nex_math/backend/routers/llm.py

319 lines
9.8 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

"""大模型多通道配置: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,
)