Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions astrbot/core/provider/sources/anthropic_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -520,6 +520,8 @@ async def _query(
**payloads, stream=False, extra_body=extra_body
),
max_attempts=request_max_retries,
provider_id=self.provider_config.get("id"),
model=payloads.get("model", self.get_model()),
)
except httpx.RequestError as e:
proxy = self.provider_config.get("proxy", "")
Expand Down Expand Up @@ -619,6 +621,8 @@ async def _query_stream(
"Anthropic",
lambda: self.client.messages.stream(**payloads, extra_body=extra_body),
max_attempts=request_max_retries,
provider_id=self.provider_config.get("id"),
model=payloads.get("model", self.get_model()),
) as stream:
assert isinstance(stream, anthropic.AsyncMessageStream)
async for event in stream:
Expand Down Expand Up @@ -996,6 +1000,7 @@ async def get_models(self) -> list[str]:
models = await retry_provider_request(
"Anthropic",
lambda: self.client.models.list(),
provider_id=self.provider_config.get("id"),
)
models = sorted(models.data, key=lambda x: x.id)
for model in models:
Expand Down
5 changes: 5 additions & 0 deletions astrbot/core/provider/sources/gemini_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -619,6 +619,8 @@ async def _query(
config=config,
),
max_attempts=request_max_retries,
provider_id=self.provider_config.get("id"),
model=model,
)
logger.debug(f"genai result: {result}")

Expand Down Expand Up @@ -711,6 +713,8 @@ async def _query_stream(
config=config,
),
max_attempts=request_max_retries,
provider_id=self.provider_config.get("id"),
model=model,
)
break
except APIError as e:
Expand Down Expand Up @@ -952,6 +956,7 @@ async def get_models(self):
models = await retry_provider_request(
"Gemini",
lambda: self.client.models.list(),
provider_id=self.provider_config.get("id"),
)
return [
m.name.replace("models/", "")
Expand Down
4 changes: 4 additions & 0 deletions astrbot/core/provider/sources/openai_responses_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -353,6 +353,8 @@ async def _query(
extra_body=extra_body,
),
max_attempts=request_max_retries,
provider_id=self.provider_config.get("id"),
model=payloads.get("model", ""),
)
if not isinstance(response, Response):
raise TypeError(
Expand Down Expand Up @@ -422,6 +424,8 @@ async def _query_stream(
extra_body=extra_body,
),
max_attempts=request_max_retries,
provider_id=self.provider_config.get("id"),
model=payloads.get("model", ""),
)

response_id: str | None = None
Expand Down
5 changes: 5 additions & 0 deletions astrbot/core/provider/sources/openai_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -440,6 +440,7 @@ async def get_models(self):
models = await retry_provider_request(
"OpenAI",
lambda: self.client.models.list(),
provider_id=self.provider_config.get("id"),
)
models = sorted(models.data, key=lambda x: x.id)
for model in models:
Expand Down Expand Up @@ -573,6 +574,8 @@ async def _query(
extra_body=extra_body,
),
max_attempts=request_max_retries,
provider_id=self.provider_config.get("id"),
model=payloads.get("model", ""),
)

if not isinstance(completion, ChatCompletion):
Expand Down Expand Up @@ -632,6 +635,8 @@ async def _query_stream(
stream_options={"include_usage": True},
),
max_attempts=request_max_retries,
provider_id=self.provider_config.get("id"),
model=payloads.get("model", ""),
)

llm_response = LLMResponse("assistant", is_chunk=True)
Expand Down
23 changes: 22 additions & 1 deletion astrbot/core/provider/sources/request_retry.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,10 +63,19 @@ def _log_retry(
provider_label: str,
retry_state: RetryCallState,
max_attempts: int,
*,
provider_id: str | None = None,
model: str | None = None,
) -> None:
error = retry_state.outcome.exception() if retry_state.outcome else None
identity_parts = []
if provider_id:
identity_parts.append(f"provider={provider_id}")
if model:
identity_parts.append(f"model={model}")
identity = f" ({', '.join(identity_parts)})" if identity_parts else ""
logger.warning(
f"[{provider_label}] Request failed with retryable error; "
f"[{provider_label}]{identity} Request failed with retryable error; "
f"retrying ({retry_state.attempt_number + 1}/{max_attempts}): "
f"{error}"
)
Expand All @@ -77,6 +86,8 @@ def _build_retrying(
*,
retry_rate_limits: bool,
max_attempts: int | None = None,
provider_id: str | None = None,
model: str | None = None,
) -> AsyncRetrying:
max_attempts = coerce_int_config(
max_attempts if max_attempts is not None else REQUEST_RETRY_ATTEMPTS,
Expand All @@ -103,6 +114,8 @@ def _build_retrying(
provider_label,
retry_state,
max_attempts,
provider_id=provider_id,
model=model,
),
reraise=True,
)
Expand All @@ -114,11 +127,15 @@ async def retry_provider_request(
*,
retry_rate_limits: bool = True,
max_attempts: int | None = None,
provider_id: str | None = None,
model: str | None = None,
) -> T:
retrying = _build_retrying(
provider_label,
retry_rate_limits=retry_rate_limits,
max_attempts=max_attempts,
provider_id=provider_id,
model=model,
)

async for attempt in retrying:
Expand All @@ -135,6 +152,8 @@ async def retry_provider_request_context(
*,
retry_rate_limits: bool = True,
max_attempts: int | None = None,
provider_id: str | None = None,
model: str | None = None,
) -> AsyncIterator[T]:
manager: AbstractAsyncContextManager[T] | None = None

Expand All @@ -148,6 +167,8 @@ async def _enter_context() -> T:
_enter_context,
retry_rate_limits=retry_rate_limits,
max_attempts=max_attempts,
provider_id=provider_id,
model=model,
)

if manager is None:
Expand Down
1 change: 1 addition & 0 deletions astrbot/core/provider/sources/ssycloud_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,7 @@ async def get_models(self) -> list[str]:
response = await retry_provider_request(
"SSYCloud",
lambda: self.client.models.list(),
provider_id=self.provider_config.get("id"),
)
model_ids: list[str] = []
for model in response.data:
Expand Down
1 change: 1 addition & 0 deletions tests/test_anthropic_kimi_code_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -198,6 +198,7 @@ async def list(self):
provider = anthropic_source.ProviderAnthropic.__new__(
anthropic_source.ProviderAnthropic
)
provider.provider_config = {"id": "test-anthropic-provider"}
provider.client = SimpleNamespace(models=models)

assert await provider.get_models() == ["claude-a", "claude-b"]
Expand Down
1 change: 1 addition & 0 deletions tests/test_gemini_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,7 @@ async def list(self):

models = FakeModels()
provider = ProviderGoogleGenAI.__new__(ProviderGoogleGenAI)
provider.provider_config = {"id": "test-gemini-provider"}
provider.client = SimpleNamespace(models=models)

assert await provider.get_models() == ["gemini-a"]
Expand Down
1 change: 1 addition & 0 deletions tests/test_openai_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -150,6 +150,7 @@ async def list(self):

models = FakeModels()
provider = ProviderOpenAIOfficial.__new__(ProviderOpenAIOfficial)
provider.provider_config = {"id": "test-openai-provider"}
provider.client = SimpleNamespace(models=models)

assert await provider.get_models() == ["gpt-a", "gpt-b"]
Expand Down
90 changes: 90 additions & 0 deletions tests/test_request_retry.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
import logging

import httpx
import pytest

Expand Down Expand Up @@ -25,3 +27,91 @@ async def request():
)

assert calls == 2


@pytest.mark.asyncio
async def test_retry_log_includes_provider_id_and_model(monkeypatch, caplog):
monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MIN_S", 0)
monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MAX_S", 0)

async def request():
raise httpx.ConnectError("temporary connection failure")

with caplog.at_level(logging.WARNING, logger="astrbot"):
with pytest.raises(httpx.ConnectError):
Comment thread
sourcery-ai[bot] marked this conversation as resolved.
await retry_provider_request(
"OpenAI",
request,
max_attempts=2,
provider_id="my-openai-instance",
model="gpt-4o",
)

assert "[OpenAI]" in caplog.text
assert "provider=my-openai-instance" in caplog.text
assert "model=gpt-4o" in caplog.text


@pytest.mark.asyncio
async def test_retry_log_omits_details_when_not_provided(monkeypatch, caplog):
monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MIN_S", 0)
monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MAX_S", 0)

async def request():
raise httpx.ConnectError("temporary connection failure")

with caplog.at_level(logging.WARNING, logger="astrbot"):
with pytest.raises(httpx.ConnectError):
await retry_provider_request(
"OpenAI",
request,
max_attempts=2,
)

assert "[OpenAI] Request failed with retryable error" in caplog.text
assert "provider=" not in caplog.text
assert "model=" not in caplog.text


@pytest.mark.asyncio
async def test_retry_log_includes_only_provider_id(monkeypatch, caplog):
monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MIN_S", 0)
monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MAX_S", 0)

async def request():
raise httpx.ConnectError("temporary connection failure")

with caplog.at_level(logging.WARNING, logger="astrbot"):
with pytest.raises(httpx.ConnectError):
await retry_provider_request(
"OpenAI",
request,
max_attempts=2,
provider_id="my-openai-instance",
)

assert "[OpenAI]" in caplog.text
assert "provider=my-openai-instance" in caplog.text
assert "model=" not in caplog.text


@pytest.mark.asyncio
async def test_retry_log_includes_only_model(monkeypatch, caplog):
monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MIN_S", 0)
monkeypatch.setattr(request_retry, "REQUEST_RETRY_WAIT_MAX_S", 0)

async def request():
raise httpx.ConnectError("temporary connection failure")

with caplog.at_level(logging.WARNING, logger="astrbot"):
with pytest.raises(httpx.ConnectError):
await retry_provider_request(
"OpenAI",
request,
max_attempts=2,
model="gpt-4o",
)

assert "[OpenAI]" in caplog.text
assert "provider=" not in caplog.text
assert "model=gpt-4o" in caplog.text
Loading