116 lines
4.1 KiB
Python
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)}
|