67 lines
2.1 KiB
Python
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()
|