Skip to content
Closed
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
18 changes: 17 additions & 1 deletion services/ai/providers/anthropic.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,14 @@
from collections.abc import AsyncIterator, Iterable
from typing import Any, ClassVar, cast

from anthropic import APIStatusError, AsyncAnthropic, AsyncStream, MessageStreamEvent
import httpx
from anthropic import (
APIConnectionError,
APIStatusError,
AsyncAnthropic,
AsyncStream,
MessageStreamEvent,
)
from anthropic.types import MessageParam, ToolParam

from . import LLMProvider, TokenUsage
Expand All @@ -32,6 +39,13 @@ def _anthropic_context_overflow(e: BaseException) -> bool:
return "prompt is too long" in str(e).lower()


def _anthropic_is_retryable(e: BaseException) -> bool:
"""Transport-level failure (connect/read/timeout, dropped stream) that a
caller may safely retry once. Status errors are excluded: the SDK already
retries those at request-creation time."""
return isinstance(e, (httpx.TransportError, APIConnectionError))


logger = logging.getLogger(__name__)


Expand Down Expand Up @@ -189,6 +203,7 @@ async def stream_response(
status_code=_anthropic_status_code(e),
cause=e,
is_context_overflow=_anthropic_context_overflow(e),
is_retryable=_anthropic_is_retryable(e),
) from e

async def generate_response(
Expand Down Expand Up @@ -238,6 +253,7 @@ async def generate_response(
status_code=_anthropic_status_code(e),
cause=e,
is_context_overflow=_anthropic_context_overflow(e),
is_retryable=_anthropic_is_retryable(e),
) from e

async def health_check(
Expand Down
2 changes: 2 additions & 0 deletions services/ai/providers/azure_foundry.py
Original file line number Diff line number Diff line change
Expand Up @@ -117,6 +117,7 @@ async def stream_response(
status_code=e.status_code,
cause=e,
is_context_overflow=e.is_context_overflow,
is_retryable=e.is_retryable,
) from e

async def generate_response(
Expand All @@ -140,6 +141,7 @@ async def generate_response(
status_code=e.status_code,
cause=e,
is_context_overflow=e.is_context_overflow,
is_retryable=e.is_retryable,
) from e

async def health_check(
Expand Down
36 changes: 31 additions & 5 deletions services/ai/providers/bedrock.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,8 @@
from typing import Any, ClassVar, cast

import boto3
import httpx
from anthropic import APIConnectionError as AnthropicConnectionError
from anthropic import APIStatusError, AnthropicBedrock
from anthropic.types import (
MessageParam,
Expand Down Expand Up @@ -36,7 +38,13 @@
)
from anthropic.types.message_stream_event import MessageStreamEvent
from anthropic.types.raw_message_delta_event import Delta
from botocore.exceptions import ClientError
from botocore.exceptions import (
ClientError,
ConnectionError as BotoConnectionError,
HTTPClientError,
SSLError as BotoSSLError,
)


from . import LLMProvider, TokenUsage
from .anthropic_message_adapter import (
Expand All @@ -48,6 +56,23 @@
logger = logging.getLogger(__name__)


def _bedrock_is_retryable(e: BaseException) -> bool:
"""Transport-level failure (dropped/errored HTTP connection, mid-stream
disconnect) that a caller may safely retry once. ``ClientError`` (API
status errors) is excluded: boto3 already retries those at request-creation
time."""
return isinstance(
e,
(
httpx.TransportError,
AnthropicConnectionError,
HTTPClientError,
BotoConnectionError,
BotoSSLError,
),
)


def sanitize_document_name(name: str) -> str:
"""
Sanitize document name for AWS Bedrock requirements.
Expand Down Expand Up @@ -113,7 +138,9 @@ def __init__(self, model_id: str, region_name: str | None = None):
else:
self.client = boto3.client("bedrock-runtime", region_name=region_name)

def _to_provider_error(self, e: Exception, model: str | None = None) -> ProviderError:
def _to_provider_error(
self, e: Exception, model: str | None = None
) -> ProviderError:
is_context_overflow = False
if isinstance(e, ClientError):
error = e.response.get("Error", {})
Expand All @@ -128,16 +155,15 @@ def _to_provider_error(self, e: Exception, model: str | None = None) -> Provider
else:
body = str(e)
status = e.status_code if isinstance(e, APIStatusError) else None
is_context_overflow = (
isinstance(e, APIStatusError) and e.status_code == 413
)
is_context_overflow = isinstance(e, APIStatusError) and e.status_code == 413
return ProviderError(
body,
provider_type=self.provider_type,
model=model or self.model_name,
status_code=status,
cause=e,
is_context_overflow=is_context_overflow,
is_retryable=_bedrock_is_retryable(e),
)

def _determine_model_family(self, model_id: str) -> str:
Expand Down
11 changes: 11 additions & 0 deletions services/ai/providers/gemini.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
from collections.abc import AsyncIterator
from typing import Any, ClassVar

import httpx
from google import genai
from google.genai import types
from google.genai.errors import APIError
Expand Down Expand Up @@ -47,6 +48,14 @@ def _gemini_status_code(e: BaseException) -> int | None:
return e.code if isinstance(e, APIError) else None


def _gemini_is_retryable(e: BaseException) -> bool:
"""Transport-level failure (connect/read/timeout, dropped stream) that a
caller may safely retry once. The genai client is httpx-based; status
errors surface as ``APIError`` and are already retried by the client at
request-creation time."""
return isinstance(e, httpx.TransportError)


logger = logging.getLogger(__name__)

THOUGHT_SIGNATURE_KEY = "_gemini_thought_signature"
Expand Down Expand Up @@ -460,6 +469,7 @@ async def stream_response(
model=model or self.model_name,
status_code=_gemini_status_code(e),
cause=e,
is_retryable=_gemini_is_retryable(e),
) from e

async def generate_response(
Expand Down Expand Up @@ -504,6 +514,7 @@ async def generate_response(
model=model or self.model_name,
status_code=_gemini_status_code(e),
cause=e,
is_retryable=_gemini_is_retryable(e),
) from e

async def health_check(
Expand Down
16 changes: 14 additions & 2 deletions services/ai/providers/openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,8 @@
from collections.abc import AsyncIterator
from typing import Any, ClassVar

from openai import APIStatusError, AsyncOpenAI
import httpx
from openai import APIConnectionError, APIStatusError, AsyncOpenAI
from anthropic.types import (
Message,
MessageDeltaUsage,
Expand Down Expand Up @@ -55,6 +56,13 @@ def _openai_context_overflow(e: BaseException) -> bool:
return _openai_error_code(e) == "context_length_exceeded"


def _openai_is_retryable(e: BaseException) -> bool:
"""Transport-level failure (connect/read/timeout, dropped stream) that a
caller may safely retry once. Status errors are excluded: the SDK already
retries those at request-creation time."""
return isinstance(e, (httpx.TransportError, APIConnectionError))


logger = logging.getLogger(__name__)

MIN_REASONING_OUTPUT_TOKENS = 1024
Expand Down Expand Up @@ -295,7 +303,9 @@ async def stream_response(
)

elif event_type == "error":
raise RuntimeError(getattr(event, "message", "Unknown stream error"))
raise RuntimeError(
getattr(event, "message", "Unknown stream error")
)

yield RawMessageStopEvent(type="message_stop")

Expand All @@ -310,6 +320,7 @@ async def stream_response(
status_code=_openai_status_code(e),
cause=e,
is_context_overflow=_openai_context_overflow(e),
is_retryable=_openai_is_retryable(e),
) from e

def _convert_messages(self, messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
Expand Down Expand Up @@ -481,6 +492,7 @@ async def generate_response(
status_code=_openai_status_code(e),
cause=e,
is_context_overflow=_openai_context_overflow(e),
is_retryable=_openai_is_retryable(e),
) from e

async def health_check(
Expand Down
12 changes: 11 additions & 1 deletion services/ai/providers/openai_compatible.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,8 @@
from collections.abc import AsyncIterator
from typing import Any, ClassVar, cast

from openai import APIStatusError, AsyncOpenAI
import httpx
from openai import APIConnectionError, APIStatusError, AsyncOpenAI
from openai.types.chat import (
ChatCompletionAssistantMessageParam,
ChatCompletionChunk,
Expand Down Expand Up @@ -75,6 +76,13 @@ def _openai_compat_context_overflow(e: BaseException) -> bool:
return _openai_compat_error_code(e) == "context_length_exceeded"


def _openai_compat_is_retryable(e: BaseException) -> bool:
"""Transport-level failure (connect/read/timeout, dropped stream) that a
caller may safely retry once. Status errors are excluded: the SDK already
retries those at request-creation time."""
return isinstance(e, (httpx.TransportError, APIConnectionError))


logger = logging.getLogger(__name__)

# Some OpenAI-compatible providers expose non-standard assistant-message fields
Expand Down Expand Up @@ -536,6 +544,7 @@ async def stream_response(
status_code=_openai_compat_status_code(e),
cause=e,
is_context_overflow=_openai_compat_context_overflow(e),
is_retryable=_openai_compat_is_retryable(e),
) from e

async def generate_response(
Expand Down Expand Up @@ -592,6 +601,7 @@ async def generate_response(
status_code=_openai_compat_status_code(e),
cause=e,
is_context_overflow=_openai_compat_context_overflow(e),
is_retryable=_openai_compat_is_retryable(e),
) from e

async def health_check(
Expand Down
6 changes: 6 additions & 0 deletions services/ai/providers/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,12 +36,18 @@ def __init__(
status_code: int | None = None,
cause: BaseException | None = None,
is_context_overflow: bool | None = None,
is_retryable: bool = False,
):
super().__init__(message)
self.message = message
self.provider_type = provider_type
self.model = model
self.status_code = status_code
self.is_context_overflow = bool(is_context_overflow)
# True only for transport-level failures (connection reset, mid-stream
# read error, remote disconnect) that a caller may safely retry once.
# Providers set it from the underlying exception type; a missing
# status_code alone does not imply transport failure.
self.is_retryable = is_retryable
if cause is not None:
self.__cause__ = cause
2 changes: 2 additions & 0 deletions services/ai/providers/vertex_ai.py
Original file line number Diff line number Diff line change
Expand Up @@ -92,6 +92,7 @@ async def stream_response(
status_code=e.status_code,
cause=e,
is_context_overflow=e.is_context_overflow,
is_retryable=e.is_retryable,
) from e

async def generate_response(
Expand All @@ -115,6 +116,7 @@ async def generate_response(
status_code=e.status_code,
cause=e,
is_context_overflow=e.is_context_overflow,
is_retryable=e.is_retryable,
) from e

async def health_check(
Expand Down
Loading