142 lines
5.3 KiB
Python
Executable File
142 lines
5.3 KiB
Python
Executable File
"""AI模型配置接口"""
|
|
import secrets
|
|
|
|
from fastapi import APIRouter, Depends, Header, HTTPException
|
|
from sqlalchemy import select, update
|
|
|
|
from app.core.database import get_db
|
|
from app.core.config import settings
|
|
from app.schemas import ApiResponse, AIModelCreateRequest, AIModelUpdateRequest, AIModelTestRequest
|
|
from app.models import AIModelConfig
|
|
from app.utils.crypto import encrypt, decrypt
|
|
from app.services.ai_service import ai_service
|
|
|
|
router = APIRouter()
|
|
|
|
|
|
@router.get("")
|
|
async def list_models(db=Depends(get_db)):
|
|
result = await db.execute(select(AIModelConfig).order_by(AIModelConfig.created_at.desc()))
|
|
models = result.scalars().all()
|
|
items = [_format_model(m) for m in models]
|
|
return ApiResponse(data=items)
|
|
|
|
|
|
@router.post("")
|
|
async def create_model(req: AIModelCreateRequest, db=Depends(get_db)):
|
|
if req.is_default:
|
|
await db.execute(
|
|
update(AIModelConfig)
|
|
.where(AIModelConfig.usage_scope == req.usage_scope)
|
|
.values(is_default=0)
|
|
)
|
|
model = AIModelConfig(
|
|
model_name=req.model_name,
|
|
provider=req.provider,
|
|
usage_scope=req.usage_scope,
|
|
api_base_url=req.api_base_url,
|
|
api_key_enc=encrypt(req.api_key) if req.api_key else None,
|
|
model_version=req.model_version,
|
|
vision_model_version=req.vision_model_version,
|
|
ocr_model_version=req.ocr_model_version,
|
|
temperature=req.temperature,
|
|
max_tokens=req.max_tokens,
|
|
timeout_seconds=req.timeout_seconds,
|
|
is_default=req.is_default,
|
|
is_enabled=1,
|
|
)
|
|
db.add(model)
|
|
await db.commit()
|
|
await db.refresh(model)
|
|
return ApiResponse(data=_format_model(model), message="模型添加成功")
|
|
|
|
|
|
@router.put("/{model_id}")
|
|
async def update_model(model_id: int, req: AIModelUpdateRequest, db=Depends(get_db)):
|
|
result = await db.execute(select(AIModelConfig).where(AIModelConfig.id == model_id))
|
|
model = result.scalar_one_or_none()
|
|
if not model:
|
|
raise HTTPException(status_code=404, detail="模型不存在")
|
|
target_scope = req.usage_scope or model.usage_scope
|
|
if req.is_default or (req.usage_scope and model.is_default):
|
|
await db.execute(
|
|
update(AIModelConfig)
|
|
.where(
|
|
AIModelConfig.id != model_id,
|
|
AIModelConfig.usage_scope == target_scope,
|
|
)
|
|
.values(is_default=0)
|
|
)
|
|
for field, val in req.model_dump(exclude_none=True).items():
|
|
if field == "api_key":
|
|
model.api_key_enc = encrypt(val) if val else None
|
|
else:
|
|
setattr(model, field, val)
|
|
await db.commit()
|
|
await db.refresh(model)
|
|
return ApiResponse(data=_format_model(model), message="更新成功")
|
|
|
|
|
|
@router.get("/runtime/digital-avatar")
|
|
async def get_digital_avatar_runtime_model(
|
|
x_avatar_config_token: str | None = Header(default=None),
|
|
db=Depends(get_db),
|
|
):
|
|
expected = settings.AVATAR_MODEL_CONFIG_TOKEN
|
|
if not expected:
|
|
raise HTTPException(status_code=503, detail="数字分身模型配置服务未启用")
|
|
if not x_avatar_config_token or not secrets.compare_digest(x_avatar_config_token, expected):
|
|
raise HTTPException(status_code=401, detail="无权读取数字分身模型配置")
|
|
|
|
result = await db.execute(
|
|
select(AIModelConfig).where(
|
|
AIModelConfig.usage_scope == "digital_avatar",
|
|
AIModelConfig.is_default == 1,
|
|
AIModelConfig.is_enabled == 1,
|
|
)
|
|
)
|
|
model = result.scalar_one_or_none()
|
|
if not model:
|
|
raise HTTPException(status_code=404, detail="尚未配置启用的数字分身专用模型")
|
|
return ApiResponse(data={
|
|
"api_base_url": model.api_base_url or "https://api.openai.com/v1",
|
|
"api_key": decrypt(model.api_key_enc) if model.api_key_enc else "",
|
|
"model": model.model_version or model.model_name,
|
|
"vision_model": model.vision_model_version or "qwen3.6-flash",
|
|
"ocr_model": model.ocr_model_version or "qwen-vl-ocr",
|
|
"temperature": model.temperature,
|
|
"max_tokens": model.max_tokens,
|
|
"timeout_seconds": model.timeout_seconds,
|
|
})
|
|
|
|
|
|
@router.delete("/{model_id}")
|
|
async def delete_model(model_id: int, db=Depends(get_db)):
|
|
result = await db.execute(select(AIModelConfig).where(AIModelConfig.id == model_id))
|
|
model = result.scalar_one_or_none()
|
|
if not model:
|
|
raise HTTPException(status_code=404, detail="模型不存在")
|
|
await db.delete(model)
|
|
await db.commit()
|
|
return ApiResponse(message="删除成功")
|
|
|
|
|
|
@router.post("/test")
|
|
async def test_model(req: AIModelTestRequest, db=Depends(get_db)):
|
|
result = await ai_service.test_model(db, req.model_id, req.test_prompt)
|
|
return ApiResponse(data=result)
|
|
|
|
|
|
def _format_model(m: AIModelConfig) -> dict:
|
|
return {
|
|
"id": m.id, "model_name": m.model_name, "provider": m.provider,
|
|
"usage_scope": m.usage_scope,
|
|
"api_base_url": m.api_base_url, "has_api_key": bool(m.api_key_enc),
|
|
"model_version": m.model_version, "temperature": m.temperature,
|
|
"vision_model_version": m.vision_model_version,
|
|
"ocr_model_version": m.ocr_model_version,
|
|
"max_tokens": m.max_tokens, "timeout_seconds": m.timeout_seconds,
|
|
"is_default": m.is_default, "is_enabled": m.is_enabled,
|
|
"created_at": m.created_at.isoformat(),
|
|
}
|