Skip to content
Merged
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
132 changes: 131 additions & 1 deletion amplifier_module_provider_anthropic/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -929,10 +929,140 @@ async def list_models(self) -> list[ModelInfo]:
model family (e.g. fable, opus, sonnet, haiku). When filtered=False,
returns all available Claude models.

The query is retried with the same shared retry_with_backoff()/
_retry_config machinery used by complete() on transient failures
(5xx, connection errors, timeouts, rate limits). Raises the
translated kernel error once retries are exhausted, or immediately
for non-retryable errors (401/403/404) -- no fallback; caller
handles empty lists.

Returns:
List of ModelInfo for available Claude models.
"""
response = await self.client.models.list()

async def _do_list_models():
"""Single API call attempt with SDK -> kernel error translation.

Mirrors the error-translation branches used by _do_complete()
(rate limit, authentication, status errors 403/404/5xx, and a
catch-all for connection/timeout errors) so list_models() shares
the same retry policy as complete().
"""
try:
return await self.client.models.list()
except AnthropicRateLimitError as e:
rate_info = self._parse_rate_limit_info(e)
retry_after = rate_info.get("retry_after_seconds")
body = getattr(e, "body", None)
msg = json.dumps(body) if body is not None else str(e)
raise KernelRateLimitError(
msg,
provider="anthropic",
status_code=429,
retryable=True,
retry_after=retry_after,
) from e
except AnthropicAuthenticationError as e:
body = getattr(e, "body", None)
msg = json.dumps(body) if body is not None else str(e)
raise KernelAuthenticationError(
msg,
provider="anthropic",
status_code=getattr(e, "status_code", 401),
) from e
except AnthropicAPIStatusError as e:
# Also catches AnthropicOverloadedError (529), a subclass of
# AnthropicAPIStatusError -- it falls into the status >= 500
# branch below and is retried like any other 5xx.
status = getattr(e, "status_code", 500)
body = getattr(e, "body", None)
error_msg = json.dumps(body) if body is not None else str(e)
if status == 403:
# Distinguish Cloudflare bot challenges (transient) from
# real API 403s (permanent), same detection _do_complete uses.
if self._is_cloudflare_challenge(e):
logger.warning(
"[PROVIDER] Cloudflare challenge detected (HTTP 403 "
"with HTML body) on list_models(). Treating as "
"transient -- will retry."
)
raise KernelProviderUnavailableError(
"Cloudflare bot challenge (transient 403 with HTML "
"body). This typically resolves on retry.",
provider="anthropic",
status_code=403,
retryable=True,
) from e
raise KernelAccessDeniedError(
error_msg,
provider="anthropic",
status_code=403,
) from e
if status == 404:
raise KernelNotFoundError(
error_msg,
provider="anthropic",
status_code=404,
) from e
if status >= 500:
raise KernelProviderUnavailableError(
error_msg,
provider="anthropic",
status_code=status,
retryable=True,
) from e
raise KernelLLMError(
error_msg,
provider="anthropic",
status_code=status,
retryable=False,
) from e
except KernelLLMError:
raise # Already translated, don't double-wrap
except Exception as e:
# Connection errors, timeouts, and anything unforeseen land
# here -- same catch-all _do_complete() uses, treated as
# transient and retryable.
body = getattr(e, "body", None)
error_msg = (
json.dumps(body)
if body is not None
else (str(e) or f"{type(e).__name__}: (no message)")
)
raise KernelLLMError(
error_msg,
provider="anthropic",
retryable=True,
) from e

async def _on_retry(attempt: int, delay: float, error: KernelLLMError):
"""Callback invoked before each retry sleep."""
error_type = type(error).__name__
logger.warning(
"[PROVIDER] Retry %d/%d for list_models(): %s, sleeping %.1fs",
attempt,
self._retry_config.max_retries,
error_type,
delay,
)
if self.coordinator and hasattr(self.coordinator, "hooks"):
await self.coordinator.hooks.emit(
PROVIDER_RETRY,
{
"provider": "anthropic",
"attempt": attempt,
"max_retries": self._retry_config.max_retries,
"delay": delay,
"error_type": error_type,
"error_message": str(error),
},
)

response = await retry_with_backoff(
_do_list_models,
self._retry_config,
on_retry=_on_retry,
)
api_models = list(response.data)

# Group models by family using _detect_family() as the single source of
Expand Down
139 changes: 139 additions & 0 deletions tests/test_list_models_retry.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,139 @@
"""Retry behavior tests for list_models().

Verifies that list_models() uses the same shared retry_with_backoff()/
_retry_config machinery as complete(): transient failures (5xx) are
retried with backoff, non-retryable failures (401) raise immediately,
and persistent transient failures raise the translated kernel error
once retries are exhausted.

See test_retry.py for the equivalent tests on the complete() path --
this file mirrors that call shape for list_models().
"""

import asyncio
from types import SimpleNamespace
from typing import cast
from unittest.mock import AsyncMock, MagicMock, patch

import anthropic
import pytest
from amplifier_core import ModuleCoordinator
from amplifier_core.llm_errors import AuthenticationError as KernelAuthenticationError
from amplifier_core.llm_errors import (
ProviderUnavailableError as KernelProviderUnavailableError,
)

from amplifier_module_provider_anthropic import AnthropicProvider
from tests._helpers import FakeCoordinator

# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------


def _make_provider(
max_retries: int = 3, max_retry_delay: float = 60.0
) -> AnthropicProvider:
provider = AnthropicProvider(
api_key="test-key",
config={
"use_streaming": False,
"max_retries": max_retries,
"min_retry_delay": 0.01, # Fast for tests
"max_retry_delay": max_retry_delay,
"retry_jitter": False, # Deterministic for tests
},
)
provider.coordinator = cast(ModuleCoordinator, FakeCoordinator())
return provider


def _fake_models_response(model_ids: list[str]) -> SimpleNamespace:
"""Create a fake Anthropic models.list() response."""
data = [
SimpleNamespace(id=mid, display_name=mid, created_at="2026-01-01")
for mid in model_ids
]
return SimpleNamespace(data=data)


def _make_sdk_server_error() -> anthropic.InternalServerError:
mock_response = MagicMock()
mock_response.status_code = 500
mock_response.headers = {}
return anthropic.InternalServerError(
"server error", response=mock_response, body=None
)


def _make_sdk_auth_error() -> anthropic.AuthenticationError:
mock_response = MagicMock()
mock_response.status_code = 401
mock_response.headers = {}
return anthropic.AuthenticationError("bad key", response=mock_response, body=None)


# ---------------------------------------------------------------------------
# Tests
# ---------------------------------------------------------------------------


def test_list_models_succeeds_first_try():
"""No transient failure: exactly one API call, result unchanged."""
provider = _make_provider()
response = _fake_models_response(["claude-sonnet-4-5-20250929"])
provider.client.models.list = AsyncMock(return_value=response)

with patch("asyncio.sleep", new_callable=AsyncMock) as mock_sleep:
models = asyncio.run(provider.list_models())

assert provider.client.models.list.await_count == 1
mock_sleep.assert_not_awaited()
assert len(models) == 1
assert models[0].id == "claude-sonnet-4-5-20250929"


def test_list_models_recovers_from_transient_500():
"""A single transient 500 is retried, then the call succeeds."""
provider = _make_provider()
response = _fake_models_response(["claude-sonnet-4-5-20250929"])
provider.client.models.list = AsyncMock(
side_effect=[_make_sdk_server_error(), response]
)

with patch("asyncio.sleep", new_callable=AsyncMock):
models = asyncio.run(provider.list_models())

assert provider.client.models.list.await_count == 2
assert len(models) == 1
assert models[0].id == "claude-sonnet-4-5-20250929"


def test_list_models_raises_after_retries_exhausted():
"""Persistent transient failure raises the kernel error after retries."""
provider = _make_provider(max_retries=2)
provider.client.models.list = AsyncMock(side_effect=_make_sdk_server_error())

with (
patch("asyncio.sleep", new_callable=AsyncMock),
pytest.raises(KernelProviderUnavailableError),
):
asyncio.run(provider.list_models())

# 1 initial + 2 retries = 3 total attempts
assert provider.client.models.list.await_count == 3


def test_list_models_non_retryable_error_raised_immediately():
"""A non-retryable error (401) raises immediately without retrying."""
provider = _make_provider()
provider.client.models.list = AsyncMock(side_effect=_make_sdk_auth_error())

with (
patch("asyncio.sleep", new_callable=AsyncMock) as mock_sleep,
pytest.raises(KernelAuthenticationError),
):
asyncio.run(provider.list_models())

assert provider.client.models.list.await_count == 1
mock_sleep.assert_not_awaited()