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]