Skip to content

Commit fa029a7

Browse files
fix(providers): normalize remaining native SDK transport failures (#184)
* fix(providers): normalize remaining native SDK transport failures Signed-off-by: Rudy Celekli <47457359+rudycelekli@users.noreply.github.com> * test(gemini): keep deadline transports local across client lifecycles Signed-off-by: Rudy Celekli <47457359+rudycelekli@users.noreply.github.com> --------- Signed-off-by: Rudy Celekli <47457359+rudycelekli@users.noreply.github.com> Co-authored-by: n-papaioannou <n.papaioannou@ime.life>
1 parent 24c0aca commit fa029a7

4 files changed

Lines changed: 170 additions & 7 deletions

File tree

‎ifixai/providers/bedrock.py‎

Lines changed: 13 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -87,11 +87,18 @@ async def send_message(
8787
),
8888
timeout=float(config.timeout),
8989
)
90-
except asyncio.TimeoutError as exc:
90+
except (
91+
asyncio.TimeoutError,
92+
botocore.exceptions.ReadTimeoutError,
93+
botocore.exceptions.ConnectTimeoutError,
94+
) as exc:
9195
raise ProviderTimeoutError(
9296
provider="bedrock",
9397
endpoint=endpoint,
94-
details=f"Request timed out after {config.timeout}s",
98+
details=(
99+
f"Request timed out after {config.timeout}s"
100+
if isinstance(exc, asyncio.TimeoutError) else str(exc)
101+
),
95102
) from exc
96103
except botocore.exceptions.NoCredentialsError as exc:
97104
raise ProviderAuthError(
@@ -138,7 +145,10 @@ async def send_message(
138145
endpoint=endpoint,
139146
details=str(exc),
140147
) from exc
141-
except botocore.exceptions.EndpointConnectionError as exc:
148+
except (
149+
botocore.exceptions.EndpointConnectionError,
150+
botocore.exceptions.ConnectionClosedError,
151+
) as exc:
142152
raise ProviderConnectionError(
143153
provider="bedrock",
144154
endpoint=endpoint,

‎ifixai/providers/gemini.py‎

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -107,11 +107,14 @@ async def send_message(
107107
)
108108
return "\n".join(text_parts)
109109

110-
except asyncio.TimeoutError as exc:
110+
except (asyncio.TimeoutError, google_exceptions.DeadlineExceeded) as exc:
111111
raise ProviderTimeoutError(
112112
provider="gemini",
113113
endpoint=endpoint,
114-
details=f"Request timed out after {config.timeout}s",
114+
details=(
115+
f"Request timed out after {config.timeout}s"
116+
if isinstance(exc, asyncio.TimeoutError) else str(exc)
117+
),
115118
) from exc
116119
except google_exceptions.Unauthenticated as exc:
117120
raise ProviderAuthError(

‎ifixai/providers/huggingface.py‎

Lines changed: 23 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,27 @@
2828
INITIAL_BACKOFF_SECONDS = 1.0
2929
BACKOFF_MULTIPLIER = 2.0
3030

31+
# Hugging Face switched from requests to httpx in 1.x. Keep the optional
32+
# provider compatible with both SDK transports without requiring either one
33+
# when the Hugging Face extra is not installed.
34+
_TRANSPORT_TIMEOUTS: tuple[type[Exception], ...] = (asyncio.TimeoutError,)
35+
_TRANSPORT_CONNECTION_ERRORS: tuple[type[Exception], ...] = (ConnectionError,)
36+
try:
37+
from httpx import NetworkError, ProxyError, RemoteProtocolError, TimeoutException
38+
except ImportError:
39+
pass
40+
else:
41+
_TRANSPORT_TIMEOUTS += (TimeoutException,)
42+
_TRANSPORT_CONNECTION_ERRORS += (NetworkError, ProxyError, RemoteProtocolError)
43+
try:
44+
from requests.exceptions import ConnectionError as RequestsConnectionError
45+
from requests.exceptions import Timeout as RequestsTimeout
46+
except ImportError:
47+
pass
48+
else:
49+
_TRANSPORT_TIMEOUTS += (RequestsTimeout,)
50+
_TRANSPORT_CONNECTION_ERRORS += (RequestsConnectionError,)
51+
3152

3253
class HuggingFaceProvider(ChatProvider):
3354
def __init__(self) -> None:
@@ -70,7 +91,7 @@ async def send_message(
7091
),
7192
timeout=float(config.timeout),
7293
)
73-
except asyncio.TimeoutError as exc:
94+
except _TRANSPORT_TIMEOUTS as exc:
7495
raise ProviderTimeoutError(
7596
provider="huggingface",
7697
endpoint=endpoint,
@@ -109,7 +130,7 @@ async def send_message(
109130
endpoint=endpoint,
110131
details=str(exc),
111132
) from exc
112-
except ConnectionError as exc:
133+
except _TRANSPORT_CONNECTION_ERRORS as exc:
113134
raise ProviderConnectionError(
114135
provider="huggingface",
115136
endpoint=endpoint,
Lines changed: 129 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,129 @@
1+
"""Normalize real SDK transport faults before the harness classifies them."""
2+
3+
import json
4+
import socket
5+
import threading
6+
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
7+
8+
import pytest
9+
10+
from ifixai.core.types import ChatMessage, ProviderConfig
11+
from ifixai.providers.base import ProviderConnectionError, ProviderTimeoutError
12+
13+
14+
@pytest.mark.asyncio
15+
@pytest.mark.parametrize("module,error_name,expected", [
16+
("httpx", "ReadTimeout", ProviderTimeoutError),
17+
("httpx", "ConnectError", ProviderConnectionError),
18+
("requests.exceptions", "ReadTimeout", ProviderTimeoutError),
19+
("requests.exceptions", "ConnectionError", ProviderConnectionError),
20+
])
21+
async def test_huggingface_transport_versions_are_normalized(monkeypatch, module, error_name, expected):
22+
pytest.importorskip("huggingface_hub")
23+
errors = pytest.importorskip(module)
24+
from ifixai.providers import huggingface
25+
26+
def fail(*args):
27+
raise getattr(errors, error_name)("owned transport failure")
28+
29+
monkeypatch.setattr(huggingface, "_call_chat_completion", fail)
30+
with pytest.raises(expected):
31+
await huggingface.HuggingFaceProvider().send_message([ChatMessage(content="hello")], ProviderConfig(provider="huggingface", model="owned-model", api_key="synthetic-local-key", max_retries=0))
32+
33+
34+
@pytest.mark.asyncio
35+
@pytest.mark.parametrize("provider_name", ["huggingface", "bedrock"])
36+
async def test_owned_http_disconnect_is_connection_error(monkeypatch, tmp_path, provider_name):
37+
seen = []
38+
39+
class Handler(BaseHTTPRequestHandler):
40+
def do_POST(self):
41+
seen.append(json.loads(self.rfile.read(int(self.headers["Content-Length"]))))
42+
self.connection.shutdown(socket.SHUT_RDWR)
43+
self.connection.close()
44+
45+
def log_message(self, *args):
46+
pass
47+
48+
server = ThreadingHTTPServer(("127.0.0.1", 0), Handler)
49+
thread = threading.Thread(target=server.serve_forever, daemon=True)
50+
thread.start()
51+
endpoint = f"http://127.0.0.1:{server.server_port}"
52+
try:
53+
if provider_name == "huggingface":
54+
pytest.importorskip("huggingface_hub")
55+
from ifixai.providers.huggingface import HuggingFaceProvider
56+
57+
provider = HuggingFaceProvider()
58+
# Use the existing model-URL route, independently of the endpoint fix.
59+
model = endpoint + "/v1/chat/completions"
60+
else:
61+
boto3 = pytest.importorskip("boto3")
62+
from botocore.config import Config
63+
64+
from ifixai.providers.bedrock import BedrockProvider
65+
66+
session = boto3.Session(aws_access_key_id="synthetic-local-key", aws_secret_access_key="synthetic-local-secret", region_name="us-east-1")
67+
native_client = session.client("bedrock-runtime", endpoint_url=endpoint, config=Config(retries={"total_max_attempts": 1}))
68+
69+
class OwnedSession:
70+
def client(self, **kwargs):
71+
assert kwargs["endpoint_url"] == endpoint
72+
return native_client
73+
74+
monkeypatch.setattr(boto3, "Session", lambda **kwargs: OwnedSession())
75+
provider = BedrockProvider()
76+
model = "owned-model"
77+
config = ProviderConfig(provider=provider_name, endpoint=endpoint, model=model, api_key="synthetic-local-key", max_retries=0)
78+
with pytest.raises(ProviderConnectionError):
79+
await provider.send_message([ChatMessage(content="hello")], config)
80+
assert len(seen) == 1
81+
finally:
82+
server.shutdown()
83+
server.server_close()
84+
thread.join()
85+
86+
87+
@pytest.mark.asyncio
88+
async def test_gemini_native_deadline_is_timeout(monkeypatch):
89+
pytest.importorskip("google.generativeai")
90+
from google.api_core.exceptions import DeadlineExceeded
91+
from google.generativeai import client
92+
93+
from ifixai.providers import gemini
94+
95+
class OwnedTransport:
96+
async def __aenter__(self):
97+
return self
98+
99+
async def __aexit__(self, *args):
100+
pass
101+
102+
async def generate_content(self, request, **kwargs):
103+
raise DeadlineExceeded("owned transport deadline")
104+
105+
monkeypatch.setattr(client, "get_default_generative_async_client", OwnedTransport)
106+
monkeypatch.setattr(gemini, "GenerativeServiceAsyncClient", lambda **kwargs: OwnedTransport(), raising=False)
107+
with pytest.raises(ProviderTimeoutError):
108+
await gemini.GeminiProvider().send_message([ChatMessage(content="hello")], ProviderConfig(provider="gemini", api_key="synthetic-local-key", max_retries=0))
109+
110+
111+
@pytest.mark.asyncio
112+
@pytest.mark.parametrize("error_name", ["ReadTimeoutError", "ConnectTimeoutError"])
113+
async def test_bedrock_sdk_deadlines_are_timeouts(monkeypatch, error_name):
114+
boto3 = pytest.importorskip("boto3")
115+
import botocore.exceptions
116+
117+
from ifixai.providers.bedrock import BedrockProvider
118+
119+
class OwnedClient:
120+
def converse(self, **kwargs):
121+
raise getattr(botocore.exceptions, error_name)(endpoint_url="http://owned.invalid")
122+
123+
class OwnedSession:
124+
def client(self, **kwargs):
125+
return OwnedClient()
126+
127+
monkeypatch.setattr(boto3, "Session", lambda **kwargs: OwnedSession())
128+
with pytest.raises(ProviderTimeoutError):
129+
await BedrockProvider().send_message([ChatMessage(content="hello")], ProviderConfig(provider="bedrock", model="owned-model", max_retries=0))

0 commit comments

Comments
 (0)