Files
langgraph/libs/sdk-py/tests/test_client_stream.py
T
efe78447b3 feat(sdk-py): emit id as part of stream events (#6581)
- **Description:** Provide the id of the event for routes that use SSE
streams. This will allow for more custom retry logic when streams
disconnect if needed.
  - **Issue:** N/A
  - **Dependencies:** N/A
  - **Twitter handle:** N/A

---------

Co-authored-by: William FH <13333726+hinthornw@users.noreply.github.com>
2025-12-12 17:16:48 -05:00

249 lines
8.0 KiB
Python

from __future__ import annotations
from collections.abc import Iterator, Sequence
from pathlib import Path
import httpx
import pytest
from langgraph_sdk.client import HttpClient, SyncHttpClient
from langgraph_sdk.schema import StreamPart
from langgraph_sdk.sse import BytesLike, BytesLineDecoder, SSEDecoder
with open(Path(__file__).parent / "fixtures" / "response.txt", "rb") as f:
RESPONSE_PAYLOAD = f.read()
class AsyncListByteStream(httpx.AsyncByteStream):
def __init__(self, chunks: Sequence[bytes], exc: Exception | None = None) -> None:
self._chunks = list(chunks)
self._exc = exc
async def __aiter__(self): # type: ignore[override]
for chunk in self._chunks:
yield chunk
if self._exc is not None:
raise self._exc
async def aclose(self) -> None:
return None
class ListByteStream(httpx.ByteStream):
def __init__(self, chunks: Sequence[bytes], exc: Exception | None = None) -> None:
self._chunks = list(chunks)
self._exc = exc
def __iter__(self): # type: ignore[override]
yield from self._chunks
if self._exc is not None:
raise self._exc
def close(self) -> None:
return None
def iter_lines_raw(payload: list[bytes]) -> Iterator[BytesLike]:
decoder = BytesLineDecoder()
for part in payload:
yield from decoder.decode(part)
yield from decoder.flush()
def test_stream_sse():
for groups in (
[RESPONSE_PAYLOAD],
RESPONSE_PAYLOAD.splitlines(keepends=True),
):
parts: list[StreamPart] = []
decoder = SSEDecoder()
for line in iter_lines_raw(groups):
sse = decoder.decode(line=line.rstrip(b"\n")) # type: ignore
if sse is not None:
parts.append(sse)
if sse := decoder.decode(b""):
parts.append(sse)
assert decoder.decode(b"") is None
assert len(parts) == 79
@pytest.mark.asyncio
async def test_http_client_stream_flushes_trailing_event():
payload = b'event: foo\ndata: {"bar": 1}\n'
async def handler(request: httpx.Request) -> httpx.Response:
assert request.headers["accept"] == "text/event-stream"
assert request.headers["cache-control"] == "no-store"
return httpx.Response(
200,
headers={"Content-Type": "text/event-stream"},
content=payload,
)
transport = httpx.MockTransport(handler)
async with httpx.AsyncClient(
transport=transport, base_url="https://example.com"
) as client:
http_client = HttpClient(client)
parts = [part async for part in http_client.stream("/stream", "GET")]
assert parts == [StreamPart(event="foo", data={"bar": 1})]
def test_sync_http_client_stream_recovers_after_disconnect():
reconnect_path = "/reconnect"
first_chunks = [
b"id: 1\n",
b"event: values\n",
b'data: {"step": 1}\n\n',
]
second_chunks = [
b"id: 2\n",
b"event: values\n",
b'data: {"step": 2}\n\n',
b"event: end\n",
b"data: null\n\n",
]
call_count = 0
def handler(request: httpx.Request) -> httpx.Response:
nonlocal call_count
call_count += 1
if call_count == 1:
assert request.method == "POST"
assert request.url.path == "/stream"
assert request.headers["accept"] == "text/event-stream"
assert request.headers["cache-control"] == "no-store"
assert "last-event-id" not in {
k.lower(): v for k, v in request.headers.items()
}
assert request.read()
return httpx.Response(
200,
headers={
"Content-Type": "text/event-stream",
"Location": reconnect_path,
},
stream=ListByteStream(
first_chunks,
httpx.RemoteProtocolError("incomplete chunked read"),
),
)
if call_count == 2:
assert request.method == "GET"
assert request.url.path == reconnect_path
assert request.headers["Last-Event-ID"] == "1"
assert request.read() == b""
return httpx.Response(
200,
headers={"Content-Type": "text/event-stream"},
stream=ListByteStream(second_chunks),
)
raise AssertionError("unexpected request")
transport = httpx.MockTransport(handler)
with httpx.Client(transport=transport, base_url="https://example.com") as client:
http_client = SyncHttpClient(client)
parts = list(http_client.stream("/stream", "POST", json={"payload": "value"}))
assert call_count == 2
assert parts == [
StreamPart(event="values", data={"step": 1}, id="1"),
StreamPart(event="values", data={"step": 2}, id="2"),
StreamPart(event="end", data=None, id="2"),
]
@pytest.mark.asyncio
async def test_http_client_stream_recovers_after_disconnect():
reconnect_path = "/reconnect"
first_chunks = [
b"id: 1\n",
b"event: values\n",
b'data: {"step": 1}\n\n',
]
second_chunks = [
b"id: 2\n",
b"event: values\n",
b'data: {"step": 2}\n\n',
b"event: end\n",
b"data: null\n\n",
]
call_count = 0
async def handler(request: httpx.Request) -> httpx.Response:
nonlocal call_count
call_count += 1
if call_count == 1:
assert request.method == "POST"
assert request.url.path == "/stream"
assert request.headers["accept"] == "text/event-stream"
assert request.headers["cache-control"] == "no-store"
assert "last-event-id" not in {
k.lower(): v for k, v in request.headers.items()
}
assert await request.aread()
return httpx.Response(
200,
headers={
"Content-Type": "text/event-stream",
"Location": reconnect_path,
},
stream=AsyncListByteStream(
first_chunks,
httpx.RemoteProtocolError("incomplete chunked read"),
),
)
if call_count == 2:
assert request.method == "GET"
assert request.url.path == reconnect_path
assert request.headers["Last-Event-ID"] == "1"
assert await request.aread() == b""
return httpx.Response(
200,
headers={"Content-Type": "text/event-stream"},
stream=AsyncListByteStream(second_chunks),
)
raise AssertionError("unexpected request")
transport = httpx.MockTransport(handler)
async with httpx.AsyncClient(
transport=transport, base_url="https://example.com"
) as client:
http_client = HttpClient(client)
parts = [
part
async for part in http_client.stream(
"/stream", "POST", json={"payload": "value"}
)
]
assert call_count == 2
assert parts == [
StreamPart(event="values", data={"step": 1}, id="1"),
StreamPart(event="values", data={"step": 2}, id="2"),
StreamPart(event="end", data=None, id="2"),
]
def test_sync_http_client_stream_flushes_trailing_event():
payload = b'event: foo\ndata: {"bar": 1}\n'
def handler(request: httpx.Request) -> httpx.Response:
assert request.headers["accept"] == "text/event-stream"
assert request.headers["cache-control"] == "no-store"
return httpx.Response(
200,
headers={"Content-Type": "text/event-stream"},
content=payload,
)
transport = httpx.MockTransport(handler)
with httpx.Client(transport=transport, base_url="https://example.com") as client:
http_client = SyncHttpClient(client)
parts = list(http_client.stream("/stream", "GET"))
assert parts == [StreamPart(event="foo", data={"bar": 1})]