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:
Rat
2026-07-08 11:28:16 +08:00
committed by GitHub
co-authored by Rat0323 SXP-Simon
parent d5b0c26289
commit e53016eb01
2 changed files with 114 additions and 100 deletions
+1 -1
View File
@@ -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",
+113 -99
View File
@@ -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