mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 09:50:21 +08:00
Initial commit
This commit is contained in:
16
src/services/provider/__init__.py
Normal file
16
src/services/provider/__init__.py
Normal file
@@ -0,0 +1,16 @@
|
||||
"""
|
||||
Provider 服务模块
|
||||
|
||||
包含 Provider 管理、格式处理、传输层等功能。
|
||||
"""
|
||||
|
||||
from src.services.provider.format import normalize_api_format
|
||||
from src.services.provider.service import ProviderService
|
||||
from src.services.provider.transport import build_provider_headers, build_provider_url
|
||||
|
||||
__all__ = [
|
||||
"ProviderService",
|
||||
"normalize_api_format",
|
||||
"build_provider_headers",
|
||||
"build_provider_url",
|
||||
]
|
||||
21
src/services/provider/format.py
Normal file
21
src/services/provider/format.py
Normal file
@@ -0,0 +1,21 @@
|
||||
"""
|
||||
API 格式辅助函数,确保在调度/编排链路中使用统一的枚举值。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Optional, Union
|
||||
|
||||
from src.core.api_format_metadata import resolve_api_format
|
||||
from src.core.enums import APIFormat
|
||||
|
||||
|
||||
def normalize_api_format(
|
||||
value: Union[str, APIFormat, None], default: APIFormat = APIFormat.CLAUDE
|
||||
) -> APIFormat:
|
||||
"""
|
||||
将任意字符串/枚举值归一化为 APIFormat。
|
||||
未识别的值回退到默认枚举(默认 CLAUDE)。
|
||||
"""
|
||||
resolved = resolve_api_format(value)
|
||||
return resolved or default
|
||||
61
src/services/provider/response_normalizer.py
Normal file
61
src/services/provider/response_normalizer.py
Normal file
@@ -0,0 +1,61 @@
|
||||
"""响应标准化服务,用于 STANDARD 模式下的响应格式验证和补全"""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from src.core.logger import logger
|
||||
from src.models.claude import ClaudeResponse
|
||||
|
||||
|
||||
|
||||
class ResponseNormalizer:
|
||||
"""响应标准化器 - 用于标准模式下验证和补全响应字段"""
|
||||
|
||||
@staticmethod
|
||||
def normalize_claude_response(
|
||||
response_data: Dict[str, Any], request_id: Optional[str] = None
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
标准化 Claude API 响应
|
||||
|
||||
Args:
|
||||
response_data: 原始响应数据
|
||||
request_id: 请求ID(用于日志)
|
||||
|
||||
Returns:
|
||||
标准化后的响应数据(失败时返回原始数据)
|
||||
"""
|
||||
if "error" in response_data:
|
||||
logger.debug(f"[ResponseNormalizer] 检测到错误响应,跳过标准化 | ID:{request_id}")
|
||||
return response_data
|
||||
|
||||
try:
|
||||
validated = ClaudeResponse.model_validate(response_data)
|
||||
normalized = validated.model_dump(mode="json", exclude_none=False)
|
||||
|
||||
logger.debug(f"[ResponseNormalizer] 响应标准化成功 | ID:{request_id}")
|
||||
return normalized
|
||||
|
||||
except Exception as e:
|
||||
logger.debug(f"[ResponseNormalizer] 响应验证失败,透传原始数据 | ID:{request_id}")
|
||||
return response_data
|
||||
|
||||
@staticmethod
|
||||
def should_normalize(response_data: Dict[str, Any]) -> bool:
|
||||
"""
|
||||
判断是否需要进行标准化
|
||||
|
||||
Args:
|
||||
response_data: 响应数据
|
||||
|
||||
Returns:
|
||||
是否需要标准化
|
||||
"""
|
||||
# 错误响应不需要标准化
|
||||
if "error" in response_data:
|
||||
return False
|
||||
|
||||
# 已经包含新字段的响应不需要再次标准化
|
||||
if "context_management" in response_data and "container" in response_data:
|
||||
return False
|
||||
|
||||
return True
|
||||
159
src/services/provider/service.py
Normal file
159
src/services/provider/service.py
Normal file
@@ -0,0 +1,159 @@
|
||||
"""
|
||||
提供商服务
|
||||
负责提供商选择、模型映射和请求处理
|
||||
"""
|
||||
|
||||
from typing import Dict
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.logger import logger
|
||||
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):
|
||||
"""
|
||||
检查模型是否可用(严格白名单模式)
|
||||
|
||||
Args:
|
||||
model_name: 模型名称
|
||||
|
||||
Returns:
|
||||
Model对象如果存在且激活,否则None
|
||||
"""
|
||||
# 首先检查是否有直接的模型记录
|
||||
model = (
|
||||
self.db.query(Model)
|
||||
.filter(Model.provider_model_name == model_name, Model.is_active == True)
|
||||
.first()
|
||||
)
|
||||
|
||||
if model:
|
||||
return model
|
||||
|
||||
# 方案 A:检查是否是别名(全局别名系统)
|
||||
from src.services.model.mapping_resolver import resolve_model_to_global_name
|
||||
|
||||
global_model_name = await resolve_model_to_global_name(self.db, model_name)
|
||||
|
||||
# 查找 GlobalModel
|
||||
global_model = (
|
||||
self.db.query(GlobalModel)
|
||||
.filter(GlobalModel.name == global_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):
|
||||
"""
|
||||
检查特定提供商是否支持特定模型
|
||||
|
||||
Args:
|
||||
provider_id: 提供商ID
|
||||
model_name: 模型名称
|
||||
|
||||
Returns:
|
||||
Model对象如果该提供商支持该模型且激活,否则None
|
||||
"""
|
||||
# 首先检查该提供商下是否有直接的模型记录
|
||||
model = (
|
||||
self.db.query(Model)
|
||||
.filter(
|
||||
Model.provider_id == provider_id,
|
||||
Model.provider_model_name == model_name,
|
||||
Model.is_active == True,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
|
||||
if model:
|
||||
return model
|
||||
|
||||
# 方案 A:检查是否是别名
|
||||
from src.services.model.mapping_resolver import resolve_model_to_global_name
|
||||
|
||||
global_model_name = await resolve_model_to_global_name(self.db, model_name, provider_id)
|
||||
|
||||
# 查找 GlobalModel
|
||||
global_model = (
|
||||
self.db.query(GlobalModel)
|
||||
.filter(GlobalModel.name == global_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 clear_cache(self):
|
||||
"""清空缓存"""
|
||||
self.mapper.clear_cache()
|
||||
self.cost_service.clear_cache()
|
||||
logger.info("Provider service cache cleared")
|
||||
146
src/services/provider/transport.py
Normal file
146
src/services/provider/transport.py
Normal file
@@ -0,0 +1,146 @@
|
||||
"""
|
||||
统一的 Provider 请求构建工具。
|
||||
|
||||
负责:
|
||||
- 根据 endpoint/key 构建标准请求头
|
||||
- 根据 API 格式或端点配置生成请求 URL
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Optional
|
||||
from urllib.parse import urlencode
|
||||
|
||||
from src.core.api_format_metadata import get_auth_config, get_default_path, resolve_api_format
|
||||
from src.core.crypto import crypto_service
|
||||
from src.core.enums import APIFormat
|
||||
from src.core.logger import logger
|
||||
|
||||
|
||||
|
||||
def build_provider_headers(
|
||||
endpoint,
|
||||
key,
|
||||
original_headers: Optional[Dict[str, str]] = None,
|
||||
*,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> Dict[str, str]:
|
||||
"""
|
||||
根据 endpoint/key 构建请求头,并透传客户端自定义头。
|
||||
"""
|
||||
headers: Dict[str, str] = {}
|
||||
|
||||
decrypted_key = crypto_service.decrypt(key.api_key)
|
||||
|
||||
# 根据 API 格式自动选择认证头
|
||||
api_format = getattr(endpoint, "api_format", None)
|
||||
resolved_format = resolve_api_format(api_format)
|
||||
auth_header, auth_type = (
|
||||
get_auth_config(resolved_format) if resolved_format else ("Authorization", "bearer")
|
||||
)
|
||||
|
||||
if auth_type == "bearer":
|
||||
headers[auth_header] = f"Bearer {decrypted_key}"
|
||||
else:
|
||||
headers[auth_header] = decrypted_key
|
||||
|
||||
if endpoint.headers:
|
||||
headers.update(endpoint.headers)
|
||||
|
||||
excluded_headers = {
|
||||
"host",
|
||||
"authorization",
|
||||
"x-api-key",
|
||||
"x-goog-api-key",
|
||||
"content-length",
|
||||
"transfer-encoding",
|
||||
}
|
||||
|
||||
if original_headers:
|
||||
for name, value in original_headers.items():
|
||||
if name.lower() not in excluded_headers:
|
||||
headers[name] = value
|
||||
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
|
||||
if "Content-Type" not in headers and "content-type" not in headers:
|
||||
headers["Content-Type"] = "application/json"
|
||||
|
||||
return headers
|
||||
|
||||
|
||||
def build_provider_url(
|
||||
endpoint,
|
||||
*,
|
||||
query_params: Optional[Dict[str, Any]] = None,
|
||||
path_params: Optional[Dict[str, Any]] = None,
|
||||
is_stream: bool = False,
|
||||
) -> str:
|
||||
"""
|
||||
根据 endpoint 配置生成请求 URL
|
||||
|
||||
优先级:
|
||||
1. endpoint.custom_path - 自定义路径(支持模板变量如 {model})
|
||||
2. API 格式默认路径 - 根据 api_format 自动选择
|
||||
|
||||
Args:
|
||||
endpoint: 端点配置
|
||||
query_params: 查询参数
|
||||
path_params: 路径模板参数 (如 {model})
|
||||
is_stream: 是否为流式请求,用于 Gemini API 选择正确的操作方法
|
||||
"""
|
||||
base = endpoint.base_url.rstrip("/")
|
||||
|
||||
# 准备路径参数,添加 Gemini API 所需的 action 参数
|
||||
effective_path_params = dict(path_params) if path_params else {}
|
||||
|
||||
# 为 Gemini API 格式自动添加 action 参数
|
||||
resolved_format = resolve_api_format(endpoint.api_format)
|
||||
if resolved_format in (APIFormat.GEMINI, APIFormat.GEMINI_CLI):
|
||||
if "action" not in effective_path_params:
|
||||
effective_path_params["action"] = (
|
||||
"streamGenerateContent" if is_stream else "generateContent"
|
||||
)
|
||||
|
||||
# 优先使用 custom_path 字段
|
||||
if endpoint.custom_path:
|
||||
path = endpoint.custom_path
|
||||
if effective_path_params:
|
||||
try:
|
||||
path = path.format(**effective_path_params)
|
||||
except KeyError:
|
||||
# 如果模板变量不匹配,保持原路径
|
||||
pass
|
||||
else:
|
||||
# 使用 API 格式的默认路径
|
||||
path = _resolve_default_path(endpoint.api_format)
|
||||
if effective_path_params:
|
||||
try:
|
||||
path = path.format(**effective_path_params)
|
||||
except KeyError:
|
||||
# 如果模板变量不匹配,保持原路径
|
||||
pass
|
||||
|
||||
if not path.startswith("/"):
|
||||
path = f"/{path}"
|
||||
|
||||
url = f"{base}{path}"
|
||||
|
||||
# 添加查询参数
|
||||
if query_params:
|
||||
query_string = urlencode(query_params, doseq=True)
|
||||
if query_string:
|
||||
url = f"{url}?{query_string}"
|
||||
|
||||
return url
|
||||
|
||||
|
||||
def _resolve_default_path(api_format) -> str:
|
||||
"""
|
||||
根据 API 格式返回默认路径
|
||||
"""
|
||||
resolved = resolve_api_format(api_format)
|
||||
if resolved:
|
||||
return get_default_path(resolved)
|
||||
|
||||
logger.warning(f"Unknown api_format '{api_format}' for endpoint, fallback to '/'")
|
||||
return "/"
|
||||
Reference in New Issue
Block a user