mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
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:
@@ -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()),
|
||||
)
|
||||
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user