diff --git a/libs/sdk-py/langgraph_sdk/client.py b/libs/sdk-py/langgraph_sdk/client.py index 2605d5bef..2d1b5e082 100644 --- a/libs/sdk-py/langgraph_sdk/client.py +++ b/libs/sdk-py/langgraph_sdk/client.py @@ -455,36 +455,94 @@ class HttpClient: if headers: request_headers.update(headers) - async with self.client.stream( - method, path, headers=request_headers, content=content, params=params - ) as res: - if on_response: - on_response(res) - # check status - try: - res.raise_for_status() - except httpx.HTTPStatusError as e: - body = (await res.aread()).decode() - if sys.version_info >= (3, 11): - e.add_note(body) + reconnect_headers = { + key: value + for key, value in request_headers.items() + if key.lower() not in {"content-length", "content-type"} + } + + last_event_id: str | None = None + reconnect_path: str | None = None + reconnect_attempts = 0 + max_reconnect_attempts = 5 + + while True: + current_headers = dict( + request_headers if reconnect_path is None else reconnect_headers + ) + if last_event_id is not None: + current_headers["Last-Event-ID"] = last_event_id + + current_method = method if reconnect_path is None else "GET" + current_content = content if reconnect_path is None else None + current_params = params if reconnect_path is None else None + + retry = False + async with self.client.stream( + current_method, + reconnect_path or path, + headers=current_headers, + content=current_content, + params=current_params, + ) as res: + if reconnect_path is None and on_response: + on_response(res) + # check status + try: + res.raise_for_status() + except httpx.HTTPStatusError as e: + body = (await res.aread()).decode() + if sys.version_info >= (3, 11): + e.add_note(body) + else: + logger.error(f"Error from langgraph-api: {body}", exc_info=e) + raise e + # check content type + content_type = res.headers.get("content-type", "").partition(";")[0] + if "text/event-stream" not in content_type: + raise httpx.TransportError( + "Expected response header Content-Type to contain 'text/event-stream', " + f"got {content_type!r}" + ) + + reconnect_location = res.headers.get("location") + if reconnect_location: + reconnect_path = reconnect_location + + # parse SSE + decoder = SSEDecoder() + try: + async for line in aiter_lines_raw(res): + sse = decoder.decode(line=line.rstrip(b"\n")) + if sse is not None: + if decoder.last_event_id is not None: + last_event_id = decoder.last_event_id + if sse.event or sse.data is not None: + yield sse + except httpx.HTTPError: + # httpx.TransportError inherits from HTTPError, so transient + # disconnects during streaming land here. + if reconnect_path is None: + raise + retry = True else: - logger.error(f"Error from langgraph-api: {body}", exc_info=e) - raise e - # check content type - content_type = res.headers.get("content-type", "").partition(";")[0] - if "text/event-stream" not in content_type: - raise httpx.TransportError( - "Expected response header Content-Type to contain 'text/event-stream', " - f"got {content_type!r}" - ) - # parse SSE - decoder = SSEDecoder() - async for line in aiter_lines_raw(res): - sse = decoder.decode(line=line.rstrip(b"\n")) - if sse is not None: - yield sse - if sse := decoder.decode(b""): - yield sse + if sse := decoder.decode(b""): + if decoder.last_event_id is not None: + last_event_id = decoder.last_event_id + if sse.event or sse.data is not None: + # decoder.decode(b"") flushes the in-flight event and may + # return an empty placeholder when there is no pending + # message. Skip these no-op events so the stream doesn't + # emit a trailing blank item after reconnects. + yield sse + if retry: + reconnect_attempts += 1 + if reconnect_attempts > max_reconnect_attempts: + raise httpx.TransportError( + "Exceeded maximum SSE reconnection attempts" + ) + continue + break async def _aencode_json(json: Any) -> tuple[dict[str, str], bytes | None]: @@ -3642,41 +3700,100 @@ class SyncHttpClient: on_response: Callable[[httpx.Response], None] | None = None, ) -> Iterator[StreamPart]: """Stream the results of a request using SSE.""" - request_headers, content = _encode_json(json) + if json is not None: + request_headers, content = _encode_json(json) + else: + request_headers, content = {}, None request_headers["Accept"] = "text/event-stream" request_headers["Cache-Control"] = "no-store" if headers: request_headers.update(headers) - with self.client.stream( - method, path, headers=request_headers, content=content, params=params - ) as res: - if on_response: - on_response(res) - # check status - try: - res.raise_for_status() - except httpx.HTTPStatusError as e: - body = (res.read()).decode() - if sys.version_info >= (3, 11): - e.add_note(body) + + reconnect_headers = { + key: value + for key, value in request_headers.items() + if key.lower() not in {"content-length", "content-type"} + } + + last_event_id: str | None = None + reconnect_path: str | None = None + reconnect_attempts = 0 + max_reconnect_attempts = 5 + + while True: + current_headers = dict( + request_headers if reconnect_path is None else reconnect_headers + ) + if last_event_id is not None: + current_headers["Last-Event-ID"] = last_event_id + + current_method = method if reconnect_path is None else "GET" + current_content = content if reconnect_path is None else None + current_params = params if reconnect_path is None else None + + retry = False + with self.client.stream( + current_method, + reconnect_path or path, + headers=current_headers, + content=current_content, + params=current_params, + ) as res: + if reconnect_path is None and on_response: + on_response(res) + # check status + try: + res.raise_for_status() + except httpx.HTTPStatusError as e: + body = (res.read()).decode() + if sys.version_info >= (3, 11): + e.add_note(body) + else: + logger.error(f"Error from langgraph-api: {body}", exc_info=e) + raise e + # check content type + content_type = res.headers.get("content-type", "").partition(";")[0] + if "text/event-stream" not in content_type: + raise httpx.TransportError( + "Expected response header Content-Type to contain 'text/event-stream', " + f"got {content_type!r}" + ) + + reconnect_location = res.headers.get("location") + if reconnect_location: + reconnect_path = reconnect_location + + decoder = SSEDecoder() + try: + for line in iter_lines_raw(res): + sse = decoder.decode(line.rstrip(b"\n")) + if sse is not None: + if decoder.last_event_id is not None: + last_event_id = decoder.last_event_id + if sse.event or sse.data is not None: + yield sse + except httpx.HTTPError: + # httpx.TransportError inherits from HTTPError, so transient + # disconnects during streaming land here. + if reconnect_path is None: + raise + retry = True else: - logger.error(f"Error from langgraph-api: {body}", exc_info=e) - raise e - # check content type - content_type = res.headers.get("content-type", "").partition(";")[0] - if "text/event-stream" not in content_type: - raise httpx.TransportError( - "Expected response header Content-Type to contain 'text/event-stream', " - f"got {content_type!r}" - ) - # parse SSE - decoder = SSEDecoder() - for line in iter_lines_raw(res): - sse = decoder.decode(line.rstrip(b"\n")) - if sse is not None: - yield sse - if sse := decoder.decode(b""): - yield sse + if sse := decoder.decode(b""): + if decoder.last_event_id is not None: + last_event_id = decoder.last_event_id + if sse.event or sse.data is not None: + # See async stream implementation for rationale on + # skipping empty flush events. + yield sse + if retry: + reconnect_attempts += 1 + if reconnect_attempts > max_reconnect_attempts: + raise httpx.TransportError( + "Exceeded maximum SSE reconnection attempts" + ) + continue + break def _encode_json(json: Any) -> tuple[dict[str, str], bytes]: diff --git a/libs/sdk-py/langgraph_sdk/sse.py b/libs/sdk-py/langgraph_sdk/sse.py index 42484df45..01bb333ac 100644 --- a/libs/sdk-py/langgraph_sdk/sse.py +++ b/libs/sdk-py/langgraph_sdk/sse.py @@ -81,6 +81,12 @@ class SSEDecoder: self._last_event_id = "" self._retry: int | None = None + @property + def last_event_id(self) -> str | None: + """Return the last event identifier that was seen.""" + + return self._last_event_id or None + def decode(self, line: bytes) -> StreamPart | None: # See: https://html.spec.whatwg.org/multipage/server-sent-events.html#event-stream-interpretation # noqa: E501 diff --git a/libs/sdk-py/tests/test_client_stream.py b/libs/sdk-py/tests/test_client_stream.py index 228ecf286..4237f89af 100644 --- a/libs/sdk-py/tests/test_client_stream.py +++ b/libs/sdk-py/tests/test_client_stream.py @@ -1,4 +1,6 @@ -from collections.abc import Iterator +from __future__ import annotations + +from collections.abc import Iterator, Sequence from pathlib import Path import httpx @@ -12,6 +14,35 @@ 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: @@ -61,6 +92,142 @@ async def test_http_client_stream_flushes_trailing_event(): 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}), + StreamPart(event="values", data={"step": 2}), + StreamPart(event="end", data=None), + ] + + +@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}), + StreamPart(event="values", data={"step": 2}), + StreamPart(event="end", data=None), + ] + + def test_sync_http_client_stream_flushes_trailing_event(): payload = b'event: foo\ndata: {"bar": 1}\n'