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
146 changes: 145 additions & 1 deletion src/archastro/platform/runtime/http_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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()

Expand Down Expand Up @@ -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):
Expand Down
125 changes: 125 additions & 0 deletions tests/test_http_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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={}))
Loading