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:
fawney19
2026-04-03 16:26:16 +08:00
parent 8f26e1a31f
commit 1d9c77522a
868 changed files with 1734 additions and 2432 deletions
@@ -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/nullnull 表示不过滤)
- `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 IDGlobalModel 与 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,
)