From 2ce2021c39bf5530c5d329754a3a5cbb329c2fc8 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Mon, 2 Dec 2024 11:29:06 -0800 Subject: [PATCH 1/3] Revert "Revert "sdk-py: Fix SSE parsing to split lines only \n \r \r\n per SSE spec"" This reverts commit 53ec7c41b2bd4261ba3790f417724a420132d863. --- libs/sdk-py/langgraph_sdk/client.py | 5 +- libs/sdk-py/langgraph_sdk/sse.py | 106 ++++++++++++++++++++++++++++ 2 files changed, 109 insertions(+), 2 deletions(-) create mode 100644 libs/sdk-py/langgraph_sdk/sse.py diff --git a/libs/sdk-py/langgraph_sdk/client.py b/libs/sdk-py/langgraph_sdk/client.py index 81e8a8506..ccf652b66 100644 --- a/libs/sdk-py/langgraph_sdk/client.py +++ b/libs/sdk-py/langgraph_sdk/client.py @@ -60,6 +60,7 @@ from langgraph_sdk.schema import ( ThreadStatus, ThreadUpdateStateResponse, ) +from langgraph_sdk.sse import EventSource logger = logging.getLogger(__name__) @@ -293,7 +294,7 @@ class HttpClient: else: logger.error(f"Error from langgraph-api: {body}", exc_info=e) raise e - async for event in sse.aiter_sse(): + async for event in EventSource(sse.response).aiter_sse(): yield StreamPart( event.event, orjson.loads(event.data) if event.data else None ) @@ -2432,7 +2433,7 @@ class SyncHttpClient: else: logger.error(f"Error from langgraph-api: {body}", exc_info=e) raise e - for event in sse.iter_sse(): + for event in EventSource(sse.response).iter_sse(): yield StreamPart( event.event, orjson.loads(event.data) if event.data else None ) diff --git a/libs/sdk-py/langgraph_sdk/sse.py b/libs/sdk-py/langgraph_sdk/sse.py new file mode 100644 index 000000000..8b019e4a6 --- /dev/null +++ b/libs/sdk-py/langgraph_sdk/sse.py @@ -0,0 +1,106 @@ +"""Adapted from httpx_sse to split lines on \n, \r, \r\n per the SSE spec.""" + +import io +from typing import AsyncIterator, Iterator + +import httpx +import httpx_sse +import httpx_sse._decoders + + +class BytesLineDecoder: + """ + Handles incrementally reading lines from text. + + Has the same behaviour as the stdllib bytes splitlines, + but handling the input iteratively. + """ + + def __init__(self) -> None: + self.buffer = io.BytesIO() + self.trailing_cr: bool = False + + def decode(self, text: bytes) -> list[bytes]: + # See https://docs.python.org/3/glossary.html#term-universal-newlines + NEWLINE_CHARS = b"\n\r" + + # We always push a trailing `\r` into the next decode iteration. + if self.trailing_cr: + text = b"\r" + text + self.trailing_cr = False + if text.endswith(b"\r"): + self.trailing_cr = True + text = text[:-1] + + if not text: + # NOTE: the edge case input of empty text doesn't occur in practice, + # because other httpx internals filter out this value + return [] # pragma: no cover + + trailing_newline = text[-1] in NEWLINE_CHARS + lines = text.splitlines() + + if len(lines) == 1 and not trailing_newline: + # No new lines, buffer the input and continue. + self.buffer.write(lines[0]) + return [] + + if self.buffer: + # Include any existing buffer in the first portion of the + # splitlines result. + lines = [self.buffer.getvalue() + lines[0]] + lines[1:] + self.buffer.truncate(0) + + if not trailing_newline: + # If the last segment of splitlines is not newline terminated, + # then drop it from our output and start a new buffer. + self.buffer.write(lines.pop()) + + return lines + + def flush(self) -> list[bytes]: + if not self.buffer and not self.trailing_cr: + return [] + + lines = [self.buffer.getvalue()] if self.buffer else [] + self.buffer.truncate(0) + self.trailing_cr = False + return lines + + +async def aiter_lines_raw(response: httpx.Response) -> AsyncIterator[bytes]: + decoder = BytesLineDecoder() + async for chunk in response.aiter_bytes(): + for line in decoder.decode(chunk): + yield line + for line in decoder.flush(): + yield line + + +def iter_lines_raw(response: httpx.Response) -> Iterator[bytes]: + decoder = BytesLineDecoder() + for chunk in response.iter_bytes(): + for line in decoder.decode(chunk): + yield line + for line in decoder.flush(): + yield line + + +class EventSource(httpx_sse.EventSource): + async def aiter_sse(self) -> AsyncIterator[httpx_sse.ServerSentEvent]: + self._check_content_type() + decoder = httpx_sse._decoders.SSEDecoder() + async for line in aiter_lines_raw(self._response): + line = line.rstrip(b"\n") + sse = decoder.decode(line.decode()) + if sse is not None: + yield sse + + def iter_sse(self) -> Iterator[httpx_sse.ServerSentEvent]: + self._check_content_type() + decoder = httpx_sse._decoders.SSEDecoder() + for line in iter_lines_raw(self._response): + line = line.rstrip(b"\n") + sse = decoder.decode(line.decode()) + if sse is not None: + yield sse From 3cee1d5087c92160f0b8acf9929b03ab64e2bb9a Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Mon, 2 Dec 2024 12:03:31 -0800 Subject: [PATCH 2/3] Remove httpx_sse, fix missing flush of sse decoder --- libs/sdk-py/langgraph_sdk/client.py | 60 ++++++++++----- libs/sdk-py/langgraph_sdk/sse.py | 112 +++++++++++++++++++--------- libs/sdk-py/poetry.lock | 13 +--- libs/sdk-py/pyproject.toml | 1 - 4 files changed, 120 insertions(+), 66 deletions(-) diff --git a/libs/sdk-py/langgraph_sdk/client.py b/libs/sdk-py/langgraph_sdk/client.py index ccf652b66..3fbf61041 100644 --- a/libs/sdk-py/langgraph_sdk/client.py +++ b/libs/sdk-py/langgraph_sdk/client.py @@ -26,7 +26,6 @@ from typing import ( ) import httpx -import httpx_sse import orjson from httpx._types import QueryParamTypes @@ -60,7 +59,7 @@ from langgraph_sdk.schema import ( ThreadStatus, ThreadUpdateStateResponse, ) -from langgraph_sdk.sse import EventSource +from langgraph_sdk.sse import SSEDecoder, aiter_lines_raw, iter_lines_raw logger = logging.getLogger(__name__) @@ -282,22 +281,37 @@ class HttpClient: ) -> AsyncIterator[StreamPart]: """Stream results using SSE.""" headers, content = await aencode_json(json) - async with httpx_sse.aconnect_sse( - self.client, method, path, headers=headers, content=content - ) as sse: + headers["Accept"] = "text/event-stream" + headers["Cache-Control"] = "no-store" + + async with self.client.stream( + method, path, headers=headers, content=content + ) as res: + # check status try: - sse.response.raise_for_status() + res.raise_for_status() except httpx.HTTPStatusError as e: - body = (await sse.response.aread()).decode() + 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 - async for event in EventSource(sse.response).aiter_sse(): - yield StreamPart( - event.event, orjson.loads(event.data) if event.data else None + # check content type + content_type = self._response.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 async def aencode_json(json: Any) -> tuple[dict[str, str], bytes]: @@ -2421,22 +2435,32 @@ class SyncHttpClient: ) -> Iterator[StreamPart]: """Stream the results of a request using SSE.""" headers, content = encode_json(json) - with httpx_sse.connect_sse( - self.client, method, path, headers=headers, content=content - ) as sse: + with self.client.stream(method, path, headers=headers, content=content) as res: + # check status try: - sse.response.raise_for_status() + res.raise_for_status() except httpx.HTTPStatusError as e: - body = sse.response.read().decode() + 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 - for event in EventSource(sse.response).iter_sse(): - yield StreamPart( - event.event, orjson.loads(event.data) if event.data else None + # check content type + content_type = self._response.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 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 8b019e4a6..6460b363c 100644 --- a/libs/sdk-py/langgraph_sdk/sse.py +++ b/libs/sdk-py/langgraph_sdk/sse.py @@ -1,11 +1,13 @@ """Adapted from httpx_sse to split lines on \n, \r, \r\n per the SSE spec.""" -import io -from typing import AsyncIterator, Iterator +from typing import AsyncIterator, Iterator, Optional, Union import httpx -import httpx_sse -import httpx_sse._decoders +import orjson + +from langgraph_sdk.schema import StreamPart + +BytesLike = Union[bytes, bytearray, memoryview] class BytesLineDecoder: @@ -17,10 +19,10 @@ class BytesLineDecoder: """ def __init__(self) -> None: - self.buffer = io.BytesIO() + self.buffer = bytearray() self.trailing_cr: bool = False - def decode(self, text: bytes) -> list[bytes]: + def decode(self, text: bytes) -> list[BytesLike]: # See https://docs.python.org/3/glossary.html#term-universal-newlines NEWLINE_CHARS = b"\n\r" @@ -42,33 +44,93 @@ class BytesLineDecoder: if len(lines) == 1 and not trailing_newline: # No new lines, buffer the input and continue. - self.buffer.write(lines[0]) + self.buffer.extend(lines[0]) return [] if self.buffer: # Include any existing buffer in the first portion of the # splitlines result. - lines = [self.buffer.getvalue() + lines[0]] + lines[1:] - self.buffer.truncate(0) + self.buffer.extend(lines[0]) + lines = [self.buffer] + lines[1:] + self.buffer = bytearray() if not trailing_newline: # If the last segment of splitlines is not newline terminated, # then drop it from our output and start a new buffer. - self.buffer.write(lines.pop()) + self.buffer.extend(lines.pop()) return lines - def flush(self) -> list[bytes]: + def flush(self) -> list[BytesLike]: if not self.buffer and not self.trailing_cr: return [] - lines = [self.buffer.getvalue()] if self.buffer else [] - self.buffer.truncate(0) + lines = [self.buffer] + self.buffer = bytearray() self.trailing_cr = False return lines -async def aiter_lines_raw(response: httpx.Response) -> AsyncIterator[bytes]: +class SSEDecoder: + def __init__(self) -> None: + self._event = "" + self._data = bytearray() + self._last_event_id = "" + self._retry: Optional[int] = None + + def decode(self, line: bytes) -> Optional[StreamPart]: + # See: https://html.spec.whatwg.org/multipage/server-sent-events.html#event-stream-interpretation # noqa: E501 + + if not line: + if ( + not self._event + and not self._data + and not self._last_event_id + and self._retry is None + ): + return None + + sse = StreamPart( + event=self._event, + data=orjson.loads(self._data) if self._data else None, + ) + + # NOTE: as per the SSE spec, do not reset last_event_id. + self._event = "" + self._data = bytearray() + self._retry = None + + return sse + + if line.startswith(b":"): + return None + + fieldname, _, value = line.partition(b":") + + if value.startswith(b" "): + value = value[1:] + + if fieldname == b"event": + self._event = value.decode() + elif fieldname == b"data": + self._data.extend(value) + elif fieldname == b"id": + if b"\0" in value: + pass + else: + self._last_event_id = value.decode() + elif fieldname == b"retry": + try: + self._retry = int(value) + except (TypeError, ValueError): + pass + else: + pass # Field is ignored. + + return None + + +async def aiter_lines_raw(response: httpx.Response) -> AsyncIterator[BytesLike]: decoder = BytesLineDecoder() async for chunk in response.aiter_bytes(): for line in decoder.decode(chunk): @@ -77,30 +139,10 @@ async def aiter_lines_raw(response: httpx.Response) -> AsyncIterator[bytes]: yield line -def iter_lines_raw(response: httpx.Response) -> Iterator[bytes]: +def iter_lines_raw(response: httpx.Response) -> Iterator[BytesLike]: decoder = BytesLineDecoder() for chunk in response.iter_bytes(): for line in decoder.decode(chunk): yield line for line in decoder.flush(): yield line - - -class EventSource(httpx_sse.EventSource): - async def aiter_sse(self) -> AsyncIterator[httpx_sse.ServerSentEvent]: - self._check_content_type() - decoder = httpx_sse._decoders.SSEDecoder() - async for line in aiter_lines_raw(self._response): - line = line.rstrip(b"\n") - sse = decoder.decode(line.decode()) - if sse is not None: - yield sse - - def iter_sse(self) -> Iterator[httpx_sse.ServerSentEvent]: - self._check_content_type() - decoder = httpx_sse._decoders.SSEDecoder() - for line in iter_lines_raw(self._response): - line = line.rstrip(b"\n") - sse = decoder.decode(line.decode()) - if sse is not None: - yield sse diff --git a/libs/sdk-py/poetry.lock b/libs/sdk-py/poetry.lock index 5bff56b98..1024f9799 100644 --- a/libs/sdk-py/poetry.lock +++ b/libs/sdk-py/poetry.lock @@ -141,17 +141,6 @@ cli = ["click (==8.*)", "pygments (==2.*)", "rich (>=10,<14)"] http2 = ["h2 (>=3,<5)"] socks = ["socksio (==1.*)"] -[[package]] -name = "httpx-sse" -version = "0.4.0" -description = "Consume Server-Sent Event (SSE) messages with HTTPX." -optional = false -python-versions = ">=3.8" -files = [ - {file = "httpx-sse-0.4.0.tar.gz", hash = "sha256:1e81a3a3070ce322add1d3529ed42eb5f70817f45ed6ec915ab753f961139721"}, - {file = "httpx_sse-0.4.0-py3-none-any.whl", hash = "sha256:f329af6eae57eaa2bdfd962b42524764af68075ea87370a2de920af5341e318f"}, -] - [[package]] name = "idna" version = "3.7" @@ -490,4 +479,4 @@ watchmedo = ["PyYAML (>=3.10)"] [metadata] lock-version = "2.0" python-versions = "^3.9.0,<4.0" -content-hash = "832acea0ad21ce71ae74edef225a1ad6f8fb166f6bf1531d876fe80fac7495f0" +content-hash = "1262a6148df18cc44ade00466b6e0f8305897a460eea370c8de649d8d20cd7a2" diff --git a/libs/sdk-py/pyproject.toml b/libs/sdk-py/pyproject.toml index c7776a6cc..678925c5b 100644 --- a/libs/sdk-py/pyproject.toml +++ b/libs/sdk-py/pyproject.toml @@ -11,7 +11,6 @@ packages = [{ include = "langgraph_sdk" }] [tool.poetry.dependencies] python = "^3.9.0,<4.0" httpx = ">=0.25.2" -httpx-sse = ">=0.4.0" orjson = ">=3.10.1" [tool.poetry.group.dev.dependencies] From 3bf92d0b031e7f2a54cf51d28d51ebcb3f04cf2e Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Tue, 3 Dec 2024 11:04:28 -0800 Subject: [PATCH 3/3] Fix --- libs/sdk-py/langgraph_sdk/client.py | 8 ++------ 1 file changed, 2 insertions(+), 6 deletions(-) diff --git a/libs/sdk-py/langgraph_sdk/client.py b/libs/sdk-py/langgraph_sdk/client.py index 3fbf61041..d6c5b6f05 100644 --- a/libs/sdk-py/langgraph_sdk/client.py +++ b/libs/sdk-py/langgraph_sdk/client.py @@ -298,9 +298,7 @@ class HttpClient: logger.error(f"Error from langgraph-api: {body}", exc_info=e) raise e # check content type - content_type = self._response.headers.get("content-type", "").partition( - ";" - )[0] + 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', " @@ -2447,9 +2445,7 @@ class SyncHttpClient: logger.error(f"Error from langgraph-api: {body}", exc_info=e) raise e # check content type - content_type = self._response.headers.get("content-type", "").partition( - ";" - )[0] + 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', "