From b7300ecd50ea55d77cf549278e82b64961e014cd Mon Sep 17 00:00:00 2001 From: Min Liu Date: Mon, 14 Sep 2026 16:01:48 -0700 Subject: [PATCH] fix: retry early provider stream disconnects and record diagnostics --- src/ava/transport/http.py | 97 +++++++++++++++++++++++++++++++++++---- tests/test_transport.py | 96 +++++++++++++++++++++++++++++++++++++- 2 files changed, 183 insertions(+), 10 deletions(-) diff --git a/src/ava/transport/http.py b/src/ava/transport/http.py index b945e5b..a7e7bfd 100644 --- a/src/ava/transport/http.py +++ b/src/ava/transport/http.py @@ -6,6 +6,10 @@ from __future__ import annotations +import asyncio +import json +import logging +import time from collections.abc import Callable from dataclasses import dataclass, field @@ -20,6 +24,8 @@ MODEL_DISCOVERY_TIMEOUT_SECONDS = 30.0 SseSink = Callable[[SseEvent], None] +_LOG = logging.getLogger(__name__) +STREAM_RETRY_DELAYS = (1.0, 2.0) @dataclass(slots=True) @@ -36,10 +42,15 @@ class Response: body: str = "" -def _transport_error(exception: Exception) -> AvaError: +class _HttpTransferError(AvaError): + # Internal retry eligibility, distinct from AvaError.recoverable (drive semantics). + retryable: bool = False + + +def _transport_error(exception: Exception) -> _HttpTransferError: if isinstance(exception, httpx.TimeoutException): - return AvaError(ErrorKind.timeout, f"provider request failed: {exception}") - return AvaError(ErrorKind.network, f"provider request failed: {exception}") + return _HttpTransferError(ErrorKind.timeout, f"provider request failed: {exception}") + return _HttpTransferError(ErrorKind.network, f"provider request failed: {exception}") class Client: @@ -54,7 +65,35 @@ async def aclose(self) -> None: async def post_sse( self, request: Request, sink: SseSink, cancel: CancelToken = NEVER ) -> Response: - return await cancel.guard(self._transfer(request, "POST", sink, stream_timeout=True)) + cancel.raise_if_cancelled() + return await cancel.guard(self._post_sse_with_retries(request, sink)) + + async def _post_sse_with_retries(self, request: Request, sink: SseSink) -> Response: + delivered = False + + def deliver(event: SseEvent) -> None: + nonlocal delivered + # Even metadata can mutate a provider adapter. Never replay after delivery. + delivered = True + sink(event) + + for attempt in range(len(STREAM_RETRY_DELAYS) + 1): + try: + return await self._transfer( + request, "POST", deliver, stream_timeout=True, attempt=attempt + 1 + ) + except _HttpTransferError as error: + transient = isinstance( + error.__cause__, + (httpx.NetworkError, httpx.RemoteProtocolError, httpx.TimeoutException), + ) + if ( + not transient or not error.retryable or delivered + or attempt == len(STREAM_RETRY_DELAYS) + ): + raise + await asyncio.sleep(STREAM_RETRY_DELAYS[attempt]) + raise AssertionError("unreachable") async def post(self, request: Request, cancel: CancelToken = NEVER) -> Response: return await cancel.guard(self._transfer(request, "POST", None, stream_timeout=False)) @@ -63,7 +102,8 @@ async def get(self, request: Request, cancel: CancelToken = NEVER) -> Response: return await cancel.guard(self._transfer(request, "GET", None, stream_timeout=False)) async def _transfer( - self, request: Request, method: str, sink: SseSink | None, *, stream_timeout: bool + self, request: Request, method: str, sink: SseSink | None, *, stream_timeout: bool, + attempt: int = 1, ) -> Response: if stream_timeout: # Idle timeout only: a total timeout kills healthy long streams. @@ -82,10 +122,36 @@ async def _transfer( parser = SseParser() received = 0 body_parts: list[bytes] = [] + started = time.monotonic() + diagnostics: dict[str, object] = {"attempt": attempt} + last_event: str | None = None + terminal_received = False + + def deliver(event: SseEvent) -> None: + nonlocal last_event, terminal_received + # Never retain event payloads, URLs, or arbitrary response headers. + kind = event.event + try: + value = json.loads(event.data) + except (ValueError, TypeError): + value = None + if isinstance(value, dict) and isinstance(value.get("type"), str): + kind = value["type"] + last_event = kind[:80] if kind else None + terminal_received |= kind in ( + "response.completed", "response.incomplete", "response.failed", "message_stop" + ) or event.data.strip() == "[DONE]" + assert sink is not None + sink(event) + try: async with self._client.stream( method, request.url, headers=headers, content=content, timeout=timeout ) as response: + diagnostics.update(status=response.status_code, http_version=response.http_version) + request_id = response.headers.get("x-request-id") or response.headers.get("request-id") + if request_id: + diagnostics["request_id"] = request_id[:200] streaming = sink is not None and 200 <= response.status_code < 300 decoder: _Utf8Decoder | None = _Utf8Decoder() if streaming else None async for chunk in response.aiter_bytes(): @@ -96,15 +162,15 @@ async def _transfer( ) if decoder is not None and sink is not None: for event in parser.feed(decoder.feed(chunk)): - sink(event) + deliver(event) else: body_parts.append(chunk) if decoder is not None and sink is not None: tail = decoder.finish() for event in parser.feed(tail) if tail else []: - sink(event) + deliver(event) for event in parser.finish(): - sink(event) + deliver(event) return Response( status=response.status_code, body=b"".join(body_parts).decode("utf-8", "replace"), @@ -112,7 +178,20 @@ async def _transfer( except AvaError: raise except httpx.HTTPError as exception: - raise _transport_error(exception) from exception + diagnostics.update( + elapsed_ms=int((time.monotonic() - started) * 1000), + bytes_received=received, + last_sse_event=last_event, + terminal_received=terminal_received, + exception_type=type(exception).__name__, + ) + detail = json.dumps(diagnostics, sort_keys=True) + _LOG.warning("Provider HTTP attempt failed: %s", detail) + error = _transport_error(exception) + error.detail = detail + status = diagnostics.get("status") + error.retryable = status is None or (isinstance(status, int) and 200 <= status < 300) + raise error from exception class _Utf8Decoder: diff --git a/tests/test_transport.py b/tests/test_transport.py index ef415df..25e471a 100644 --- a/tests/test_transport.py +++ b/tests/test_transport.py @@ -1,4 +1,98 @@ -from ava.transport import SseParser +import asyncio +import json +from contextlib import asynccontextmanager + +import pytest + +from ava.base import AvaError, CancelToken, ErrorKind +from ava.transport import Client, Request, SseParser + + +@asynccontextmanager +async def chunked_peer(responses): + requests = [] + + async def handle(reader, writer): + try: + headers = await reader.readuntil(b"\r\n\r\n") + length = next(int(line.split(b":", 1)[1]) for line in headers.split(b"\r\n") + if line.lower().startswith(b"content-length:")) + requests.append(await reader.readexactly(length)) + status, payload, complete = responses[min(len(requests) - 1, len(responses) - 1)] + writer.write( + f"HTTP/1.1 {status} Test\r\n".encode() + + b"Content-Type: text/event-stream\r\nTransfer-Encoding: chunked\r\n" + + b"X-Request-ID: fixture-request\r\nConnection: close\r\n\r\n" + ) + if payload: + writer.write(f"{len(payload):x}\r\n".encode() + payload + b"\r\n") + if complete: + writer.write(b"0\r\n\r\n") + await writer.drain() + finally: + writer.close() + await writer.wait_closed() + + server = await asyncio.start_server(handle, "127.0.0.1", 0) + client = Client() + try: + yield client, Request( + url=f"http://127.0.0.1:{server.sockets[0].getsockname()[1]}/responses?secret=query", + headers=[("authorization", "Bearer secret-token")], body="secret-body", + ), requests + finally: + await client.aclose() + server.close() + await server.wait_closed() + + +@pytest.mark.parametrize("mode", ["recover", "exhaust", "partial", "http_error", "terminal"]) +async def test_stream_disconnect_recovery_and_diagnostics(mode, monkeypatch, caplog): + monkeypatch.setattr("ava.transport.http.STREAM_RETRY_DELAYS", (0, 0)) + payload = b'data: {"type":"response.created"}\n\n' + first = { + "recover": (200, b": heartbeat\n\n", False), + "exhaust": (200, b"", False), + "partial": (200, payload, False), + "http_error": (401, b"secret-response", False), + "terminal": (200, b'data: {"type":"response.completed"}\n\n', False), + }[mode] + responses = [first, (200, payload, True)] if mode == "recover" else [first] + async with chunked_peer(responses) as (client, request, requests): + events = [] + if mode == "recover": + assert (await client.post_sse(request, events.append)).status == 200 + assert len(requests) == 2 + assert len(events) == 1 + assert requests == [b"secret-body"] * 2 + else: + with pytest.raises(AvaError, match="incomplete chunked read") as caught: + await client.post_sse(request, events.append) + assert len(requests) == (3 if mode == "exhaust" else 1) + assert len(events) == (1 if mode in ("partial", "terminal") else 0) + detail = json.loads(caught.value.detail) + assert detail["request_id"] == "fixture-request" + assert detail["http_version"] == "HTTP/1.1" + assert detail["bytes_received"] == len(first[1]) + assert detail["terminal_received"] == (mode == "terminal") + assert detail["elapsed_ms"] >= 0 + assert "secret" not in caplog.text + + +async def test_cancel_during_stream_retry_backoff(monkeypatch): + monkeypatch.setattr("ava.transport.http.STREAM_RETRY_DELAYS", (60, 60)) + async with chunked_peer([(200, b"", False)]) as (client, request, requests): + token = CancelToken() + task = asyncio.create_task(client.post_sse(request, lambda event: None, token)) + async with asyncio.timeout(2): + while not requests: + await asyncio.sleep(0.001) + await asyncio.sleep(0.05) + token.cancel() + with pytest.raises(AvaError) as caught: + await task + assert caught.value.kind == ErrorKind.cancelled + assert len(requests) == 1 def _feed_in_chunks(text: str, size: int) -> list: