Initial commit

This commit is contained in:
fawney19
2025-12-10 20:52:44 +08:00
commit f784106826
485 changed files with 110993 additions and 0 deletions

View 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",
]

View 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

View 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

View 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")

View 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 "/"