""" LLM 模型 API """ import logging from typing import List, Optional from fastapi import APIRouter, Depends, HTTPException, Query, Body from sqlalchemy import select, func from sqlalchemy.ext.asyncio import AsyncSession from app.database import get_db from app.base_schema import PaginatedResponse, ResponseModel from ai_platform.models import LLMModel, LLMProvider from ai_platform.schemas.model_schema import ( ModelCreate, ModelUpdate, ModelResponse, ModelListResponse, ) logger = logging.getLogger(__name__) router = APIRouter(prefix="/model", tags=["AI-模型"]) @router.get("/list", response_model=PaginatedResponse[ModelListResponse], summary="模型列表") async def list_models( provider_id: Optional[str] = Query(None, description="提供商 ID"), model_type: Optional[str] = Query(None, description="模型类型"), is_active: Optional[bool] = Query(None, description="是否启用"), page: int = Query(1, ge=1, description="页码"), page_size: int = Query(20, ge=1, le=100, alias="pageSize", description="每页数量"), db: AsyncSession = Depends(get_db), ): """获取模型列表""" query = select(LLMModel).where(LLMModel.is_deleted == False) if provider_id: query = query.where(LLMModel.provider_id == provider_id) if model_type: query = query.where(LLMModel.model_type == model_type) if is_active is not None: query = query.where(LLMModel.is_active == is_active) # 获取总数 count_result = await db.execute(select(func.count()).select_from(query.subquery())) total = count_result.scalar() or 0 # 分页 offset = (page - 1) * page_size query = query.order_by(LLMModel.sort.desc(), LLMModel.sys_create_datetime.desc()) query = query.offset(offset).limit(page_size) result = await db.execute(query) models = result.scalars().all() # 获取提供商名称 items = [] for model in models: provider_result = await db.execute( select(LLMProvider).where(LLMProvider.id == model.provider_id) ) provider = provider_result.scalar_one_or_none() items.append({ "id": model.id, "provider_id": model.provider_id, "provider_name": provider.name if provider else "", "model_name": model.model_name, "display_name": model.display_name, "model_type": model.model_type or "chat", "is_active": model.is_active, }) return PaginatedResponse(items=items, total=total) @router.get("/active", response_model=List[ModelListResponse], summary="获取可用模型列表") async def list_active_models(db: AsyncSession = Depends(get_db)): """获取所有可用的模型(用于选择器)""" query = select(LLMModel).where( LLMModel.is_deleted == False, LLMModel.is_active == True ).order_by(LLMModel.display_name) result = await db.execute(query) models = result.scalars().all() items = [] for model in models: provider_result = await db.execute( select(LLMProvider).where( LLMProvider.id == model.provider_id, LLMProvider.is_deleted == False, LLMProvider.is_active == True ) ) provider = provider_result.scalar_one_or_none() if provider: items.append({ "id": model.id, "provider_id": model.provider_id, "provider_name": provider.name, "model_name": model.model_name, "display_name": model.display_name, "model_type": model.model_type or "chat", "is_active": model.is_active, }) return items @router.get("/{model_id}", response_model=ModelResponse, summary="模型详情") async def get_model(model_id: str, db: AsyncSession = Depends(get_db)): """获取模型详情""" result = await db.execute( select(LLMModel).where( LLMModel.id == model_id, LLMModel.is_deleted == False ) ) model = result.scalar_one_or_none() if not model: raise HTTPException(status_code=404, detail="模型不存在") # 获取提供商名称 provider_result = await db.execute( select(LLMProvider).where(LLMProvider.id == model.provider_id) ) provider = provider_result.scalar_one_or_none() return _build_model_response(model, provider) @router.post("", response_model=ModelResponse, summary="创建模型") async def create_model(data: ModelCreate, db: AsyncSession = Depends(get_db)): """创建模型""" # 验证提供商 provider_result = await db.execute( select(LLMProvider).where( LLMProvider.id == data.provider_id, LLMProvider.is_deleted == False ) ) provider = provider_result.scalar_one_or_none() if not provider: raise HTTPException(status_code=400, detail="提供商不存在") # 检查模型名称是否重复 exists_result = await db.execute( select(LLMModel).where( LLMModel.provider_id == data.provider_id, LLMModel.model_name == data.model_name, LLMModel.is_deleted == False ) ) if exists_result.scalar_one_or_none(): raise HTTPException(status_code=400, detail=f"模型 {data.model_name} 已存在") model = LLMModel(**data.model_dump()) db.add(model) await db.commit() await db.refresh(model) return _build_model_response(model, provider) @router.post("/batch", response_model=ResponseModel, summary="批量创建模型") async def batch_create_models( provider_id: str = Query(..., description="提供商 ID"), models: List[dict] = Body(..., description="模型列表"), db: AsyncSession = Depends(get_db), ): """批量创建模型(从默认模型列表)""" provider_result = await db.execute( select(LLMProvider).where( LLMProvider.id == provider_id, LLMProvider.is_deleted == False ) ) provider = provider_result.scalar_one_or_none() if not provider: raise HTTPException(status_code=400, detail="提供商不存在") created_count = 0 for model_data in models: model_name = model_data.get("model_name", "") if not model_name: continue # 跳过已存在的 exists_result = await db.execute( select(LLMModel).where( LLMModel.provider_id == provider_id, LLMModel.model_name == model_name, LLMModel.is_deleted == False ) ) if exists_result.scalar_one_or_none(): continue model = LLMModel( provider_id=provider_id, model_name=model_name, display_name=model_data.get("display_name", model_name), model_type=model_data.get("model_type", "chat"), max_tokens=model_data.get("max_tokens", 4096), context_window=model_data.get("context_window", 4096), supports_vision=model_data.get("supports_vision", False), supports_function_call=model_data.get("supports_function_call", False), input_price=model_data.get("input_price", 0), output_price=model_data.get("output_price", 0), is_active=True, ) db.add(model) created_count += 1 await db.commit() return ResponseModel(message=f"成功创建 {created_count} 个模型") @router.put("/{model_id}", response_model=ModelResponse, summary="更新模型") async def update_model( model_id: str, data: ModelUpdate, db: AsyncSession = Depends(get_db), ): """更新模型""" result = await db.execute( select(LLMModel).where( LLMModel.id == model_id, LLMModel.is_deleted == False ) ) model = result.scalar_one_or_none() if not model: raise HTTPException(status_code=404, detail="模型不存在") update_data = data.model_dump(exclude_unset=True) for key, value in update_data.items(): setattr(model, key, value) await db.commit() await db.refresh(model) # 获取提供商名称 provider_result = await db.execute( select(LLMProvider).where(LLMProvider.id == model.provider_id) ) provider = provider_result.scalar_one_or_none() return _build_model_response(model, provider) @router.delete("/{model_id}", response_model=ResponseModel, summary="删除模型") async def delete_model(model_id: str, db: AsyncSession = Depends(get_db)): """删除模型""" result = await db.execute( select(LLMModel).where( LLMModel.id == model_id, LLMModel.is_deleted == False ) ) model = result.scalar_one_or_none() if not model: raise HTTPException(status_code=404, detail="模型不存在") model.is_deleted = True await db.commit() return ResponseModel(message="删除成功") def _build_model_response(model: LLMModel, provider: Optional[LLMProvider] = None) -> dict: """构建模型输出""" return { "id": model.id, "provider_id": model.provider_id, "provider_name": provider.name if provider else "", "model_name": model.model_name, "display_name": model.display_name, "model_type": model.model_type or "chat", "max_tokens": model.max_tokens or 4096, "context_window": model.context_window or 4096, "default_temperature": model.default_temperature or 0.7, "default_top_p": model.default_top_p or 1.0, "input_price": model.input_price or 0, "output_price": model.output_price or 0, "supports_vision": model.supports_vision or False, "supports_function_call": model.supports_function_call or False, "supports_streaming": model.supports_streaming if model.supports_streaming is not None else True, "is_active": model.is_active, "sort": model.sort or 0, "sys_create_datetime": model.sys_create_datetime, "sys_update_datetime": model.sys_update_datetime, }