From c6f341e1a4dd4e9f7886b42abf21c3c40482090b Mon Sep 17 00:00:00 2001 From: Calvin Grunewald Date: Mon, 29 Jun 2026 16:14:39 -0700 Subject: [PATCH] feat(runtime): add stream_sse / stream_sse_sync for SSE endpoints MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Adds the SSE primitives the generated stream() methods target for x-sdk-streaming endpoints: HttpClient.stream_sse (async) and SyncHttpClient.stream_sse_sync (sync). Each opens the request (any method + JSON body), sets Accept: text/event-stream, raises ApiError on a non-2xx before the stream opens, and yields parsed {"event", "data"} records from the event stream via httpx streaming. (Token refresh is not retried mid-stream.) Runtime-only — no codegen changes. Tests cover async + sync ordered events + body send and non-2xx -> ApiError. 42 pass; ruff clean. Co-Authored-By: Claude Opus 4.8 (1M context) --- src/archastro/platform/runtime/http_client.py | 146 +++++++++++++++++- tests/test_http_client.py | 125 +++++++++++++++ 2 files changed, 270 insertions(+), 1 deletion(-) diff --git a/src/archastro/platform/runtime/http_client.py b/src/archastro/platform/runtime/http_client.py index 4676e11..2ba290d 100644 --- a/src/archastro/platform/runtime/http_client.py +++ b/src/archastro/platform/runtime/http_client.py @@ -4,8 +4,9 @@ from __future__ import annotations import asyncio +import json import threading -from collections.abc import Callable, Coroutine +from collections.abc import AsyncIterator, Callable, Coroutine, Iterator from functools import cache from typing import Any, TypeVar, overload @@ -246,6 +247,72 @@ async def request_raw( "mime_type": response.headers.get("content-type", "text/plain"), } + async def stream_sse( + self, + path: str, + *, + method: str = "GET", + body: Any = None, + headers: dict[str, str] | None = None, + query: dict[str, Any] | None = None, + ) -> AsyncIterator[dict[str, Any]]: + """Open a Server-Sent Events stream, yielding parsed ``{"event", "data"}``. + + Backs the generated async ``stream()`` methods for ``x-sdk-streaming`` + endpoints. Raises :class:`ApiError` on a non-2xx response before the + stream opens. (Token refresh is not retried mid-stream.) + """ + auth_prefix = f"{DEFAULT_API_PREFIX}/auth/" + if self._refresh_only and not path.startswith(auth_prefix): + raise RuntimeError( + f"Refresh-only HTTP client cannot make requests outside {auth_prefix}" + ) + + url = f"{self._base_url}{self._transform_path(path)}" + sends_body = body is not None and method not in ("GET", "HEAD") + req_headers = {**self._default_headers, "Accept": "text/event-stream"} + if sends_body: + req_headers["Content-Type"] = "application/json" + token = self._get_token() + if token: + req_headers["Authorization"] = f"Bearer {token}" + if headers: + req_headers.update(headers) + params = {k: v for k, v in (query or {}).items() if v is not None} or None + + async with self._client.stream( + method, + url, + json=body if sends_body else None, + headers=req_headers, + params=params, + ) as response: + if response.status_code >= 400: + await response.aread() + raw: dict[str, Any] = {} + try: + raw = response.json() + except Exception: + pass + code, message = _parse_error(raw, response.status_code) + raise ApiError(response.status_code, code, message, raw) + + event: str | None = None + data_lines: list[str] = [] + async for line in response.aiter_lines(): + if line == "": + parsed = _build_sse_event(event, data_lines) + if parsed is not None: + yield parsed + event, data_lines = None, [] + elif line.startswith("event:"): + event = line[6:].strip() + elif line.startswith("data:"): + data_lines.append(line[5:].strip()) + parsed = _build_sse_event(event, data_lines) + if parsed is not None: + yield parsed + async def close(self) -> None: await self._client.aclose() @@ -440,10 +507,87 @@ def request_raw( "mime_type": response.headers.get("content-type", "text/plain"), } + def stream_sse_sync( + self, + path: str, + *, + method: str = "GET", + body: Any = None, + headers: dict[str, str] | None = None, + query: dict[str, Any] | None = None, + ) -> Iterator[dict[str, Any]]: + """Synchronous counterpart to :meth:`HttpClient.stream_sse`. + + Backs the generated sync ``stream()`` methods. Raises :class:`ApiError` + on a non-2xx response before the stream opens. + """ + auth_prefix = f"{DEFAULT_API_PREFIX}/auth/" + if self._refresh_only and not path.startswith(auth_prefix): + raise RuntimeError( + f"Refresh-only HTTP client cannot make requests outside {auth_prefix}" + ) + + url = f"{self._base_url}{self._transform_path(path)}" + sends_body = body is not None and method not in ("GET", "HEAD") + req_headers = {**self._default_headers, "Accept": "text/event-stream"} + if sends_body: + req_headers["Content-Type"] = "application/json" + token = self._get_token() + if token: + req_headers["Authorization"] = f"Bearer {token}" + if headers: + req_headers.update(headers) + params = {k: v for k, v in (query or {}).items() if v is not None} or None + + with self._client.stream( + method, + url, + json=body if sends_body else None, + headers=req_headers, + params=params, + ) as response: + if response.status_code >= 400: + response.read() + raw: dict[str, Any] = {} + try: + raw = response.json() + except Exception: + pass + code, message = _parse_error(raw, response.status_code) + raise ApiError(response.status_code, code, message, raw) + + event: str | None = None + data_lines: list[str] = [] + for line in response.iter_lines(): + if line == "": + parsed = _build_sse_event(event, data_lines) + if parsed is not None: + yield parsed + event, data_lines = None, [] + elif line.startswith("event:"): + event = line[6:].strip() + elif line.startswith("data:"): + data_lines.append(line[5:].strip()) + parsed = _build_sse_event(event, data_lines) + if parsed is not None: + yield parsed + def close(self) -> None: self._client.close() +def _build_sse_event(event: str | None, data_lines: list[str]) -> dict[str, Any] | None: + """Assemble one SSE frame into ``{"event", "data"}``; ``None`` if empty.""" + if event is None and not data_lines: + return None + raw = "\n".join(data_lines) + try: + data: Any = json.loads(raw) + except Exception: + data = raw + return {"event": event or "message", "data": data} + + def _parse_error(raw_data: dict[str, Any], status: int) -> tuple[str, str]: error = raw_data.get("error") if isinstance(error, dict): diff --git a/tests/test_http_client.py b/tests/test_http_client.py index 11c8b18..602898b 100644 --- a/tests/test_http_client.py +++ b/tests/test_http_client.py @@ -731,3 +731,128 @@ def test_non_json_error_body_falls_back_to_http_status_message(): assert exc_info.value.error_code == "unknown_error" assert str(exc_info.value) == "HTTP 500" + + +# ─── SSE streaming (stream_sse / stream_sse_sync) ────────────────── + + +class _FakeAsyncResponse: + def __init__(self, status_code, lines, json_body=None): + self.status_code = status_code + self._lines = lines + self._json = json_body or {} + + async def aiter_lines(self): + for line in self._lines: + yield line + + async def aread(self): + return b"" + + def json(self): + return self._json + + +class _FakeAsyncStream: + def __init__(self, response): + self._response = response + + async def __aenter__(self): + return self._response + + async def __aexit__(self, *exc): + return False + + +class _FakeSyncResponse: + def __init__(self, status_code, lines, json_body=None): + self.status_code = status_code + self._lines = lines + self._json = json_body or {} + + def iter_lines(self): + yield from self._lines + + def read(self): + return b"" + + def json(self): + return self._json + + +class _FakeSyncStream: + def __init__(self, response): + self._response = response + + def __enter__(self): + return self._response + + def __exit__(self, *exc): + return False + + +_SSE_LINES = [ + "event: chunk", + 'data: {"text": "He"}', + "", + "event: chunk", + 'data: {"text": "llo"}', + "", + "event: done", + 'data: {"ok": true}', + "", +] + + +async def test_stream_sse_yields_parsed_events_and_sends_body(): + client = HttpClient(base_url="https://api.test") + resp = _FakeAsyncResponse(200, _SSE_LINES) + with patch.object(client._client, "stream", return_value=_FakeAsyncStream(resp)) as m: + events = [ + ev + async for ev in client.stream_sse( + "/api/v1/echo/stream", method="POST", body={"prompt": "hi"} + ) + ] + + assert events == [ + {"event": "chunk", "data": {"text": "He"}}, + {"event": "chunk", "data": {"text": "llo"}}, + {"event": "done", "data": {"ok": True}}, + ] + _, kwargs = m.call_args + assert kwargs["json"] == {"prompt": "hi"} + assert kwargs["headers"]["Accept"] == "text/event-stream" + + +async def test_stream_sse_raises_apierror_on_non_2xx(): + client = HttpClient(base_url="https://api.test") + resp = _FakeAsyncResponse( + 402, [], json_body={"error": {"code": "plan_not_entitled", "message": "no"}} + ) + with patch.object(client._client, "stream", return_value=_FakeAsyncStream(resp)): + with pytest.raises(ApiError) as exc_info: + [ev async for ev in client.stream_sse("/api/v1/x/stream", method="POST", body={})] + assert exc_info.value.status == 402 + + +def test_stream_sse_sync_yields_parsed_events(): + client = SyncHttpClient(base_url="https://api.test") + resp = _FakeSyncResponse(200, _SSE_LINES) + with patch.object(client._client, "stream", return_value=_FakeSyncStream(resp)): + events = list( + client.stream_sse_sync("/api/v1/echo/stream", method="POST", body={"prompt": "hi"}) + ) + assert events == [ + {"event": "chunk", "data": {"text": "He"}}, + {"event": "chunk", "data": {"text": "llo"}}, + {"event": "done", "data": {"ok": True}}, + ] + + +def test_stream_sse_sync_raises_apierror_on_non_2xx(): + client = SyncHttpClient(base_url="https://api.test") + resp = _FakeSyncResponse(401, [], json_body={"error": "unauthenticated"}) + with patch.object(client._client, "stream", return_value=_FakeSyncStream(resp)): + with pytest.raises(ApiError): + list(client.stream_sse_sync("/api/v1/x/stream", method="POST", body={}))