Files
ai-agent-admin/backend-fastapi/ai_platform/api/provider_api.py
T

378 lines
13 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
LLM 提供商 API
"""
import logging
from typing import Any, 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-提供商"])
def _pick_message_from_payload(payload: Any) -> Optional[str]:
if isinstance(payload, dict):
for key in ('message', 'detail', 'error_description'):
value = payload.get(key)
if isinstance(value, str) and value:
return value
error = payload.get('error')
if isinstance(error, str) and error:
return error
if isinstance(error, dict):
return _pick_message_from_payload(error)
if isinstance(payload, list):
for item in payload:
message = _pick_message_from_payload(item)
if message:
return message
return None
def _extract_upstream_error(exc: Exception) -> str:
response = getattr(exc, 'response', None)
if response is not None:
status_code = getattr(response, 'status_code', None)
try:
message = _pick_message_from_payload(response.json())
except Exception:
message = None
if not message:
message = (getattr(response, 'text', '') or '').strip()
if status_code and message:
return f'HTTP {status_code}: {message[:240]}'
if status_code:
return f'HTTP {status_code}'
if message:
return message[:240]
return str(exc)[:240]
def _format_provider_test_error(exc: Exception, provider: LLMProvider, api_base: str) -> str:
error_text = f'{type(exc).__name__}: {str(exc)} {_extract_upstream_error(exc)}'
lower_text = error_text.lower()
if (
'401' in lower_text
or '403' in lower_text
or 'unauthorized' in lower_text
or 'forbidden' in lower_text
or 'invalid api key' in lower_text
or 'incorrect api key' in lower_text
or '无效的令牌' in error_text
or '鉴权' in error_text
):
reason = '上游模型鉴权失败,API Key 无效或已过期'
elif '404' in lower_text or 'not found' in lower_text:
reason = '模型接口不存在或 Base URL 路径不正确'
elif 'timeout' in lower_text or 'timed out' in lower_text or '超时' in error_text:
reason = '上游模型接口请求超时'
elif (
'connect' in lower_text
or 'network' in lower_text
or 'name or service not known' in lower_text
or 'connection' in lower_text
):
reason = '无法连接到上游模型接口'
elif '429' in lower_text or 'rate limit' in lower_text or 'quota' in lower_text:
reason = '上游模型限流或额度不足'
else:
reason = '上游模型接口调用失败'
upstream_error = _extract_upstream_error(exc) or ''
return (
f'{reason}:提供商 {provider.name},类型 {provider.provider_type}'
f'Base URL {api_base or "未配置"}'
f'请检查 Base URL、API Key 和上游账号权限。'
f'上游返回:{upstream_error}'
)
@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="提供商不存在")
api_base = provider.api_base or ''
instance_kwargs = {}
if provider.provider_type == 'ollama' and provider.ollama_host:
api_base = provider.ollama_host
instance_kwargs['ollama_host'] = provider.ollama_host
provider_instance = ProviderRegistry.create_instance(
provider_type=provider.provider_type,
api_key=provider.api_key,
api_base=api_base,
**instance_kwargs,
)
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 = await provider_instance.fetch_models_from_api_strict()
except Exception as e:
logger.warning(
"Provider online test failed: upstream error "
f"[provider={provider.name}, type={provider.provider_type}, api_base={api_base or '(empty)'}]: {e}"
)
raise HTTPException(
status_code=400,
detail=_format_provider_test_error(e, provider, api_base),
)
if not models:
logger.warning(
"Provider online test failed: no models returned "
f"[provider={provider.name}, type={provider.provider_type}, api_base={api_base or '(empty)'}]"
)
raise HTTPException(
status_code=400,
detail=(
f"连接失败:模型接口返回为空。提供商 {provider.name}"
f"类型 {provider.provider_type}Base URL {api_base or '未配置'}"
"请检查 Base URL、API Key 和上游账号权限。"
),
)
return {
"success": True,
"message": f"连接成功,已从模型接口拉取到 {len(models)} 个模型",
"source": "api",
"model_count": len(models),
"models": models[:20],
}
@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="提供商不存在")
# 构建 kwargsOllama 需要 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,
}