mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-12 06:00:20 +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:
@@ -0,0 +1,18 @@
|
||||
"""
|
||||
模型管理相关 Admin API
|
||||
"""
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
from .catalog import router as catalog_router
|
||||
from .external import router as external_router
|
||||
from .global_models import router as global_models_router
|
||||
from .routing import router as routing_router
|
||||
|
||||
router = APIRouter(prefix="/api/admin/models", tags=["Admin - Models"])
|
||||
|
||||
# 挂载子路由
|
||||
router.include_router(catalog_router)
|
||||
router.include_router(global_models_router)
|
||||
router.include_router(external_router)
|
||||
router.include_router(routing_router)
|
||||
@@ -0,0 +1,172 @@
|
||||
"""
|
||||
统一模型目录 Admin API
|
||||
|
||||
基于 GlobalModel 的聚合视图
|
||||
"""
|
||||
|
||||
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.pipeline import get_pipeline
|
||||
from src.database import get_db
|
||||
from src.models.database import GlobalModel, Model
|
||||
from src.models.pydantic_models import (
|
||||
ModelCapabilities,
|
||||
ModelCatalogItem,
|
||||
ModelCatalogProviderDetail,
|
||||
ModelCatalogResponse,
|
||||
ModelPriceRange,
|
||||
)
|
||||
|
||||
router = APIRouter(prefix="/catalog", tags=["Admin - Model Catalog"])
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
@router.get("", response_model=ModelCatalogResponse)
|
||||
async def get_model_catalog(
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
) -> ModelCatalogResponse:
|
||||
"""
|
||||
获取统一模型目录
|
||||
|
||||
基于 GlobalModel 聚合所有活跃模型及其关联提供商的信息,返回完整的模型目录视图。
|
||||
|
||||
**返回字段**:
|
||||
- `models`: 模型列表,每个模型包含:
|
||||
- `global_model_name`: GlobalModel 名称
|
||||
- `display_name`: 显示名称
|
||||
- `description`: 模型描述
|
||||
- `providers`: 提供商列表,包含提供商名称、价格、能力等详细信息
|
||||
- `price_range`: 价格区间(基于 GlobalModel 第一阶梯价格)
|
||||
- `total_providers`: 关联提供商数量
|
||||
- `capabilities`: 模型能力标志(视觉、函数调用、流式输出)
|
||||
- `total`: 模型总数
|
||||
"""
|
||||
adapter = AdminGetModelCatalogAdapter()
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminGetModelCatalogAdapter(AdminApiAdapter):
|
||||
"""管理员查询统一模型目录
|
||||
|
||||
架构说明:
|
||||
1. 以 GlobalModel 为中心聚合数据
|
||||
2. Model 表提供关联提供商和价格
|
||||
"""
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
db: Session = context.db
|
||||
|
||||
# 1. 获取所有活跃的 GlobalModel
|
||||
global_models: list[GlobalModel] = (
|
||||
db.query(GlobalModel).filter(GlobalModel.is_active == True).all()
|
||||
)
|
||||
|
||||
# 2. 获取所有活跃的 Model 实现(包含 global_model 以便计算有效价格)
|
||||
models: list[Model] = (
|
||||
db.query(Model)
|
||||
.options(joinedload(Model.provider), joinedload(Model.global_model))
|
||||
.filter(Model.is_active == True)
|
||||
.all()
|
||||
)
|
||||
|
||||
# 按 GlobalModel ID 组织关联提供商
|
||||
models_by_global_model: dict[str, list[Model]] = {}
|
||||
for model in models:
|
||||
if model.global_model_id:
|
||||
models_by_global_model.setdefault(model.global_model_id, []).append(model)
|
||||
|
||||
# 3. 为每个 GlobalModel 构建 catalog item
|
||||
catalog_items: list[ModelCatalogItem] = []
|
||||
|
||||
for gm in global_models:
|
||||
gm_id = gm.id
|
||||
provider_entries: list[ModelCatalogProviderDetail] = []
|
||||
# 从 config JSON 读取能力标志
|
||||
gm_config = gm.config or {}
|
||||
capability_flags = {
|
||||
"supports_vision": gm_config.get("vision", False),
|
||||
"supports_function_calling": gm_config.get("function_calling", False),
|
||||
"supports_streaming": gm_config.get("streaming", True),
|
||||
}
|
||||
|
||||
# 遍历该 GlobalModel 的所有关联提供商
|
||||
for model in models_by_global_model.get(gm_id, []):
|
||||
provider = model.provider
|
||||
if not provider:
|
||||
continue
|
||||
|
||||
# 使用有效价格(考虑 GlobalModel 默认值)
|
||||
effective_input = model.get_effective_input_price()
|
||||
effective_output = model.get_effective_output_price()
|
||||
effective_tiered = model.get_effective_tiered_pricing()
|
||||
tier_count = len(effective_tiered.get("tiers", [])) if effective_tiered else 1
|
||||
|
||||
# 使用有效能力值
|
||||
capability_flags["supports_vision"] = (
|
||||
capability_flags["supports_vision"] or model.get_effective_supports_vision()
|
||||
)
|
||||
capability_flags["supports_function_calling"] = (
|
||||
capability_flags["supports_function_calling"]
|
||||
or model.get_effective_supports_function_calling()
|
||||
)
|
||||
capability_flags["supports_streaming"] = (
|
||||
capability_flags["supports_streaming"]
|
||||
or model.get_effective_supports_streaming()
|
||||
)
|
||||
|
||||
provider_entries.append(
|
||||
ModelCatalogProviderDetail(
|
||||
provider_id=provider.id,
|
||||
provider_name=provider.name,
|
||||
model_id=model.id,
|
||||
target_model=model.provider_model_name,
|
||||
# 显示有效价格
|
||||
input_price_per_1m=effective_input,
|
||||
output_price_per_1m=effective_output,
|
||||
cache_creation_price_per_1m=model.get_effective_cache_creation_price(),
|
||||
cache_read_price_per_1m=model.get_effective_cache_read_price(),
|
||||
cache_1h_creation_price_per_1m=model.get_effective_1h_cache_creation_price(),
|
||||
price_per_request=model.get_effective_price_per_request(),
|
||||
effective_tiered_pricing=effective_tiered,
|
||||
tier_count=tier_count,
|
||||
supports_vision=model.get_effective_supports_vision(),
|
||||
supports_function_calling=model.get_effective_supports_function_calling(),
|
||||
supports_streaming=model.get_effective_supports_streaming(),
|
||||
is_active=bool(model.is_active),
|
||||
)
|
||||
)
|
||||
|
||||
# 模型目录显示 GlobalModel 的第一个阶梯价格(不是 Provider 聚合价格)
|
||||
tiered = gm.default_tiered_pricing or {}
|
||||
first_tier = tiered.get("tiers", [{}])[0] if tiered.get("tiers") else {}
|
||||
price_range = ModelPriceRange(
|
||||
min_input=first_tier.get("input_price_per_1m", 0),
|
||||
max_input=first_tier.get("input_price_per_1m", 0),
|
||||
min_output=first_tier.get("output_price_per_1m", 0),
|
||||
max_output=first_tier.get("output_price_per_1m", 0),
|
||||
)
|
||||
|
||||
catalog_items.append(
|
||||
ModelCatalogItem(
|
||||
global_model_name=gm.name,
|
||||
display_name=gm.display_name,
|
||||
description=gm_config.get("description"),
|
||||
providers=provider_entries,
|
||||
price_range=price_range,
|
||||
total_providers=len(provider_entries),
|
||||
capabilities=ModelCapabilities(**capability_flags),
|
||||
)
|
||||
)
|
||||
|
||||
return ModelCatalogResponse(
|
||||
models=catalog_items,
|
||||
total=len(catalog_items),
|
||||
)
|
||||
@@ -0,0 +1,179 @@
|
||||
"""
|
||||
models.dev 外部模型数据代理
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.clients import get_redis_client
|
||||
from src.core.logger import logger
|
||||
from src.database import get_db
|
||||
from src.models.database import User
|
||||
from src.utils.auth_utils import require_admin
|
||||
|
||||
router = APIRouter()
|
||||
pipeline = get_pipeline()
|
||||
|
||||
CACHE_KEY = "aether:external:models_dev"
|
||||
CACHE_TTL = 15 * 60 # 15 分钟
|
||||
|
||||
# 标记官方/一手提供商,前端可据此过滤第三方转售商
|
||||
OFFICIAL_PROVIDERS = {
|
||||
"anthropic", # Claude 官方
|
||||
"openai", # OpenAI 官方
|
||||
"google", # Gemini 官方
|
||||
"google-vertex", # Google Vertex AI
|
||||
"azure", # Azure OpenAI
|
||||
"amazon-bedrock", # AWS Bedrock
|
||||
"xai", # Grok 官方
|
||||
"meta", # Llama 官方
|
||||
"deepseek", # DeepSeek 官方
|
||||
"mistral", # Mistral 官方
|
||||
"cohere", # Cohere 官方
|
||||
"zhipuai", # 智谱 AI 官方
|
||||
"alibaba", # 阿里云(通义千问)
|
||||
"minimax", # MiniMax 官方
|
||||
"moonshot", # 月之暗面(Kimi)
|
||||
"baichuan", # 百川智能
|
||||
"ai21", # AI21 Labs
|
||||
}
|
||||
|
||||
|
||||
async def _get_cached_data() -> dict[str, Any] | None:
|
||||
"""从 Redis 获取缓存数据"""
|
||||
redis = await get_redis_client()
|
||||
if redis is None:
|
||||
return None
|
||||
try:
|
||||
cached = await redis.get(CACHE_KEY)
|
||||
if cached:
|
||||
result: dict[str, Any] = json.loads(cached)
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.warning(f"读取 models.dev 缓存失败: {e}")
|
||||
return None
|
||||
|
||||
|
||||
async def _set_cached_data(data: dict) -> None:
|
||||
"""将数据写入 Redis 缓存"""
|
||||
redis = await get_redis_client()
|
||||
if redis is None:
|
||||
return
|
||||
try:
|
||||
await redis.setex(CACHE_KEY, CACHE_TTL, json.dumps(data, ensure_ascii=False))
|
||||
except Exception as e:
|
||||
logger.warning(f"写入 models.dev 缓存失败: {e}")
|
||||
|
||||
|
||||
def _mark_official_providers(data: dict[str, Any]) -> dict[str, Any]:
|
||||
"""为每个提供商标记是否为官方"""
|
||||
result = {}
|
||||
for provider_id, provider_data in data.items():
|
||||
result[provider_id] = {
|
||||
**provider_data,
|
||||
"official": provider_id in OFFICIAL_PROVIDERS,
|
||||
}
|
||||
return result
|
||||
|
||||
|
||||
async def _get_external_models_response() -> JSONResponse:
|
||||
"""
|
||||
获取外部模型数据
|
||||
|
||||
从 models.dev 获取第三方模型数据,用于导入新模型或参考定价信息。
|
||||
该接口作为代理请求解决跨域问题,并提供缓存优化。
|
||||
|
||||
**功能特性**:
|
||||
- 代理 models.dev API,解决前端跨域问题
|
||||
- 使用 Redis 缓存 15 分钟,多 worker 共享缓存
|
||||
- 自动标记官方提供商(official 字段),前端可据此过滤第三方转售商
|
||||
|
||||
**返回字段**:
|
||||
- 键为提供商 ID(如 "anthropic"、"openai")
|
||||
- 值为提供商详细信息,包含:
|
||||
- `official`: 是否为官方提供商(true/false)
|
||||
- 其他 models.dev 提供的原始字段(模型列表、定价等)
|
||||
"""
|
||||
# 检查缓存
|
||||
cached = await _get_cached_data()
|
||||
if cached is not None:
|
||||
# 兼容旧缓存:如果没有 official 字段则补全并回写
|
||||
try:
|
||||
needs_mark = False
|
||||
for provider_data in cached.values():
|
||||
if not isinstance(provider_data, dict) or "official" not in provider_data:
|
||||
needs_mark = True
|
||||
break
|
||||
if needs_mark:
|
||||
marked_cached = _mark_official_providers(cached)
|
||||
await _set_cached_data(marked_cached)
|
||||
return JSONResponse(content=marked_cached)
|
||||
except Exception as e:
|
||||
logger.warning(f"处理 models.dev 缓存数据失败,将直接返回原缓存: {e}")
|
||||
return JSONResponse(content=cached)
|
||||
|
||||
raise HTTPException(status_code=503, detail="External models catalog requires Rust admin backend")
|
||||
|
||||
|
||||
async def _clear_external_models_cache_response() -> dict[str, Any]:
|
||||
"""
|
||||
清除外部模型数据缓存
|
||||
|
||||
手动清除 models.dev 的 Redis 缓存,强制下次请求重新获取最新数据。
|
||||
通常用于需要立即更新外部模型数据的场景。
|
||||
|
||||
**返回字段**:
|
||||
- `cleared`: 是否成功清除缓存(true/false)
|
||||
- `message`: 提示信息(仅在 Redis 未启用时返回)
|
||||
"""
|
||||
redis = await get_redis_client()
|
||||
if redis is None:
|
||||
return {"cleared": False, "message": "Redis 未启用"}
|
||||
try:
|
||||
await redis.delete(CACHE_KEY)
|
||||
return {"cleared": True}
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=500, detail=f"清除缓存失败: {str(e)}")
|
||||
|
||||
|
||||
class ExternalModelsAdminAdapter(AdminApiAdapter):
|
||||
"""models.dev 外部模型管理基类。"""
|
||||
|
||||
|
||||
class AdminGetExternalModelsAdapter(ExternalModelsAdminAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> JSONResponse: # type: ignore[override]
|
||||
del context
|
||||
return await _get_external_models_response()
|
||||
|
||||
|
||||
class AdminClearExternalModelsCacheAdapter(ExternalModelsAdminAdapter):
|
||||
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
|
||||
del context
|
||||
return await _clear_external_models_cache_response()
|
||||
|
||||
|
||||
@router.get("/external")
|
||||
async def get_external_models(
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
_: User = Depends(require_admin),
|
||||
) -> JSONResponse:
|
||||
adapter = AdminGetExternalModelsAdapter()
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.delete("/external/cache")
|
||||
async def clear_external_models_cache(
|
||||
request: Request,
|
||||
db: Session = Depends(get_db),
|
||||
_: User = Depends(require_admin),
|
||||
) -> dict[str, Any]:
|
||||
adapter = AdminClearExternalModelsCacheAdapter()
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
@@ -0,0 +1,671 @@
|
||||
"""
|
||||
GlobalModel Admin API
|
||||
|
||||
提供 GlobalModel 的 CRUD 操作接口
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Body, Depends, Query, Request, Response
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
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.logger import logger
|
||||
from src.database import get_db
|
||||
from src.models.pydantic_models import (
|
||||
BatchAssignToProvidersRequest,
|
||||
BatchAssignToProvidersResponse,
|
||||
GlobalModelCreate,
|
||||
GlobalModelListResponse,
|
||||
GlobalModelProvidersResponse,
|
||||
GlobalModelResponse,
|
||||
GlobalModelUpdate,
|
||||
GlobalModelWithStats,
|
||||
ModelCatalogProviderDetail,
|
||||
)
|
||||
from src.services.model.global_model import GlobalModelService
|
||||
|
||||
router = APIRouter(prefix="/global", tags=["Admin - Global Models"])
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
@router.get("", response_model=GlobalModelListResponse)
|
||||
async def list_global_models(
|
||||
request: Request,
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(100, ge=1, le=1000),
|
||||
is_active: bool | None = Query(None),
|
||||
search: str | None = Query(None),
|
||||
db: Session = Depends(get_db),
|
||||
) -> GlobalModelListResponse:
|
||||
"""
|
||||
获取 GlobalModel 列表
|
||||
|
||||
查询系统中的全局模型列表,支持分页、过滤和搜索功能。
|
||||
|
||||
**查询参数**:
|
||||
- `skip`: 跳过记录数,用于分页(默认 0)
|
||||
- `limit`: 返回记录数,用于分页(默认 100,最大 1000)
|
||||
- `is_active`: 过滤活跃状态(true/false/null,null 表示不过滤)
|
||||
- `search`: 搜索关键词,支持按名称或显示名称模糊搜索
|
||||
|
||||
**返回字段**:
|
||||
- `models`: GlobalModel 列表,每个包含:
|
||||
- `id`: GlobalModel ID
|
||||
- `name`: 模型名称(唯一)
|
||||
- `display_name`: 显示名称
|
||||
- `is_active`: 是否活跃
|
||||
- `provider_count`: 关联提供商数量
|
||||
- 定价和能力配置等其他字段
|
||||
- `total`: 返回的模型总数
|
||||
"""
|
||||
adapter = AdminListGlobalModelsAdapter(
|
||||
skip=skip,
|
||||
limit=limit,
|
||||
is_active=is_active,
|
||||
search=search,
|
||||
)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get("/{global_model_id}", response_model=GlobalModelWithStats)
|
||||
async def get_global_model(
|
||||
request: Request,
|
||||
global_model_id: str,
|
||||
db: Session = Depends(get_db),
|
||||
) -> GlobalModelWithStats:
|
||||
"""
|
||||
获取单个 GlobalModel 详情
|
||||
|
||||
查询指定 GlobalModel 的详细信息,包含关联的提供商和价格统计数据。
|
||||
|
||||
**路径参数**:
|
||||
- `global_model_id`: GlobalModel ID
|
||||
|
||||
**返回字段**:
|
||||
- 基础字段:`id`, `name`, `display_name`, `is_active` 等
|
||||
- 统计字段:
|
||||
- `total_models`: 关联的 Model 实现数量
|
||||
- `total_providers`: 关联的提供商数量
|
||||
- `price_range`: 价格区间统计(最低/最高输入输出价格)
|
||||
"""
|
||||
adapter = AdminGetGlobalModelAdapter(global_model_id=global_model_id)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.post("", response_model=GlobalModelResponse, status_code=201)
|
||||
async def create_global_model(
|
||||
request: Request,
|
||||
payload: GlobalModelCreate,
|
||||
db: Session = Depends(get_db),
|
||||
) -> GlobalModelResponse:
|
||||
"""
|
||||
创建 GlobalModel
|
||||
|
||||
创建一个新的全局模型定义,作为多个提供商实现的统一抽象。
|
||||
|
||||
**请求体字段**:
|
||||
- `name`: 模型名称(唯一标识,如 "claude-3-5-sonnet-20241022")
|
||||
- `display_name`: 显示名称(如 "Claude 3.5 Sonnet")
|
||||
- `is_active`: 是否活跃(默认 true)
|
||||
- `default_price_per_request`: 默认按次计费价格(可选)
|
||||
- `default_tiered_pricing`: 默认阶梯定价配置(包含多个价格阶梯)
|
||||
- `supported_capabilities`: 支持的能力标志(vision、function_calling、streaming)
|
||||
- `config`: 额外配置(JSON 格式,如 description、context_window 等)
|
||||
|
||||
**返回字段**:
|
||||
- `id`: 创建的 GlobalModel ID
|
||||
- 其他请求体中的所有字段
|
||||
"""
|
||||
adapter = AdminCreateGlobalModelAdapter(payload=payload)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.patch("/{global_model_id}", response_model=GlobalModelResponse)
|
||||
async def update_global_model(
|
||||
request: Request,
|
||||
global_model_id: str,
|
||||
payload: GlobalModelUpdate,
|
||||
db: Session = Depends(get_db),
|
||||
) -> GlobalModelResponse:
|
||||
"""
|
||||
更新 GlobalModel
|
||||
|
||||
更新指定 GlobalModel 的配置信息,支持部分字段更新。
|
||||
更新后会自动失效相关缓存。
|
||||
|
||||
**路径参数**:
|
||||
- `global_model_id`: GlobalModel ID
|
||||
|
||||
**请求体字段**(均为可选):
|
||||
- `display_name`: 显示名称
|
||||
- `is_active`: 是否活跃
|
||||
- `default_price_per_request`: 默认按次计费价格
|
||||
- `default_tiered_pricing`: 默认阶梯定价配置
|
||||
- `supported_capabilities`: 支持的能力标志
|
||||
- `config`: 额外配置
|
||||
|
||||
**返回字段**:
|
||||
- 更新后的完整 GlobalModel 信息
|
||||
"""
|
||||
adapter = AdminUpdateGlobalModelAdapter(global_model_id=global_model_id, payload=payload)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.delete("/{global_model_id}", status_code=204, response_class=Response)
|
||||
async def delete_global_model(
|
||||
request: Request,
|
||||
global_model_id: str,
|
||||
db: Session = Depends(get_db),
|
||||
) -> Response:
|
||||
"""
|
||||
删除 GlobalModel
|
||||
|
||||
删除指定的 GlobalModel,会级联删除所有关联的 Provider 模型实现。
|
||||
删除后会自动失效相关缓存。
|
||||
|
||||
**路径参数**:
|
||||
- `global_model_id`: GlobalModel ID
|
||||
|
||||
**返回**:
|
||||
- 成功删除返回 204 状态码,无响应体
|
||||
"""
|
||||
adapter = AdminDeleteGlobalModelAdapter(global_model_id=global_model_id)
|
||||
await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
return Response(status_code=204)
|
||||
|
||||
|
||||
@router.post("/batch-delete")
|
||||
async def batch_delete_global_models(
|
||||
request: Request,
|
||||
ids: list[str] = Body(..., embed=True, max_length=100),
|
||||
db: Session = Depends(get_db),
|
||||
) -> dict:
|
||||
"""
|
||||
批量删除 GlobalModel
|
||||
|
||||
顺序删除多个 GlobalModel(每个独立提交),避免并行删除导致的锁竞争。
|
||||
"""
|
||||
adapter = AdminBatchDeleteGlobalModelsAdapter(ids=ids)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/{global_model_id}/assign-to-providers", response_model=BatchAssignToProvidersResponse
|
||||
)
|
||||
async def batch_assign_to_providers(
|
||||
request: Request,
|
||||
global_model_id: str,
|
||||
payload: BatchAssignToProvidersRequest,
|
||||
db: Session = Depends(get_db),
|
||||
) -> BatchAssignToProvidersResponse:
|
||||
"""
|
||||
批量为提供商添加模型实现
|
||||
|
||||
为指定的 GlobalModel 批量创建多个 Provider 的模型实现(Model 记录)。
|
||||
用于快速将一个统一模型分配给多个提供商。
|
||||
|
||||
**路径参数**:
|
||||
- `global_model_id`: GlobalModel ID
|
||||
|
||||
**请求体字段**:
|
||||
- `provider_ids`: 提供商 ID 列表
|
||||
- `create_models`: Model 创建配置列表,每个包含:
|
||||
- `provider_id`: 提供商 ID
|
||||
- `provider_model_name`: 提供商侧的模型名称(如 "claude-3-5-sonnet-20241022")
|
||||
- 其他可选字段(价格覆盖、能力覆盖等)
|
||||
|
||||
**返回字段**:
|
||||
- `success`: 成功创建的 Model 列表
|
||||
- `errors`: 失败的提供商及错误信息列表
|
||||
- `total_requested`: 请求处理的总数
|
||||
- `total_success`: 成功创建的数量
|
||||
- `total_errors`: 失败的数量
|
||||
"""
|
||||
adapter = AdminBatchAssignToProvidersAdapter(global_model_id=global_model_id, payload=payload)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
@router.get("/{global_model_id}/providers", response_model=GlobalModelProvidersResponse)
|
||||
async def get_global_model_providers(
|
||||
request: Request,
|
||||
global_model_id: str,
|
||||
db: Session = Depends(get_db),
|
||||
) -> GlobalModelProvidersResponse:
|
||||
"""
|
||||
获取 GlobalModel 的关联提供商
|
||||
|
||||
查询指定 GlobalModel 的所有关联提供商及其模型实现详情,包括非活跃的提供商。
|
||||
用于查看某个统一模型在各个提供商上的具体配置。
|
||||
|
||||
**路径参数**:
|
||||
- `global_model_id`: GlobalModel ID
|
||||
|
||||
**返回字段**:
|
||||
- `providers`: 提供商列表,每个包含:
|
||||
- `provider_id`: 提供商 ID
|
||||
- `provider_name`: 提供商名称
|
||||
- `provider_display_name`: 提供商显示名称
|
||||
- `model_id`: Model 实现 ID
|
||||
- `target_model`: 提供商侧的模型名称
|
||||
- 价格信息(input_price_per_1m、output_price_per_1m 等)
|
||||
- 能力标志(supports_vision、supports_function_calling、supports_streaming)
|
||||
- `is_active`: 是否活跃
|
||||
- `total`: 关联提供商总数
|
||||
"""
|
||||
adapter = AdminGetGlobalModelProvidersAdapter(global_model_id=global_model_id)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
# ========== Adapters ==========
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminListGlobalModelsAdapter(AdminApiAdapter):
|
||||
"""列出 GlobalModel"""
|
||||
|
||||
skip: int
|
||||
limit: int
|
||||
is_active: bool | None
|
||||
search: str | None
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
from sqlalchemy import and_, case, func, or_
|
||||
|
||||
from src.models.database import GlobalModel, Model, Provider
|
||||
|
||||
query = context.db.query(GlobalModel)
|
||||
if self.is_active is not None:
|
||||
query = query.filter(GlobalModel.is_active == self.is_active)
|
||||
if self.search:
|
||||
search_pattern = f"%{self.search}%"
|
||||
query = query.filter(
|
||||
or_(
|
||||
GlobalModel.name.ilike(search_pattern),
|
||||
GlobalModel.display_name.ilike(search_pattern),
|
||||
)
|
||||
)
|
||||
|
||||
total = int(query.with_entities(func.count(GlobalModel.id)).scalar() or 0)
|
||||
models = query.order_by(GlobalModel.name).offset(self.skip).limit(self.limit).all()
|
||||
|
||||
# 一次性查询所有 GlobalModel 的 provider_count(优化 N+1 问题)
|
||||
# 用条件聚合同时获取总数和活跃数,减少一次 DB 往返
|
||||
model_ids = [gm.id for gm in models]
|
||||
provider_counts = {}
|
||||
active_provider_counts = {}
|
||||
if model_ids:
|
||||
count_results = (
|
||||
context.db.query(
|
||||
Model.global_model_id,
|
||||
func.count(func.distinct(Model.provider_id)),
|
||||
func.count(
|
||||
func.distinct(
|
||||
case(
|
||||
(
|
||||
and_(
|
||||
Model.is_active.is_(True),
|
||||
Provider.is_active.is_(True),
|
||||
),
|
||||
Model.provider_id,
|
||||
),
|
||||
else_=None,
|
||||
)
|
||||
)
|
||||
),
|
||||
)
|
||||
.join(Provider, Model.provider_id == Provider.id)
|
||||
.filter(Model.global_model_id.in_(model_ids))
|
||||
.group_by(Model.global_model_id)
|
||||
.all()
|
||||
)
|
||||
provider_counts = {gm_id: total for gm_id, total, _ in count_results}
|
||||
active_provider_counts = {gm_id: active for gm_id, _, active in count_results}
|
||||
|
||||
# 构建响应
|
||||
model_responses = []
|
||||
for gm in models:
|
||||
response = GlobalModelResponse.model_validate(gm)
|
||||
response.provider_count = provider_counts.get(gm.id, 0)
|
||||
response.active_provider_count = active_provider_counts.get(gm.id, 0)
|
||||
model_responses.append(response)
|
||||
|
||||
return GlobalModelListResponse(
|
||||
models=model_responses,
|
||||
total=total,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminGetGlobalModelAdapter(AdminApiAdapter):
|
||||
"""获取单个 GlobalModel"""
|
||||
|
||||
global_model_id: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
from sqlalchemy import and_, case, func
|
||||
|
||||
from src.models.database import Model, Provider
|
||||
|
||||
global_model = GlobalModelService.get_global_model(context.db, self.global_model_id)
|
||||
stats = GlobalModelService.get_global_model_stats(context.db, self.global_model_id)
|
||||
|
||||
# total_providers 已由 stats 提供,这里只查询活跃 provider 数量
|
||||
active_count = (
|
||||
context.db.query(
|
||||
func.count(
|
||||
func.distinct(
|
||||
case(
|
||||
(
|
||||
and_(
|
||||
Model.is_active.is_(True),
|
||||
Provider.is_active.is_(True),
|
||||
),
|
||||
Model.provider_id,
|
||||
),
|
||||
else_=None,
|
||||
)
|
||||
)
|
||||
)
|
||||
)
|
||||
.join(Provider, Model.provider_id == Provider.id)
|
||||
.filter(Model.global_model_id == global_model.id)
|
||||
.scalar()
|
||||
or 0
|
||||
)
|
||||
|
||||
response = GlobalModelResponse.model_validate(global_model)
|
||||
response.provider_count = stats["total_providers"]
|
||||
response.active_provider_count = int(active_count)
|
||||
|
||||
return GlobalModelWithStats(
|
||||
**response.model_dump(),
|
||||
total_models=stats["total_models"],
|
||||
total_providers=stats["total_providers"],
|
||||
price_range=stats["price_range"],
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminCreateGlobalModelAdapter(AdminApiAdapter):
|
||||
"""创建 GlobalModel"""
|
||||
|
||||
payload: GlobalModelCreate
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
from src.core.exceptions import InvalidRequestException
|
||||
from src.core.model_permissions import validate_and_extract_model_mappings
|
||||
|
||||
# 验证 model_mappings(如果有)
|
||||
is_valid, error, _ = validate_and_extract_model_mappings(self.payload.config)
|
||||
if not is_valid:
|
||||
raise InvalidRequestException(f"映射规则验证失败: {error}", "model_mappings")
|
||||
|
||||
# 将 TieredPricingConfig 转换为 dict
|
||||
tiered_pricing_dict = self.payload.default_tiered_pricing.model_dump()
|
||||
|
||||
global_model = GlobalModelService.create_global_model(
|
||||
db=context.db,
|
||||
name=self.payload.name,
|
||||
display_name=self.payload.display_name,
|
||||
is_active=self.payload.is_active,
|
||||
# 按次计费配置
|
||||
default_price_per_request=self.payload.default_price_per_request,
|
||||
# 阶梯计费配置
|
||||
default_tiered_pricing=tiered_pricing_dict,
|
||||
# Key 能力配置
|
||||
supported_capabilities=self.payload.supported_capabilities,
|
||||
# 模型配置(JSON)
|
||||
config=self.payload.config,
|
||||
)
|
||||
|
||||
logger.info(f"GlobalModel 已创建: id={global_model.id} name={global_model.name}")
|
||||
|
||||
# 创建成功后失效缓存(避免 mapping-preview 在 TTL 内读到旧结果)
|
||||
from src.services.cache.invalidation import get_cache_invalidation_service
|
||||
|
||||
cache_service = get_cache_invalidation_service()
|
||||
await cache_service.on_global_model_changed(global_model.name, str(global_model.id))
|
||||
|
||||
return GlobalModelResponse.model_validate(global_model)
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminUpdateGlobalModelAdapter(AdminApiAdapter):
|
||||
"""更新 GlobalModel"""
|
||||
|
||||
global_model_id: str
|
||||
payload: GlobalModelUpdate
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
from src.core.exceptions import InvalidRequestException
|
||||
from src.core.model_permissions import validate_and_extract_model_mappings
|
||||
|
||||
# 验证 model_mappings(如果有)
|
||||
is_valid, error, _ = validate_and_extract_model_mappings(self.payload.config)
|
||||
if not is_valid:
|
||||
raise InvalidRequestException(f"映射规则验证失败: {error}", "model_mappings")
|
||||
|
||||
# 使用行级锁获取旧的 GlobalModel 信息,防止并发更新导致的竞态条件
|
||||
# 设置 2 秒锁超时,允许短暂等待而非立即失败,提升并发操作的成功率
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.exc import OperationalError
|
||||
|
||||
from src.models.database import GlobalModel
|
||||
|
||||
try:
|
||||
# 设置会话级别的锁超时(仅影响当前事务)
|
||||
context.db.execute(text("SET LOCAL lock_timeout = '2s'"))
|
||||
old_global_model = (
|
||||
context.db.query(GlobalModel)
|
||||
.filter(GlobalModel.id == self.global_model_id)
|
||||
.with_for_update()
|
||||
.first()
|
||||
)
|
||||
except OperationalError as e:
|
||||
# 锁超时或锁冲突时返回友好的错误提示
|
||||
error_msg = str(e).lower()
|
||||
if "lock" in error_msg or "timeout" in error_msg:
|
||||
raise InvalidRequestException("该模型正在被其他操作更新,请稍后重试")
|
||||
raise
|
||||
old_model_name = old_global_model.name if old_global_model else None
|
||||
|
||||
# 执行更新(此时仍持有行锁)
|
||||
global_model = GlobalModelService.update_global_model(
|
||||
db=context.db,
|
||||
global_model_id=self.global_model_id,
|
||||
update_data=self.payload,
|
||||
)
|
||||
|
||||
logger.info(f"GlobalModel 已更新: id={global_model.id} name={global_model.name}")
|
||||
|
||||
# 更新成功后才失效缓存(避免回滚时缓存已被清除的竞态问题)
|
||||
# 注意:此时事务已提交(由 pipeline 管理),数据已持久化
|
||||
from src.services.cache.invalidation import get_cache_invalidation_service
|
||||
|
||||
cache_service = get_cache_invalidation_service()
|
||||
if old_model_name:
|
||||
await cache_service.on_global_model_changed(old_model_name, self.global_model_id)
|
||||
|
||||
return GlobalModelResponse.model_validate(global_model)
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminDeleteGlobalModelAdapter(AdminApiAdapter):
|
||||
"""删除 GlobalModel(级联删除所有关联的 Provider 模型实现)"""
|
||||
|
||||
global_model_id: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
# 使用行级锁获取 GlobalModel 信息,防止并发操作导致的竞态条件
|
||||
# 设置 2 秒锁超时,允许短暂等待而非立即失败
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.exc import OperationalError
|
||||
|
||||
from src.core.exceptions import InvalidRequestException
|
||||
from src.models.database import GlobalModel
|
||||
|
||||
try:
|
||||
# 设置会话级别的锁超时(仅影响当前事务)
|
||||
context.db.execute(text("SET LOCAL lock_timeout = '2s'"))
|
||||
global_model = (
|
||||
context.db.query(GlobalModel)
|
||||
.filter(GlobalModel.id == self.global_model_id)
|
||||
.with_for_update()
|
||||
.first()
|
||||
)
|
||||
except OperationalError as e:
|
||||
# 锁超时或锁冲突时返回友好的错误提示
|
||||
error_msg = str(e).lower()
|
||||
if "lock" in error_msg or "timeout" in error_msg:
|
||||
raise InvalidRequestException("该模型正在被其他操作处理,请稍后重试")
|
||||
raise
|
||||
model_name = global_model.name if global_model else None
|
||||
model_id = global_model.id if global_model else self.global_model_id
|
||||
|
||||
# 执行删除(此时仍持有行锁)
|
||||
GlobalModelService.delete_global_model(context.db, self.global_model_id)
|
||||
|
||||
logger.info(f"GlobalModel 已删除: id={self.global_model_id}")
|
||||
|
||||
# 删除成功后才失效缓存(避免回滚时缓存已被清除的竞态问题)
|
||||
from src.services.cache.invalidation import get_cache_invalidation_service
|
||||
|
||||
cache_service = get_cache_invalidation_service()
|
||||
if model_name:
|
||||
await cache_service.on_global_model_changed(model_name, model_id)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminBatchDeleteGlobalModelsAdapter(AdminApiAdapter):
|
||||
"""批量删除多个 GlobalModel(顺序执行,每个删除独立提交)"""
|
||||
|
||||
ids: list[str]
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
from src.core.exceptions import NotFoundException
|
||||
from src.models.database import GlobalModel
|
||||
|
||||
success_count = 0
|
||||
failed: list[dict] = []
|
||||
deleted_names: list[tuple[str, str]] = [] # (name, id)
|
||||
|
||||
for gm_id in self.ids:
|
||||
try:
|
||||
gm = context.db.query(GlobalModel).filter(GlobalModel.id == gm_id).first()
|
||||
if gm:
|
||||
name = gm.name
|
||||
mid = gm.id
|
||||
GlobalModelService.delete_global_model(context.db, gm_id)
|
||||
deleted_names.append((name, mid))
|
||||
success_count += 1
|
||||
else:
|
||||
failed.append({"id": gm_id, "error": "not found"})
|
||||
except NotFoundException:
|
||||
failed.append({"id": gm_id, "error": "not found"})
|
||||
except Exception as e:
|
||||
context.db.rollback()
|
||||
failed.append({"id": gm_id, "error": str(e)})
|
||||
|
||||
# 批量失效缓存
|
||||
if deleted_names:
|
||||
from src.services.cache.invalidation import get_cache_invalidation_service
|
||||
|
||||
cache_service = get_cache_invalidation_service()
|
||||
for name, mid in deleted_names:
|
||||
await cache_service.on_global_model_changed(name, mid)
|
||||
|
||||
logger.info("批量删除 GlobalModel: success={}, failed={}", success_count, len(failed))
|
||||
|
||||
return {"success_count": success_count, "failed": failed}
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminBatchAssignToProvidersAdapter(AdminApiAdapter):
|
||||
"""批量为 Provider 添加 GlobalModel 实现"""
|
||||
|
||||
global_model_id: str
|
||||
payload: BatchAssignToProvidersRequest
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
result = GlobalModelService.batch_assign_to_providers(
|
||||
db=context.db,
|
||||
global_model_id=self.global_model_id,
|
||||
provider_ids=self.payload.provider_ids,
|
||||
create_models=self.payload.create_models,
|
||||
)
|
||||
|
||||
# 如果有成功创建的关联,清除 /v1/models 列表缓存
|
||||
if result["success"]:
|
||||
await invalidate_models_list_cache()
|
||||
|
||||
logger.info(
|
||||
f"批量为 Provider 添加 GlobalModel: global_model_id={self.global_model_id} success={len(result['success'])} errors={len(result['errors'])}"
|
||||
)
|
||||
|
||||
return BatchAssignToProvidersResponse(**result)
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminGetGlobalModelProvidersAdapter(AdminApiAdapter):
|
||||
"""获取 GlobalModel 的所有关联提供商(包括非活跃的)"""
|
||||
|
||||
global_model_id: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
|
||||
from sqlalchemy.orm import joinedload
|
||||
|
||||
from src.models.database import Model
|
||||
|
||||
global_model = GlobalModelService.get_global_model(context.db, self.global_model_id)
|
||||
|
||||
# 获取所有关联的 Model(包括非活跃的)
|
||||
models = (
|
||||
context.db.query(Model)
|
||||
.options(joinedload(Model.provider), joinedload(Model.global_model))
|
||||
.filter(Model.global_model_id == global_model.id)
|
||||
.all()
|
||||
)
|
||||
|
||||
provider_entries = []
|
||||
for model in models:
|
||||
provider = model.provider
|
||||
if not provider:
|
||||
continue
|
||||
|
||||
effective_tiered = model.get_effective_tiered_pricing()
|
||||
tier_count = len(effective_tiered.get("tiers", [])) if effective_tiered else 1
|
||||
|
||||
provider_entries.append(
|
||||
ModelCatalogProviderDetail(
|
||||
provider_id=provider.id,
|
||||
provider_name=provider.name,
|
||||
model_id=model.id,
|
||||
target_model=model.provider_model_name,
|
||||
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(),
|
||||
cache_1h_creation_price_per_1m=model.get_effective_1h_cache_creation_price(),
|
||||
price_per_request=model.get_effective_price_per_request(),
|
||||
effective_tiered_pricing=effective_tiered,
|
||||
tier_count=tier_count,
|
||||
supports_vision=model.get_effective_supports_vision(),
|
||||
supports_function_calling=model.get_effective_supports_function_calling(),
|
||||
supports_streaming=model.get_effective_supports_streaming(),
|
||||
is_active=bool(model.is_active),
|
||||
)
|
||||
)
|
||||
|
||||
return GlobalModelProvidersResponse(
|
||||
providers=provider_entries,
|
||||
total=len(provider_entries),
|
||||
)
|
||||
@@ -0,0 +1,541 @@
|
||||
"""
|
||||
GlobalModel 请求链路预览 API
|
||||
|
||||
提供模型的请求链路信息,包括:
|
||||
- 请求会流向哪些提供商
|
||||
- 每个提供商的优先级和负载均衡配置
|
||||
- 模型名称映射关系
|
||||
- Key 的并发配置和健康状态
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastapi import APIRouter, Depends, Request
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from sqlalchemy.orm import Session, selectinload
|
||||
|
||||
from src.api.base.admin_adapter import AdminApiAdapter
|
||||
from src.api.base.context import ApiRequestContext
|
||||
from src.api.base.pipeline import get_pipeline
|
||||
from src.core.crypto import CryptoService
|
||||
from src.core.model_permissions import (
|
||||
check_model_allowed_with_mappings,
|
||||
parse_allowed_models_to_list,
|
||||
)
|
||||
from src.database import get_db
|
||||
from src.models.database import (
|
||||
GlobalModel,
|
||||
Model,
|
||||
Provider,
|
||||
ProviderAPIKey,
|
||||
ProviderEndpoint,
|
||||
)
|
||||
from src.services.scheduling.aware_scheduler import CacheAwareScheduler
|
||||
from src.services.system.config import SystemConfigService
|
||||
|
||||
router = APIRouter(prefix="/global", tags=["Admin - Global Models"])
|
||||
pipeline = get_pipeline()
|
||||
|
||||
|
||||
# ========== Response Models ==========
|
||||
|
||||
|
||||
class RoutingKeyInfo(BaseModel):
|
||||
"""Key 路由信息"""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
masked_key: str = Field("", description="脱敏的 API Key")
|
||||
internal_priority: int = Field(..., description="Key 内部优先级")
|
||||
global_priority_by_format: dict[str, int] | None = Field(
|
||||
None, description="按 API 格式的全局优先级"
|
||||
)
|
||||
rpm_limit: int | None = Field(None, description="RPM 限制,null 表示自适应")
|
||||
is_adaptive: bool = Field(False, description="是否为自适应 RPM 模式")
|
||||
effective_rpm: int | None = Field(None, description="有效 RPM 限制")
|
||||
cache_ttl_minutes: int = Field(0, description="缓存 TTL(分钟)")
|
||||
health_score: float = Field(1.0, description="健康度分数(0-1 小数格式)")
|
||||
is_active: bool
|
||||
api_formats: list[str] = Field(default_factory=list, description="支持的 API 格式")
|
||||
# 模型白名单
|
||||
allowed_models: list[str] | None = Field(None, description="允许的模型列表,null 表示不限制")
|
||||
# 熔断状态
|
||||
circuit_breaker_open: bool = Field(False, description="熔断器是否打开")
|
||||
circuit_breaker_formats: list[str] = Field(
|
||||
default_factory=list, description="熔断的 API 格式列表"
|
||||
)
|
||||
next_probe_at: str | None = Field(None, description="下次探测时间(ISO格式)")
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
|
||||
class RoutingEndpointInfo(BaseModel):
|
||||
"""Endpoint 路由信息"""
|
||||
|
||||
id: str
|
||||
api_format: str
|
||||
base_url: str
|
||||
custom_path: str | None = None
|
||||
is_active: bool
|
||||
keys: list[RoutingKeyInfo] = Field(default_factory=list)
|
||||
total_keys: int = 0
|
||||
active_keys: int = 0
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
|
||||
class RoutingModelMapping(BaseModel):
|
||||
"""模型名称映射信息"""
|
||||
|
||||
name: str = Field(..., description="映射名称")
|
||||
priority: int = Field(..., description="优先级(数字越小优先级越高)")
|
||||
api_formats: list[str] | None = Field(None, description="作用域(适用的 API 格式)")
|
||||
|
||||
|
||||
class RoutingProviderInfo(BaseModel):
|
||||
"""Provider 路由信息"""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
model_id: str = Field(..., description="Model ID(GlobalModel 与 Provider 的关联记录 ID)")
|
||||
provider_priority: int = Field(..., description="提供商优先级(数字越小优先级越高)")
|
||||
billing_type: str | None = Field(None, description="计费类型")
|
||||
monthly_quota_usd: float | None = Field(None, description="月额度(美元)")
|
||||
monthly_used_usd: float | None = Field(None, description="已用额度(美元)")
|
||||
is_active: bool
|
||||
# 模型映射信息
|
||||
provider_model_name: str = Field(..., description="提供商侧的模型名称")
|
||||
model_mappings: list[RoutingModelMapping] = Field(
|
||||
default_factory=list, description="模型名称映射列表"
|
||||
)
|
||||
model_is_active: bool = Field(True, description="Model 是否活跃")
|
||||
# Endpoint 和 Key 信息
|
||||
endpoints: list[RoutingEndpointInfo] = Field(default_factory=list)
|
||||
total_endpoints: int = 0
|
||||
active_endpoints: int = 0
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
|
||||
class GlobalKeyWhitelistItem(BaseModel):
|
||||
"""全局 Key 白名单项(用于前端实时匹配)"""
|
||||
|
||||
key_id: str = Field(..., description="Key ID")
|
||||
key_name: str = Field(..., description="Key 名称")
|
||||
masked_key: str = Field(..., description="脱敏的 API Key")
|
||||
provider_id: str = Field(..., description="Provider ID")
|
||||
provider_name: str = Field(..., description="Provider 名称")
|
||||
allowed_models: list[str] = Field(default_factory=list, description="Key 白名单模型列表")
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
|
||||
class ModelRoutingPreviewResponse(BaseModel):
|
||||
"""模型请求链路预览响应"""
|
||||
|
||||
global_model_id: str
|
||||
global_model_name: str
|
||||
display_name: str
|
||||
is_active: bool
|
||||
# GlobalModel 的模型映射(用于前端匹配 Key 白名单)
|
||||
global_model_mappings: list[str] = Field(
|
||||
default_factory=list, description="GlobalModel 的模型映射规则(正则模式)"
|
||||
)
|
||||
# 链路信息
|
||||
providers: list[RoutingProviderInfo] = Field(
|
||||
default_factory=list, description="按优先级排序的提供商列表"
|
||||
)
|
||||
total_providers: int = 0
|
||||
active_providers: int = 0
|
||||
# 调度配置
|
||||
scheduling_mode: str = Field("cache_affinity", description="调度模式")
|
||||
priority_mode: str = Field("provider", description="优先级模式")
|
||||
# 全局 Key 白名单数据(供前端实时匹配,包含所有 Provider 的 Key)
|
||||
all_keys_whitelist: list[GlobalKeyWhitelistItem] = Field(
|
||||
default_factory=list, description="所有 Provider 的 Key 白名单数据"
|
||||
)
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
|
||||
# ========== API Endpoints ==========
|
||||
|
||||
|
||||
@router.get("/{global_model_id}/routing", response_model=ModelRoutingPreviewResponse)
|
||||
async def get_model_routing_preview(
|
||||
request: Request,
|
||||
global_model_id: str,
|
||||
db: Session = Depends(get_db),
|
||||
) -> ModelRoutingPreviewResponse:
|
||||
"""
|
||||
获取模型请求链路预览
|
||||
|
||||
查看指定 GlobalModel 的完整请求链路信息,包括:
|
||||
- 关联的所有提供商及其优先级
|
||||
- 每个提供商的模型名称映射配置
|
||||
- Endpoint 和 Key 的详细配置
|
||||
- 负载均衡和调度策略
|
||||
|
||||
**路径参数**:
|
||||
- `global_model_id`: GlobalModel ID
|
||||
|
||||
**返回字段**:
|
||||
- `global_model_id`: GlobalModel ID
|
||||
- `global_model_name`: 模型名称
|
||||
- `display_name`: 显示名称
|
||||
- `is_active`: 是否活跃
|
||||
- `providers`: 按优先级排序的提供商列表,每个包含:
|
||||
- `id`: Provider ID
|
||||
- `name`: Provider 名称
|
||||
- `provider_priority`: 提供商优先级
|
||||
- `provider_model_name`: 提供商侧的模型名称
|
||||
- `model_mappings`: 模型名称映射列表
|
||||
- `endpoints`: Endpoint 列表,每个包含 Key 信息
|
||||
- `scheduling_mode`: 调度模式(cache_affinity, fixed_order, load_balance)
|
||||
- `priority_mode`: 优先级模式(provider, global_key)
|
||||
"""
|
||||
adapter = AdminGetModelRoutingPreviewAdapter(global_model_id=global_model_id)
|
||||
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
|
||||
|
||||
|
||||
# ========== Adapters ==========
|
||||
|
||||
|
||||
@dataclass
|
||||
class AdminGetModelRoutingPreviewAdapter(AdminApiAdapter):
|
||||
"""获取模型请求链路预览"""
|
||||
|
||||
global_model_id: str
|
||||
|
||||
async def handle(self, context: ApiRequestContext) -> ModelRoutingPreviewResponse: # type: ignore[override]
|
||||
db = context.db
|
||||
|
||||
# 获取 GlobalModel
|
||||
global_model = db.query(GlobalModel).filter(GlobalModel.id == self.global_model_id).first()
|
||||
if not global_model:
|
||||
from fastapi import HTTPException
|
||||
|
||||
raise HTTPException(status_code=404, detail="GlobalModel not found")
|
||||
|
||||
# 获取所有关联的 Model(包含 Provider 信息)
|
||||
models = (
|
||||
db.query(Model)
|
||||
.options(selectinload(Model.provider))
|
||||
.filter(Model.global_model_id == global_model.id)
|
||||
.all()
|
||||
)
|
||||
|
||||
# 获取所有相关的 Provider ID
|
||||
provider_ids = [m.provider_id for m in models if m.provider_id]
|
||||
|
||||
# 批量获取 Provider 的 Endpoints
|
||||
endpoints_by_provider: dict[str, list[ProviderEndpoint]] = {}
|
||||
if provider_ids:
|
||||
endpoints = (
|
||||
db.query(ProviderEndpoint)
|
||||
.filter(ProviderEndpoint.provider_id.in_(provider_ids))
|
||||
.all()
|
||||
)
|
||||
for ep in endpoints:
|
||||
if ep.provider_id not in endpoints_by_provider:
|
||||
endpoints_by_provider[ep.provider_id] = []
|
||||
endpoints_by_provider[ep.provider_id].append(ep)
|
||||
|
||||
# 批量获取 Provider 的 Keys
|
||||
keys_by_provider: dict[str, list[ProviderAPIKey]] = {}
|
||||
if provider_ids:
|
||||
keys = (
|
||||
db.query(ProviderAPIKey).filter(ProviderAPIKey.provider_id.in_(provider_ids)).all()
|
||||
)
|
||||
for key in keys:
|
||||
if key.provider_id not in keys_by_provider:
|
||||
keys_by_provider[key.provider_id] = []
|
||||
keys_by_provider[key.provider_id].append(key)
|
||||
|
||||
# 提取 GlobalModel 的 model_mappings(用于 Key 白名单匹配)
|
||||
global_model_mappings: list[str] = []
|
||||
if global_model.config and isinstance(global_model.config, dict):
|
||||
mappings = global_model.config.get("model_mappings")
|
||||
if isinstance(mappings, list):
|
||||
global_model_mappings = [m for m in mappings if isinstance(m, str)]
|
||||
|
||||
# 构建 Provider 路由信息
|
||||
provider_infos: list[RoutingProviderInfo] = []
|
||||
for model in models:
|
||||
provider = model.provider
|
||||
if not provider:
|
||||
continue
|
||||
|
||||
# 获取模型映射
|
||||
model_mappings = []
|
||||
if model.provider_model_mappings:
|
||||
for mapping in model.provider_model_mappings:
|
||||
model_mappings.append(
|
||||
RoutingModelMapping(
|
||||
name=mapping.get("name", ""),
|
||||
priority=mapping.get("priority", 0),
|
||||
api_formats=mapping.get("api_formats"),
|
||||
)
|
||||
)
|
||||
|
||||
# 获取 Endpoints
|
||||
provider_endpoints = endpoints_by_provider.get(provider.id, [])
|
||||
provider_keys = keys_by_provider.get(provider.id, [])
|
||||
|
||||
# 按 api_format 组织 Keys
|
||||
keys_by_endpoint: dict[str, list[ProviderAPIKey]] = {}
|
||||
for key in provider_keys:
|
||||
# 每个 Key 可能支持多个 api_formats
|
||||
for fmt in key.api_formats or []:
|
||||
if fmt not in keys_by_endpoint:
|
||||
keys_by_endpoint[fmt] = []
|
||||
keys_by_endpoint[fmt].append(key)
|
||||
|
||||
# 定义 Key 模型权限匹配检查函数(用于过滤 Key)
|
||||
def is_key_model_allowed(key: ProviderAPIKey) -> bool:
|
||||
"""检查 Key 的白名单是否匹配当前 GlobalModel"""
|
||||
raw_allowed_models = key.allowed_models
|
||||
if not raw_allowed_models:
|
||||
# 没有白名单限制,允许所有模型
|
||||
return True
|
||||
allowed_models_list = parse_allowed_models_to_list(raw_allowed_models)
|
||||
is_allowed, _ = check_model_allowed_with_mappings(
|
||||
model_name=global_model.name,
|
||||
allowed_models=allowed_models_list,
|
||||
model_mappings=global_model_mappings,
|
||||
)
|
||||
return is_allowed
|
||||
|
||||
endpoint_infos = []
|
||||
for ep in provider_endpoints:
|
||||
# 获取该 Endpoint 格式对应的 Keys
|
||||
ep_keys = keys_by_endpoint.get(ep.api_format or "", [])
|
||||
|
||||
# 过滤:只保留白名单匹配当前 GlobalModel 的 Keys
|
||||
ep_keys = [k for k in ep_keys if is_key_model_allowed(k)]
|
||||
|
||||
# 如果该 Endpoint 没有任何匹配的 Key,跳过此 Endpoint
|
||||
if not ep_keys:
|
||||
continue
|
||||
|
||||
# 按优先级排序(使用当前格式的全局优先级)
|
||||
api_format = ep.api_format or ""
|
||||
|
||||
def get_key_priority(k: ProviderAPIKey) -> tuple[int, int]:
|
||||
format_priority = 999
|
||||
if k.global_priority_by_format and api_format in k.global_priority_by_format:
|
||||
format_priority = k.global_priority_by_format[api_format]
|
||||
return (format_priority, k.internal_priority or 0)
|
||||
|
||||
ep_keys.sort(key=get_key_priority)
|
||||
|
||||
key_infos = []
|
||||
for key in ep_keys:
|
||||
# 计算有效 RPM
|
||||
effective_rpm = key.rpm_limit
|
||||
is_adaptive = key.rpm_limit is None
|
||||
if is_adaptive and key.learned_rpm_limit:
|
||||
effective_rpm = key.learned_rpm_limit
|
||||
|
||||
# 从 health_by_format 获取健康度(0-1 小数格式)
|
||||
health_score = 1.0
|
||||
if key.health_by_format and ep.api_format:
|
||||
format_health = key.health_by_format.get(ep.api_format, {})
|
||||
health_score = format_health.get("health_score", 1.0)
|
||||
|
||||
# 生成脱敏 SK(先解密再脱敏)
|
||||
masked_key = ""
|
||||
if key.api_key:
|
||||
crypto = CryptoService()
|
||||
try:
|
||||
decrypted_key = crypto.decrypt(key.api_key, silent=True)
|
||||
except Exception:
|
||||
# 解密失败时使用加密后的值(可能是未加密的旧数据)
|
||||
decrypted_key = key.api_key
|
||||
if len(decrypted_key) > 8:
|
||||
masked_key = f"{decrypted_key[:4]}***{decrypted_key[-4:]}"
|
||||
else:
|
||||
masked_key = f"{decrypted_key[:2]}***"
|
||||
|
||||
# 检查熔断状态
|
||||
circuit_breaker_open = False
|
||||
circuit_breaker_formats: list[str] = []
|
||||
next_probe_at: str | None = None
|
||||
if key.circuit_breaker_by_format:
|
||||
for fmt, cb_state in key.circuit_breaker_by_format.items():
|
||||
if isinstance(cb_state, dict) and cb_state.get("open"):
|
||||
circuit_breaker_open = True
|
||||
circuit_breaker_formats.append(fmt)
|
||||
# 取最早的探测时间
|
||||
fmt_next_probe = cb_state.get("next_probe_at")
|
||||
if fmt_next_probe:
|
||||
if next_probe_at is None or fmt_next_probe < next_probe_at:
|
||||
next_probe_at = fmt_next_probe
|
||||
|
||||
# 解析 allowed_models
|
||||
raw_allowed_models = key.allowed_models
|
||||
allowed_models_list = (
|
||||
parse_allowed_models_to_list(raw_allowed_models)
|
||||
if raw_allowed_models
|
||||
else None
|
||||
)
|
||||
|
||||
key_infos.append(
|
||||
RoutingKeyInfo(
|
||||
id=key.id or "",
|
||||
name=key.name or "",
|
||||
masked_key=masked_key,
|
||||
internal_priority=key.internal_priority or 0,
|
||||
global_priority_by_format=key.global_priority_by_format,
|
||||
rpm_limit=key.rpm_limit,
|
||||
is_adaptive=is_adaptive,
|
||||
effective_rpm=effective_rpm,
|
||||
cache_ttl_minutes=key.cache_ttl_minutes or 0,
|
||||
health_score=health_score,
|
||||
is_active=bool(key.is_active),
|
||||
api_formats=key.api_formats or [],
|
||||
allowed_models=allowed_models_list,
|
||||
circuit_breaker_open=circuit_breaker_open,
|
||||
circuit_breaker_formats=circuit_breaker_formats,
|
||||
next_probe_at=next_probe_at,
|
||||
)
|
||||
)
|
||||
|
||||
# 计算有效 Keys 数量:is_active 即可(模型权限已在前面过滤)
|
||||
active_keys = sum(1 for k in key_infos if k.is_active)
|
||||
endpoint_infos.append(
|
||||
RoutingEndpointInfo(
|
||||
id=ep.id or "",
|
||||
api_format=ep.api_format or "",
|
||||
base_url=ep.base_url or "",
|
||||
custom_path=ep.custom_path,
|
||||
is_active=bool(ep.is_active),
|
||||
keys=key_infos,
|
||||
total_keys=len(key_infos),
|
||||
active_keys=active_keys,
|
||||
)
|
||||
)
|
||||
|
||||
# 按 endpoint signature 的推荐顺序排序 Endpoints(与前端展示保持一致)
|
||||
preferred_order = [
|
||||
"openai:chat",
|
||||
"openai:cli",
|
||||
"openai:compact",
|
||||
"openai:video",
|
||||
"claude:chat",
|
||||
"claude:cli",
|
||||
"gemini:chat",
|
||||
"gemini:cli",
|
||||
"gemini:video",
|
||||
]
|
||||
order_map = {key: i for i, key in enumerate(preferred_order)}
|
||||
endpoint_infos.sort(
|
||||
key=lambda e: order_map.get(str(e.api_format or "").strip().lower(), 999)
|
||||
)
|
||||
|
||||
active_endpoints = sum(1 for e in endpoint_infos if e.is_active)
|
||||
provider_infos.append(
|
||||
RoutingProviderInfo(
|
||||
id=provider.id,
|
||||
name=provider.name,
|
||||
model_id=model.id,
|
||||
provider_priority=provider.provider_priority,
|
||||
billing_type=provider.billing_type,
|
||||
monthly_quota_usd=provider.monthly_quota_usd,
|
||||
monthly_used_usd=provider.monthly_used_usd,
|
||||
is_active=bool(provider.is_active),
|
||||
provider_model_name=model.provider_model_name,
|
||||
model_mappings=model_mappings,
|
||||
model_is_active=bool(model.is_active),
|
||||
endpoints=endpoint_infos,
|
||||
total_endpoints=len(endpoint_infos),
|
||||
active_endpoints=active_endpoints,
|
||||
)
|
||||
)
|
||||
|
||||
# 按 provider_priority 排序
|
||||
provider_infos.sort(key=lambda p: p.provider_priority)
|
||||
|
||||
active_providers = sum(1 for p in provider_infos if p.is_active and p.model_is_active)
|
||||
|
||||
# 从数据库获取当前调度配置
|
||||
scheduling_mode = (
|
||||
SystemConfigService.get_config(
|
||||
db,
|
||||
"scheduling_mode",
|
||||
CacheAwareScheduler.SCHEDULING_MODE_CACHE_AFFINITY,
|
||||
)
|
||||
or CacheAwareScheduler.SCHEDULING_MODE_CACHE_AFFINITY
|
||||
)
|
||||
priority_mode = (
|
||||
SystemConfigService.get_config(
|
||||
db,
|
||||
"provider_priority_mode",
|
||||
CacheAwareScheduler.PRIORITY_MODE_PROVIDER,
|
||||
)
|
||||
or CacheAwareScheduler.PRIORITY_MODE_PROVIDER
|
||||
)
|
||||
|
||||
# 获取所有活跃 Provider 的 Key 白名单数据(供前端实时匹配)
|
||||
all_keys_whitelist: list[GlobalKeyWhitelistItem] = []
|
||||
crypto = CryptoService()
|
||||
|
||||
# 获取所有活跃的 Key(带白名单),使用 selectinload 避免 N+1 查询
|
||||
all_keys = (
|
||||
db.query(ProviderAPIKey)
|
||||
.join(Provider, ProviderAPIKey.provider_id == Provider.id)
|
||||
.options(selectinload(ProviderAPIKey.provider))
|
||||
.filter(ProviderAPIKey.is_active == True)
|
||||
.filter(Provider.is_active == True)
|
||||
.filter(ProviderAPIKey.allowed_models.isnot(None)) # 只获取有白名单的 Key
|
||||
.all()
|
||||
)
|
||||
|
||||
# 转换为白名单数据
|
||||
for key in all_keys:
|
||||
if not key.allowed_models:
|
||||
continue
|
||||
|
||||
# 解析白名单
|
||||
allowed_models_list = parse_allowed_models_to_list(key.allowed_models)
|
||||
if not allowed_models_list:
|
||||
continue
|
||||
|
||||
# 生成脱敏 Key
|
||||
masked = ""
|
||||
if key.api_key:
|
||||
try:
|
||||
decrypted = crypto.decrypt(key.api_key, silent=True)
|
||||
except Exception:
|
||||
decrypted = key.api_key
|
||||
if len(decrypted) > 8:
|
||||
masked = f"{decrypted[:4]}***{decrypted[-4:]}"
|
||||
else:
|
||||
masked = f"{decrypted[:2]}***"
|
||||
|
||||
all_keys_whitelist.append(
|
||||
GlobalKeyWhitelistItem(
|
||||
key_id=key.id or "",
|
||||
key_name=key.name or "",
|
||||
masked_key=masked,
|
||||
provider_id=key.provider_id or "",
|
||||
provider_name=key.provider.name if key.provider else "",
|
||||
allowed_models=allowed_models_list,
|
||||
)
|
||||
)
|
||||
|
||||
return ModelRoutingPreviewResponse(
|
||||
global_model_id=global_model.id,
|
||||
global_model_name=global_model.name,
|
||||
display_name=global_model.display_name,
|
||||
is_active=bool(global_model.is_active),
|
||||
global_model_mappings=global_model_mappings,
|
||||
providers=provider_infos,
|
||||
total_providers=len(provider_infos),
|
||||
active_providers=active_providers,
|
||||
scheduling_mode=scheduling_mode,
|
||||
priority_mode=priority_mode,
|
||||
all_keys_whitelist=all_keys_whitelist,
|
||||
)
|
||||
Reference in New Issue
Block a user