refactor: 共享请求管道、按需懒加载、流式内存护栏与连接池治理

- 抽取 ApiRequestPipeline 单例,44 个路由文件共享同一实例
- Handler/Adapter 模块级 __getattr__ 延迟导入,减少启动时间
- 新增 ensure_stream_buffer_limit() 流式内存护栏(16MB 单行 / 32MB 总量)
- HTTP 空闲连接清理与 curl_cffi LRU 会话池
- ensure_providers_bootstrapped 按需引导指定 provider_types
- Usage 事件序列化迁移至 msgpack,Redis codec 隔离
- 启动预热任务(/readyz 就绪门控)与优雅关闭
- 通知邮件模块独立开关与 SMTP 配置校验
- CryptoService DCL 线程安全修复
- 通知模块开关 DB 查询 30s 内存缓存
- /readyz 对 unknown 状态返回 503
- 预热关闭 5s 超时保护
- 预热适配器逐个 try-except 容错
- FormatConversionRegistry 哨兵模式防并发重复物化
- 流式缓冲检查无条件执行

Closes #230

Co-authored-by: AAEE86 <ppk0227@hotmail.com>
This commit is contained in:
fawney19
2026-03-14 11:59:07 +08:00
parent 45985f1c04
commit e0286aebe3
111 changed files with 2775 additions and 1102 deletions

View File

@@ -9,10 +9,14 @@ source -> internal -> target
- 转换失败将抛出 `FormatConversionError`(不再静默回退)。
"""
import ast
import importlib
import inspect
import threading
import time
from collections.abc import Generator
from contextlib import contextmanager
from pathlib import Path
from typing import Any
from src.core.api_format.conversion.exceptions import FormatConversionError
@@ -43,32 +47,136 @@ def _track_conversion_metrics(
)
_MATERIALIZING: tuple[str, str] = ("__materializing__", "")
class FormatConversionRegistry:
"""基于 Normalizer 的格式转换注册表"""
def __init__(self) -> None:
self._normalizers: dict[str, FormatNormalizer] = {}
self._lazy_normalizers: dict[str, tuple[str, str]] = {}
self._lock = threading.RLock()
def register(self, normalizer: FormatNormalizer) -> None:
self._normalizers[str(normalizer.FORMAT_ID).upper()] = normalizer
key = str(normalizer.FORMAT_ID).upper()
with self._lock:
self._normalizers[key] = normalizer
self._lazy_normalizers.pop(key, None)
logger.info(f"[FormatConversionRegistry] 注册 normalizer: {normalizer.FORMAT_ID}")
def register_lazy(self, format_id: str, module_path: str, class_name: str) -> None:
key = str(format_id).upper()
with self._lock:
if key in self._normalizers:
logger.debug(
"[FormatConversionRegistry] 跳过 lazy 注册normalizer 已实例化): {}",
key,
)
return
existing = self._lazy_normalizers.get(key)
if existing and existing != (module_path, class_name):
logger.warning(
"[FormatConversionRegistry] FORMAT_ID '{}' 重复 lazy 注册,{}.{}, 将覆盖 {}.{}",
key,
module_path,
class_name,
existing[0],
existing[1],
)
self._lazy_normalizers[key] = (module_path, class_name)
logger.info(
"[FormatConversionRegistry] 注册 lazy normalizer: {} -> {}.{}",
key,
module_path,
class_name,
)
def _materialize_lazy_normalizer(self, key: str) -> FormatNormalizer | None:
with self._lock:
existing = self._normalizers.get(key)
if existing is not None:
return existing
lazy_spec = self._lazy_normalizers.get(key)
if lazy_spec is None or lazy_spec is _MATERIALIZING:
return None
# 标记为正在加载,防止其他线程重复 materialize
self._lazy_normalizers[key] = _MATERIALIZING
module_path, class_name = lazy_spec
try:
mod = importlib.import_module(module_path)
obj = getattr(mod, class_name, None)
if not inspect.isclass(obj) or not issubclass(obj, FormatNormalizer):
raise TypeError(f"{module_path}.{class_name} 不是有效的 FormatNormalizer")
normalizer = obj()
except Exception as e:
logger.error(
"[FormatConversionRegistry] lazy 加载 {}.{} 失败: {}",
module_path,
class_name,
e,
)
# 恢复 lazy_spec 以便后续重试
with self._lock:
if self._lazy_normalizers.get(key) is _MATERIALIZING:
self._lazy_normalizers[key] = lazy_spec
return None
self.register(normalizer)
key_upper = str(normalizer.FORMAT_ID).upper()
with self._lock:
return self._normalizers.get(key) or self._normalizers.get(key_upper)
def _find_registered_by_data_format_id(self, target_dfid: str) -> FormatNormalizer | None:
from src.core.api_format.metadata import get_data_format_id_for_endpoint
with self._lock:
registered_items = list(self._normalizers.items())
for reg_key, reg_normalizer in registered_items:
if get_data_format_id_for_endpoint(reg_key) == target_dfid:
return reg_normalizer
return None
def _find_lazy_key_by_data_format_id(self, target_dfid: str) -> str | None:
from src.core.api_format.metadata import get_data_format_id_for_endpoint
with self._lock:
lazy_keys = list(self._lazy_normalizers.keys())
for lazy_key in lazy_keys:
if get_data_format_id_for_endpoint(lazy_key) == target_dfid:
return lazy_key
return None
def get_normalizer(self, format_id: str) -> FormatNormalizer | None:
key = str(format_id).upper()
# 1. 精确匹配
normalizer = self._normalizers.get(key)
with self._lock:
normalizer = self._normalizers.get(key)
if normalizer is not None:
return normalizer
# 2. lazy 精确匹配
normalizer = self._materialize_lazy_normalizer(key)
if normalizer is not None:
return normalizer
# 2. data_format_id 回退:如 "claude:cli" (dfid="claude") -> ClaudeNormalizer (dfid="claude")
from src.core.api_format.metadata import get_data_format_id_for_endpoint
target_dfid = get_data_format_id_for_endpoint(format_id)
if not target_dfid:
return None
for reg_key, reg_normalizer in self._normalizers.items():
reg_dfid = get_data_format_id_for_endpoint(reg_key)
if reg_dfid == target_dfid:
return reg_normalizer
# 3. data_format_id 在已实例化 normalizer 中回退
normalizer = self._find_registered_by_data_format_id(target_dfid)
if normalizer is not None:
return normalizer
# 4. data_format_id 在 lazy normalizer 中回退
lazy_key = self._find_lazy_key_by_data_format_id(target_dfid)
if lazy_key:
return self._materialize_lazy_normalizer(lazy_key)
return None
def _require_normalizer(self, format_id: str) -> FormatNormalizer:
@@ -467,13 +575,17 @@ class FormatConversionRegistry:
return True
def list_normalizers(self) -> list[str]:
return sorted(self._normalizers.keys())
with self._lock:
all_keys = set(self._normalizers.keys()) | set(self._lazy_normalizers.keys())
return sorted(all_keys)
def get_supported_targets(self, source_format: str) -> list[str]:
src = str(source_format).upper()
if src not in self._normalizers:
with self._lock:
all_keys = set(self._normalizers.keys()) | set(self._lazy_normalizers.keys())
if src not in all_keys:
return []
return [k for k in self._normalizers.keys() if k != src]
return [k for k in sorted(all_keys) if k != src]
# 全局注册表(唯一实现)
@@ -482,8 +594,91 @@ _DEFAULT_NORMALIZERS_REGISTERED = False
_REGISTRATION_LOCK = threading.Lock()
def _is_format_normalizer_base(node: ast.expr) -> bool:
if isinstance(node, ast.Name):
return node.id == "FormatNormalizer"
if isinstance(node, ast.Attribute):
return node.attr == "FormatNormalizer"
return False
def _extract_format_id_literal(class_node: ast.ClassDef) -> str | None:
for stmt in class_node.body:
if isinstance(stmt, ast.Assign):
for target in stmt.targets:
if isinstance(target, ast.Name) and target.id == "FORMAT_ID":
if isinstance(stmt.value, ast.Constant) and isinstance(stmt.value.value, str):
value = stmt.value.value.strip()
return value or None
elif isinstance(stmt, ast.AnnAssign):
target = stmt.target
if isinstance(target, ast.Name) and target.id == "FORMAT_ID":
value = stmt.value
if isinstance(value, ast.Constant) and isinstance(value.value, str):
text = value.value.strip()
return text or None
return None
def _discover_normalizer_specs(normalizers_dir: Path) -> list[tuple[str, str, str]]:
specs: list[tuple[str, str, str]] = []
for py_file in sorted(normalizers_dir.glob("*.py")):
if py_file.name.startswith("_"):
continue
module_name = py_file.stem
module_path = f"src.core.api_format.conversion.normalizers.{module_name}"
module_specs: list[tuple[str, str, str]] = []
# 优先 AST 发现,避免导入大模块
try:
source = py_file.read_text(encoding="utf-8")
tree = ast.parse(source, filename=str(py_file))
for node in tree.body:
if not isinstance(node, ast.ClassDef):
continue
if not any(_is_format_normalizer_base(base) for base in node.bases):
continue
fmt_id = _extract_format_id_literal(node)
if fmt_id:
module_specs.append((fmt_id, module_path, node.name))
except Exception as e:
logger.warning("[FormatConversionRegistry] AST 扫描 {} 失败: {}", module_path, e)
if module_specs:
specs.extend(module_specs)
continue
# AST 无法识别时,回退到反射发现(保持兼容)
try:
mod = importlib.import_module(module_path)
except Exception as e:
logger.error("[FormatConversionRegistry] 导入 {} 失败: {}", module_path, e)
continue
for _attr_name, obj in inspect.getmembers(mod, inspect.isclass):
if (
issubclass(obj, FormatNormalizer)
and obj is not FormatNormalizer
and hasattr(obj, "FORMAT_ID")
and obj.__module__ == mod.__name__
):
fmt_id = str(getattr(obj, "FORMAT_ID", "")).strip()
if fmt_id:
module_specs.append((fmt_id, module_path, obj.__name__))
if not module_specs:
logger.warning("[FormatConversionRegistry] 未在 {} 发现可注册 normalizer", module_path)
continue
specs.extend(module_specs)
return specs
def register_default_normalizers() -> None:
"""自动发现并注册 normalizers/ 目录下的所有 FormatNormalizer 实现"""
"""自动发现并注册 normalizers/ 目录下的所有 FormatNormalizer 实现"""
global _DEFAULT_NORMALIZERS_REGISTERED # noqa: PLW0603 - module-level 缓存
# 快速路径:已注册则直接返回(无锁)
@@ -495,43 +690,13 @@ def register_default_normalizers() -> None:
if _DEFAULT_NORMALIZERS_REGISTERED:
return
import importlib
import inspect
from pathlib import Path
normalizers_dir = Path(__file__).parent / "normalizers"
for py_file in sorted(normalizers_dir.glob("*.py")):
if py_file.name.startswith("_"):
continue
module_name = py_file.stem
module_path = f"src.core.api_format.conversion.normalizers.{module_name}"
try:
mod = importlib.import_module(module_path)
except Exception as e:
logger.error("[FormatConversionRegistry] 导入 {} 失败: {}", module_path, e)
continue
for _attr_name, obj in inspect.getmembers(mod, inspect.isclass):
if (
issubclass(obj, FormatNormalizer)
and obj is not FormatNormalizer
and hasattr(obj, "FORMAT_ID")
and obj.__module__ == mod.__name__
):
fmt_id = str(obj.FORMAT_ID).upper()
if format_conversion_registry.get_normalizer(fmt_id) is not None:
logger.warning(
"[FormatConversionRegistry] FORMAT_ID '{}' 重复注册,{} 将覆盖已有实现",
fmt_id,
obj.__name__,
)
try:
format_conversion_registry.register(obj())
except Exception as e:
logger.error("[FormatConversionRegistry] 注册 {} 失败: {}", obj.__name__, e)
for fmt_id, module_path, class_name in _discover_normalizer_specs(normalizers_dir):
format_conversion_registry.register_lazy(fmt_id, module_path, class_name)
_DEFAULT_NORMALIZERS_REGISTERED = True
logger.info(
"[FormatConversionRegistry] 已注册 {} 个 normalizer",
"[FormatConversionRegistry] 已注册 {} 个 normalizer",
len(format_conversion_registry.list_normalizers()),
)

View File

@@ -15,10 +15,7 @@ import hashlib
import threading
import time
from collections import OrderedDict
from cryptography.fernet import Fernet
from cryptography.hazmat.primitives import hashes
from cryptography.hazmat.primitives.kdf.pbkdf2 import PBKDF2HMAC
from typing import TYPE_CHECKING, Any, cast
from src.core.logger import logger
from src.utils.perf import PerfRecorder
@@ -26,6 +23,9 @@ from src.utils.perf import PerfRecorder
from ..config import config
from ..core.exceptions import DecryptionException
if TYPE_CHECKING:
from cryptography.fernet import Fernet
class CryptoService:
"""
@@ -36,6 +36,7 @@ class CryptoService:
"""
_instance: CryptoService | None = None
_instance_lock = threading.Lock()
_cipher: Fernet | None = None
_key_source: str = "unknown" # 记录密钥来源,用于调试
@@ -45,12 +46,17 @@ class CryptoService:
def __new__(cls) -> CryptoService:
if cls._instance is None:
cls._instance = super().__new__(cls)
cls._instance._initialize()
with cls._instance_lock:
if cls._instance is None:
inst = super().__new__(cls)
inst._initialize()
cls._instance = inst
return cls._instance
def _initialize(self) -> None:
"""初始化加密服务"""
from cryptography.fernet import Fernet
logger.info("初始化加密服务")
encryption_key = config.encryption_key
@@ -99,6 +105,10 @@ class CryptoService:
Returns:
Fernet 兼容的 base64 编码密钥
"""
from cryptography.fernet import Fernet
from cryptography.hazmat.primitives import hashes
from cryptography.hazmat.primitives.kdf.pbkdf2 import PBKDF2HMAC
# 首先尝试直接作为 Fernet 密钥使用
try:
key_bytes = (
@@ -235,5 +245,19 @@ class CryptoService:
self._decrypt_cache.popitem(last=False)
# 创建全局加密服务实例
crypto_service = CryptoService()
def get_crypto_service() -> CryptoService:
"""获取加密服务单例(首次使用时才会初始化)。"""
return CryptoService()
class _LazyCryptoServiceProxy:
"""延迟代理,避免 import 阶段触发 cryptography 重载。"""
def __getattr__(self, name: str) -> Any:
return getattr(get_crypto_service(), name)
if TYPE_CHECKING:
crypto_service = CryptoService()
else:
crypto_service = cast(CryptoService, _LazyCryptoServiceProxy())

View File

@@ -1,106 +0,0 @@
"""
优化工具类 - 包含Token计数和响应头管理
"""
from typing import Any
import tiktoken
class TokenCounter:
"""
改进的Token计数器
支持多种模型的准确计数
"""
# 模型到编码器的映射
MODEL_TO_ENCODING = {
"gpt-4": "cl100k_base",
"gpt-3.5-turbo": "cl100k_base",
"claude-3": "cl100k_base", # Claude使用类似的tokenizer
"claude-2": "cl100k_base",
}
def __init__(self) -> None:
self._encodings = {}
self._default_encoding = None
def _get_encoding(self, model: str) -> Any:
"""获取模型对应的编码器"""
# 标准化模型名称
model_base = model.lower().split("-")[0]
if model_base not in self._encodings:
encoding_name = self.MODEL_TO_ENCODING.get(model_base, "cl100k_base") # 默认编码器
try:
self._encodings[model_base] = tiktoken.get_encoding(encoding_name)
except Exception:
# 如果失败,使用默认编码器
if not self._default_encoding:
self._default_encoding = tiktoken.get_encoding("cl100k_base")
self._encodings[model_base] = self._default_encoding
return self._encodings[model_base]
def count_tokens(self, text: str, model: str = "claude-3") -> int:
"""
精确计算文本的token数量
"""
if not text:
return 0
try:
encoding = self._get_encoding(model)
return len(encoding.encode(text))
except Exception:
# 降级到简单估算
return len(text) // 4
def count_messages_tokens(self, messages: list, model: str = "claude-3") -> int:
"""
计算消息列表的总token数
"""
total = 0
for message in messages:
if isinstance(message, dict):
# 计算角色标记
total += 4 # 角色和分隔符的开销
# 计算内容
content = message.get("content", "")
if isinstance(content, str):
total += self.count_tokens(content, model)
elif isinstance(content, list):
# 处理多模态内容
for item in content:
if isinstance(item, dict) and "text" in item:
total += self.count_tokens(item["text"], model)
return total
def estimate_response_tokens(self, response: Any, model: str = "claude-3") -> int:
"""
估算响应的token数量
"""
if isinstance(response, dict):
# 尝试从响应中提取内容
if "content" in response:
content = response["content"]
if isinstance(content, list):
text = " ".join(
item.get("text", "") for item in content if isinstance(item, dict)
)
else:
text = str(content)
return self.count_tokens(text, model)
elif "choices" in response:
# OpenAI格式
total = 0
for choice in response.get("choices", []):
message = choice.get("message", {})
content = message.get("content", "")
total += self.count_tokens(content, model)
return total
# 降级到简单估算
return len(str(response)) // 4

View File

@@ -1,214 +0,0 @@
"""
提供商健康度管理
基于简单的失败计数和优先级调整
"""
import time
from collections import defaultdict
from typing import Any
class ProviderHealthTracker:
"""
追踪提供商的健康状态
根据失败率动态调整优先级
"""
def __init__(
self,
failure_window: int = 300, # 5分钟时间窗口
failure_threshold: int = 3, # 3次失败降低优先级
recovery_time: int = 600, # 10分钟后重置
):
self.failure_window = failure_window
self.failure_threshold = failure_threshold
self.recovery_time = recovery_time
# 存储每个提供商的失败记录
self.failures: dict[str, list] = defaultdict(list)
# 存储每个提供商的成功记录
self.successes: dict[str, list] = defaultdict(list)
# 存储优先级调整
self.priority_adjustments: dict[str, int] = {}
def record_success(self, provider_name: str) -> None:
"""记录成功的请求"""
current_time = time.time()
# 记录成功时间
self.successes[provider_name].append(current_time)
# 清理旧记录
self._cleanup_old_records(provider_name, current_time)
# 如果连续成功,可以恢复优先级
if len(self.successes[provider_name]) >= 5:
if self.priority_adjustments.get(provider_name, 0) < 0:
self.priority_adjustments[provider_name] += 1
def record_failure(self, provider_name: str) -> None:
"""记录失败的请求"""
current_time = time.time()
# 记录失败时间
self.failures[provider_name].append(current_time)
# 清理旧记录
self._cleanup_old_records(provider_name, current_time)
# 检查是否需要降低优先级
recent_failures = len(self.failures[provider_name])
if recent_failures >= self.failure_threshold:
# 降低优先级
current_adjustment = self.priority_adjustments.get(provider_name, 0)
self.priority_adjustments[provider_name] = current_adjustment - 1
def get_priority_adjustment(self, provider_name: str) -> int:
"""
获取优先级调整值
负数表示降低优先级,正数表示提高优先级
"""
return self.priority_adjustments.get(provider_name, 0)
def get_health_status(self, provider_name: str) -> dict:
"""
获取提供商的健康状态
"""
current_time = time.time()
self._cleanup_old_records(provider_name, current_time)
recent_failures = len(self.failures[provider_name])
recent_successes = len(self.successes[provider_name])
total_requests = recent_failures + recent_successes
failure_rate = recent_failures / total_requests if total_requests > 0 else 0
return {
"provider": provider_name,
"recent_failures": recent_failures,
"recent_successes": recent_successes,
"failure_rate": failure_rate,
"priority_adjustment": self.get_priority_adjustment(provider_name),
"status": self._get_status_label(failure_rate, recent_failures),
}
def _cleanup_old_records(self, provider_name: str, current_time: float) -> None:
"""清理超出时间窗口的记录"""
# 清理失败记录
self.failures[provider_name] = [
t for t in self.failures[provider_name] if current_time - t < self.failure_window
]
# 清理成功记录
self.successes[provider_name] = [
t for t in self.successes[provider_name] if current_time - t < self.failure_window
]
# 如果很久没有失败,重置优先级调整
if not self.failures[provider_name] and self.priority_adjustments.get(provider_name, 0) < 0:
# 检查恢复时间
if all(current_time - t > self.recovery_time for t in self.successes[provider_name]):
self.priority_adjustments[provider_name] = 0
# 清理已无记录且无优先级调整的 key防止 dict 无限增长
if (
not self.failures[provider_name]
and not self.successes[provider_name]
and self.priority_adjustments.get(provider_name, 0) == 0
):
self.failures.pop(provider_name, None)
self.successes.pop(provider_name, None)
self.priority_adjustments.pop(provider_name, None)
def _get_status_label(self, failure_rate: float, recent_failures: int) -> str:
"""根据失败率返回状态标签"""
if recent_failures >= self.failure_threshold:
return "degraded" # 降级
elif failure_rate > 0.5:
return "unstable" # 不稳定
elif failure_rate > 0.1:
return "warning" # 警告
else:
return "healthy" # 健康
def should_use_provider(self, provider_name: str) -> bool:
"""
判断是否应该使用该提供商
简单的策略:如果优先级调整低于-3暂时不使用
"""
adjustment = self.get_priority_adjustment(provider_name)
return adjustment > -3
def reset_provider_health(self, provider_name: str) -> None:
"""重置提供商的健康状态(管理员手动操作)"""
self.failures[provider_name] = []
self.successes[provider_name] = []
self.priority_adjustments[provider_name] = 0
class SimpleProviderSelector:
"""
简单的提供商选择器
基于优先级和健康状态
"""
def __init__(self, health_tracker: ProviderHealthTracker):
self.health_tracker = health_tracker
def select_provider(self, providers: list, specified_provider: str | None = None) -> Any:
"""
选择提供商
Args:
providers: 可用提供商列表(已按基础优先级排序)
specified_provider: 用户指定的提供商
Returns:
选中的提供商
"""
# 如果用户指定了提供商,直接使用(不管健康状态)
if specified_provider:
return next((p for p in providers if p.name == specified_provider), None)
# 否则,根据优先级和健康状态选择
# 对提供商列表进行动态排序
sorted_providers = sorted(
providers,
key=lambda p: (
p.priority + self.health_tracker.get_priority_adjustment(p.name),
-p.id, # 相同优先级时使用ID作为次要排序
),
reverse=True, # 优先级高的在前
)
# 选择第一个健康的提供商
for provider in sorted_providers:
if self.health_tracker.should_use_provider(provider.name):
return provider
# 如果都不健康,还是返回第一个(降级策略)
return sorted_providers[0] if sorted_providers else None
def get_provider_rankings(self, providers: list) -> list:
"""
获取提供商的当前排名(用于调试和监控)
"""
rankings = []
for provider in providers:
health_status = self.health_tracker.get_health_status(provider.name)
effective_priority = provider.priority + health_status["priority_adjustment"]
rankings.append(
{
"name": provider.name,
"base_priority": provider.priority,
"adjustment": health_status["priority_adjustment"],
"effective_priority": effective_priority,
"status": health_status["status"],
"failure_rate": health_status["failure_rate"],
}
)
# 按有效优先级排序
rankings.sort(key=lambda x: x["effective_priority"], reverse=True)
return rankings

View File

@@ -519,11 +519,13 @@ async def enrich_auth_config(
"""Enrich auth_config with non-secret metadata (email/account_id).
各 provider 的 enrichment 逻辑通过 register_auth_enricher 注册。
注意: ensure_providers_bootstrapped() 在应用启动时(main.py lifespan)已显式调用
为支持按需 bootstrap这里会尝试按 provider_type 触发插件注册
"""
from src.core.provider_types import normalize_provider_type
from src.services.provider.envelope import ensure_providers_bootstrapped
pt = normalize_provider_type(provider_type)
ensure_providers_bootstrapped(provider_types=[pt] if pt else None)
enricher = _auth_enrichers.get(pt)
if enricher:
return await enricher(auth_config, token_response, access_token, proxy_config)