Files
2026-07-08 20:02:29 +08:00

116 lines
4.1 KiB
Python

from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select
from app.core.database import get_db
from app.core.security import encrypt_api_key
from app.models.ai_model import AIModel
from app.schemas.ai_model import AIModelCreate, AIModelUpdate, AIModelToggle, AIModelResponse
from app.services.ai_adapter import call_ai_model
router = APIRouter(prefix="/models", tags=["AI模型管理"])
@router.post("", response_model=AIModelResponse)
async def create_model(data: AIModelCreate, db: AsyncSession = Depends(get_db)):
encrypted_key = encrypt_api_key(data.api_key)
model = AIModel(
name=data.name,
provider=data.provider,
endpoint=data.endpoint,
api_key=encrypted_key,
extra_params=data.extra_params,
is_enabled=data.is_enabled,
remark=data.remark,
)
db.add(model)
await db.flush()
await db.refresh(model)
return model
@router.get("", response_model=list[AIModelResponse])
async def list_models(
enabled: bool | None = Query(None, description="过滤启用/禁用"),
skip: int = Query(0, ge=0),
limit: int = Query(20, ge=1, le=100),
db: AsyncSession = Depends(get_db),
):
query = select(AIModel)
if enabled is not None:
query = query.where(AIModel.is_enabled == enabled)
query = query.offset(skip).limit(limit).order_by(AIModel.created_at.desc())
result = await db.execute(query)
models = result.scalars().all()
return models
@router.get("/{model_id}", response_model=AIModelResponse)
async def get_model(model_id: str, db: AsyncSession = Depends(get_db)):
result = await db.execute(select(AIModel).where(AIModel.id == model_id))
model = result.scalar_one_or_none()
if not model:
raise HTTPException(status_code=404, detail="模型不存在")
return model
@router.put("/{model_id}", response_model=AIModelResponse)
async def update_model(model_id: str, data: AIModelUpdate, db: AsyncSession = Depends(get_db)):
result = await db.execute(select(AIModel).where(AIModel.id == model_id))
model = result.scalar_one_or_none()
if not model:
raise HTTPException(status_code=404, detail="模型不存在")
update_data = data.model_dump(exclude_unset=True)
if "api_key" in update_data and update_data["api_key"] is not None:
update_data["api_key"] = encrypt_api_key(update_data["api_key"])
for key, value in update_data.items():
setattr(model, key, value)
await db.flush()
await db.refresh(model)
return model
@router.delete("/{model_id}")
async def delete_model(model_id: str, db: AsyncSession = Depends(get_db)):
result = await db.execute(select(AIModel).where(AIModel.id == model_id))
model = result.scalar_one_or_none()
if not model:
raise HTTPException(status_code=404, detail="模型不存在")
await db.delete(model)
return {"detail": "删除成功"}
@router.patch("/{model_id}/toggle", response_model=AIModelResponse)
async def toggle_model(model_id: str, data: AIModelToggle, db: AsyncSession = Depends(get_db)):
result = await db.execute(select(AIModel).where(AIModel.id == model_id))
model = result.scalar_one_or_none()
if not model:
raise HTTPException(status_code=404, detail="模型不存在")
model.is_enabled = data.is_enabled
await db.flush()
await db.refresh(model)
return model
@router.post("/{model_id}/test")
async def test_model(model_id: str, db: AsyncSession = Depends(get_db)):
result = await db.execute(select(AIModel).where(AIModel.id == model_id))
model = result.scalar_one_or_none()
if not model:
raise HTTPException(status_code=404, detail="模型不存在")
model_config = {
"provider": model.provider,
"endpoint": model.endpoint,
"api_key": model.api_key,
"extra_params": model.extra_params,
}
try:
response = await call_ai_model(model_config, "请用一句话介绍你自己。")
return {"success": True, "result": response}
except Exception as e:
return {"success": False, "error": str(e)}