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 1735 additions and 2433 deletions

View File

@@ -0,0 +1,429 @@
"""
模型映射中间件
根据数据库中的配置,将用户请求的模型映射到提供商的实际模型
"""
from __future__ import annotations
from sqlalchemy.orm import Session, joinedload
from src.core.cache_utils import SyncLRUCache
from src.core.logger import logger
from src.models.claude import ClaudeMessagesRequest
from src.models.database import GlobalModel, Model, Provider, ProviderEndpoint
from src.services.cache.model_cache import ModelCacheService
# 模块级共享缓存,所有 ModelMapperMiddleware 实例共用
_shared_cache = SyncLRUCache(max_size=1000, ttl=300)
class ModelMapperMiddleware:
"""
模型映射中间件
负责将用户请求的模型名映射到提供商的实际模型名
"""
def __init__(self, db: Session):
"""
初始化模型映射中间件
Args:
db: 数据库会话
"""
self.db = db
async def apply_mapping(
self, request: ClaudeMessagesRequest, provider: Provider
) -> ClaudeMessagesRequest:
"""
应用模型映射到请求
Args:
request: 原始请求
provider: 目标提供商
Returns:
应用映射后的请求
"""
# 获取请求的模型名
source_model = request.model
# 查找映射
mapping = await self.get_mapping(source_model, provider.id)
if mapping:
# 应用映射
original_model = request.model
request.model = mapping.model.select_provider_model_name()
logger.debug(
f"Applied model mapping for provider {provider.name}: "
f"{original_model} -> {request.model}"
)
else:
# 没有找到映射,使用原始模型名
logger.debug(
f"No model mapping found for {source_model} with provider {provider.name}, "
f"forwarding with original model name"
)
return request
async def get_mapping(self, source_model: str, provider_id: str) -> object | None:
"""
获取模型映射
简化后的逻辑:
1. 通过 GlobalModel.name 解析 GlobalModel
2. 找到 GlobalModel 后,查找该 Provider 的 Model 实现
Args:
source_model: 用户请求的模型名(必须是 GlobalModel.name
provider_id: 提供商ID (UUID)
Returns:
模型映射对象(包含 model 字段如果没有找到返回None
"""
# 步骤 1: 规范化模型名称
normalized_name = source_model.strip() if isinstance(source_model, str) else ""
if not normalized_name:
logger.debug("GlobalModel not found: <empty model name>")
return None
# 检查缓存(使用规范化后的名称)
cache_key = f"{provider_id}:{normalized_name}"
if cache_key in _shared_cache:
return _shared_cache[cache_key]
mapping = None
global_model = await ModelCacheService.get_global_model_by_name(self.db, normalized_name)
if not global_model or not global_model.is_active:
logger.debug(f"GlobalModel not found or inactive: {normalized_name}")
_shared_cache[cache_key] = None
return None
# 步骤 2: 查找该 Provider 是否有实现这个 GlobalModel 的 Model使用缓存
model = await ModelCacheService.get_model_by_provider_and_global_model(
self.db, provider_id, global_model.id
)
if model:
# 将 ORM Model 转为无 Session 绑定的实例,避免跨请求缓存导致 DetachedInstanceError
from sqlalchemy.orm.session import object_session
if object_session(model) is not None:
model_dict = ModelCacheService._model_to_dict(model)
model = ModelCacheService._dict_to_model(model_dict)
# 创建映射对象
mapping = type(
"obj",
(object,),
{
"source_model": source_model,
"model": model,
"is_active": True,
"provider_id": provider_id,
},
)()
logger.debug(
f"Found model mapping: {normalized_name} -> {model.provider_model_name} "
f"(provider={provider_id[:8]}...)"
)
# 缓存结果
_shared_cache[cache_key] = mapping
return mapping
def get_all_mappings(self, provider_id: str) -> list[object]:
"""
获取提供商的所有可用模型(通过 GlobalModel)
Args:
provider_id: 提供商ID (UUID)
Returns:
模型映射列表
"""
# 查询该 Provider 的所有活跃 Model使用 joinedload 避免 N+1
models = (
self.db.query(Model)
.join(GlobalModel)
.options(joinedload(Model.global_model))
.filter(
Model.provider_id == provider_id,
Model.is_active == True,
GlobalModel.is_active == True,
)
.all()
)
# 构造兼容的映射对象列表
mappings = []
for model in models:
mapping = type(
"obj",
(object,),
{
"source_model": model.global_model.name,
"model": model,
"is_active": True,
"provider_id": provider_id,
},
)()
mappings.append(mapping)
return mappings
def get_supported_models(self, provider_id: str) -> list[str]:
"""
获取提供商支持的所有源模型名
Args:
provider_id: 提供商ID (UUID)
Returns:
支持的模型名列表
"""
mappings = self.get_all_mappings(provider_id)
return [mapping.source_model for mapping in mappings]
async def validate_request(
self, request: ClaudeMessagesRequest, provider: Provider
) -> tuple[bool, str | None]:
"""
验证请求是否符合映射的限制
Args:
request: 请求对象
provider: 提供商对象
Returns:
(是否有效, 错误信息)
"""
mapping = await self.get_mapping(request.model, provider.id)
if not mapping:
# 没有映射,可能是默认支持的模型
return True, None
if not mapping.is_active:
return False, f"Model mapping for {request.model} is disabled"
return True, None
@staticmethod
def clear_cache() -> None:
"""清空共享缓存"""
_shared_cache.clear()
logger.debug("Model mapping cache cleared")
@staticmethod
def refresh_cache(provider_id: str | None = None) -> None:
"""
刷新缓存
Args:
provider_id: 如果指定,只刷新该提供商的缓存 (UUID)
"""
if provider_id:
keys_to_remove = [
key for key in _shared_cache.keys() if key.startswith(f"{provider_id}:")
]
for key in keys_to_remove:
del _shared_cache[key]
logger.debug(f"Refreshed cache for provider {provider_id}")
else:
ModelMapperMiddleware.clear_cache()
class ModelRoutingMiddleware:
"""
模型路由中间件
根据模型名选择合适的提供商
"""
def __init__(self, db: Session):
"""
初始化模型路由中间件
Args:
db: 数据库会话
"""
self.db = db
self.mapper = ModelMapperMiddleware(db)
def select_provider(
self,
model_name: str,
preferred_provider: str | None = None,
allowed_api_formats: list[str] | None = None,
request_id: str | None = None,
) -> Provider | None:
"""
根据模型名选择提供商
Args:
model_name: 请求的模型名
preferred_provider: 首选提供商名称
allowed_api_formats: 允许的API格式列表
request_id: 请求ID用于日志关联
Returns:
选中的提供商如果没有找到返回None
"""
request_prefix = f"ID:{request_id} | " if request_id else ""
allowed_norm: set[str] | None = None
if allowed_api_formats:
from src.services.provider.format import normalize_endpoint_signature
allowed_norm = {
normalize_endpoint_signature(str(fmt))
for fmt in allowed_api_formats
if isinstance(fmt, str) and fmt
}
# 1. 如果指定了提供商,直接使用
if preferred_provider:
provider = (
self.db.query(Provider)
.filter(Provider.name == preferred_provider, Provider.is_active == True)
.first()
)
if provider:
# 检查API格式 - 从 endpoints 中检查
if allowed_norm:
has_matching_endpoint = any(
ep.is_active
and ep.api_format
and str(ep.api_format).strip().lower() in allowed_norm
for ep in provider.endpoints
)
if not has_matching_endpoint:
logger.warning(
f"Specified provider {provider.name} has no active endpoints with allowed API formats ({allowed_api_formats})"
)
else:
logger.debug(
f" └─ {request_prefix}使用指定提供商: {provider.name} | 模型:{model_name}"
)
return provider
else:
logger.debug(
f" └─ {request_prefix}使用指定提供商: {provider.name} | 模型:{model_name}"
)
return provider
else:
logger.warning(f"Specified provider {preferred_provider} not found or inactive")
# 2. 查找优先级最高的活动提供商
query = self.db.query(Provider).filter(Provider.is_active == True)
if allowed_norm:
query = (
query.join(ProviderEndpoint)
.filter(
ProviderEndpoint.is_active == True,
ProviderEndpoint.api_format.in_(sorted(allowed_norm)),
)
.distinct()
)
best_provider = query.order_by(Provider.provider_priority.asc(), Provider.id.asc()).first()
if best_provider:
logger.debug(
f" └─ {request_prefix}使用优先级最高提供商: {best_provider.name} (priority:{best_provider.provider_priority}) | 模型:{model_name}"
)
return best_provider
if allowed_api_formats:
logger.error(
f"No active providers found with allowed API formats {allowed_api_formats}."
)
else:
logger.error("No active providers found.")
return None
def get_available_models(self) -> dict[str, list[str]]:
"""
获取所有可用的模型及其提供商
Returns:
字典,键为 GlobalModel.name值为支持该模型的提供商名列表
"""
result = {}
models = (
self.db.query(GlobalModel.name, Provider.name)
.join(Model, GlobalModel.id == Model.global_model_id)
.join(Provider, Model.provider_id == Provider.id)
.filter(
GlobalModel.is_active == True, Model.is_active == True, Provider.is_active == True
)
.all()
)
for global_model_name, provider_name in models:
if global_model_name not in result:
result[global_model_name] = []
if provider_name not in result[global_model_name]:
result[global_model_name].append(provider_name)
return result
async def get_cheapest_provider(self, model_name: str) -> Provider | None:
"""
获取某个模型最便宜的提供商
Args:
model_name: GlobalModel 名称
Returns:
最便宜的提供商
"""
# 直接查找 GlobalModel
global_model = (
self.db.query(GlobalModel)
.filter(GlobalModel.name == model_name, GlobalModel.is_active == True)
.first()
)
if not global_model:
return None
# 查询所有支持该模型的 Provider 及其价格
models_with_providers = (
self.db.query(Provider, Model)
.join(Model, Provider.id == Model.provider_id)
.filter(
Model.global_model_id == global_model.id,
Model.is_active == True,
Provider.is_active == True,
)
.all()
)
if not models_with_providers:
return None
# 按总价格排序
cheapest = min(
models_with_providers,
key=lambda x: x[1].get_effective_input_price() + x[1].get_effective_output_price(),
)
provider = cheapest[0]
model = cheapest[1]
logger.debug(
f"Selected cheapest provider {provider.name} for model {model_name} "
f"(input: ${model.get_effective_input_price()}/M, output: ${model.get_effective_output_price()}/M)"
)
return provider