mirror of
https://github.com/Nezumi-2711/astrbot_plugin_qq_group_daily_analysis.git
synced 2026-09-23 04:09:59 +00:00
fix: 正确的导入各重构模块
This commit is contained in:
@@ -1,29 +1 @@
|
||||
# 基础设施层
|
||||
# 持久化
|
||||
# LLM
|
||||
# 配置
|
||||
# 弹性/容错
|
||||
from . import config, llm, persistence, platform, resilience
|
||||
|
||||
__all__ = [
|
||||
"config",
|
||||
"llm",
|
||||
"persistence",
|
||||
"platform",
|
||||
"resilience",
|
||||
# 平台
|
||||
"PlatformAdapter",
|
||||
"PlatformAdapterFactory",
|
||||
"OneBotAdapter",
|
||||
# 持久化
|
||||
"HistoryRepository",
|
||||
# LLM
|
||||
"LLMClient",
|
||||
# 配置
|
||||
"ConfigManager",
|
||||
# 弹性
|
||||
"CircuitBreaker",
|
||||
"RateLimiter",
|
||||
"retry_async",
|
||||
"RetryConfig",
|
||||
]
|
||||
|
||||
@@ -6,8 +6,8 @@ LLM API请求处理工具模块
|
||||
import asyncio
|
||||
from typing import Any
|
||||
|
||||
from ...utils.logger import logger
|
||||
from ...utils.resilience import CircuitBreaker, global_llm_rate_limiter
|
||||
from ....utils.logger import logger
|
||||
from ....utils.resilience import CircuitBreaker, global_llm_rate_limiter
|
||||
|
||||
_circuit_breakers = {}
|
||||
|
||||
|
||||
@@ -1,7 +0,0 @@
|
||||
"""
|
||||
LLM Module - LLM client implementations
|
||||
"""
|
||||
|
||||
from .llm_client import LLMClient
|
||||
|
||||
__all__ = ["LLMClient"]
|
||||
@@ -1,187 +0,0 @@
|
||||
"""
|
||||
LLM 客户端 - 包装 AstrBot 的 LLM 提供商系统
|
||||
|
||||
该模块提供了一个访问 AstrBot LLM 功能的清晰接口,
|
||||
抽象了提供商管理的细节。
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from ...domain.exceptions import LLMException, LLMRateLimitException
|
||||
from ...domain.value_objects.statistics import TokenUsage
|
||||
from ...utils.logger import logger
|
||||
|
||||
|
||||
class LLMClient:
|
||||
"""
|
||||
用于与 LLM 提供商交互的客户端。
|
||||
|
||||
该类包装了 AstrBot 的提供商系统,并提供了一个
|
||||
清晰的接口来进行 LLM 调用。
|
||||
"""
|
||||
|
||||
def __init__(self, context: Any):
|
||||
"""
|
||||
初始化 LLM 客户端。
|
||||
|
||||
Args:
|
||||
context: 具有提供商访问权限的 AstrBot 插件上下文
|
||||
"""
|
||||
self.context = context
|
||||
self._provider_cache: dict[str, Any] = {}
|
||||
|
||||
def get_provider(self, provider_id: str | None = None) -> Any:
|
||||
"""
|
||||
通过 ID 获取 LLM 提供商。
|
||||
|
||||
Args:
|
||||
provider_id: 特定的提供商 ID,None 表示默认
|
||||
|
||||
Returns:
|
||||
提供商实例
|
||||
|
||||
Raises:
|
||||
LLMException: 如果未找到提供商
|
||||
"""
|
||||
try:
|
||||
if provider_id and provider_id in self._provider_cache:
|
||||
return self._provider_cache[provider_id]
|
||||
|
||||
if provider_id:
|
||||
provider = self.context.get_provider_by_id(provider_id)
|
||||
else:
|
||||
# 获取默认提供商
|
||||
providers = self.context.get_all_providers()
|
||||
if not providers:
|
||||
raise LLMException("无可用 LLM 提供商")
|
||||
provider = providers[0]
|
||||
|
||||
if provider:
|
||||
self._provider_cache[provider_id or "default"] = provider
|
||||
|
||||
return provider
|
||||
|
||||
except Exception as e:
|
||||
raise LLMException(f"获取提供商失败: {e}")
|
||||
|
||||
async def chat_completion(
|
||||
self,
|
||||
prompt: str,
|
||||
provider_id: str | None = None,
|
||||
max_tokens: int = 2000,
|
||||
temperature: float = 0.7,
|
||||
system_prompt: str | None = None,
|
||||
) -> tuple[str, TokenUsage]:
|
||||
"""
|
||||
发起聊天完成请求。
|
||||
|
||||
Args:
|
||||
prompt: 用户提示词
|
||||
provider_id: 特定的提供商 ID (可选)
|
||||
max_tokens: 响应中的最大 token 数
|
||||
temperature: 采样温度
|
||||
system_prompt: 可选的系统提示词
|
||||
|
||||
Returns:
|
||||
(response_text, token_usage) 元组
|
||||
|
||||
Raises:
|
||||
LLMException: 如果请求失败
|
||||
"""
|
||||
try:
|
||||
provider = self.get_provider(provider_id)
|
||||
if not provider:
|
||||
raise LLMException("无可用提供商", provider_id or "default")
|
||||
|
||||
# 构建消息
|
||||
messages = []
|
||||
if system_prompt:
|
||||
messages.append({"role": "system", "content": system_prompt})
|
||||
messages.append({"role": "user", "content": prompt})
|
||||
|
||||
# 发起请求
|
||||
response = await provider.text_chat(
|
||||
messages=messages,
|
||||
session_id=None, # 无状态
|
||||
)
|
||||
|
||||
# 提取响应文本
|
||||
if hasattr(response, "completion_text"):
|
||||
response_text = response.completion_text
|
||||
elif isinstance(response, dict):
|
||||
response_text = response.get(
|
||||
"completion_text", response.get("text", "")
|
||||
)
|
||||
else:
|
||||
response_text = str(response)
|
||||
|
||||
# 提取 token 使用情况
|
||||
token_usage = TokenUsage()
|
||||
if hasattr(response, "usage"):
|
||||
usage = response.usage
|
||||
if hasattr(usage, "prompt_tokens"):
|
||||
token_usage = TokenUsage(
|
||||
prompt_tokens=usage.prompt_tokens or 0,
|
||||
completion_tokens=usage.completion_tokens or 0,
|
||||
total_tokens=usage.total_tokens or 0,
|
||||
)
|
||||
|
||||
return response_text, token_usage
|
||||
|
||||
except Exception as e:
|
||||
error_msg = str(e).lower()
|
||||
if "rate limit" in error_msg or "429" in error_msg:
|
||||
raise LLMRateLimitException(str(e), provider_id or "default")
|
||||
raise LLMException(f"聊天完成请求失败: {e}", provider_id or "default")
|
||||
|
||||
async def analyze_with_json_output(
|
||||
self,
|
||||
prompt: str,
|
||||
provider_id: str | None = None,
|
||||
max_tokens: int = 2000,
|
||||
temperature: float = 0.7,
|
||||
) -> tuple[str, TokenUsage]:
|
||||
"""
|
||||
发起期望 JSON 输出的完成请求。
|
||||
|
||||
Args:
|
||||
prompt: 分析提示词
|
||||
provider_id: 特定的提供商 ID (可选)
|
||||
max_tokens: 响应中的最大 token 数
|
||||
temperature: 采样温度
|
||||
|
||||
Returns:
|
||||
(response_text, token_usage) 元组
|
||||
"""
|
||||
# 如果提示词中没有 JSON 指令,则添加
|
||||
json_instruction = "\nRespond with valid JSON only."
|
||||
if "json" not in prompt.lower():
|
||||
prompt = prompt + json_instruction
|
||||
|
||||
return await self.chat_completion(
|
||||
prompt=prompt,
|
||||
provider_id=provider_id,
|
||||
max_tokens=max_tokens,
|
||||
temperature=temperature,
|
||||
)
|
||||
|
||||
def list_available_providers(self) -> list[dict[str, str]]:
|
||||
"""
|
||||
列出所有可用的 LLM 提供商。
|
||||
|
||||
Returns:
|
||||
提供商信息字典列表
|
||||
"""
|
||||
try:
|
||||
providers = self.context.get_all_providers()
|
||||
return [
|
||||
{
|
||||
"id": getattr(p, "id", str(i)),
|
||||
"name": getattr(p, "name", f"Provider {i}"),
|
||||
"type": getattr(p, "type", "unknown"),
|
||||
}
|
||||
for i, p in enumerate(providers)
|
||||
]
|
||||
except Exception as e:
|
||||
logger.error(f"列出提供商失败: {e}")
|
||||
return []
|
||||
@@ -0,0 +1,53 @@
|
||||
"""
|
||||
消息发送器 - 基础设施层
|
||||
提供高层消息发送接口,支持跨平台智能路由。
|
||||
"""
|
||||
|
||||
from ...utils.logger import logger
|
||||
|
||||
|
||||
class MessageSender:
|
||||
"""
|
||||
消息发送器
|
||||
封装了 PlatformAdapter 的底层调用,提供更高层的发送接口
|
||||
"""
|
||||
|
||||
def __init__(self, bot_manager, config_manager, retry_manager):
|
||||
self.bot_manager = bot_manager
|
||||
self.config_manager = config_manager
|
||||
self.retry_manager = retry_manager
|
||||
|
||||
async def send_text(
|
||||
self, group_id: str, text: str, platform_id: str = None
|
||||
) -> bool:
|
||||
"""发送文本消息"""
|
||||
adapter = self.bot_manager.get_adapter(platform_id)
|
||||
if not adapter:
|
||||
logger.error(f"[MessageSender] 未找到平台 {platform_id} 的适配器")
|
||||
return False
|
||||
return await adapter.send_text(group_id, text)
|
||||
|
||||
async def send_image_smart(
|
||||
self, group_id: str, image_url: str, caption: str = "", platform_id: str = None
|
||||
) -> bool:
|
||||
"""智能发送图片,支持自动选择适配器"""
|
||||
adapter = self.bot_manager.get_adapter(platform_id)
|
||||
if not adapter:
|
||||
logger.error(f"[MessageSender] 未找到平台 {platform_id} 的适配器")
|
||||
return False
|
||||
return await adapter.send_image(group_id, image_url, caption)
|
||||
|
||||
async def send_pdf(
|
||||
self, group_id: str, pdf_path: str, caption: str = "", platform_id: str = None
|
||||
) -> bool:
|
||||
"""发送 PDF 文件"""
|
||||
adapter = self.bot_manager.get_adapter(platform_id)
|
||||
if not adapter:
|
||||
logger.error(f"[MessageSender] 未找到平台 {platform_id} 的适配器")
|
||||
return False
|
||||
return await adapter.send_file(group_id, pdf_path)
|
||||
|
||||
def _get_available_platforms(self, group_id: str):
|
||||
"""获取可用的平台列表 (Helper for Dispatcher)"""
|
||||
# 简单实现:返回所有已加载的平台
|
||||
return [(pid, None) for pid in self.bot_manager.get_platform_ids()]
|
||||
@@ -1,8 +1,8 @@
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
from ..utils.logger import logger
|
||||
from ..utils.trace_context import TraceContext
|
||||
from ...utils.logger import logger
|
||||
from ...utils.trace_context import TraceContext
|
||||
|
||||
|
||||
class ReportDispatcher:
|
||||
|
||||
@@ -1,15 +0,0 @@
|
||||
"""
|
||||
弹性模块 - 断路器、速率限制器和重试工具
|
||||
"""
|
||||
|
||||
from .circuit_breaker import CircuitBreaker, CircuitState
|
||||
from .rate_limiter import RateLimiter
|
||||
from .retry import RetryConfig, retry_async
|
||||
|
||||
__all__ = [
|
||||
"CircuitBreaker",
|
||||
"CircuitState",
|
||||
"RateLimiter",
|
||||
"retry_async",
|
||||
"RetryConfig",
|
||||
]
|
||||
@@ -1,137 +0,0 @@
|
||||
"""
|
||||
断路器 - 防止级联故障
|
||||
|
||||
实现断路器模式,防止对失败服务的重复调用。
|
||||
"""
|
||||
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
|
||||
from ...utils.logger import logger
|
||||
|
||||
|
||||
class CircuitState(Enum):
|
||||
"""断路器状态。"""
|
||||
|
||||
CLOSED = "closed" # 正常运行
|
||||
OPEN = "open" # 故障中,拒绝调用
|
||||
HALF_OPEN = "half_open" # 测试服务是否恢复
|
||||
|
||||
|
||||
@dataclass
|
||||
class CircuitBreaker:
|
||||
"""
|
||||
断路器实现。
|
||||
|
||||
通过跟踪故障率并临时阻止对故障服务的调用
|
||||
来防止级联故障。
|
||||
"""
|
||||
|
||||
name: str
|
||||
failure_threshold: int = 5
|
||||
recovery_timeout: float = 30.0
|
||||
half_open_max_calls: int = 3
|
||||
|
||||
# 内部状态
|
||||
_state: CircuitState = field(default=CircuitState.CLOSED, init=False)
|
||||
_failure_count: int = field(default=0, init=False)
|
||||
_success_count: int = field(default=0, init=False)
|
||||
_last_failure_time: float = field(default=0, init=False)
|
||||
_half_open_calls: int = field(default=0, init=False)
|
||||
|
||||
@property
|
||||
def state(self) -> CircuitState:
|
||||
"""获取当前断路器状态,检查是否恢复。"""
|
||||
if self._state == CircuitState.OPEN:
|
||||
if time.time() - self._last_failure_time >= self.recovery_timeout:
|
||||
self._transition_to(CircuitState.HALF_OPEN)
|
||||
return self._state
|
||||
|
||||
def _transition_to(self, new_state: CircuitState) -> None:
|
||||
"""转换到新状态。"""
|
||||
old_state = self._state
|
||||
self._state = new_state
|
||||
|
||||
if new_state == CircuitState.CLOSED:
|
||||
self._failure_count = 0
|
||||
self._success_count = 0
|
||||
elif new_state == CircuitState.HALF_OPEN:
|
||||
self._half_open_calls = 0
|
||||
|
||||
logger.debug(f"断路器 {self.name}: {old_state.value} -> {new_state.value}")
|
||||
|
||||
def record_success(self) -> None:
|
||||
"""记录成功调用。"""
|
||||
if self._state == CircuitState.HALF_OPEN:
|
||||
self._success_count += 1
|
||||
if self._success_count >= self.half_open_max_calls:
|
||||
self._transition_to(CircuitState.CLOSED)
|
||||
elif self._state == CircuitState.CLOSED:
|
||||
# 成功时重置故障计数
|
||||
self._failure_count = 0
|
||||
|
||||
def record_failure(self) -> None:
|
||||
"""记录失败调用。"""
|
||||
self._failure_count += 1
|
||||
self._last_failure_time = time.time()
|
||||
|
||||
if self._state == CircuitState.HALF_OPEN:
|
||||
self._transition_to(CircuitState.OPEN)
|
||||
elif self._state == CircuitState.CLOSED:
|
||||
if self._failure_count >= self.failure_threshold:
|
||||
self._transition_to(CircuitState.OPEN)
|
||||
|
||||
def can_execute(self) -> bool:
|
||||
"""检查是否可以执行调用。"""
|
||||
state = self.state # 这可能触发状态转换
|
||||
|
||||
if state == CircuitState.CLOSED:
|
||||
return True
|
||||
elif state == CircuitState.OPEN:
|
||||
return False
|
||||
elif state == CircuitState.HALF_OPEN:
|
||||
self._half_open_calls += 1
|
||||
return self._half_open_calls <= self.half_open_max_calls
|
||||
|
||||
return False
|
||||
|
||||
def reset(self) -> None:
|
||||
"""重置断路器到关闭状态。"""
|
||||
self._transition_to(CircuitState.CLOSED)
|
||||
|
||||
async def execute(
|
||||
self,
|
||||
func: Callable,
|
||||
*args,
|
||||
fallback: Callable | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
使用断路器保护执行函数。
|
||||
|
||||
参数:
|
||||
func: 要执行的异步函数
|
||||
*args: 函数参数
|
||||
fallback: 断路器打开时的可选降级函数
|
||||
**kwargs: 函数关键字参数
|
||||
|
||||
返回:
|
||||
函数结果或降级结果
|
||||
|
||||
异常:
|
||||
Exception: 如果断路器打开且没有提供降级函数
|
||||
"""
|
||||
if not self.can_execute():
|
||||
if fallback:
|
||||
return await fallback(*args, **kwargs)
|
||||
raise Exception(f"断路器 {self.name} 已打开")
|
||||
|
||||
try:
|
||||
result = await func(*args, **kwargs)
|
||||
self.record_success()
|
||||
return result
|
||||
except Exception:
|
||||
self.record_failure()
|
||||
raise
|
||||
@@ -1,139 +0,0 @@
|
||||
"""
|
||||
速率限制器 - 控制请求速率
|
||||
|
||||
实现令牌桶速率限制,防止服务过载。
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
|
||||
@dataclass
|
||||
class RateLimiter:
|
||||
"""
|
||||
令牌桶速率限制器。
|
||||
|
||||
使用令牌桶算法控制操作速率。
|
||||
"""
|
||||
|
||||
name: str
|
||||
rate: float # 每秒令牌数
|
||||
burst: int # 最大突发大小(桶容量)
|
||||
|
||||
# 内部状态
|
||||
_tokens: float = field(default=0, init=False)
|
||||
_last_update: float = field(default=0, init=False)
|
||||
_lock: asyncio.Lock = field(default_factory=asyncio.Lock, init=False)
|
||||
|
||||
def __post_init__(self):
|
||||
"""初始化令牌桶。"""
|
||||
self._tokens = float(self.burst)
|
||||
self._last_update = time.time()
|
||||
|
||||
def _refill(self) -> None:
|
||||
"""根据经过的时间补充令牌。"""
|
||||
now = time.time()
|
||||
elapsed = now - self._last_update
|
||||
self._tokens = min(self.burst, self._tokens + elapsed * self.rate)
|
||||
self._last_update = now
|
||||
|
||||
async def acquire(self, tokens: int = 1, timeout: float | None = None) -> bool:
|
||||
"""
|
||||
从桶中获取令牌。
|
||||
|
||||
参数:
|
||||
tokens: 要获取的令牌数
|
||||
timeout: 最大等待时间(None = 无限等待)
|
||||
|
||||
返回:
|
||||
如果获取到令牌返回 True,超时返回 False
|
||||
"""
|
||||
start_time = time.time()
|
||||
|
||||
async with self._lock:
|
||||
while True:
|
||||
self._refill()
|
||||
|
||||
if self._tokens >= tokens:
|
||||
self._tokens -= tokens
|
||||
return True
|
||||
|
||||
if timeout is not None:
|
||||
elapsed = time.time() - start_time
|
||||
if elapsed >= timeout:
|
||||
return False
|
||||
|
||||
# 计算获取足够令牌的等待时间
|
||||
tokens_needed = tokens - self._tokens
|
||||
wait_time = tokens_needed / self.rate
|
||||
|
||||
if timeout is not None:
|
||||
remaining = timeout - (time.time() - start_time)
|
||||
wait_time = min(wait_time, remaining)
|
||||
|
||||
if wait_time > 0:
|
||||
await asyncio.sleep(wait_time)
|
||||
|
||||
def try_acquire(self, tokens: int = 1) -> bool:
|
||||
"""
|
||||
尝试获取令牌而不等待。
|
||||
|
||||
参数:
|
||||
tokens: 要获取的令牌数
|
||||
|
||||
返回:
|
||||
如果获取到令牌返回 True,否则返回 False
|
||||
"""
|
||||
self._refill()
|
||||
|
||||
if self._tokens >= tokens:
|
||||
self._tokens -= tokens
|
||||
return True
|
||||
return False
|
||||
|
||||
@property
|
||||
def available_tokens(self) -> float:
|
||||
"""获取当前可用令牌数。"""
|
||||
self._refill()
|
||||
return self._tokens
|
||||
|
||||
def reset(self) -> None:
|
||||
"""重置速率限制器到满容量。"""
|
||||
self._tokens = float(self.burst)
|
||||
self._last_update = time.time()
|
||||
|
||||
|
||||
class RateLimiterGroup:
|
||||
"""
|
||||
不同操作的速率限制器组。
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self._limiters: dict[str, RateLimiter] = {}
|
||||
|
||||
def get_or_create(
|
||||
self,
|
||||
name: str,
|
||||
rate: float = 1.0,
|
||||
burst: int = 5,
|
||||
) -> RateLimiter:
|
||||
"""
|
||||
获取或创建速率限制器。
|
||||
|
||||
参数:
|
||||
name: 限制器名称
|
||||
rate: 每秒令牌数
|
||||
burst: 最大突发大小
|
||||
|
||||
返回:
|
||||
RateLimiter 实例
|
||||
"""
|
||||
if name not in self._limiters:
|
||||
self._limiters[name] = RateLimiter(name=name, rate=rate, burst=burst)
|
||||
return self._limiters[name]
|
||||
|
||||
def reset_all(self) -> None:
|
||||
"""重置所有速率限制器。"""
|
||||
for limiter in self._limiters.values():
|
||||
limiter.reset()
|
||||
@@ -1,176 +0,0 @@
|
||||
"""
|
||||
重试 - 带指数退避的重试工具
|
||||
|
||||
提供用于处理瞬态故障的重试装饰器和工具。
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import random
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from functools import wraps
|
||||
|
||||
from ...utils.logger import logger
|
||||
|
||||
|
||||
@dataclass
|
||||
class RetryConfig:
|
||||
"""重试行为配置。"""
|
||||
|
||||
max_attempts: int = 3
|
||||
base_delay: float = 1.0
|
||||
max_delay: float = 60.0
|
||||
exponential_base: float = 2.0
|
||||
jitter: bool = True
|
||||
retry_exceptions: tuple[type[Exception], ...] = (Exception,)
|
||||
|
||||
|
||||
def calculate_delay(
|
||||
attempt: int,
|
||||
base_delay: float,
|
||||
max_delay: float,
|
||||
exponential_base: float,
|
||||
jitter: bool,
|
||||
) -> float:
|
||||
"""
|
||||
计算重试尝试的延迟。
|
||||
|
||||
参数:
|
||||
attempt: 当前尝试次数(从 0 开始)
|
||||
base_delay: 基础延迟(秒)
|
||||
max_delay: 最大延迟(秒)
|
||||
exponential_base: 指数退避的基数
|
||||
jitter: 是否添加随机抖动
|
||||
|
||||
返回:
|
||||
延迟时间(秒)
|
||||
"""
|
||||
delay = base_delay * (exponential_base**attempt)
|
||||
delay = min(delay, max_delay)
|
||||
|
||||
if jitter:
|
||||
delay = delay * (0.5 + random.random())
|
||||
|
||||
return delay
|
||||
|
||||
|
||||
def retry_async(
|
||||
max_attempts: int = 3,
|
||||
base_delay: float = 1.0,
|
||||
max_delay: float = 60.0,
|
||||
exponential_base: float = 2.0,
|
||||
jitter: bool = True,
|
||||
retry_exceptions: tuple[type[Exception], ...] = (Exception,),
|
||||
on_retry: Callable[[Exception, int], None] | None = None,
|
||||
):
|
||||
"""
|
||||
带指数退避的异步函数重试装饰器。
|
||||
|
||||
参数:
|
||||
max_attempts: 最大尝试次数
|
||||
base_delay: 重试之间的基础延迟
|
||||
max_delay: 重试之间的最大延迟
|
||||
exponential_base: 指数退避的基数
|
||||
jitter: 是否添加随机抖动
|
||||
retry_exceptions: 要重试的异常元组
|
||||
on_retry: 重试时的可选回调(异常,尝试次数)
|
||||
|
||||
返回:
|
||||
装饰后的函数
|
||||
"""
|
||||
|
||||
def decorator(func: Callable):
|
||||
@wraps(func)
|
||||
async def wrapper(*args, **kwargs):
|
||||
last_exception = None
|
||||
|
||||
for attempt in range(max_attempts):
|
||||
try:
|
||||
return await func(*args, **kwargs)
|
||||
except retry_exceptions as e:
|
||||
last_exception = e
|
||||
|
||||
if attempt < max_attempts - 1:
|
||||
delay = calculate_delay(
|
||||
attempt, base_delay, max_delay, exponential_base, jitter
|
||||
)
|
||||
|
||||
if on_retry:
|
||||
on_retry(e, attempt + 1)
|
||||
|
||||
logger.debug(
|
||||
f"重试 {attempt + 1}/{max_attempts} {func.__name__} "
|
||||
f"延迟 {delay:.2f}s: {e}"
|
||||
)
|
||||
await asyncio.sleep(delay)
|
||||
else:
|
||||
logger.warning(
|
||||
f"{func.__name__} 的所有 {max_attempts} 次尝试均失败: {e}"
|
||||
)
|
||||
|
||||
raise last_exception
|
||||
|
||||
return wrapper
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
class RetryExecutor:
|
||||
"""
|
||||
带重试逻辑的函数执行器。
|
||||
"""
|
||||
|
||||
def __init__(self, config: RetryConfig | None = None):
|
||||
"""
|
||||
初始化重试执行器。
|
||||
|
||||
参数:
|
||||
config: 重试配置
|
||||
"""
|
||||
self.config = config or RetryConfig()
|
||||
|
||||
async def execute(
|
||||
self,
|
||||
func: Callable,
|
||||
*args,
|
||||
config: RetryConfig | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
使用重试逻辑执行函数。
|
||||
|
||||
参数:
|
||||
func: 要执行的异步函数
|
||||
*args: 函数参数
|
||||
config: 可选的覆盖配置
|
||||
**kwargs: 函数关键字参数
|
||||
|
||||
返回:
|
||||
函数结果
|
||||
|
||||
异常:
|
||||
Exception: 如果所有重试都失败
|
||||
"""
|
||||
cfg = config or self.config
|
||||
last_exception = None
|
||||
|
||||
for attempt in range(cfg.max_attempts):
|
||||
try:
|
||||
return await func(*args, **kwargs)
|
||||
except cfg.retry_exceptions as e:
|
||||
last_exception = e
|
||||
|
||||
if attempt < cfg.max_attempts - 1:
|
||||
delay = calculate_delay(
|
||||
attempt,
|
||||
cfg.base_delay,
|
||||
cfg.max_delay,
|
||||
cfg.exponential_base,
|
||||
cfg.jitter,
|
||||
)
|
||||
logger.debug(
|
||||
f"重试 {attempt + 1}/{cfg.max_attempts} 延迟 {delay:.2f}s: {e}"
|
||||
)
|
||||
await asyncio.sleep(delay)
|
||||
|
||||
raise last_exception
|
||||
@@ -9,10 +9,10 @@ import weakref
|
||||
from apscheduler.triggers.cron import CronTrigger
|
||||
|
||||
from ...utils.logger import logger
|
||||
from ..core.message_sender import MessageSender
|
||||
from ...utils.trace_context import TraceContext
|
||||
from ..messaging.message_sender import MessageSender
|
||||
from ..platform.factory import PlatformAdapterFactory
|
||||
from ..reporting.dispatcher import ReportDispatcher
|
||||
from ..utils.trace_context import TraceContext
|
||||
|
||||
|
||||
class AutoScheduler:
|
||||
|
||||
Reference in New Issue
Block a user