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
97 changes: 88 additions & 9 deletions src/ava/transport/http.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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)
Expand All @@ -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:
Expand All @@ -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))
Expand All @@ -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.
Expand All @@ -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():
Expand All @@ -96,23 +162,36 @@ 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"),
)
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:
Expand Down
96 changes: 95 additions & 1 deletion tests/test_transport.py
Original file line number Diff line number Diff line change
@@ -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:
Expand Down
Loading