mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
- 删除全部 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)
131 lines
3.7 KiB
Python
131 lines
3.7 KiB
Python
"""
|
||
提供商服务
|
||
负责提供商选择、模型映射和请求处理
|
||
"""
|
||
|
||
from typing import Any
|
||
|
||
from sqlalchemy.orm import Session
|
||
|
||
from src.models.database import GlobalModel, Model, Provider
|
||
from src.services.model.cost import ModelCostService
|
||
from src.services.model.mapper import ModelMapperMiddleware, ModelRoutingMiddleware
|
||
|
||
|
||
class ProviderService:
|
||
"""提供商服务类"""
|
||
|
||
def __init__(self, db: Session):
|
||
"""
|
||
初始化提供商服务
|
||
|
||
Args:
|
||
db: 数据库会话
|
||
"""
|
||
self.db = db
|
||
self.mapper = ModelMapperMiddleware(db)
|
||
self.router = ModelRoutingMiddleware(db)
|
||
self.cost_service = ModelCostService(db)
|
||
|
||
async def _check_model_availability(self, model_name: str) -> None:
|
||
"""
|
||
检查模型是否可用(严格白名单模式)
|
||
|
||
Args:
|
||
model_name: 模型名称(必须是 GlobalModel.name)
|
||
|
||
Returns:
|
||
Model对象如果存在且激活,否则None
|
||
"""
|
||
# 直接查找 GlobalModel
|
||
global_model = (
|
||
self.db.query(GlobalModel)
|
||
.filter(GlobalModel.name == model_name, GlobalModel.is_active == True)
|
||
.first()
|
||
)
|
||
|
||
if global_model:
|
||
# 查找任意 Provider 的 Model 实现
|
||
model_obj = (
|
||
self.db.query(Model)
|
||
.filter(Model.global_model_id == global_model.id, Model.is_active == True)
|
||
.first()
|
||
)
|
||
if model_obj:
|
||
return model_obj
|
||
|
||
return None
|
||
|
||
async def _check_provider_model_availability(self, provider_id: str, model_name: str) -> None:
|
||
"""
|
||
检查特定提供商是否支持特定模型
|
||
|
||
Args:
|
||
provider_id: 提供商ID
|
||
model_name: 模型名称(必须是 GlobalModel.name)
|
||
|
||
Returns:
|
||
Model对象如果该提供商支持该模型且激活,否则None
|
||
"""
|
||
# 直接查找 GlobalModel
|
||
global_model = (
|
||
self.db.query(GlobalModel)
|
||
.filter(GlobalModel.name == model_name, GlobalModel.is_active == True)
|
||
.first()
|
||
)
|
||
|
||
if global_model:
|
||
# 查找该 Provider 是否有实现该 GlobalModel
|
||
model_obj = (
|
||
self.db.query(Model)
|
||
.filter(
|
||
Model.provider_id == provider_id,
|
||
Model.global_model_id == global_model.id,
|
||
Model.is_active == True,
|
||
)
|
||
.first()
|
||
)
|
||
if model_obj:
|
||
return model_obj
|
||
|
||
return None
|
||
|
||
def calculate_cost(
|
||
self, provider: Provider, model: str, input_tokens: int, output_tokens: int
|
||
) -> dict[str, float]:
|
||
"""
|
||
计算使用成本
|
||
|
||
Args:
|
||
provider: 提供商对象
|
||
model: 模型名
|
||
input_tokens: 输入tokens
|
||
output_tokens: 输出tokens
|
||
|
||
Returns:
|
||
成本信息
|
||
"""
|
||
return self.mapper.calculate_cost(model, provider.id, input_tokens, output_tokens)
|
||
|
||
def get_available_models(self) -> dict[str, list]:
|
||
"""
|
||
获取所有可用的模型
|
||
|
||
Returns:
|
||
字典,键为模型名,值为提供商列表
|
||
"""
|
||
return self.router.get_available_models()
|
||
|
||
def select_provider(self, model_name: str, preferred_provider: Any | None = None) -> Any:
|
||
"""
|
||
选择提供商
|
||
|
||
Args:
|
||
model_name: 模型名
|
||
preferred_provider: 首选提供商
|
||
|
||
Returns:
|
||
Provider对象
|
||
"""
|
||
return self.router.select_provider(model_name, preferred_provider)
|