mirror of
https://github.com/Nezumi-2711/astrbot_plugin_qq_group_daily_analysis.git
synced 2026-09-22 05:31:52 +00:00
feat(llm): 重构并优化模型请求的重试与降级补偿机制 (#195)
* refactor: optimize LLM fallback strategy using a unified attempt queue (Strategy B) * refactor: apply exponential backoff, fix format state bleeding, and update typing to Any * fix: avoid false breaker failures in LLM fallback * style: format llm_utils.py with ruff * fix: resolve possibly unbound variable warning in llm_utils.py * refactor: replace loose Any type annotations with precise Context, ConfigManager, JSONValue, and LLMResponse * docs: update call_provider_with_retry docstring and timeout comments --------- Co-authored-by: Rat0323 <261020116+Rat0323@users.noreply.github.com> Co-authored-by: SXP-Simon <sxp20061207@163.com>
This commit is contained in:
+1
-1
@@ -336,7 +336,7 @@
|
|||||||
"type": "int",
|
"type": "int",
|
||||||
"description": "LLM 请求重试退避基值(秒)",
|
"description": "LLM 请求重试退避基值(秒)",
|
||||||
"default": 2,
|
"default": 2,
|
||||||
"hint": "重试之间的基准等待时间(秒),实际等待时间为基值乘以尝试次数。"
|
"hint": "重试之间的基准等待时间(秒)。实际等待时间会随尝试次数指数增长,并加入少量随机抖动;当专用 Provider 失败后,会继续按队列尝试默认 Provider。"
|
||||||
},
|
},
|
||||||
"enable_streaming_llm_call": {
|
"enable_streaming_llm_call": {
|
||||||
"type": "bool",
|
"type": "bool",
|
||||||
|
|||||||
@@ -4,12 +4,15 @@ LLM API请求处理工具模块
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import random
|
||||||
|
|
||||||
from astrbot.api.provider import LLMResponse
|
from astrbot.api.provider import LLMResponse
|
||||||
|
from astrbot.api.star import Context
|
||||||
|
|
||||||
from ....utils.logger import logger
|
from ....utils.logger import logger
|
||||||
from ....utils.resilience import CircuitBreaker, GlobalRateLimiter
|
from ....utils.resilience import CircuitBreaker, GlobalRateLimiter
|
||||||
from .structured_output_schema import JSONObject
|
from ...config.config_manager import ConfigManager
|
||||||
|
from .structured_output_schema import JSONObject, JSONValue
|
||||||
|
|
||||||
_circuit_breakers = {}
|
_circuit_breakers = {}
|
||||||
|
|
||||||
@@ -39,8 +42,8 @@ def _get_circuit_breaker(provider_id: str) -> CircuitBreaker:
|
|||||||
|
|
||||||
|
|
||||||
async def _call_provider_stream(
|
async def _call_provider_stream(
|
||||||
context, provider_id: str, llm_kwargs: dict[str, object]
|
context: Context, provider_id: str, llm_kwargs: dict[str, JSONValue]
|
||||||
):
|
) -> LLMResponse:
|
||||||
provider = context.get_provider_by_id(provider_id=provider_id)
|
provider = context.get_provider_by_id(provider_id=provider_id)
|
||||||
if provider is None:
|
if provider is None:
|
||||||
raise RuntimeError(f"Provider 不存在: {provider_id}")
|
raise RuntimeError(f"Provider 不存在: {provider_id}")
|
||||||
@@ -151,8 +154,8 @@ async def _try_get_first_available_provider_id(context) -> str | None:
|
|||||||
|
|
||||||
|
|
||||||
async def get_provider_id_with_fallback(
|
async def get_provider_id_with_fallback(
|
||||||
context,
|
context: Context,
|
||||||
config_manager,
|
config_manager: ConfigManager,
|
||||||
provider_id_key: str | None,
|
provider_id_key: str | None,
|
||||||
umo: str | None = None,
|
umo: str | None = None,
|
||||||
) -> str | None:
|
) -> str | None:
|
||||||
@@ -235,15 +238,15 @@ async def get_provider_id_with_fallback(
|
|||||||
|
|
||||||
|
|
||||||
async def call_provider_with_retry(
|
async def call_provider_with_retry(
|
||||||
context,
|
context: Context,
|
||||||
config_manager,
|
config_manager: ConfigManager,
|
||||||
prompt: str,
|
prompt: str,
|
||||||
umo: str | None = None,
|
umo: str | None = None,
|
||||||
provider_id_key: str | None = None,
|
provider_id_key: str | None = None,
|
||||||
system_prompt: str | None = None,
|
system_prompt: str | None = None,
|
||||||
response_format: JSONObject | None = None,
|
response_format: JSONObject | None = None,
|
||||||
extra_generate_kwargs: dict[str, object] | None = None,
|
extra_generate_kwargs: dict[str, JSONValue] | None = None,
|
||||||
) -> object | None:
|
) -> LLMResponse | None:
|
||||||
"""
|
"""
|
||||||
调用LLM提供者,带超时、重试与退避。支持自定义服务商和配置化 Provider 选择。
|
调用LLM提供者,带超时、重试与退避。支持自定义服务商和配置化 Provider 选择。
|
||||||
|
|
||||||
@@ -264,111 +267,122 @@ async def call_provider_with_retry(
|
|||||||
# 用户可在 AstrBot WebUI 中为每个 Provider 配置 timeout 参数
|
# 用户可在 AstrBot WebUI 中为每个 Provider 配置 timeout 参数
|
||||||
retries = config_manager.get_llm_retries()
|
retries = config_manager.get_llm_retries()
|
||||||
backoff = config_manager.get_llm_backoff()
|
backoff = config_manager.get_llm_backoff()
|
||||||
|
|
||||||
# 检查流式调用配置
|
|
||||||
enable_streaming_llm_call = config_manager.get_enable_streaming_llm_call()
|
enable_streaming_llm_call = config_manager.get_enable_streaming_llm_call()
|
||||||
|
|
||||||
last_exc = None
|
# 1. 确定我们要尝试的 Provider 队列
|
||||||
for attempt in range(1, retries + 1):
|
attempt_queue = []
|
||||||
|
|
||||||
|
# 尝试获取指定的 Provider
|
||||||
|
specific_provider_id = await get_provider_id_with_fallback(
|
||||||
|
context, config_manager, provider_id_key, umo
|
||||||
|
)
|
||||||
|
if specific_provider_id:
|
||||||
|
attempt_queue.extend([(specific_provider_id, False)] * retries)
|
||||||
|
|
||||||
|
# 尝试获取降级/默认的 Provider (传入 None 走主模型/会话模型路径)
|
||||||
|
if provider_id_key is not None:
|
||||||
|
fallback_provider_id = await get_provider_id_with_fallback(
|
||||||
|
context, config_manager, None, umo
|
||||||
|
)
|
||||||
|
if fallback_provider_id and fallback_provider_id != specific_provider_id:
|
||||||
|
attempt_queue.extend([(fallback_provider_id, True)] * retries)
|
||||||
|
|
||||||
|
if not attempt_queue:
|
||||||
|
logger.error("无可用 Provider,无法调用 llm_generate")
|
||||||
|
return None
|
||||||
|
|
||||||
|
# 2. 核心请求执行闭包
|
||||||
|
async def _execute_llm_request(
|
||||||
|
pid: str, r_format: JSONObject | None
|
||||||
|
) -> LLMResponse:
|
||||||
|
cb = _get_circuit_breaker(pid)
|
||||||
|
if not cb.allow_request():
|
||||||
|
logger.warning(f"Provider {pid} 熔断器已打开,跳过本次请求")
|
||||||
|
raise Exception("Circuit breaker open")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# 使用新的 provider 选择逻辑,获取 Provider ID
|
async with GlobalRateLimiter.get_instance().semaphore:
|
||||||
provider_id = await get_provider_id_with_fallback(
|
llm_kwargs: dict[str, JSONValue] = {
|
||||||
context, config_manager, provider_id_key, umo
|
"chat_provider_id": pid,
|
||||||
)
|
"prompt": prompt,
|
||||||
|
}
|
||||||
|
if system_prompt is not None:
|
||||||
|
llm_kwargs["system_prompt"] = system_prompt
|
||||||
|
if r_format is not None:
|
||||||
|
llm_kwargs["response_format"] = r_format
|
||||||
|
if extra_generate_kwargs:
|
||||||
|
llm_kwargs.update(extra_generate_kwargs)
|
||||||
|
|
||||||
if not provider_id:
|
if enable_streaming_llm_call:
|
||||||
logger.error("provider_id 为空,无法调用 llm_generate,直接返回 None")
|
resp = await _call_provider_stream(context, pid, llm_kwargs)
|
||||||
return None
|
else:
|
||||||
|
resp = await context.llm_generate(**llm_kwargs)
|
||||||
|
cb.record_success()
|
||||||
|
return resp
|
||||||
|
except Exception as err:
|
||||||
|
if r_format is not None and _is_response_format_unsupported_error(err):
|
||||||
|
raise err
|
||||||
|
cb.record_failure()
|
||||||
|
raise err
|
||||||
|
|
||||||
logger.info(
|
# 3. 开始执行队列
|
||||||
f"[LLM 调用] 使用 Provider ID: {provider_id} | "
|
last_exc = None
|
||||||
"max_tokens=provider-default | "
|
current_response_format = response_format
|
||||||
"temperature=provider-default | "
|
|
||||||
f"prompt长度={len(prompt) if prompt else 0}字符"
|
|
||||||
)
|
|
||||||
|
|
||||||
logger.debug(
|
# 记录上一次尝试的 Provider ID,用于判断是否发生切换
|
||||||
f"[LLM 调用] Prompt 前100字符: {prompt[:100] if prompt else 'None'}..."
|
previous_pid = None
|
||||||
)
|
|
||||||
|
|
||||||
if system_prompt:
|
for i, (current_pid, is_fallback) in enumerate(attempt_queue):
|
||||||
logger.debug(
|
attempt_num = i + 1
|
||||||
f"[LLM 调用] System Prompt 前100字符: {system_prompt[:100]}..."
|
is_last_attempt = i == len(attempt_queue) - 1
|
||||||
)
|
|
||||||
|
|
||||||
# 检查 prompt 是否为空
|
# 修复状态污染:如果切换了全新的 Provider,必须重置 response_format 约束
|
||||||
if not prompt or not prompt.strip():
|
if current_pid != previous_pid:
|
||||||
logger.error(
|
current_response_format = response_format
|
||||||
"LLM provider: prompt 为空或只包含空白字符,无法调用 llm_generate"
|
previous_pid = current_pid
|
||||||
)
|
|
||||||
return None
|
|
||||||
|
|
||||||
# 获取熔断器
|
prefix = "[降级补偿] " if is_fallback else "[LLM 调用] "
|
||||||
cb = _get_circuit_breaker(provider_id)
|
logger.info(
|
||||||
if not cb.allow_request():
|
f"{prefix}尝试 #{attempt_num} | Provider ID: {current_pid} | "
|
||||||
logger.warning(f"Provider {provider_id} 熔断器已打开,跳过本次请求")
|
f"prompt长度={len(prompt) if prompt else 0}字符"
|
||||||
return None
|
)
|
||||||
|
|
||||||
# 使用全局限流器 + 熔断器记录
|
if not prompt or not prompt.strip():
|
||||||
# 超时由 Provider 内部控制,无需外层 wait_for
|
logger.error("LLM provider: prompt 为空,无法调用")
|
||||||
try:
|
return None
|
||||||
async with GlobalRateLimiter.get_instance().semaphore:
|
|
||||||
llm_kwargs: dict[str, object] = {
|
|
||||||
"chat_provider_id": provider_id,
|
|
||||||
"prompt": prompt,
|
|
||||||
}
|
|
||||||
if system_prompt is not None:
|
|
||||||
llm_kwargs["system_prompt"] = system_prompt
|
|
||||||
if response_format is not None:
|
|
||||||
llm_kwargs["response_format"] = response_format
|
|
||||||
if extra_generate_kwargs:
|
|
||||||
llm_kwargs.update(extra_generate_kwargs)
|
|
||||||
|
|
||||||
if enable_streaming_llm_call:
|
try:
|
||||||
logger.info("[LLM 调用] 使用流式 Provider 调用")
|
return await _execute_llm_request(current_pid, current_response_format)
|
||||||
|
|
||||||
async def _invoke_llm(pid: str):
|
|
||||||
if enable_streaming_llm_call:
|
|
||||||
return await _call_provider_stream(context, pid, llm_kwargs)
|
|
||||||
return await context.llm_generate(**llm_kwargs)
|
|
||||||
|
|
||||||
try:
|
|
||||||
llm_resp = await _invoke_llm(provider_id)
|
|
||||||
except Exception as e:
|
|
||||||
if (
|
|
||||||
response_format is not None
|
|
||||||
and _is_response_format_unsupported_error(e)
|
|
||||||
):
|
|
||||||
logger.warning(
|
|
||||||
"[LLM 调用] 当前 Provider 可能不支持 response_format,"
|
|
||||||
"已自动降级为无 schema 约束重试本次请求。"
|
|
||||||
)
|
|
||||||
llm_kwargs.pop("response_format", None)
|
|
||||||
llm_resp = await _invoke_llm(provider_id)
|
|
||||||
else:
|
|
||||||
raise
|
|
||||||
|
|
||||||
# 成功记录
|
|
||||||
cb.record_success()
|
|
||||||
return llm_resp
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
# 失败记录
|
|
||||||
cb.record_failure()
|
|
||||||
raise e
|
|
||||||
|
|
||||||
except asyncio.TimeoutError as e:
|
|
||||||
last_exc = e
|
|
||||||
logger.warning(f"LLM请求超时: 第{attempt}次 (Provider 内部超时)")
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
last_exc = e
|
last_exc = e
|
||||||
logger.warning(f"LLM请求失败: 第{attempt}次, 错误: {last_exc}")
|
|
||||||
# 若非最后一次,等待退避后重试
|
|
||||||
if attempt < retries:
|
|
||||||
await asyncio.sleep(backoff * attempt)
|
|
||||||
|
|
||||||
# 最终仍失败,记录错误并返回 None 由调用方处理降级,避免抛出异常
|
# 处理不支持 response_format 的情况
|
||||||
logger.error(f"LLM请求全部重试失败: {last_exc}")
|
if (
|
||||||
|
current_response_format is not None
|
||||||
|
and _is_response_format_unsupported_error(e)
|
||||||
|
):
|
||||||
|
logger.warning(
|
||||||
|
f"{prefix}当前 Provider 可能不支持 response_format,已自动降级为无 schema 约束。"
|
||||||
|
)
|
||||||
|
current_response_format = None
|
||||||
|
# 在当前尝试额度内立即再试一次剥离了 schema 的请求
|
||||||
|
try:
|
||||||
|
return await _execute_llm_request(
|
||||||
|
current_pid, current_response_format
|
||||||
|
)
|
||||||
|
except Exception as inner_e:
|
||||||
|
last_exc = inner_e
|
||||||
|
|
||||||
|
logger.warning(f"{prefix}请求失败: {last_exc}")
|
||||||
|
|
||||||
|
if not is_last_attempt:
|
||||||
|
# Exponential backoff with jitter: backoff * (2 ^ (attempt_num - 1)) + random jitter
|
||||||
|
sleep_time = backoff * (2 ** (attempt_num - 1)) + random.uniform(0, 1)
|
||||||
|
logger.debug(f"等待 {sleep_time:.2f} 秒后重试...")
|
||||||
|
await asyncio.sleep(sleep_time)
|
||||||
|
|
||||||
|
logger.error(f"LLM请求队列全部耗尽,最终失败: {last_exc}")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user