nex_math/backend/services/llm_config.py

67 lines
2.1 KiB
Python

"""模型通道(多通道)配置辅助:默认命名、幂等初始化与默认通道管理。"""
from __future__ import annotations
import os
from sqlalchemy.orm import Session
from models import LlmSetting
from services.llm_client import PROVIDER_PRESETS
def provider_label(provider: str) -> str:
preset = PROVIDER_PRESETS.get(provider)
return preset["label"] if preset else provider
def default_channel_name(provider: str, model: str) -> str:
label = provider_label(provider)
return f"{label} · {model}" if model else label
def ensure_llm_channels(db: Session) -> None:
"""老库升级:给历史单通道配置补名字并保证存在默认通道。"""
rows = db.query(LlmSetting).order_by(LlmSetting.id.asc()).all()
if not rows:
api_key = os.getenv("OPENAI_API_KEY", "").strip()
if api_key:
preset = PROVIDER_PRESETS["openai"]
db.add(
LlmSetting(
name=default_channel_name("openai", preset["model"]),
provider="openai",
base_url=preset["base_url"],
model=preset["model"],
temperature=0.3,
api_key=api_key,
is_default=True,
)
)
db.commit()
return
for row in rows:
if not row.name:
row.name = default_channel_name(row.provider, row.model)
if not any(row.is_default for row in rows):
rows[0].is_default = True
db.commit()
def get_default_channel(db: Session) -> LlmSetting | None:
return (
db.query(LlmSetting)
.filter(LlmSetting.is_default.is_(True))
.order_by(LlmSetting.id.asc())
.first()
)
def promote_default(db: Session, channel_id: int | None = None) -> None:
"""把指定通道设为默认;未指定时选择第一个可用通道。"""
rows = db.query(LlmSetting).order_by(LlmSetting.id.asc()).all()
for row in rows:
row.is_default = channel_id is not None and row.id == channel_id
if channel_id is None and rows:
rows[0].is_default = True
db.commit()