""" LLM 提供商 API """ import logging from typing import List, Optional from fastapi import APIRouter, Depends, HTTPException, Query 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 LLMProvider from ai_platform.schemas.provider_schema import ( ProviderCreate, ProviderUpdate, ProviderResponse, ProviderListResponse, ProviderTypeResponse, ) from ai_platform.providers import ProviderRegistry logger = logging.getLogger(__name__) router = APIRouter(prefix="/provider", tags=["AI-提供商"]) @router.get("/types", response_model=List[ProviderTypeResponse], summary="获取提供商类型列表") async def get_provider_types(): """获取所有支持的提供商类型""" return ProviderRegistry.get_all_types() @router.get("/list", response_model=PaginatedResponse[ProviderListResponse], summary="提供商列表") async def list_providers( name: Optional[str] = Query(None, description="名称"), provider_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(LLMProvider).where(LLMProvider.is_deleted == False) if name: query = query.where(LLMProvider.name.ilike(f"%{name}%")) if provider_type: query = query.where(LLMProvider.provider_type == provider_type) if is_active is not None: query = query.where(LLMProvider.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(LLMProvider.sort.desc(), LLMProvider.sys_create_datetime.desc()) query = query.offset(offset).limit(page_size) result = await db.execute(query) items = result.scalars().all() return PaginatedResponse(items=items, total=total) @router.get("/{provider_id}", response_model=ProviderResponse, summary="提供商详情") async def get_provider(provider_id: str, db: AsyncSession = Depends(get_db)): """获取提供商详情""" result = await db.execute( select(LLMProvider).where( LLMProvider.id == provider_id, LLMProvider.is_deleted == False ) ) provider = result.scalar_one_or_none() if not provider: raise HTTPException(status_code=404, detail="提供商不存在") return _build_provider_response(provider) @router.post("", response_model=ProviderResponse, summary="创建提供商") async def create_provider(data: ProviderCreate, db: AsyncSession = Depends(get_db)): """创建提供商""" provider = LLMProvider(**data.model_dump()) db.add(provider) await db.commit() await db.refresh(provider) return _build_provider_response(provider) @router.put("/{provider_id}", response_model=ProviderResponse, summary="更新提供商") async def update_provider( provider_id: str, data: ProviderUpdate, db: AsyncSession = Depends(get_db), ): """更新提供商""" result = await db.execute( select(LLMProvider).where( LLMProvider.id == provider_id, LLMProvider.is_deleted == False ) ) provider = result.scalar_one_or_none() if not provider: raise HTTPException(status_code=404, detail="提供商不存在") update_data = data.model_dump(exclude_unset=True) if update_data.get("api_key") == "": update_data.pop("api_key") for key, value in update_data.items(): setattr(provider, key, value) await db.commit() await db.refresh(provider) return _build_provider_response(provider) @router.delete("/{provider_id}", response_model=ResponseModel, summary="删除提供商") async def delete_provider(provider_id: str, db: AsyncSession = Depends(get_db)): """删除提供商""" result = await db.execute( select(LLMProvider).where( LLMProvider.id == provider_id, LLMProvider.is_deleted == False ) ) provider = result.scalar_one_or_none() if not provider: raise HTTPException(status_code=404, detail="提供商不存在") provider.is_deleted = True await db.commit() return ResponseModel(message="删除成功") @router.post("/{provider_id}/test", summary="测试提供商连接") async def test_provider(provider_id: str, db: AsyncSession = Depends(get_db)): """测试提供商连接""" result = await db.execute( select(LLMProvider).where( LLMProvider.id == provider_id, LLMProvider.is_deleted == False ) ) provider = result.scalar_one_or_none() if not provider: raise HTTPException(status_code=404, detail="提供商不存在") # 创建提供商实例 provider_instance = ProviderRegistry.create_instance( provider_type=provider.provider_type, api_key=provider.api_key, api_base=provider.api_base, ollama_host=provider.ollama_host, ) if not provider_instance: raise HTTPException(status_code=400, detail=f"不支持的提供商类型: {provider.provider_type}") # 验证配置 if not provider_instance.validate_config(): raise HTTPException(status_code=400, detail="配置无效,请检查 API Key") # 尝试获取模型列表(如果支持) try: models = provider_instance.get_available_models() return { "success": True, "message": "连接成功", "models": models, } except Exception as e: logger.warning(f"获取模型列表失败: {e}") return { "success": True, "message": "连接成功(无法获取模型列表)", "models": [], } @router.get("/{provider_id}/default-models", summary="获取默认模型列表") async def get_default_models(provider_id: str, db: AsyncSession = Depends(get_db)): """获取提供商的默认模型列表""" result = await db.execute( select(LLMProvider).where( LLMProvider.id == provider_id, LLMProvider.is_deleted == False ) ) provider = result.scalar_one_or_none() if not provider: raise HTTPException(status_code=404, detail="提供商不存在") return ProviderRegistry.get_default_models(provider.provider_type) @router.get("/{provider_id}/fetch-models", summary="在线拉取提供商模型列表") async def fetch_provider_models(provider_id: str, db: AsyncSession = Depends(get_db)): """ 通过提供商 API 在线拉取最新模型列表。 如果在线拉取失败,自动 fallback 到硬编码的默认模型列表。 """ result = await db.execute( select(LLMProvider).where( LLMProvider.id == provider_id, LLMProvider.is_deleted == False ) ) provider = result.scalar_one_or_none() if not provider: raise HTTPException(status_code=404, detail="提供商不存在") # 构建 kwargs(Ollama 需要 ollama_host) kwargs = {} api_base = provider.api_base or '' if provider.provider_type == 'ollama' and provider.ollama_host: api_base = provider.ollama_host # 尝试在线拉取 try: models = await ProviderRegistry.fetch_models_from_api( provider_type=provider.provider_type, api_key=provider.api_key or '', api_base=api_base, **kwargs, ) except Exception as e: logger.warning( f'在线拉取模型列表异常 [provider={provider.name}, type={provider.provider_type}]: {e}' ) models = [] source = 'api' if not models: logger.warning( f'在线拉取模型列表为空,fallback 到默认列表 ' f'[provider={provider.name}, type={provider.provider_type}, api_base={api_base or "(empty)"}]' ) models = ProviderRegistry.get_default_models(provider.provider_type) source = 'default' return { "source": source, "models": models, } def _build_provider_response(provider: LLMProvider) -> dict: """构建提供商输出""" return { "id": provider.id, "name": provider.name, "provider_type": provider.provider_type, "api_key_masked": provider.get_api_key_masked(), "api_base": provider.api_base or "", "api_version": provider.api_version or "", "ollama_host": provider.ollama_host or "", "description": provider.description or "", "is_active": provider.is_active, "quota_limit": provider.quota_limit or 0, "quota_used": provider.quota_used or 0, "sort": provider.sort or 0, "sys_create_datetime": provider.sys_create_datetime, "sys_update_datetime": provider.sys_update_datetime, }