feat: implement request retry mechanism for provider requests (#8893)

* feat: implement request retry mechanism for provider requests

* feat: add request max retries configuration and implement retry logic for provider requests

* feat: update fake_query function to accept request_max_retries parameter

* feat: remove retry_rate_limits from provider request calls
This commit is contained in:
Weilong Liao
2026-06-19 17:13:40 +08:00
committed by GitHub
parent 143f846b92
commit dd36979eca
15 changed files with 467 additions and 35 deletions
@@ -224,6 +224,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
custom_compressor: ContextCompressor | None = None,
tool_schema_mode: str | None = "full",
fallback_providers: list[Provider] | None = None,
request_max_retries: int | None = None,
tool_result_overflow_dir: str | None = None,
read_tool: FunctionTool | None = None,
**kwargs: T.Any,
@@ -237,6 +238,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
self.truncate_turns = truncate_turns
self.custom_token_counter = custom_token_counter
self.custom_compressor = custom_compressor
self.request_max_retries = request_max_retries
self.tool_result_overflow_dir = tool_result_overflow_dir
self.read_tool = read_tool
self._tool_result_token_counter = EstimateTokenCounter()
@@ -463,6 +465,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
"session_id": self.req.session_id,
"extra_user_content_parts": self.req.extra_user_content_parts, # list[ContentPart]
"abort_signal": self._abort_signal,
"request_max_retries": self.request_max_retries,
}
if include_model:
# For primary provider we keep explicit model selection if provided.
@@ -1305,6 +1308,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
extra_user_content_parts=self.req.extra_user_content_parts,
# tool_choice="required",
abort_signal=self._abort_signal,
request_max_retries=self.request_max_retries,
)
if requery_resp:
llm_resp = requery_resp
@@ -1331,6 +1335,7 @@ class ToolLoopAgentRunner(BaseAgentRunner[TContext]):
extra_user_content_parts=self.req.extra_user_content_parts,
# tool_choice="required",
abort_signal=self._abort_signal,
request_max_retries=self.request_max_retries,
)
if repair_resp:
llm_resp = repair_resp
+1
View File
@@ -1629,6 +1629,7 @@ async def build_main_agent(
enforce_max_turns=config.max_context_length,
tool_schema_mode=config.tool_schema_mode,
fallback_providers=fallback_providers,
request_max_retries=config.provider_settings.get("request_max_retries", 5),
tool_result_overflow_dir=(
get_astrbot_system_tmp_path()
if req.func_tool and req.func_tool.get_tool("astrbot_file_read_tool")
+9
View File
@@ -101,6 +101,7 @@ DEFAULT_CONFIG = {
"enable": True,
"default_provider_id": "",
"fallback_chat_models": [],
"request_max_retries": 5,
"default_image_caption_provider_id": "",
"image_caption_prompt": "Please describe the image using Chinese.",
"provider_pool": ["*"], # "*" 表示使用所有可用的提供者
@@ -2808,6 +2809,9 @@ CONFIG_METADATA_2 = {
"type": "list",
"items": {"type": "string"},
},
"request_max_retries": {
"type": "int",
},
"wake_prefix": {
"type": "string",
},
@@ -3167,6 +3171,11 @@ CONFIG_METADATA_3 = {
"_special": "select_providers",
"hint": "主聊天模型请求失败时,按顺序切换到这些模型。",
},
"provider_settings.request_max_retries": {
"description": "请求最大重试次数",
"type": "int",
"hint": "单次模型请求遇到可重试错误时的最大尝试次数。",
},
"provider_settings.default_image_caption_provider_id": {
"description": "默认图片转述模型",
"type": "string",
+4
View File
@@ -106,6 +106,7 @@ class Provider(AbstractProvider):
model: str | None = None,
extra_user_content_parts: list[ContentPart] | None = None,
tool_choice: Literal["auto", "required"] = "auto",
request_max_retries: int | None = None,
**kwargs,
) -> LLMResponse:
"""获得 LLM 的文本对话结果。会使用当前的模型进行对话。
@@ -120,6 +121,7 @@ class Provider(AbstractProvider):
contexts: 上下文,和 prompt 二选一使用
tool_calls_result: 回传给 LLM 的工具调用结果。参考: https://platform.openai.com/docs/guides/function-calling
extra_user_content_parts: 额外的内容块列表,用于在用户消息后添加额外的文本块(如系统提醒、指令等)
request_max_retries: 可重试请求错误的最大尝试次数,包含首次请求。
kwargs: 其他参数
Notes:
@@ -142,6 +144,7 @@ class Provider(AbstractProvider):
tool_calls_result: ToolCallsResult | list[ToolCallsResult] | None = None,
model: str | None = None,
tool_choice: Literal["auto", "required"] = "auto",
request_max_retries: int | None = None,
**kwargs,
) -> AsyncGenerator[LLMResponse, None]:
"""获得 LLM 的流式文本对话结果。会使用当前的模型进行对话。在生成的最后会返回一次完整的结果。
@@ -155,6 +158,7 @@ class Provider(AbstractProvider):
tool_choice: 工具调用策略,`auto` 表示由模型自行决定,`required` 表示要求模型必须调用工具
contexts: 上下文,和 prompt 二选一使用
tool_calls_result: 回传给 LLM 的工具调用结果。参考: https://platform.openai.com/docs/guides/function-calling
request_max_retries: 可重试请求错误的最大尝试次数,包含首次请求。
kwargs: 其他参数
Notes:
@@ -27,6 +27,7 @@ from astrbot.core.utils.network_utils import (
)
from ..register import register_provider_adapter
from .request_retry import retry_provider_request, retry_provider_request_context
@register_provider_adapter(
@@ -353,7 +354,13 @@ class ProviderAnthropic(Provider):
logger.warning(f"未知的 tool_choice 值: {tool_choice},已回退为 'auto'")
return {"type": "auto"}
async def _query(self, payloads: dict, tools: ToolSet | None) -> LLMResponse:
async def _query(
self,
payloads: dict,
tools: ToolSet | None,
*,
request_max_retries: int | None = None,
) -> LLMResponse:
if tools:
if tool_list := tools.get_func_desc_anthropic_style():
payloads["tools"] = tool_list
@@ -368,8 +375,12 @@ class ProviderAnthropic(Provider):
self._apply_thinking_config(payloads)
try:
completion = await self.client.messages.create(
**payloads, stream=False, extra_body=extra_body
completion = await retry_provider_request(
"Anthropic",
lambda: self.client.messages.create(
**payloads, stream=False, extra_body=extra_body
),
max_attempts=request_max_retries,
)
except httpx.RequestError as e:
proxy = self.provider_config.get("proxy", "")
@@ -438,6 +449,8 @@ class ProviderAnthropic(Provider):
self,
payloads: dict,
tools: ToolSet | None,
*,
request_max_retries: int | None = None,
) -> AsyncGenerator[LLMResponse, None]:
if tools:
if tool_list := tools.get_func_desc_anthropic_style():
@@ -461,8 +474,10 @@ class ProviderAnthropic(Provider):
payloads["max_tokens"] = 65536
self._apply_thinking_config(payloads)
async with self.client.messages.stream(
**payloads, extra_body=extra_body
async with retry_provider_request_context(
"Anthropic",
lambda: self.client.messages.stream(**payloads, extra_body=extra_body),
max_attempts=request_max_retries,
) as stream:
assert isinstance(stream, anthropic.AsyncMessageStream)
async for event in stream:
@@ -601,6 +616,7 @@ class ProviderAnthropic(Provider):
model=None,
extra_user_content_parts=None,
tool_choice: Literal["auto", "any", "tool", "none"] | dict[str, str] = "auto",
request_max_retries: int | None = None,
**kwargs,
) -> LLMResponse:
if contexts is None:
@@ -650,7 +666,11 @@ class ProviderAnthropic(Provider):
llm_response = None
try:
llm_response = await self._query(payloads, func_tool)
llm_response = await self._query(
payloads,
func_tool,
request_max_retries=request_max_retries,
)
except Exception as e:
raise e
@@ -669,6 +689,7 @@ class ProviderAnthropic(Provider):
model=None,
extra_user_content_parts=None,
tool_choice: Literal["auto", "any", "tool", "none"] | dict[str, str] = "auto",
request_max_retries: int | None = None,
**kwargs,
):
if contexts is None:
@@ -715,7 +736,11 @@ class ProviderAnthropic(Provider):
else system_prompt
)
async for llm_response in self._query_stream(payloads, func_tool):
async for llm_response in self._query_stream(
payloads,
func_tool,
request_max_retries=request_max_retries,
):
yield llm_response
def _detect_image_mime_type(self, data: bytes) -> str:
@@ -827,7 +852,10 @@ class ProviderAnthropic(Provider):
async def get_models(self) -> list[str]:
models_str = []
models = await self.client.models.list()
models = await retry_provider_request(
"Anthropic",
lambda: self.client.models.list(),
)
models = sorted(models.data, key=lambda x: x.id)
for model in models:
models_str.append(model.id)
+42 -12
View File
@@ -26,6 +26,7 @@ from astrbot.core.utils.media_utils import (
from astrbot.core.utils.network_utils import is_connection_error, log_connection_failure
from ..register import register_provider_adapter
from .request_retry import retry_provider_request
class SuppressNonTextPartsWarning(logging.Filter):
@@ -577,7 +578,13 @@ class ProviderGoogleGenAI(Provider):
)
return chain_result
async def _query(self, payloads: dict, tools: ToolSet | None) -> LLMResponse:
async def _query(
self,
payloads: dict,
tools: ToolSet | None,
*,
request_max_retries: int | None = None,
) -> LLMResponse:
"""非流式请求 Gemini API"""
system_instruction = next(
(msg["content"] for msg in payloads["messages"] if msg["role"] == "system"),
@@ -604,10 +611,14 @@ class ProviderGoogleGenAI(Provider):
modalities,
temperature,
)
result = await self.client.models.generate_content(
model=model,
contents=cast(types.ContentListUnion, conversation),
config=config,
result = await retry_provider_request(
"Gemini",
lambda: self.client.models.generate_content(
model=model,
contents=cast(types.ContentListUnion, conversation),
config=config,
),
max_attempts=request_max_retries,
)
logger.debug(f"genai result: {result}")
@@ -672,6 +683,8 @@ class ProviderGoogleGenAI(Provider):
self,
payloads: dict,
tools: ToolSet | None,
*,
request_max_retries: int | None = None,
) -> AsyncGenerator[LLMResponse, None]:
"""流式请求 Gemini API"""
system_instruction = next(
@@ -690,10 +703,14 @@ class ProviderGoogleGenAI(Provider):
payloads.get("tool_choice", "auto"),
system_instruction,
)
result = await self.client.models.generate_content_stream(
model=model,
contents=cast(types.ContentListUnion, conversation),
config=config,
result = await retry_provider_request(
"Gemini",
lambda: self.client.models.generate_content_stream(
model=model,
contents=cast(types.ContentListUnion, conversation),
config=config,
),
max_attempts=request_max_retries,
)
break
except APIError as e:
@@ -809,6 +826,7 @@ class ProviderGoogleGenAI(Provider):
model=None,
extra_user_content_parts=None,
tool_choice: Literal["auto", "required"] = "auto",
request_max_retries: int | None = None,
**kwargs,
) -> LLMResponse:
if contexts is None:
@@ -850,7 +868,11 @@ class ProviderGoogleGenAI(Provider):
for _ in range(retry):
try:
return await self._query(payloads, func_tool)
return await self._query(
payloads,
func_tool,
request_max_retries=request_max_retries,
)
except APIError as e:
if await self._handle_api_error(e, keys):
continue
@@ -871,6 +893,7 @@ class ProviderGoogleGenAI(Provider):
model=None,
extra_user_content_parts=None,
tool_choice: Literal["auto", "required"] = "auto",
request_max_retries: int | None = None,
**kwargs,
) -> AsyncGenerator[LLMResponse, None]:
if contexts is None:
@@ -912,7 +935,11 @@ class ProviderGoogleGenAI(Provider):
for _ in range(retry):
try:
async for response in self._query_stream(payloads, func_tool):
async for response in self._query_stream(
payloads,
func_tool,
request_max_retries=request_max_retries,
):
yield response
break
except APIError as e:
@@ -922,7 +949,10 @@ class ProviderGoogleGenAI(Provider):
async def get_models(self):
try:
models = await self.client.models.list()
models = await retry_provider_request(
"Gemini",
lambda: self.client.models.list(),
)
return [
m.name.replace("models/", "")
for m in models
+43 -13
View File
@@ -41,6 +41,7 @@ from astrbot.core.utils.network_utils import (
from astrbot.core.utils.string_utils import normalize_and_dedupe_strings
from ..register import register_provider_adapter
from .request_retry import retry_provider_request
@register_provider_adapter(
@@ -420,7 +421,10 @@ class ProviderOpenAIOfficial(Provider):
async def get_models(self):
try:
models_str = []
models = await self.client.models.list()
models = await retry_provider_request(
"OpenAI",
lambda: self.client.models.list(),
)
models = sorted(models.data, key=lambda x: x.id)
for model in models:
models_str.append(model.id)
@@ -465,7 +469,13 @@ class ProviderOpenAIOfficial(Provider):
payloads["messages"] = cleaned
async def _query(self, payloads: dict, tools: ToolSet | None) -> LLMResponse:
async def _query(
self,
payloads: dict,
tools: ToolSet | None,
*,
request_max_retries: int | None = None,
) -> LLMResponse:
if tools:
model = payloads.get("model", "").lower()
omit_empty_param_field = "gemini" in model
@@ -496,10 +506,14 @@ class ProviderOpenAIOfficial(Provider):
self._sanitize_assistant_messages(payloads)
completion = await self.client.chat.completions.create(
**payloads,
stream=False,
extra_body=extra_body,
completion = await retry_provider_request(
"OpenAI",
lambda: self.client.chat.completions.create(
**payloads,
stream=False,
extra_body=extra_body,
),
max_attempts=request_max_retries,
)
if not isinstance(completion, ChatCompletion):
@@ -517,6 +531,8 @@ class ProviderOpenAIOfficial(Provider):
self,
payloads: dict,
tools: ToolSet | None,
*,
request_max_retries: int | None = None,
) -> AsyncGenerator[LLMResponse, None]:
"""流式查询API,逐步返回结果"""
if tools:
@@ -548,11 +564,15 @@ class ProviderOpenAIOfficial(Provider):
self._sanitize_assistant_messages(payloads)
stream = await self.client.chat.completions.create(
**payloads,
stream=True,
extra_body=extra_body,
stream_options={"include_usage": True},
stream = await retry_provider_request(
"OpenAI",
lambda: self.client.chat.completions.create(
**payloads,
stream=True,
extra_body=extra_body,
stream_options={"include_usage": True},
),
max_attempts=request_max_retries,
)
llm_response = LLMResponse("assistant", is_chunk=True)
@@ -1104,6 +1124,7 @@ class ProviderOpenAIOfficial(Provider):
model=None,
extra_user_content_parts=None,
tool_choice: Literal["auto", "required"] = "auto",
request_max_retries: int | None = None,
**kwargs,
) -> LLMResponse:
payloads, context_query = await self._prepare_chat_payload(
@@ -1131,7 +1152,11 @@ class ProviderOpenAIOfficial(Provider):
for retry_cnt in range(max_retries):
try:
self.client.api_key = chosen_key
llm_response = await self._query(payloads, func_tool)
llm_response = await self._query(
payloads,
func_tool,
request_max_retries=request_max_retries,
)
break
except Exception as e:
last_exception = e
@@ -1176,6 +1201,7 @@ class ProviderOpenAIOfficial(Provider):
tool_calls_result=None,
model=None,
tool_choice: Literal["auto", "required"] = "auto",
request_max_retries: int | None = None,
**kwargs,
) -> AsyncGenerator[LLMResponse, None]:
"""流式对话,与服务商交互并逐步返回结果"""
@@ -1202,7 +1228,11 @@ class ProviderOpenAIOfficial(Provider):
for retry_cnt in range(max_retries):
try:
self.client.api_key = chosen_key
async for response in self._query_stream(payloads, func_tool):
async for response in self._query_stream(
payloads,
func_tool,
request_max_retries=request_max_retries,
):
yield response
break
except Exception as e:
@@ -0,0 +1,163 @@
from collections.abc import AsyncIterator, Awaitable, Callable
from contextlib import AbstractAsyncContextManager, asynccontextmanager
from typing import TypeVar
from tenacity import (
AsyncRetrying,
RetryCallState,
retry_if_exception,
stop_after_attempt,
wait_exponential,
)
from astrbot import logger
from astrbot.core.utils.config_number import coerce_int_config
from astrbot.core.utils.network_utils import is_connection_error
T = TypeVar("T")
REQUEST_RETRY_ATTEMPTS = 5 # default value
REQUEST_RETRY_WAIT_MIN_S = 0.2
REQUEST_RETRY_WAIT_MAX_S = 30
REQUEST_RETRY_STATUS_CODES = {408, 409, 429, 500, 502, 503, 504, 529}
def _get_status_code(error: BaseException) -> int | None:
for attr in ("status_code", "status", "code"):
value = getattr(error, attr, None)
if isinstance(value, int):
return value
response = getattr(error, "response", None)
if response is not None:
status_code = getattr(response, "status_code", None)
if isinstance(status_code, int):
return status_code
return None
def _is_retryable_provider_request_error(
error: BaseException,
*,
retry_rate_limits: bool,
) -> bool:
if is_connection_error(error):
return True
error_type_name = type(error).__name__
if error_type_name in {"APIConnectionError", "APITimeoutError"}:
return True
status_code = _get_status_code(error)
if status_code is None:
return False
if status_code == 429 and not retry_rate_limits:
return False
return status_code in REQUEST_RETRY_STATUS_CODES or 500 <= status_code <= 599
def _log_retry(
provider_label: str,
retry_state: RetryCallState,
max_attempts: int,
) -> None:
error = retry_state.outcome.exception() if retry_state.outcome else None
logger.warning(
f"[{provider_label}] Request failed with retryable error; "
f"retrying ({retry_state.attempt_number + 1}/{max_attempts}): "
f"{error}"
)
def _build_retrying(
provider_label: str,
*,
retry_rate_limits: bool,
max_attempts: int | None = None,
) -> AsyncRetrying:
max_attempts = coerce_int_config(
max_attempts if max_attempts is not None else REQUEST_RETRY_ATTEMPTS,
default=REQUEST_RETRY_ATTEMPTS,
min_value=1,
field_name="request_max_retries",
source=provider_label,
)
return AsyncRetrying(
retry=retry_if_exception(
lambda error: _is_retryable_provider_request_error(
error,
retry_rate_limits=retry_rate_limits,
)
),
stop=stop_after_attempt(max_attempts),
wait=wait_exponential(
multiplier=1,
min=REQUEST_RETRY_WAIT_MIN_S,
max=REQUEST_RETRY_WAIT_MAX_S,
),
before_sleep=lambda retry_state: _log_retry(
provider_label,
retry_state,
max_attempts,
),
reraise=True,
)
async def retry_provider_request(
provider_label: str,
request_factory: Callable[[], Awaitable[T]],
*,
retry_rate_limits: bool = True,
max_attempts: int | None = None,
) -> T:
retrying = _build_retrying(
provider_label,
retry_rate_limits=retry_rate_limits,
max_attempts=max_attempts,
)
async for attempt in retrying:
with attempt:
return await request_factory()
raise RuntimeError("Provider request retry loop exited unexpectedly.")
@asynccontextmanager
async def retry_provider_request_context(
provider_label: str,
context_manager_factory: Callable[[], AbstractAsyncContextManager[T]],
*,
retry_rate_limits: bool = True,
max_attempts: int | None = None,
) -> AsyncIterator[T]:
manager: AbstractAsyncContextManager[T] | None = None
async def _enter_context() -> T:
nonlocal manager
manager = context_manager_factory()
return await manager.__aenter__()
value = await retry_provider_request(
provider_label,
_enter_context,
retry_rate_limits=retry_rate_limits,
max_attempts=max_attempts,
)
if manager is None:
raise RuntimeError("Provider request context was not created.")
try:
yield value
except BaseException as error:
if await manager.__aexit__(type(error), error, error.__traceback__):
return
raise
else:
await manager.__aexit__(None, None, None)
@@ -45,6 +45,10 @@
"description": "Fallback chat model IDs",
"hint": "When the primary chat model request fails, fallback to these chat models in order."
},
"request_max_retries": {
"description": "Request Max Retries",
"hint": "Maximum attempts for a single model request when retryable errors occur."
},
"default_image_caption_provider_id": {
"description": "Default Image Caption Model",
"hint": "Leave empty to disable; useful for non-multimodal models"
@@ -45,6 +45,10 @@
"description": "Резервные модели чата (ID)",
"hint": "Если текущая модель недоступна, запрос будет перенаправлен на эти модели по порядку."
},
"request_max_retries": {
"description": "Максимум повторов запроса",
"hint": "Максимальное число попыток для одного запроса модели при повторяемых ошибках."
},
"default_image_caption_provider_id": {
"description": "Модель описания изображений",
"hint": "Оставьте пустым для отключения; полезно для моделей без поддержки мультимодальности"
@@ -45,6 +45,10 @@
"description": "回退对话模型列表",
"hint": "主对话模型请求失败时,按顺序切换到这些对话模型。"
},
"request_max_retries": {
"description": "请求最大重试次数",
"hint": "单次模型请求遇到可重试错误时的最大尝试次数。"
},
"default_image_caption_provider_id": {
"description": "默认图片转述模型",
"hint": "留空代表不使用,可用于非多模态模型"
+35 -2
View File
@@ -1,9 +1,12 @@
import builtins
from types import SimpleNamespace
import httpx
import pytest
import astrbot.core.provider.sources.anthropic_source as anthropic_source
import astrbot.core.provider.sources.kimi_code_source as kimi_code_source
import astrbot.core.provider.sources.request_retry as request_retry
from astrbot.core.exceptions import EmptyModelOutputError
from astrbot.core.provider.entities import LLMResponse
@@ -171,6 +174,36 @@ def test_create_http_client_falls_back_to_global_httpx_module(monkeypatch):
assert captured["httpx_module"] is anthropic_source.httpx
@pytest.mark.asyncio
async def test_anthropic_get_models_retries_transient_request_error(monkeypatch):
monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MIN_S", 0)
monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MAX_S", 0)
class FakeModels:
def __init__(self):
self.calls = 0
async def list(self):
self.calls += 1
if self.calls == 1:
raise httpx.ConnectError("temporary connection failure")
return SimpleNamespace(
data=[
SimpleNamespace(id="claude-b"),
SimpleNamespace(id="claude-a"),
]
)
models = FakeModels()
provider = anthropic_source.ProviderAnthropic.__new__(
anthropic_source.ProviderAnthropic
)
provider.client = SimpleNamespace(models=models)
assert await provider.get_models() == ["claude-a", "claude-b"]
assert models.calls == 2
@pytest.mark.asyncio
async def test_text_chat_wraps_string_system_prompt_as_list(monkeypatch):
monkeypatch.setattr(anthropic_source, "AsyncAnthropic", _FakeAsyncAnthropic)
@@ -187,7 +220,7 @@ async def test_text_chat_wraps_string_system_prompt_as_list(monkeypatch):
captured_payloads: dict[str, object] = {}
async def fake_query(payloads, tools):
async def fake_query(payloads, tools, *, request_max_retries=None):
captured_payloads.update(payloads)
return LLMResponse(role="assistant", completion_text="ok")
@@ -214,7 +247,7 @@ async def test_text_chat_passes_through_list_system_prompt(monkeypatch):
captured_payloads: dict[str, object] = {}
async def fake_query(payloads, tools):
async def fake_query(payloads, tools, *, request_max_retries=None):
captured_payloads.update(payloads)
return LLMResponse(role="assistant", completion_text="ok")
+36
View File
@@ -1,6 +1,10 @@
from types import SimpleNamespace
import httpx
import pytest
from astrbot.core.exceptions import EmptyModelOutputError
import astrbot.core.provider.sources.request_retry as request_retry
from astrbot.core.provider.entities import LLMResponse
from astrbot.core.provider.sources.gemini_source import ProviderGoogleGenAI
@@ -27,3 +31,35 @@ def test_gemini_reasoning_only_output_is_allowed():
response_id="resp_reasoning",
finish_reason="STOP",
)
@pytest.mark.asyncio
async def test_gemini_get_models_retries_transient_request_error(monkeypatch):
monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MIN_S", 0)
monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MAX_S", 0)
class FakeModels:
def __init__(self):
self.calls = 0
async def list(self):
self.calls += 1
if self.calls == 1:
raise httpx.ConnectError("temporary connection failure")
return [
SimpleNamespace(
name="models/gemini-a",
supported_actions=["generateContent"],
),
SimpleNamespace(
name="models/gemini-b",
supported_actions=["embedContent"],
),
]
models = FakeModels()
provider = ProviderGoogleGenAI.__new__(ProviderGoogleGenAI)
provider.client = SimpleNamespace(models=models)
assert await provider.get_models() == ["gemini-a"]
assert models.calls == 2
+54
View File
@@ -3,13 +3,16 @@ import builtins
from io import BytesIO
from types import SimpleNamespace
import httpx
import pytest
from openai.types.chat.chat_completion import ChatCompletion
from openai.types.chat.chat_completion_chunk import ChatCompletionChunk
from PIL import Image as PILImage
import astrbot.core.provider.sources.openai_source as openai_source_module
import astrbot.core.provider.sources.request_retry as request_retry
from astrbot.core.exceptions import EmptyModelOutputError
from astrbot.core.provider.entities import LLMResponse
from astrbot.core.provider.sources.groq_source import ProviderGroq
from astrbot.core.provider.sources.openai_source import ProviderOpenAIOfficial
from astrbot.core.utils.media_utils import ResolvedMediaData, file_uri_to_path
@@ -117,6 +120,57 @@ def test_create_http_client_falls_back_to_global_httpx_module(monkeypatch):
assert captured["httpx_module"] is openai_source_module.httpx
@pytest.mark.asyncio
async def test_get_models_retries_transient_request_error(monkeypatch):
monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MIN_S", 0)
monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MAX_S", 0)
class FakeModels:
def __init__(self):
self.calls = 0
async def list(self):
self.calls += 1
if self.calls == 1:
raise httpx.ConnectError("temporary connection failure")
return SimpleNamespace(
data=[
SimpleNamespace(id="gpt-b"),
SimpleNamespace(id="gpt-a"),
]
)
models = FakeModels()
provider = ProviderOpenAIOfficial.__new__(ProviderOpenAIOfficial)
provider.client = SimpleNamespace(models=models)
assert await provider.get_models() == ["gpt-a", "gpt-b"]
assert models.calls == 2
@pytest.mark.asyncio
async def test_text_chat_passes_request_max_retries_to_query():
captured: dict[str, object] = {}
provider = ProviderOpenAIOfficial.__new__(ProviderOpenAIOfficial)
provider.api_keys = ["test-key"]
provider.client = SimpleNamespace(api_key=None)
async def fake_prepare_chat_payload(*args, **kwargs):
return {"messages": [], "model": "gpt-4o-mini"}, []
async def fake_query(payloads, func_tool, *, request_max_retries=None):
captured["request_max_retries"] = request_max_retries
return LLMResponse(role="assistant", completion_text="ok")
provider._prepare_chat_payload = fake_prepare_chat_payload
provider._query = fake_query
await provider.text_chat(prompt="hello", request_max_retries=2)
assert captured["request_max_retries"] == 2
@pytest.mark.asyncio
async def test_handle_api_error_content_moderated_removes_images():
provider = _make_provider(
+27
View File
@@ -0,0 +1,27 @@
import httpx
import pytest
import astrbot.core.provider.sources.request_retry as request_retry
from astrbot.core.provider.sources.request_retry import retry_provider_request
@pytest.mark.asyncio
async def test_retry_provider_request_uses_configured_max_retries(monkeypatch):
monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MIN_S", 0)
monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MAX_S", 0)
calls = 0
async def request():
nonlocal calls
calls += 1
raise httpx.ConnectError("temporary connection failure")
with pytest.raises(httpx.ConnectError):
await retry_provider_request(
"Test",
request,
max_attempts=2,
)
assert calls == 2