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