mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
refactor: 移除 Python 后端源码,全面迁移至 Rust gateway 架构
- 删除全部 Python 源码 (src/) 及 Alembic 迁移脚本,归档至 _deprecated_py_src/ - 重构 Rust gateway ai_pipeline: 拆分 planner/finalize 模块,新增 contracts/adaptation 层 - 重组 handlers 模块为 admin/public/proxy/internal/shared 子模块结构 - 新增 executor 模块,引入 Rust 原生数据库迁移 (aether-data/migrations) - 简化 CI/Docker 构建流程,移除 base image 二级构建,统一为单一 app image - 移除 Python 相关基础设施文件 (entrypoint.sh, gunicorn_conf.py, Dockerfile.base)
This commit is contained in:
20
_deprecated_py_src/api/admin/providers/__init__.py
Normal file
20
_deprecated_py_src/api/admin/providers/__init__.py
Normal file
@@ -0,0 +1,20 @@
|
||||
"""Provider admin routes export."""
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
from .models import router as models_router
|
||||
from .routes import router as routes_router
|
||||
from .summary import router as summary_router
|
||||
|
||||
router = APIRouter(prefix="/api/admin/providers", tags=["Admin - Providers"])
|
||||
|
||||
# Provider CRUD
|
||||
router.include_router(routes_router)
|
||||
|
||||
# Provider summary & health monitor
|
||||
router.include_router(summary_router)
|
||||
|
||||
# Provider models management
|
||||
router.include_router(models_router)
|
||||
|
||||
__all__ = ["router"]
|
||||
806
_deprecated_py_src/api/admin/providers/models.py
Normal file
806
_deprecated_py_src/api/admin/providers/models.py
Normal file
@@ -0,0 +1,806 @@
|
||||
"""
|
||||
Provider 模型管理 API
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from sqlalchemy.orm import Session, joinedload
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.models_service import invalidate_models_list_cache
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.exceptions import InvalidRequestException, NotFoundException
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
from src.models.api import (
|
||||
ModelCreate,
|
||||
ModelResponse,
|
||||
ModelUpdate,
|
||||
)
|
||||
from src.models.database import (
|
||||
GlobalModel,
|
||||
Model,
|
||||
Provider,
|
||||
)
|
||||
from src.models.pydantic_models import (
|
||||
BatchAssignModelsToProviderRequest,
|
||||
BatchAssignModelsToProviderResponse,
|
||||
ImportFromUpstreamErrorItem,
|
||||
ImportFromUpstreamRequest,
|
||||
ImportFromUpstreamResponse,
|
||||
ImportFromUpstreamSuccessItem,
|
||||
ProviderAvailableSourceModel,
|
||||
ProviderAvailableSourceModelsResponse,
|
||||
)
|
||||
from src.services.model.service import ModelService
|
||||
|
||||
router = APIRouter(tags=["Model Management"])
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
@router.get("/{provider_id}/models", response_model=list[ModelResponse])
|
||||
async def list_provider_models(
|
||||
provider_id: str,
|
||||
request: Request,
|
||||
is_active: bool | None = None,
|
||||
skip: int = 0,
|
||||
limit: int = 100,
|
||||
db: Session = Depends(get_db),
|
||||
) -> list[ModelResponse]:
|
||||
"""
|
||||
获取提供商的所有模型
|
||||
|
||||
获取指定提供商的模型列表,支持分页和状态过滤。
|
||||
|
||||
**路径参数**:
|
||||
- `provider_id`: 提供商 ID
|
||||
|
||||
**查询参数**:
|
||||
- `is_active`: 可选的活跃状态过滤,true 仅返回活跃模型,false 返回禁用模型,不传则返回全部
|
||||
- `skip`: 跳过的记录数,默认为 0
|
||||
- `limit`: 返回的最大记录数,默认为 100
|
||||
|
||||
**返回字段**(数组,每项包含):
|
||||
- `id`: 模型 ID
|
||||
- `provider_id`: 提供商 ID
|
||||
- `global_model_id`: 全局模型 ID
|
||||
- `provider_model_name`: 提供商模型名称
|
||||
- `is_active`: 是否启用
|
||||
- `input_price_per_1m`: 输入价格(每百万 token)
|
||||
- `output_price_per_1m`: 输出价格(每百万 token)
|
||||
- `cache_creation_price_per_1m`: 缓存创建价格(每百万 token)
|
||||
- `cache_read_price_per_1m`: 缓存读取价格(每百万 token)
|
||||
- `price_per_request`: 每次请求价格
|
||||
- `supports_vision`: 是否支持视觉
|
||||
- `supports_function_calling`: 是否支持函数调用
|
||||
- `supports_streaming`: 是否支持流式输出
|
||||
- `created_at`: 创建时间
|
||||
- `updated_at`: 更新时间
|
||||
"""
|
||||
adapter = AdminListProviderModelsAdapter(
|
||||
provider_id=provider_id,
|
||||
is_active=is_active,
|
||||
skip=skip,
|
||||
limit=limit,
|
||||
)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.post("/{provider_id}/models", response_model=ModelResponse)
|
||||
async def create_provider_model(
|
||||
provider_id: str,
|
||||
model_data: ModelCreate,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> ModelResponse:
|
||||
"""
|
||||
创建模型
|
||||
|
||||
为指定提供商创建一个新的模型配置。
|
||||
|
||||
**路径参数**:
|
||||
- `provider_id`: 提供商 ID
|
||||
|
||||
**请求体字段**:
|
||||
- `provider_model_name`: 提供商模型名称(必填)
|
||||
- `global_model_id`: 全局模型 ID(可选,关联到全局模型)
|
||||
- `is_active`: 是否启用(默认 true)
|
||||
- `input_price_per_1m`: 输入价格(每百万 token)(可选)
|
||||
- `output_price_per_1m`: 输出价格(每百万 token)(可选)
|
||||
- `cache_creation_price_per_1m`: 缓存创建价格(每百万 token)(可选)
|
||||
- `cache_read_price_per_1m`: 缓存读取价格(每百万 token)(可选)
|
||||
- `price_per_request`: 每次请求价格(可选)
|
||||
- `supports_vision`: 是否支持视觉(可选)
|
||||
- `supports_function_calling`: 是否支持函数调用(可选)
|
||||
- `supports_streaming`: 是否支持流式输出(可选)
|
||||
|
||||
**返回字段**: 返回创建的模型详细信息(与 GET 单个模型接口返回格式相同)
|
||||
"""
|
||||
adapter = AdminCreateProviderModelAdapter(provider_id=provider_id, model_data=model_data)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get("/{provider_id}/models/{model_id}", response_model=ModelResponse)
|
||||
async def get_provider_model(
|
||||
provider_id: str,
|
||||
model_id: str,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> ModelResponse:
|
||||
"""
|
||||
获取模型详情
|
||||
|
||||
获取指定模型的详细配置信息。
|
||||
|
||||
**路径参数**:
|
||||
- `provider_id`: 提供商 ID
|
||||
- `model_id`: 模型 ID
|
||||
|
||||
**返回字段**:
|
||||
- `id`: 模型 ID
|
||||
- `provider_id`: 提供商 ID
|
||||
- `global_model_id`: 全局模型 ID
|
||||
- `provider_model_name`: 提供商模型名称
|
||||
- `is_active`: 是否启用
|
||||
- `input_price_per_1m`: 输入价格(每百万 token)
|
||||
- `output_price_per_1m`: 输出价格(每百万 token)
|
||||
- `cache_creation_price_per_1m`: 缓存创建价格(每百万 token)
|
||||
- `cache_read_price_per_1m`: 缓存读取价格(每百万 token)
|
||||
- `price_per_request`: 每次请求价格
|
||||
- `supports_vision`: 是否支持视觉
|
||||
- `supports_function_calling`: 是否支持函数调用
|
||||
- `supports_streaming`: 是否支持流式输出
|
||||
- `created_at`: 创建时间
|
||||
- `updated_at`: 更新时间
|
||||
"""
|
||||
adapter = AdminGetProviderModelAdapter(provider_id=provider_id, model_id=model_id)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.patch("/{provider_id}/models/{model_id}", response_model=ModelResponse)
|
||||
async def update_provider_model(
|
||||
provider_id: str,
|
||||
model_id: str,
|
||||
model_data: ModelUpdate,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> ModelResponse:
|
||||
"""
|
||||
更新模型配置
|
||||
|
||||
更新指定模型的配置信息。只需传入需要更新的字段,未传入的字段保持不变。
|
||||
|
||||
**路径参数**:
|
||||
- `provider_id`: 提供商 ID
|
||||
- `model_id`: 模型 ID
|
||||
|
||||
**请求体字段**(所有字段可选):
|
||||
- `provider_model_name`: 提供商模型名称
|
||||
- `global_model_id`: 全局模型 ID
|
||||
- `is_active`: 是否启用
|
||||
- `input_price_per_1m`: 输入价格(每百万 token)
|
||||
- `output_price_per_1m`: 输出价格(每百万 token)
|
||||
- `cache_creation_price_per_1m`: 缓存创建价格(每百万 token)
|
||||
- `cache_read_price_per_1m`: 缓存读取价格(每百万 token)
|
||||
- `price_per_request`: 每次请求价格
|
||||
- `supports_vision`: 是否支持视觉
|
||||
- `supports_function_calling`: 是否支持函数调用
|
||||
- `supports_streaming`: 是否支持流式输出
|
||||
|
||||
**返回字段**: 返回更新后的模型详细信息(与 GET 单个模型接口返回格式相同)
|
||||
"""
|
||||
adapter = AdminUpdateProviderModelAdapter(
|
||||
provider_id=provider_id,
|
||||
model_id=model_id,
|
||||
model_data=model_data,
|
||||
)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.delete("/{provider_id}/models/{model_id}")
|
||||
async def delete_provider_model(
|
||||
provider_id: str,
|
||||
model_id: str,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> Any:
|
||||
"""
|
||||
删除模型
|
||||
|
||||
删除指定的模型配置。注意:此操作不可逆。
|
||||
|
||||
**路径参数**:
|
||||
- `provider_id`: 提供商 ID
|
||||
- `model_id`: 模型 ID
|
||||
|
||||
**返回字段**:
|
||||
- `message`: 删除成功提示信息
|
||||
"""
|
||||
adapter = AdminDeleteProviderModelAdapter(provider_id=provider_id, model_id=model_id)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.post("/{provider_id}/models/batch", response_model=list[ModelResponse])
|
||||
async def batch_create_provider_models(
|
||||
provider_id: str,
|
||||
models_data: list[ModelCreate],
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> list[ModelResponse]:
|
||||
"""
|
||||
批量创建模型
|
||||
|
||||
为指定提供商批量创建多个模型配置。
|
||||
|
||||
**路径参数**:
|
||||
- `provider_id`: 提供商 ID
|
||||
|
||||
**请求体**: 模型数据数组,每项包含:
|
||||
- `provider_model_name`: 提供商模型名称(必填)
|
||||
- `global_model_id`: 全局模型 ID(可选)
|
||||
- `is_active`: 是否启用(默认 true)
|
||||
- `input_price_per_1m`: 输入价格(每百万 token)(可选)
|
||||
- `output_price_per_1m`: 输出价格(每百万 token)(可选)
|
||||
- `cache_creation_price_per_1m`: 缓存创建价格(每百万 token)(可选)
|
||||
- `cache_read_price_per_1m`: 缓存读取价格(每百万 token)(可选)
|
||||
- `price_per_request`: 每次请求价格(可选)
|
||||
- `supports_vision`: 是否支持视觉(可选)
|
||||
- `supports_function_calling`: 是否支持函数调用(可选)
|
||||
- `supports_streaming`: 是否支持流式输出(可选)
|
||||
|
||||
**返回字段**: 返回创建的模型列表(与 GET 模型列表接口返回格式相同)
|
||||
"""
|
||||
adapter = AdminBatchCreateModelsAdapter(provider_id=provider_id, models_data=models_data)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/{provider_id}/available-source-models",
|
||||
response_model=ProviderAvailableSourceModelsResponse,
|
||||
)
|
||||
async def get_provider_available_source_models(
|
||||
provider_id: str,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> Any:
|
||||
"""
|
||||
获取提供商支持的可用源模型
|
||||
|
||||
获取该提供商支持的所有统一模型名(source_model),包含价格和能力信息。
|
||||
|
||||
**路径参数**:
|
||||
- `provider_id`: 提供商 ID
|
||||
|
||||
**返回字段**:
|
||||
- `models`: 可用源模型数组,每项包含:
|
||||
- `global_model_name`: 全局模型名称
|
||||
- `display_name`: 显示名称
|
||||
- `provider_model_name`: 提供商模型名称
|
||||
- `model_id`: 模型 ID
|
||||
- `price`: 价格信息(包含 input_price_per_1m, output_price_per_1m, cache_creation_price_per_1m, cache_read_price_per_1m, price_per_request)
|
||||
- `capabilities`: 能力信息(包含 supports_vision, supports_function_calling, supports_streaming)
|
||||
- `is_active`: 是否启用
|
||||
- `total`: 总数
|
||||
"""
|
||||
adapter = AdminGetProviderAvailableSourceModelsAdapter(provider_id=provider_id)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/{provider_id}/assign-global-models",
|
||||
response_model=BatchAssignModelsToProviderResponse,
|
||||
)
|
||||
async def batch_assign_global_models_to_provider(
|
||||
provider_id: str,
|
||||
payload: BatchAssignModelsToProviderRequest,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> BatchAssignModelsToProviderResponse:
|
||||
"""
|
||||
批量关联全局模型
|
||||
|
||||
批量为提供商关联全局模型,自动继承全局模型的价格和能力配置。
|
||||
|
||||
**路径参数**:
|
||||
- `provider_id`: 提供商 ID
|
||||
|
||||
**请求体字段**:
|
||||
- `global_model_ids`: 全局模型 ID 数组(必填)
|
||||
|
||||
**返回字段**:
|
||||
- `success`: 成功关联的模型数组,每项包含:
|
||||
- `global_model_id`: 全局模型 ID
|
||||
- `global_model_name`: 全局模型名称
|
||||
- `model_id`: 新创建的模型 ID
|
||||
- `errors`: 失败的模型数组,每项包含:
|
||||
- `global_model_id`: 全局模型 ID
|
||||
- `global_model_name`: 全局模型名称(如果可用)
|
||||
- `error`: 错误信息
|
||||
"""
|
||||
adapter = AdminBatchAssignModelsToProviderAdapter(provider_id=provider_id, payload=payload)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/{provider_id}/import-from-upstream",
|
||||
response_model=ImportFromUpstreamResponse,
|
||||
)
|
||||
async def import_models_from_upstream(
|
||||
provider_id: str,
|
||||
payload: ImportFromUpstreamRequest,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> ImportFromUpstreamResponse:
|
||||
"""
|
||||
从上游提供商导入模型
|
||||
|
||||
从上游提供商导入模型列表。自动匹配已有的 GlobalModel,如果不存在则自动创建。
|
||||
|
||||
**流程说明**:
|
||||
1. 检查模型是否已存在于当前 Provider(按 provider_model_name 匹配)
|
||||
2. 尝试按名称精确匹配已有的 GlobalModel
|
||||
3. 如果没有匹配到,自动创建新的 GlobalModel
|
||||
4. 创建 Model 记录并关联到 GlobalModel
|
||||
|
||||
**路径参数**:
|
||||
- `provider_id`: 提供商 ID
|
||||
|
||||
**请求体字段**:
|
||||
- `model_ids`: 模型 ID 数组(必填,每个 ID 长度 1-100 字符)
|
||||
- `tiered_pricing`: 可选的阶梯计费配置(应用于所有导入的模型和新创建的 GlobalModel)
|
||||
- `price_per_request`: 可选的按次计费价格(应用于所有导入的模型和新创建的 GlobalModel)
|
||||
|
||||
**返回字段**:
|
||||
- `success`: 成功导入的模型数组,每项包含:
|
||||
- `model_id`: 模型 ID
|
||||
- `provider_model_id`: 提供商模型 ID
|
||||
- `global_model_id`: 全局模型 ID
|
||||
- `global_model_name`: 全局模型名称
|
||||
- `created_global_model`: 是否新创建了全局模型
|
||||
- `errors`: 失败的模型数组,每项包含:
|
||||
- `model_id`: 模型 ID
|
||||
- `error`: 错误信息
|
||||
"""
|
||||
adapter = AdminImportFromUpstreamAdapter(provider_id=provider_id, payload=payload)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
# -------- Adapters --------
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminListProviderModelsAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
is_active: bool | None
|
||||
skip: int
|
||||
limit: int
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db = context.db
|
||||
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
|
||||
if not provider:
|
||||
raise NotFoundException("Provider not found", "provider")
|
||||
|
||||
models = ModelService.get_models_by_provider(
|
||||
db, self.provider_id, self.skip, self.limit, self.is_active
|
||||
)
|
||||
return [ModelService.convert_to_response(model) for model in models]
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminCreateProviderModelAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
model_data: ModelCreate
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db = context.db
|
||||
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
|
||||
if not provider:
|
||||
raise NotFoundException("Provider not found", "provider")
|
||||
|
||||
try:
|
||||
model = ModelService.create_model(db, self.provider_id, self.model_data)
|
||||
logger.info(
|
||||
f"Model created: {model.provider_model_name} for provider {provider.name} by {context.user.username}"
|
||||
)
|
||||
# 缓存失效已在 ModelService.create_model 中处理
|
||||
return ModelService.convert_to_response(model)
|
||||
except Exception as exc:
|
||||
raise InvalidRequestException(str(exc))
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminGetProviderModelAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
model_id: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db = context.db
|
||||
model = (
|
||||
db.query(Model)
|
||||
.filter(Model.id == self.model_id, Model.provider_id == self.provider_id)
|
||||
.first()
|
||||
)
|
||||
if not model:
|
||||
raise NotFoundException("Model not found", "model")
|
||||
|
||||
return ModelService.convert_to_response(model)
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminUpdateProviderModelAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
model_id: str
|
||||
model_data: ModelUpdate
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db = context.db
|
||||
model = (
|
||||
db.query(Model)
|
||||
.filter(Model.id == self.model_id, Model.provider_id == self.provider_id)
|
||||
.first()
|
||||
)
|
||||
if not model:
|
||||
raise NotFoundException("Model not found", "model")
|
||||
|
||||
try:
|
||||
updated_model = ModelService.update_model(db, self.model_id, self.model_data)
|
||||
logger.info(
|
||||
f"Model updated: {updated_model.provider_model_name} by {context.user.username}"
|
||||
)
|
||||
# 缓存失效已在 ModelService.update_model 中处理
|
||||
return ModelService.convert_to_response(updated_model)
|
||||
except Exception as exc:
|
||||
raise InvalidRequestException(str(exc))
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminDeleteProviderModelAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
model_id: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db = context.db
|
||||
model = (
|
||||
db.query(Model)
|
||||
.filter(Model.id == self.model_id, Model.provider_id == self.provider_id)
|
||||
.first()
|
||||
)
|
||||
if not model:
|
||||
raise NotFoundException("Model not found", "model")
|
||||
|
||||
model_name = model.provider_model_name
|
||||
try:
|
||||
ModelService.delete_model(db, self.model_id)
|
||||
logger.info(f"Model deleted: {model_name} by {context.user.username}")
|
||||
# 缓存失效已在 ModelService.delete_model 中处理
|
||||
return {"message": f"Model '{model_name}' deleted successfully"}
|
||||
except Exception as exc:
|
||||
raise InvalidRequestException(str(exc))
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminBatchCreateModelsAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
models_data: list[ModelCreate]
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db = context.db
|
||||
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
|
||||
if not provider:
|
||||
raise NotFoundException("Provider not found", "provider")
|
||||
|
||||
try:
|
||||
models = ModelService.batch_create_models(db, self.provider_id, self.models_data)
|
||||
logger.info(
|
||||
f"Batch created {len(models)} models for provider {provider.name} by {context.user.username}"
|
||||
)
|
||||
# 缓存失效已在 ModelService.batch_create_models 中处理
|
||||
return [ModelService.convert_to_response(model) for model in models]
|
||||
except Exception as exc:
|
||||
raise InvalidRequestException(str(exc))
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminGetProviderAvailableSourceModelsAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
"""
|
||||
返回 Provider 支持的所有 GlobalModel
|
||||
|
||||
逻辑:
|
||||
1. 查询该 Provider 的所有 Model
|
||||
2. 通过 Model.global_model_id 获取 GlobalModel
|
||||
"""
|
||||
db = context.db
|
||||
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
|
||||
if not provider:
|
||||
raise NotFoundException("Provider not found", "provider")
|
||||
|
||||
# 1. 查询该 Provider 的所有活跃 Model(预加载 GlobalModel)
|
||||
models = (
|
||||
db.query(Model)
|
||||
.options(joinedload(Model.global_model))
|
||||
.filter(Model.provider_id == self.provider_id, Model.is_active == True)
|
||||
.all()
|
||||
)
|
||||
|
||||
# 2. 构建以 GlobalModel 为主键的字典
|
||||
global_models_dict: dict[str, dict[str, Any]] = {}
|
||||
|
||||
for model in models:
|
||||
global_model = model.global_model
|
||||
if not global_model or not global_model.is_active:
|
||||
continue
|
||||
|
||||
global_model_name = global_model.name
|
||||
|
||||
# 如果该 GlobalModel 还未处理,初始化
|
||||
if global_model_name not in global_models_dict:
|
||||
global_models_dict[global_model_name] = {
|
||||
"global_model_name": global_model_name,
|
||||
"display_name": global_model.display_name,
|
||||
"provider_model_name": model.provider_model_name,
|
||||
"model_id": model.id,
|
||||
"price": {
|
||||
"input_price_per_1m": model.get_effective_input_price(),
|
||||
"output_price_per_1m": model.get_effective_output_price(),
|
||||
"cache_creation_price_per_1m": model.get_effective_cache_creation_price(),
|
||||
"cache_read_price_per_1m": model.get_effective_cache_read_price(),
|
||||
"price_per_request": model.get_effective_price_per_request(),
|
||||
},
|
||||
"capabilities": {
|
||||
"supports_vision": bool(model.supports_vision),
|
||||
"supports_function_calling": bool(model.supports_function_calling),
|
||||
"supports_streaming": bool(model.supports_streaming),
|
||||
},
|
||||
"is_active": bool(model.is_active),
|
||||
}
|
||||
|
||||
models_list = [
|
||||
ProviderAvailableSourceModel(**global_models_dict[name])
|
||||
for name in sorted(global_models_dict.keys())
|
||||
]
|
||||
|
||||
return ProviderAvailableSourceModelsResponse(models=models_list, total=len(models_list))
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminBatchAssignModelsToProviderAdapter(AdminApiAdapter):
|
||||
"""批量为 Provider 关联 GlobalModels"""
|
||||
|
||||
provider_id: str
|
||||
payload: BatchAssignModelsToProviderRequest
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db = context.db
|
||||
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
|
||||
if not provider:
|
||||
raise NotFoundException("Provider not found", "provider")
|
||||
|
||||
success = []
|
||||
errors = []
|
||||
|
||||
for global_model_id in self.payload.global_model_ids:
|
||||
try:
|
||||
global_model = (
|
||||
db.query(GlobalModel).filter(GlobalModel.id == global_model_id).first()
|
||||
)
|
||||
if not global_model:
|
||||
errors.append(
|
||||
{"global_model_id": global_model_id, "error": "GlobalModel not found"}
|
||||
)
|
||||
continue
|
||||
|
||||
# 检查是否已存在关联
|
||||
existing = (
|
||||
db.query(Model)
|
||||
.filter(
|
||||
Model.provider_id == self.provider_id,
|
||||
Model.global_model_id == global_model_id,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if existing:
|
||||
errors.append(
|
||||
{
|
||||
"global_model_id": global_model_id,
|
||||
"global_model_name": global_model.name,
|
||||
"error": "Already associated",
|
||||
}
|
||||
)
|
||||
continue
|
||||
|
||||
# 创建新的 Model 记录,继承 GlobalModel 的配置
|
||||
new_model = Model(
|
||||
provider_id=self.provider_id,
|
||||
global_model_id=global_model_id,
|
||||
provider_model_name=global_model.name,
|
||||
is_active=True,
|
||||
)
|
||||
db.add(new_model)
|
||||
db.flush()
|
||||
|
||||
success.append(
|
||||
{
|
||||
"global_model_id": global_model_id,
|
||||
"global_model_name": global_model.name,
|
||||
"model_id": new_model.id,
|
||||
}
|
||||
)
|
||||
except Exception as e:
|
||||
errors.append({"global_model_id": global_model_id, "error": str(e)})
|
||||
|
||||
db.commit()
|
||||
logger.info(
|
||||
f"Batch assigned {len(success)} GlobalModels to provider {provider.name} by {context.user.username}"
|
||||
)
|
||||
|
||||
# 清除 /v1/models 列表缓存
|
||||
if success:
|
||||
# Provider 新增模型实现后,清除同进程的 ModelMapper 缓存,避免 TTL 内仍返回 None
|
||||
from src.services.cache.invalidation import get_cache_invalidation_service
|
||||
|
||||
cache_service = get_cache_invalidation_service()
|
||||
cache_service.on_model_changed(self.provider_id, success[0].get("global_model_id", ""))
|
||||
|
||||
await invalidate_models_list_cache()
|
||||
|
||||
return BatchAssignModelsToProviderResponse(success=success, errors=errors)
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminImportFromUpstreamAdapter(AdminApiAdapter):
|
||||
"""从上游提供商导入模型(自动匹配或创建 GlobalModel)"""
|
||||
|
||||
provider_id: str
|
||||
payload: ImportFromUpstreamRequest
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db = context.db
|
||||
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
|
||||
if not provider:
|
||||
raise NotFoundException("Provider not found", "provider")
|
||||
|
||||
success: list[ImportFromUpstreamSuccessItem] = []
|
||||
errors: list[ImportFromUpstreamErrorItem] = []
|
||||
|
||||
# 获取价格覆盖配置
|
||||
tiered_pricing = None
|
||||
price_per_request = None
|
||||
if hasattr(self.payload, "tiered_pricing") and self.payload.tiered_pricing:
|
||||
tiered_pricing = self.payload.tiered_pricing
|
||||
if (
|
||||
hasattr(self.payload, "price_per_request")
|
||||
and self.payload.price_per_request is not None
|
||||
):
|
||||
price_per_request = self.payload.price_per_request
|
||||
|
||||
# 默认价格配置(用于自动创建的 GlobalModel)
|
||||
default_pricing = {
|
||||
"tiers": [
|
||||
{
|
||||
"up_to": None,
|
||||
"input_price_per_1m": 0.0,
|
||||
"output_price_per_1m": 0.0,
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
for model_id in self.payload.model_ids:
|
||||
# 输入验证:检查 model_id 长度
|
||||
if not model_id or len(model_id) > 100:
|
||||
errors.append(
|
||||
ImportFromUpstreamErrorItem(
|
||||
model_id=(
|
||||
model_id[:50] + "..."
|
||||
if model_id and len(model_id) > 50
|
||||
else model_id or "<empty>"
|
||||
),
|
||||
error="Invalid model_id: must be 1-100 characters",
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
try:
|
||||
# 使用 savepoint 确保单个模型导入的原子性
|
||||
savepoint = db.begin_nested()
|
||||
try:
|
||||
# 1. 检查是否已存在同名的 ProviderModel
|
||||
existing = (
|
||||
db.query(Model)
|
||||
.options(joinedload(Model.global_model))
|
||||
.filter(
|
||||
Model.provider_id == self.provider_id,
|
||||
Model.provider_model_name == model_id,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if existing:
|
||||
# 已存在,提交 savepoint 并记录成功
|
||||
savepoint.commit()
|
||||
success.append(
|
||||
ImportFromUpstreamSuccessItem(
|
||||
model_id=model_id,
|
||||
global_model_id=existing.global_model_id,
|
||||
global_model_name=existing.global_model.name,
|
||||
provider_model_id=existing.id,
|
||||
created_global_model=False,
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
# 2. 尝试匹配已有的 GlobalModel(按名称精确匹配)
|
||||
global_model = (
|
||||
db.query(GlobalModel).filter(GlobalModel.name == model_id).first()
|
||||
)
|
||||
created_global_model = False
|
||||
|
||||
# 3. 如果没有匹配到,自动创建新的 GlobalModel
|
||||
if not global_model:
|
||||
global_model = GlobalModel(
|
||||
name=model_id,
|
||||
display_name=model_id,
|
||||
default_tiered_pricing=tiered_pricing or default_pricing,
|
||||
default_price_per_request=price_per_request,
|
||||
is_active=True,
|
||||
)
|
||||
db.add(global_model)
|
||||
db.flush()
|
||||
created_global_model = True
|
||||
logger.info(
|
||||
f"Auto-created GlobalModel: {model_id} for provider {provider.name} "
|
||||
f"by {context.user.username}"
|
||||
)
|
||||
|
||||
# 4. 创建新的 Model 记录(关联到 GlobalModel)
|
||||
new_model = Model(
|
||||
provider_id=self.provider_id,
|
||||
global_model_id=global_model.id,
|
||||
provider_model_name=model_id,
|
||||
is_active=True,
|
||||
tiered_pricing=tiered_pricing,
|
||||
price_per_request=price_per_request,
|
||||
)
|
||||
db.add(new_model)
|
||||
db.flush()
|
||||
|
||||
# 提交 savepoint
|
||||
savepoint.commit()
|
||||
success.append(
|
||||
ImportFromUpstreamSuccessItem(
|
||||
model_id=model_id,
|
||||
global_model_id=global_model.id,
|
||||
global_model_name=global_model.name,
|
||||
provider_model_id=new_model.id,
|
||||
created_global_model=created_global_model,
|
||||
)
|
||||
)
|
||||
logger.info(
|
||||
f"Imported model: {model_id} -> GlobalModel: {global_model.name} "
|
||||
f"(created={created_global_model}) for provider {provider.name}"
|
||||
)
|
||||
except Exception as e:
|
||||
# 回滚到 savepoint
|
||||
savepoint.rollback()
|
||||
raise e
|
||||
except Exception as e:
|
||||
logger.error(f"Error importing model {model_id}: {e}")
|
||||
errors.append(ImportFromUpstreamErrorItem(model_id=model_id, error=str(e)))
|
||||
|
||||
db.commit()
|
||||
logger.info(
|
||||
f"Imported {len(success)} models to provider {provider.name} by {context.user.username}"
|
||||
)
|
||||
|
||||
# 清除 /v1/models 列表缓存(导入的模型现在参与路由)
|
||||
if success:
|
||||
await invalidate_models_list_cache()
|
||||
|
||||
return ImportFromUpstreamResponse(success=success, errors=errors)
|
||||
1245
_deprecated_py_src/api/admin/providers/routes.py
Normal file
1245
_deprecated_py_src/api/admin/providers/routes.py
Normal file
File diff suppressed because it is too large
Load Diff
932
_deprecated_py_src/api/admin/providers/summary.py
Normal file
932
_deprecated_py_src/api/admin/providers/summary.py
Normal file
@@ -0,0 +1,932 @@
|
||||
"""
|
||||
Provider 摘要与健康监控 API
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any, Protocol
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy import case, func
|
||||
from sqlalchemy.orm import Session, load_only
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.models_service import invalidate_models_list_cache
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.config.constants import CacheTTL
|
||||
from src.core.enums import ProviderBillingType
|
||||
from src.core.exceptions import InvalidRequestException, NotFoundException
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
from src.models.admin_requests import (
|
||||
ClaudeCodeAdvancedConfig,
|
||||
FailoverRulesConfig,
|
||||
PoolAdvancedConfig,
|
||||
)
|
||||
from src.models.database import (
|
||||
Model,
|
||||
Provider,
|
||||
ProviderAPIKey,
|
||||
ProviderEndpoint,
|
||||
RequestCandidate,
|
||||
)
|
||||
from src.models.endpoint_models import (
|
||||
EndpointHealthEvent,
|
||||
EndpointHealthMonitor,
|
||||
ProviderEndpointHealthMonitorResponse,
|
||||
ProviderSummaryPageResponse,
|
||||
ProviderUpdateRequest,
|
||||
ProviderWithEndpointsSummary,
|
||||
)
|
||||
from src.services.cache.model_cache import ModelCacheService
|
||||
from src.services.cache.provider_cache import ProviderCacheService
|
||||
from src.utils.cache_decorator import cache_result
|
||||
|
||||
|
||||
class _HasProviderSortFields(Protocol):
|
||||
is_active: Any
|
||||
provider_priority: Any
|
||||
created_at: Any
|
||||
|
||||
|
||||
router = APIRouter(tags=["Provider Summary"])
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
def _provider_summary_ordering(provider_model: _HasProviderSortFields) -> tuple[Any, Any, Any]:
|
||||
"""Provider 摘要列表排序:启用在前,其次按优先级与创建时间。"""
|
||||
return (
|
||||
case((provider_model.is_active == True, 0), else_=1).asc(),
|
||||
provider_model.provider_priority.asc(),
|
||||
provider_model.created_at.asc(),
|
||||
)
|
||||
|
||||
|
||||
@router.get("/summary", response_model=ProviderSummaryPageResponse)
|
||||
async def get_providers_summary(
|
||||
request: Request,
|
||||
page: int = Query(1, ge=1),
|
||||
page_size: int = Query(20, ge=1, le=10000),
|
||||
search: str = Query("", description="按名称搜索"),
|
||||
status: str = Query("all", description="all/active/inactive"),
|
||||
api_format: str = Query("all", description="API 格式筛选"),
|
||||
model_id: str = Query("all", description="全局模型 ID 筛选"),
|
||||
db: Session = Depends(get_db),
|
||||
) -> ProviderSummaryPageResponse:
|
||||
"""获取提供商摘要信息(分页)"""
|
||||
adapter = AdminProviderSummaryAdapter(
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
search=search,
|
||||
status=status,
|
||||
api_format=api_format,
|
||||
model_id=model_id,
|
||||
)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get("/{provider_id}/summary", response_model=ProviderWithEndpointsSummary)
|
||||
async def get_provider_summary(
|
||||
provider_id: str,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> ProviderWithEndpointsSummary:
|
||||
"""
|
||||
获取单个提供商摘要信息
|
||||
|
||||
获取指定提供商的详细摘要信息,包含端点、密钥、模型统计和健康状态。
|
||||
|
||||
**路径参数**:
|
||||
- `provider_id`: 提供商 ID
|
||||
|
||||
**返回字段**:
|
||||
- `id`: 提供商 ID
|
||||
- `name`: 提供商名称
|
||||
- `description`: 描述信息
|
||||
- `website`: 官网地址
|
||||
- `provider_priority`: 优先级
|
||||
- `is_active`: 是否启用
|
||||
- `billing_type`: 计费类型
|
||||
- `monthly_quota_usd`: 月度配额(美元)
|
||||
- `monthly_used_usd`: 本月已使用金额(美元)
|
||||
- `quota_reset_day`: 配额重置日期
|
||||
- `quota_last_reset_at`: 上次配额重置时间
|
||||
- `quota_expires_at`: 配额过期时间
|
||||
- `timeout`: 默认请求超时(秒)
|
||||
- `max_retries`: 默认最大重试次数
|
||||
- `proxy`: 默认代理配置
|
||||
- `total_endpoints`: 端点总数
|
||||
- `active_endpoints`: 活跃端点数
|
||||
- `total_keys`: 密钥总数
|
||||
- `active_keys`: 活跃密钥数
|
||||
- `total_models`: 模型总数
|
||||
- `active_models`: 活跃模型数
|
||||
- `avg_health_score`: 平均健康分数(0-1)
|
||||
- `unhealthy_endpoints`: 不健康端点数(健康分数 < 0.5)
|
||||
- `api_formats`: 支持的 API 格式列表
|
||||
- `endpoint_health_details`: 端点健康详情(包含 api_format, health_score, is_active, active_keys)
|
||||
- `created_at`: 创建时间
|
||||
- `updated_at`: 更新时间
|
||||
"""
|
||||
adapter = AdminProviderDetailAdapter(provider_id=provider_id)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get("/{provider_id}/health-monitor", response_model=ProviderEndpointHealthMonitorResponse)
|
||||
async def get_provider_health_monitor(
|
||||
provider_id: str,
|
||||
request: Request,
|
||||
lookback_hours: int = Query(6, ge=1, le=72, description="回溯的小时数"),
|
||||
per_endpoint_limit: int = Query(48, ge=10, le=200, description="每个端点的事件数量"),
|
||||
db: Session = Depends(get_db),
|
||||
) -> ProviderEndpointHealthMonitorResponse:
|
||||
"""
|
||||
获取提供商健康监控数据
|
||||
|
||||
获取指定提供商下所有端点的健康监控时间线,包含请求成功率、延迟、错误信息等。
|
||||
|
||||
**路径参数**:
|
||||
- `provider_id`: 提供商 ID
|
||||
|
||||
**查询参数**:
|
||||
- `lookback_hours`: 回溯的小时数,范围 1-72,默认为 6
|
||||
- `per_endpoint_limit`: 每个端点返回的事件数量,范围 10-200,默认为 48
|
||||
|
||||
**返回字段**:
|
||||
- `provider_id`: 提供商 ID
|
||||
- `provider_name`: 提供商名称
|
||||
- `generated_at`: 生成时间
|
||||
- `endpoints`: 端点健康监控数据数组,每项包含:
|
||||
- `endpoint_id`: 端点 ID
|
||||
- `api_format`: API 格式
|
||||
- `is_active`: 是否活跃
|
||||
- `total_attempts`: 总请求次数
|
||||
- `success_count`: 成功次数
|
||||
- `failed_count`: 失败次数
|
||||
- `skipped_count`: 跳过次数
|
||||
- `success_rate`: 成功率(0-1)
|
||||
- `last_event_at`: 最后事件时间
|
||||
- `events`: 事件详情数组(包含 timestamp, status, status_code, latency_ms, error_type, error_message)
|
||||
"""
|
||||
|
||||
adapter = AdminProviderHealthMonitorAdapter(
|
||||
provider_id=provider_id,
|
||||
lookback_hours=lookback_hours,
|
||||
per_endpoint_limit=per_endpoint_limit,
|
||||
)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.patch("/{provider_id}", response_model=ProviderWithEndpointsSummary)
|
||||
async def update_provider_settings(
|
||||
provider_id: str,
|
||||
update_data: ProviderUpdateRequest,
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> ProviderWithEndpointsSummary:
|
||||
"""
|
||||
更新提供商基础配置
|
||||
|
||||
更新提供商的基础配置信息,如名称、描述、优先级等。只需传入需要更新的字段。
|
||||
|
||||
**路径参数**:
|
||||
- `provider_id`: 提供商 ID
|
||||
|
||||
**请求体字段**(所有字段可选):
|
||||
- `name`: 提供商名称
|
||||
- `description`: 描述信息
|
||||
- `website`: 官网地址
|
||||
- `provider_priority`: 优先级
|
||||
- `is_active`: 是否启用
|
||||
- `billing_type`: 计费类型
|
||||
- `monthly_quota_usd`: 月度配额(美元)
|
||||
- `quota_reset_day`: 配额重置日期
|
||||
- `quota_expires_at`: 配额过期时间
|
||||
- `timeout`: 默认请求超时(秒)
|
||||
- `max_retries`: 默认最大重试次数
|
||||
- `proxy`: 默认代理配置
|
||||
|
||||
**返回字段**: 返回更新后的提供商摘要信息(与 GET /summary 接口返回格式相同)
|
||||
"""
|
||||
|
||||
adapter = AdminUpdateProviderSettingsAdapter(provider_id=provider_id, update_data=update_data)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
def _extract_pool_advanced_from_config(
|
||||
provider_config: dict[str, Any] | None,
|
||||
*,
|
||||
provider_id: str,
|
||||
) -> PoolAdvancedConfig | None:
|
||||
"""从 Provider.config 中安全提取通用号池配置。
|
||||
|
||||
优先查找 ``pool_advanced``,回退查找 ``claude_code_advanced`` 中的号池字段。
|
||||
"""
|
||||
cfg = provider_config or {}
|
||||
raw = cfg.get("pool_advanced")
|
||||
if raw is None:
|
||||
return None
|
||||
|
||||
if isinstance(raw, PoolAdvancedConfig):
|
||||
return raw
|
||||
|
||||
if not isinstance(raw, dict):
|
||||
logger.warning(
|
||||
"Provider {} 的 pool_advanced 类型无效: {},已忽略",
|
||||
provider_id,
|
||||
type(raw).__name__,
|
||||
)
|
||||
return None
|
||||
|
||||
try:
|
||||
return PoolAdvancedConfig.model_validate(raw)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Provider {} 的 pool_advanced 配置无效,已忽略: {}",
|
||||
provider_id,
|
||||
str(exc),
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
def _extract_claude_code_advanced_from_config(
|
||||
provider_config: dict[str, Any] | None,
|
||||
*,
|
||||
provider_id: str,
|
||||
) -> ClaudeCodeAdvancedConfig | None:
|
||||
"""从 Provider.config 中安全提取 Claude Code 高级配置。"""
|
||||
raw_config = (provider_config or {}).get("claude_code_advanced")
|
||||
if raw_config is None:
|
||||
return None
|
||||
|
||||
if isinstance(raw_config, ClaudeCodeAdvancedConfig):
|
||||
return raw_config
|
||||
|
||||
if not isinstance(raw_config, dict):
|
||||
logger.warning(
|
||||
"Provider {} 的 claude_code_advanced 类型无效: {},已忽略",
|
||||
provider_id,
|
||||
type(raw_config).__name__,
|
||||
)
|
||||
return None
|
||||
|
||||
try:
|
||||
return ClaudeCodeAdvancedConfig.model_validate(raw_config)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Provider {} 的 claude_code_advanced 配置无效,已忽略: {}",
|
||||
provider_id,
|
||||
str(exc),
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
def _extract_failover_rules_from_config(
|
||||
provider_config: dict[str, Any] | None,
|
||||
*,
|
||||
provider_id: str,
|
||||
) -> FailoverRulesConfig | None:
|
||||
"""从 Provider.config 中安全提取故障转移规则配置。"""
|
||||
raw = (provider_config or {}).get("failover_rules")
|
||||
if raw is None:
|
||||
return None
|
||||
|
||||
if isinstance(raw, FailoverRulesConfig):
|
||||
return raw
|
||||
|
||||
if not isinstance(raw, dict):
|
||||
logger.warning(
|
||||
"Provider {} 的 failover_rules 类型无效: {},已忽略",
|
||||
provider_id,
|
||||
type(raw).__name__,
|
||||
)
|
||||
return None
|
||||
|
||||
try:
|
||||
return FailoverRulesConfig.model_validate(raw)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Provider {} 的 failover_rules 配置无效,已忽略: {}",
|
||||
provider_id,
|
||||
str(exc),
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
def _build_provider_summary(db: Session, provider: Provider) -> ProviderWithEndpointsSummary:
|
||||
endpoints = (
|
||||
db.query(ProviderEndpoint)
|
||||
.options(
|
||||
load_only(
|
||||
ProviderEndpoint.id,
|
||||
ProviderEndpoint.provider_id,
|
||||
ProviderEndpoint.api_format,
|
||||
ProviderEndpoint.is_active,
|
||||
)
|
||||
)
|
||||
.filter(ProviderEndpoint.provider_id == provider.id)
|
||||
.all()
|
||||
)
|
||||
|
||||
key_stats = (
|
||||
db.query(
|
||||
func.count(ProviderAPIKey.id).label("total"),
|
||||
func.sum(case((ProviderAPIKey.is_active == True, 1), else_=0)).label("active"),
|
||||
)
|
||||
.filter(ProviderAPIKey.provider_id == provider.id)
|
||||
.first()
|
||||
)
|
||||
total_keys = int(key_stats.total or 0)
|
||||
active_keys = int(key_stats.active or 0)
|
||||
|
||||
model_stats = (
|
||||
db.query(
|
||||
func.count(Model.id).label("total"),
|
||||
func.sum(case((Model.is_active == True, 1), else_=0)).label("active"),
|
||||
)
|
||||
.filter(Model.provider_id == provider.id)
|
||||
.first()
|
||||
)
|
||||
total_models = int(model_stats.total or 0)
|
||||
active_models = int(model_stats.active or 0)
|
||||
|
||||
global_model_ids = [
|
||||
row[0]
|
||||
for row in db.query(Model.global_model_id)
|
||||
.filter(
|
||||
Model.provider_id == provider.id,
|
||||
Model.is_active == True,
|
||||
Model.global_model_id.isnot(None),
|
||||
)
|
||||
.distinct()
|
||||
.all()
|
||||
]
|
||||
|
||||
all_keys = (
|
||||
db.query(ProviderAPIKey)
|
||||
.options(
|
||||
load_only(
|
||||
ProviderAPIKey.id,
|
||||
ProviderAPIKey.provider_id,
|
||||
ProviderAPIKey.is_active,
|
||||
ProviderAPIKey.api_formats,
|
||||
ProviderAPIKey.health_by_format,
|
||||
)
|
||||
)
|
||||
.filter(ProviderAPIKey.provider_id == provider.id)
|
||||
.all()
|
||||
)
|
||||
|
||||
return _compose_provider_summary(
|
||||
provider=provider,
|
||||
endpoints=endpoints,
|
||||
all_keys=all_keys,
|
||||
total_keys=total_keys,
|
||||
active_keys=active_keys,
|
||||
total_models=total_models,
|
||||
active_models=active_models,
|
||||
global_model_ids=global_model_ids,
|
||||
)
|
||||
|
||||
|
||||
def _compose_provider_summary(
|
||||
*,
|
||||
provider: Provider,
|
||||
endpoints: list[ProviderEndpoint],
|
||||
all_keys: list[ProviderAPIKey],
|
||||
total_keys: int,
|
||||
active_keys: int,
|
||||
total_models: int,
|
||||
active_models: int,
|
||||
global_model_ids: list[Any],
|
||||
) -> ProviderWithEndpointsSummary:
|
||||
total_endpoints = len(endpoints)
|
||||
active_endpoints = sum(1 for e in endpoints if e.is_active)
|
||||
api_formats = [e.api_format for e in endpoints]
|
||||
|
||||
# 按 api_formats 分组 keys(通过 api_formats 关联)
|
||||
format_to_endpoint_id: dict[str, str] = {e.api_format: e.id for e in endpoints}
|
||||
keys_by_endpoint: dict[str, list[ProviderAPIKey]] = {e.id: [] for e in endpoints}
|
||||
for key in all_keys:
|
||||
formats = key.api_formats or []
|
||||
for fmt in formats:
|
||||
endpoint_id = format_to_endpoint_id.get(fmt)
|
||||
if endpoint_id:
|
||||
keys_by_endpoint[endpoint_id].append(key)
|
||||
|
||||
endpoint_health_map: dict[str, float] = {}
|
||||
for endpoint in endpoints:
|
||||
keys = keys_by_endpoint.get(endpoint.id, [])
|
||||
if keys:
|
||||
api_fmt = endpoint.api_format
|
||||
health_scores: list[float] = []
|
||||
for k in keys:
|
||||
health_by_format = k.health_by_format or {}
|
||||
if api_fmt in health_by_format:
|
||||
score = health_by_format[api_fmt].get("health_score")
|
||||
if score is not None:
|
||||
health_scores.append(float(score))
|
||||
else:
|
||||
health_scores.append(1.0)
|
||||
avg_health = sum(health_scores) / len(health_scores) if health_scores else 1.0
|
||||
endpoint_health_map[endpoint.id] = avg_health
|
||||
else:
|
||||
endpoint_health_map[endpoint.id] = 1.0
|
||||
|
||||
all_health_scores = list(endpoint_health_map.values())
|
||||
avg_health_score = sum(all_health_scores) / len(all_health_scores) if all_health_scores else 1.0
|
||||
unhealthy_endpoints = sum(1 for score in all_health_scores if score < 0.5)
|
||||
|
||||
active_keys_by_endpoint: dict[str, int] = {}
|
||||
for endpoint_id, keys in keys_by_endpoint.items():
|
||||
active_keys_by_endpoint[endpoint_id] = sum(1 for k in keys if k.is_active)
|
||||
|
||||
endpoint_health_details = [
|
||||
{
|
||||
"api_format": e.api_format,
|
||||
"health_score": endpoint_health_map.get(e.id, 1.0),
|
||||
"is_active": e.is_active,
|
||||
"total_keys": len(keys_by_endpoint.get(e.id, [])),
|
||||
"active_keys": active_keys_by_endpoint.get(e.id, 0),
|
||||
}
|
||||
for e in endpoints
|
||||
]
|
||||
|
||||
provider_config_raw = provider.config
|
||||
provider_config = provider_config_raw if isinstance(provider_config_raw, dict) else {}
|
||||
if provider_config_raw is not None and not isinstance(provider_config_raw, dict):
|
||||
logger.warning(
|
||||
"Provider {} 的 config 类型无效: {},按空配置处理",
|
||||
provider.id,
|
||||
type(provider_config_raw).__name__,
|
||||
)
|
||||
|
||||
# 检查是否配置了 Provider Ops(余额监控等)
|
||||
provider_ops_config = provider_config.get("provider_ops")
|
||||
ops_configured = bool(provider_ops_config)
|
||||
ops_architecture_id = (
|
||||
provider_ops_config.get("architecture_id") if provider_ops_config else None
|
||||
)
|
||||
claude_code_advanced = _extract_claude_code_advanced_from_config(
|
||||
provider_config,
|
||||
provider_id=str(provider.id),
|
||||
)
|
||||
pool_advanced = _extract_pool_advanced_from_config(
|
||||
provider_config,
|
||||
provider_id=str(provider.id),
|
||||
)
|
||||
failover_rules = _extract_failover_rules_from_config(
|
||||
provider_config,
|
||||
provider_id=str(provider.id),
|
||||
)
|
||||
|
||||
return ProviderWithEndpointsSummary(
|
||||
id=provider.id,
|
||||
name=provider.name,
|
||||
provider_type=getattr(provider, "provider_type", None),
|
||||
description=provider.description,
|
||||
website=provider.website,
|
||||
provider_priority=provider.provider_priority,
|
||||
keep_priority_on_conversion=provider.keep_priority_on_conversion,
|
||||
enable_format_conversion=provider.enable_format_conversion,
|
||||
is_active=provider.is_active,
|
||||
billing_type=provider.billing_type.value if provider.billing_type else None,
|
||||
monthly_quota_usd=provider.monthly_quota_usd,
|
||||
monthly_used_usd=provider.monthly_used_usd,
|
||||
quota_reset_day=provider.quota_reset_day,
|
||||
quota_last_reset_at=provider.quota_last_reset_at,
|
||||
quota_expires_at=provider.quota_expires_at,
|
||||
max_retries=provider.max_retries,
|
||||
proxy=provider.proxy,
|
||||
stream_first_byte_timeout=provider.stream_first_byte_timeout,
|
||||
request_timeout=provider.request_timeout,
|
||||
claude_code_advanced=claude_code_advanced,
|
||||
pool_advanced=pool_advanced,
|
||||
failover_rules=failover_rules,
|
||||
total_endpoints=total_endpoints,
|
||||
active_endpoints=active_endpoints,
|
||||
total_keys=total_keys,
|
||||
active_keys=active_keys,
|
||||
total_models=total_models,
|
||||
active_models=active_models,
|
||||
global_model_ids=global_model_ids,
|
||||
avg_health_score=avg_health_score,
|
||||
unhealthy_endpoints=unhealthy_endpoints,
|
||||
api_formats=api_formats,
|
||||
endpoint_health_details=endpoint_health_details,
|
||||
ops_configured=ops_configured,
|
||||
ops_architecture_id=ops_architecture_id,
|
||||
created_at=provider.created_at,
|
||||
updated_at=provider.updated_at,
|
||||
)
|
||||
|
||||
|
||||
def _build_provider_summaries_batch(
|
||||
db: Session, providers: list[Provider]
|
||||
) -> list[ProviderWithEndpointsSummary]:
|
||||
if not providers:
|
||||
return []
|
||||
|
||||
provider_ids = [provider.id for provider in providers]
|
||||
|
||||
endpoint_rows = (
|
||||
db.query(ProviderEndpoint)
|
||||
.options(
|
||||
load_only(
|
||||
ProviderEndpoint.id,
|
||||
ProviderEndpoint.provider_id,
|
||||
ProviderEndpoint.api_format,
|
||||
ProviderEndpoint.is_active,
|
||||
)
|
||||
)
|
||||
.filter(ProviderEndpoint.provider_id.in_(provider_ids))
|
||||
.all()
|
||||
)
|
||||
endpoints_by_provider: dict[str, list[ProviderEndpoint]] = {}
|
||||
for endpoint in endpoint_rows:
|
||||
endpoints_by_provider.setdefault(str(endpoint.provider_id), []).append(endpoint)
|
||||
|
||||
key_rows = (
|
||||
db.query(ProviderAPIKey)
|
||||
.options(
|
||||
load_only(
|
||||
ProviderAPIKey.id,
|
||||
ProviderAPIKey.provider_id,
|
||||
ProviderAPIKey.is_active,
|
||||
ProviderAPIKey.api_formats,
|
||||
ProviderAPIKey.health_by_format,
|
||||
)
|
||||
)
|
||||
.filter(ProviderAPIKey.provider_id.in_(provider_ids))
|
||||
.all()
|
||||
)
|
||||
keys_by_provider: dict[str, list[ProviderAPIKey]] = {}
|
||||
for key in key_rows:
|
||||
keys_by_provider.setdefault(str(key.provider_id), []).append(key)
|
||||
|
||||
key_stats_rows = (
|
||||
db.query(
|
||||
ProviderAPIKey.provider_id.label("provider_id"),
|
||||
func.count(ProviderAPIKey.id).label("total"),
|
||||
func.sum(case((ProviderAPIKey.is_active == True, 1), else_=0)).label("active"),
|
||||
)
|
||||
.filter(ProviderAPIKey.provider_id.in_(provider_ids))
|
||||
.group_by(ProviderAPIKey.provider_id)
|
||||
.all()
|
||||
)
|
||||
key_stats_by_provider: dict[str, dict[str, int]] = {
|
||||
str(row.provider_id): {
|
||||
"total": int(row.total or 0),
|
||||
"active": int(row.active or 0),
|
||||
}
|
||||
for row in key_stats_rows
|
||||
}
|
||||
|
||||
model_stats_rows = (
|
||||
db.query(
|
||||
Model.provider_id.label("provider_id"),
|
||||
func.count(Model.id).label("total"),
|
||||
func.sum(case((Model.is_active == True, 1), else_=0)).label("active"),
|
||||
)
|
||||
.filter(Model.provider_id.in_(provider_ids))
|
||||
.group_by(Model.provider_id)
|
||||
.all()
|
||||
)
|
||||
model_stats_by_provider: dict[str, dict[str, int]] = {
|
||||
str(row.provider_id): {
|
||||
"total": int(row.total or 0),
|
||||
"active": int(row.active or 0),
|
||||
}
|
||||
for row in model_stats_rows
|
||||
}
|
||||
|
||||
global_model_rows = (
|
||||
db.query(Model.provider_id, Model.global_model_id)
|
||||
.filter(
|
||||
Model.provider_id.in_(provider_ids),
|
||||
Model.is_active == True,
|
||||
Model.global_model_id.isnot(None),
|
||||
)
|
||||
.distinct()
|
||||
.all()
|
||||
)
|
||||
global_model_ids_by_provider: dict[str, list[Any]] = {}
|
||||
for provider_id, global_model_id in global_model_rows:
|
||||
global_model_ids_by_provider.setdefault(str(provider_id), []).append(global_model_id)
|
||||
|
||||
summaries: list[ProviderWithEndpointsSummary] = []
|
||||
for provider in providers:
|
||||
pid = str(provider.id)
|
||||
key_stats = key_stats_by_provider.get(pid, {"total": 0, "active": 0})
|
||||
model_stats = model_stats_by_provider.get(pid, {"total": 0, "active": 0})
|
||||
summaries.append(
|
||||
_compose_provider_summary(
|
||||
provider=provider,
|
||||
endpoints=endpoints_by_provider.get(pid, []),
|
||||
all_keys=keys_by_provider.get(pid, []),
|
||||
total_keys=key_stats["total"],
|
||||
active_keys=key_stats["active"],
|
||||
total_models=model_stats["total"],
|
||||
active_models=model_stats["active"],
|
||||
global_model_ids=global_model_ids_by_provider.get(pid, []),
|
||||
)
|
||||
)
|
||||
return summaries
|
||||
|
||||
|
||||
# -------- Adapters --------
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminProviderHealthMonitorAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
lookback_hours: int
|
||||
per_endpoint_limit: int
|
||||
|
||||
@cache_result(
|
||||
key_prefix="admin:providers:health-monitor",
|
||||
ttl=CacheTTL.ADMIN_USAGE_RECORDS,
|
||||
user_specific=False,
|
||||
vary_by=["provider_id", "lookback_hours", "per_endpoint_limit"],
|
||||
)
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db = context.db
|
||||
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
|
||||
if not provider:
|
||||
raise NotFoundException(f"Provider {self.provider_id} 不存在")
|
||||
|
||||
endpoints = (
|
||||
db.query(ProviderEndpoint)
|
||||
.options(
|
||||
load_only(
|
||||
ProviderEndpoint.id,
|
||||
ProviderEndpoint.provider_id,
|
||||
ProviderEndpoint.api_format,
|
||||
ProviderEndpoint.is_active,
|
||||
)
|
||||
)
|
||||
.filter(ProviderEndpoint.provider_id == self.provider_id)
|
||||
.all()
|
||||
)
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
since = now - timedelta(hours=self.lookback_hours)
|
||||
|
||||
endpoint_ids = [str(endpoint.id) for endpoint in endpoints]
|
||||
if not endpoint_ids:
|
||||
response = ProviderEndpointHealthMonitorResponse(
|
||||
provider_id=provider.id,
|
||||
provider_name=provider.name,
|
||||
generated_at=now,
|
||||
endpoints=[],
|
||||
)
|
||||
context.add_audit_metadata(
|
||||
action="provider_health_monitor",
|
||||
provider_id=self.provider_id,
|
||||
endpoint_count=0,
|
||||
lookback_hours=self.lookback_hours,
|
||||
)
|
||||
return response.model_dump()
|
||||
|
||||
ranked_attempts_subq = (
|
||||
db.query(RequestCandidate)
|
||||
.with_entities(
|
||||
RequestCandidate.endpoint_id.label("endpoint_id"),
|
||||
RequestCandidate.status.label("status"),
|
||||
RequestCandidate.status_code.label("status_code"),
|
||||
RequestCandidate.latency_ms.label("latency_ms"),
|
||||
RequestCandidate.error_type.label("error_type"),
|
||||
RequestCandidate.error_message.label("error_message"),
|
||||
func.coalesce(
|
||||
RequestCandidate.finished_at,
|
||||
RequestCandidate.started_at,
|
||||
RequestCandidate.created_at,
|
||||
).label("event_timestamp"),
|
||||
func.row_number()
|
||||
.over(
|
||||
partition_by=RequestCandidate.endpoint_id,
|
||||
order_by=RequestCandidate.created_at.desc(),
|
||||
)
|
||||
.label("rn"),
|
||||
)
|
||||
.filter(
|
||||
RequestCandidate.endpoint_id.in_(endpoint_ids),
|
||||
RequestCandidate.created_at >= since,
|
||||
)
|
||||
.subquery()
|
||||
)
|
||||
attempt_rows = (
|
||||
db.query(ranked_attempts_subq)
|
||||
.filter(ranked_attempts_subq.c.rn <= self.per_endpoint_limit)
|
||||
.order_by(
|
||||
ranked_attempts_subq.c.endpoint_id.asc(),
|
||||
ranked_attempts_subq.c.event_timestamp.asc(),
|
||||
)
|
||||
.all()
|
||||
)
|
||||
|
||||
events_by_endpoint: dict[str, list[EndpointHealthEvent]] = {eid: [] for eid in endpoint_ids}
|
||||
for row in attempt_rows:
|
||||
endpoint_id = str(row.endpoint_id) if row.endpoint_id is not None else ""
|
||||
if not endpoint_id or endpoint_id not in events_by_endpoint:
|
||||
continue
|
||||
events_by_endpoint[endpoint_id].append(
|
||||
EndpointHealthEvent(
|
||||
timestamp=row.event_timestamp,
|
||||
status=row.status,
|
||||
status_code=row.status_code,
|
||||
latency_ms=row.latency_ms,
|
||||
error_type=row.error_type,
|
||||
error_message=row.error_message,
|
||||
)
|
||||
)
|
||||
|
||||
endpoint_monitors: list[EndpointHealthMonitor] = []
|
||||
for endpoint in endpoints:
|
||||
endpoint_id = str(endpoint.id)
|
||||
events = events_by_endpoint.get(endpoint_id, [])
|
||||
|
||||
success_count = sum(1 for event in events if event.status == "success")
|
||||
failed_count = sum(1 for event in events if event.status == "failed")
|
||||
skipped_count = sum(1 for event in events if event.status == "skipped")
|
||||
total_attempts = len(events)
|
||||
success_rate = success_count / total_attempts if total_attempts else 1.0
|
||||
last_event_at = events[-1].timestamp if events else None
|
||||
|
||||
endpoint_monitors.append(
|
||||
EndpointHealthMonitor(
|
||||
endpoint_id=endpoint.id,
|
||||
api_format=endpoint.api_format,
|
||||
is_active=endpoint.is_active,
|
||||
total_attempts=total_attempts,
|
||||
success_count=success_count,
|
||||
failed_count=failed_count,
|
||||
skipped_count=skipped_count,
|
||||
success_rate=success_rate,
|
||||
last_event_at=last_event_at,
|
||||
events=events,
|
||||
)
|
||||
)
|
||||
|
||||
response = ProviderEndpointHealthMonitorResponse(
|
||||
provider_id=provider.id,
|
||||
provider_name=provider.name,
|
||||
generated_at=now,
|
||||
endpoints=endpoint_monitors,
|
||||
)
|
||||
context.add_audit_metadata(
|
||||
action="provider_health_monitor",
|
||||
provider_id=self.provider_id,
|
||||
endpoint_count=len(endpoint_monitors),
|
||||
lookback_hours=self.lookback_hours,
|
||||
per_endpoint_limit=self.per_endpoint_limit,
|
||||
)
|
||||
return response.model_dump()
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminProviderSummaryAdapter(AdminApiAdapter):
|
||||
page: int = 1
|
||||
page_size: int = 20
|
||||
search: str = ""
|
||||
status: str = "all"
|
||||
api_format: str = "all"
|
||||
model_id: str = "all"
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db = context.db
|
||||
|
||||
query = db.query(Provider)
|
||||
|
||||
# 搜索筛选
|
||||
if self.search.strip():
|
||||
keywords = self.search.strip().lower().split()
|
||||
for kw in keywords:
|
||||
query = query.filter(func.lower(Provider.name).contains(kw))
|
||||
|
||||
# 状态筛选
|
||||
if self.status == "active":
|
||||
query = query.filter(Provider.is_active == True)
|
||||
elif self.status == "inactive":
|
||||
query = query.filter(Provider.is_active == False)
|
||||
|
||||
# API 格式筛选
|
||||
if self.api_format != "all":
|
||||
query = query.filter(
|
||||
Provider.id.in_(
|
||||
db.query(ProviderEndpoint.provider_id)
|
||||
.filter(ProviderEndpoint.api_format == self.api_format)
|
||||
.distinct()
|
||||
)
|
||||
)
|
||||
|
||||
# 全局模型 ID 筛选
|
||||
if self.model_id != "all":
|
||||
query = query.filter(
|
||||
Provider.id.in_(
|
||||
db.query(Model.provider_id)
|
||||
.filter(
|
||||
Model.global_model_id == self.model_id,
|
||||
Model.is_active == True,
|
||||
)
|
||||
.distinct()
|
||||
)
|
||||
)
|
||||
|
||||
total = query.count()
|
||||
|
||||
providers = (
|
||||
query.order_by(*_provider_summary_ordering(Provider))
|
||||
.offset((self.page - 1) * self.page_size)
|
||||
.limit(self.page_size)
|
||||
.all()
|
||||
)
|
||||
|
||||
items = _build_provider_summaries_batch(db, providers)
|
||||
return ProviderSummaryPageResponse(
|
||||
total=total,
|
||||
page=self.page,
|
||||
page_size=self.page_size,
|
||||
items=items,
|
||||
).model_dump()
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminProviderDetailAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
|
||||
@cache_result(
|
||||
key_prefix="admin:providers:summary:detail",
|
||||
ttl=CacheTTL.ADMIN_USAGE_RECORDS,
|
||||
user_specific=False,
|
||||
vary_by=["provider_id"],
|
||||
)
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db = context.db
|
||||
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
|
||||
if not provider:
|
||||
raise NotFoundException(f"Provider {self.provider_id} not found")
|
||||
return _build_provider_summary(db, provider).model_dump()
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminUpdateProviderSettingsAdapter(AdminApiAdapter):
|
||||
provider_id: str
|
||||
update_data: ProviderUpdateRequest
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db = context.db
|
||||
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
|
||||
if not provider:
|
||||
raise NotFoundException("Provider not found", "provider")
|
||||
|
||||
update_dict = self.update_data.model_dump(exclude_unset=True)
|
||||
if "claude_code_advanced" in update_dict:
|
||||
claude_advanced = update_dict.pop("claude_code_advanced")
|
||||
provider_type = str(getattr(provider, "provider_type", "") or "").strip().lower()
|
||||
if claude_advanced is not None and provider_type != "claude_code":
|
||||
raise InvalidRequestException(
|
||||
"claude_code_advanced 仅适用于 provider_type=claude_code"
|
||||
)
|
||||
|
||||
provider_config = dict(provider.config or {})
|
||||
if claude_advanced is None:
|
||||
provider_config.pop("claude_code_advanced", None)
|
||||
else:
|
||||
provider_config["claude_code_advanced"] = dict(claude_advanced)
|
||||
update_dict["config"] = provider_config or None
|
||||
|
||||
if "pool_advanced" in update_dict:
|
||||
pool_advanced = update_dict.pop("pool_advanced")
|
||||
provider_config = dict(update_dict.get("config") or provider.config or {})
|
||||
if pool_advanced is None:
|
||||
provider_config.pop("pool_advanced", None)
|
||||
else:
|
||||
provider_config["pool_advanced"] = dict(pool_advanced)
|
||||
update_dict["config"] = provider_config or None
|
||||
|
||||
if "billing_type" in update_dict and update_dict["billing_type"] is not None:
|
||||
update_dict["billing_type"] = ProviderBillingType(update_dict["billing_type"])
|
||||
|
||||
for key, value in update_dict.items():
|
||||
setattr(provider, key, value)
|
||||
|
||||
provider.updated_at = datetime.now(timezone.utc)
|
||||
db.commit()
|
||||
db.refresh(provider)
|
||||
|
||||
admin_name = context.user.username if context.user else "admin"
|
||||
logger.info(f"Provider {provider.name} updated by {admin_name}: {update_dict}")
|
||||
|
||||
# 缓存失效
|
||||
affects_model_visibility = {"is_active", "enable_format_conversion"} & update_dict.keys()
|
||||
if affects_model_visibility:
|
||||
await invalidate_models_list_cache()
|
||||
if "is_active" in update_dict:
|
||||
await ModelCacheService.invalidate_all_resolve_cache()
|
||||
|
||||
if "billing_type" in update_dict:
|
||||
await ProviderCacheService.invalidate_provider_cache(provider.id)
|
||||
|
||||
return _build_provider_summary(db, provider)
|
||||
Reference in New Issue
Block a user