378 lines
13 KiB
Python
378 lines
13 KiB
Python
"""
|
||
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="提供商不存在")
|
||
|
||
# 构建 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,
|
||
}
|