"""大模型多通道配置:OpenAI / DeepSeek / 阿里千问 / OpenAI 兼容。 每个通道保存一份独立的 provider / base_url / model / api_key / temperature, 通道名建议为“服务商 · 模型”,例如“阿里千问(百炼) · qwen-plus”; 生成题目时默认使用 is_default 通道,也可在调用方显式指定 channel_id。 """ from __future__ import annotations import json from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException from sqlalchemy.orm import Session from database import get_db from dependencies import require_permission from models import LlmSetting, LlmTask, User from schemas import ( LlmChannelCreate, LlmChannelOut, LlmChannelUpdate, LlmProviderOut, LlmSettingsOut, LlmTestResult, LlmTaskCreate, LlmTaskOut, ) from services.llm_client import PROVIDER_PRESETS, chat_completion, mask_key from services.llm_config import ( default_channel_name, ensure_llm_channels, get_default_channel, promote_default, provider_label, ) from services.task_runner import run_llm_task router = APIRouter(prefix="/llm", tags=["llm"]) def _providers() -> list[LlmProviderOut]: return [ LlmProviderOut( id=provider_id, label=preset["label"], base_url=preset["base_url"], model=preset["model"], ) for provider_id, preset in PROVIDER_PRESETS.items() ] def _channel_out(channel: LlmSetting) -> LlmChannelOut: return LlmChannelOut( id=channel.id, name=channel.name or default_channel_name(channel.provider, channel.model), provider=channel.provider, provider_label=provider_label(channel.provider), base_url=channel.base_url, model=channel.model, temperature=channel.temperature, max_tokens=channel.max_tokens, has_api_key=bool(channel.api_key), api_key_preview=mask_key(channel.api_key), is_default=bool(channel.is_default), ) def _settings_out(db: Session) -> LlmSettingsOut: channels = ( db.query(LlmSetting) .order_by(LlmSetting.is_default.desc(), LlmSetting.id.asc()) .all() ) return LlmSettingsOut( channels=[_channel_out(channel) for channel in channels], providers=_providers(), ) def _load_channel(db: Session, channel_id: int) -> LlmSetting: channel = db.get(LlmSetting, channel_id) if channel is None: raise HTTPException(status_code=404, detail="模型通道不存在") return channel def _ensure_provider(provider: str) -> dict: preset = PROVIDER_PRESETS.get(provider) if preset is None: raise HTTPException(status_code=400, detail="不支持的模型服务商") return preset def _ensure_unique_name(db: Session, name: str, exclude_id: int | None = None) -> None: query = db.query(LlmSetting).filter(LlmSetting.name == name) if exclude_id is not None: query = query.filter(LlmSetting.id != exclude_id) if query.first() is not None: raise HTTPException(status_code=400, detail="通道名称已存在,请换一个名称") def _preset_fill(provider: str, base_url: str, model: str) -> tuple[str, str]: preset = _ensure_provider(provider) return ( base_url.strip() or preset["base_url"], model.strip() or preset["model"], ) @router.get("/settings", response_model=LlmSettingsOut) def get_settings( db: Session = Depends(get_db), _: User = Depends(require_permission("llm:manage")), ): ensure_llm_channels(db) return _settings_out(db) @router.post("/channels", response_model=LlmChannelOut) def create_channel( payload: LlmChannelCreate, db: Session = Depends(get_db), _: User = Depends(require_permission("llm:manage")), ): _ensure_provider(payload.provider) name = payload.name.strip() or default_channel_name(payload.provider, payload.model) _ensure_unique_name(db, name) base_url, model = _preset_fill(payload.provider, payload.base_url, payload.model) count = db.query(LlmSetting).count() channel = LlmSetting( name=name, provider=payload.provider, base_url=base_url, model=model, temperature=payload.temperature, max_tokens=payload.max_tokens, api_key=(payload.api_key or "").strip(), is_default=count == 0, ) db.add(channel) db.commit() db.refresh(channel) return _channel_out(channel) @router.put("/channels/{channel_id}", response_model=LlmChannelOut) def update_channel( channel_id: int, payload: LlmChannelUpdate, db: Session = Depends(get_db), _: User = Depends(require_permission("llm:manage")), ): channel = _load_channel(db, channel_id) provider = payload.provider or channel.provider _ensure_provider(provider) base_url, model = _preset_fill( provider, payload.base_url if payload.base_url is not None else channel.base_url, payload.model if payload.model is not None else channel.model, ) if payload.name is not None and payload.name.strip(): name = payload.name.strip() else: name = channel.name or default_channel_name(provider, model) _ensure_unique_name(db, name, exclude_id=channel.id) channel.name = name channel.provider = provider channel.base_url = base_url channel.model = model if payload.temperature is not None: channel.temperature = payload.temperature if payload.clear_max_tokens: channel.max_tokens = None elif payload.max_tokens is not None: channel.max_tokens = payload.max_tokens if payload.clear_api_key: channel.api_key = "" elif payload.api_key and payload.api_key.strip(): channel.api_key = payload.api_key.strip() db.commit() return _channel_out(channel) @router.delete("/channels/{channel_id}") def delete_channel( channel_id: int, db: Session = Depends(get_db), _: User = Depends(require_permission("llm:manage")), ): channel = _load_channel(db, channel_id) was_default = bool(channel.is_default) db.delete(channel) db.flush() if was_default: promote_default(db, None) else: db.commit() return {"deleted": channel_id} @router.put("/channels/{channel_id}/default", response_model=LlmSettingsOut) def set_default_channel( channel_id: int, db: Session = Depends(get_db), _: User = Depends(require_permission("llm:manage")), ): _load_channel(db, channel_id) promote_default(db, channel_id) return _settings_out(db) def _run_test(channel: LlmSetting) -> str: if not channel.base_url or not channel.model: raise HTTPException(status_code=400, detail="通道缺少 Base URL 或模型名称,请先补全") if not channel.api_key: raise HTTPException(status_code=400, detail="通道尚未保存 API Key,请先配置") try: return chat_completion( base_url=channel.base_url, api_key=channel.api_key, model=channel.model, temperature=0.0, max_tokens=256, timeout=45.0, messages=[ { "role": "system", "content": "你是连通性测试助手,只回复两个字:正常", }, {"role": "user", "content": "请确认服务可用。"}, ], ) except ValueError as exc: raise HTTPException(status_code=400, detail=f"测试失败:{exc}") from exc @router.post("/channels/{channel_id}/test", response_model=LlmTestResult) def test_channel( channel_id: int, db: Session = Depends(get_db), _: User = Depends(require_permission("llm:manage")), ): channel = _load_channel(db, channel_id) reply = _run_test(channel) return LlmTestResult(ok=True, message=f"连接成功,模型已回复:“{reply[:40]}”") @router.post("/test", response_model=LlmTestResult) def test_default_channel( db: Session = Depends(get_db), _: User = Depends(require_permission("llm:manage")), ): """兼容旧前端:测试默认通道。""" ensure_llm_channels(db) channel = get_default_channel(db) if channel is None: raise HTTPException(status_code=400, detail="尚未配置任何模型通道") reply = _run_test(channel) return LlmTestResult(ok=True, message=f"连接成功,模型已回复:“{reply[:40]}”") @router.post("/tasks", response_model=LlmTaskOut, status_code=201) def create_task( payload: LlmTaskCreate, background_tasks: BackgroundTasks, db: Session = Depends(get_db), user: User = Depends(require_permission("llm:manage")), ): if payload.kind not in { "generate_questions", "generate_chapters", "organize_chapters", }: raise HTTPException(status_code=400, detail="不支持的任务类型") task = LlmTask( kind=payload.kind, params=json.dumps(payload.params, ensure_ascii=False), created_by=user.id, ) db.add(task) db.commit() db.refresh(task) background_tasks.add_task( run_llm_task, task.id, task.kind, payload.params ) return LlmTaskOut( id=task.id, kind=task.kind, status=task.status, progress=task.progress, message=task.message, result={}, ) @router.get("/tasks/{task_id}", response_model=LlmTaskOut) def get_task( task_id: int, db: Session = Depends(get_db), _: User = Depends(require_permission("llm:manage")), ): task = db.get(LlmTask, task_id) if task is None: raise HTTPException(status_code=404, detail="任务不存在") try: result = json.loads(task.result or "{}") except ValueError: result = {} return LlmTaskOut( id=task.id, kind=task.kind, status=task.status, progress=task.progress, message=task.message, result=result, )