diff --git a/litellm/llms/langgraph/chat/sse_iterator.py b/litellm/llms/langgraph/chat/sse_iterator.py index 2eb17b4d4b4..8be27e436c7 100644 --- a/litellm/llms/langgraph/chat/sse_iterator.py +++ b/litellm/llms/langgraph/chat/sse_iterator.py @@ -6,33 +6,40 @@ Handles Server-Sent Events (SSE) streaming responses from LangGraph. import json import uuid -from typing import TYPE_CHECKING, Optional +from typing import Any, Optional import httpx from litellm._logging import verbose_logger from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices -if TYPE_CHECKING: - pass - class LangGraphSSEStreamIterator: """ Iterator for LangGraph SSE streaming responses. Supports both sync and async iteration. - LangGraph stream format with stream_mode="messages-tuple": - Each SSE event is a tuple: (event_type, data) - Common event types: "messages", "metadata" + Supports two LangGraph SSE formats: + + 1. Standard SSE (``stream_mode="messages"``): + ``event: messages\\ndata: [{message_obj}, {metadata}]`` + + 2. Legacy tuple format (``stream_mode="messages-tuple"``): + ``data: ["messages", payload]`` """ + # Event types we act on + _MESSAGE_EVENTS = frozenset({"messages", "messages/partial", "messages/complete"}) + _METADATA_EVENTS = frozenset({"metadata"}) + def __init__(self, response: httpx.Response, model: str): self.response = response self.model = model self.finished = False self.line_iterator = None self.async_line_iterator = None + # Tracks the most recent ``event:`` header across SSE lines + self._current_event_type: Optional[str] = None def __iter__(self): """Initialize sync iteration.""" @@ -46,17 +53,30 @@ class LangGraphSSEStreamIterator: def _parse_sse_line(self, line: str) -> Optional[ModelResponseStream]: """ - Parse a single SSE line and return a ModelResponse chunk if applicable. + Parse a single SSE line. - LangGraph SSE format can vary: - - data: [...] (tuple format) - - event: ...\ndata: ... + Per the SSE specification an event block looks like:: + + event: \\n + data: \\n + \\n + + The ``event:`` line is optional; when absent the event type defaults + to ``"message"``. We track event type across calls so that when + ``data:`` arrives we already know the type. """ line = line.strip() if not line: + # Blank line = end of SSE event block; reset tracked type + self._current_event_type = None return None - # Handle SSE data lines + # Capture ``event:`` header for the next ``data:`` line + if line.startswith("event:"): + self._current_event_type = line[6:].strip() + return None + + # Handle ``data:`` lines if line.startswith("data:"): json_str = line[5:].strip() if not json_str: @@ -64,70 +84,108 @@ class LangGraphSSEStreamIterator: try: data = json.loads(json_str) - return self._process_data(data) except json.JSONDecodeError: verbose_logger.debug(f"Skipping non-JSON SSE line: {line[:100]}") return None + result = self._process_data(data, self._current_event_type) + # Reset after consuming so repeated calls don't re-use a stale type + self._current_event_type = None + return result + return None - def _process_data(self, data) -> Optional[ModelResponseStream]: - """ - Process parsed data from SSE stream. + # Data processing - LangGraph uses tuple format: [event_type, payload] + def _process_data( + self, data: Any, event_type: Optional[str] = None + ) -> Optional[ModelResponseStream]: """ - # Handle tuple format: ["messages", ...] - if isinstance(data, list) and len(data) >= 2: - event_type = data[0] + Route parsed JSON *data* using *event_type* (from ``event:`` header). + + Backward-compatible: when no ``event:`` header was present and *data* + is a list whose first element is a string, we fall back to the legacy + tuple format ``["messages", payload]``. + """ + + # Legacy tuple format: ["messages", payload] + if isinstance(data, list) and len(data) >= 2 and isinstance(data[0], str): + legacy_event = data[0] payload = data[1] - - if event_type == "messages": + if legacy_event == "messages": return self._process_messages_event(payload) - elif event_type == "metadata": - # Metadata event, might contain usage info + if legacy_event == "metadata": return self._process_metadata_event(payload) + return None - # Handle dict format (alternative response format) - elif isinstance(data, dict): + # Standard SSE: event type comes from the header + if event_type is not None: + if event_type in self._MESSAGE_EVENTS: + return self._process_messages_event(data) + if event_type in self._METADATA_EVENTS: + return self._process_metadata_event(data) + # Unknown event type – fall through to heuristic handling + verbose_logger.debug(f"Ignoring unknown LangGraph event type: {event_type}") + + # Heuristic fallback for headerless dict payloads + if isinstance(data, dict): if "content" in data: - return self._create_content_chunk(data.get("content", "")) - elif "messages" in data: - messages = data.get("messages", []) - if messages: - last_msg = messages[-1] - if isinstance(last_msg, dict) and last_msg.get("type") == "ai": - return self._create_content_chunk(last_msg.get("content", "")) + return self._create_content_chunk(data["content"]) + messages = data.get("messages") + if messages and isinstance(messages, list): + last_msg = messages[-1] + if isinstance(last_msg, dict) and last_msg.get("type") == "ai": + return self._create_content_chunk(last_msg.get("content", "")) + + # Headerless list of message objects (no tuple string key) + if isinstance(data, list) and event_type is None: + # Attempt to treat as a messages payload directly + return self._process_messages_event(data) return None - def _process_messages_event(self, payload) -> Optional[ModelResponseStream]: + def _process_messages_event(self, payload: Any) -> Optional[ModelResponseStream]: """ Process a messages event from the stream. - payload format: [[message_object, metadata], ...] + Handles multiple payload shapes emitted by LangGraph: + + * Flat list of message objects: + ``[{message_obj}, {metadata}]`` + * Nested list (legacy ``messages-tuple`` second element): + ``[[message_obj, metadata], ...]`` + * Single message dict """ + if isinstance(payload, dict): + return self._extract_ai_content(payload) + if isinstance(payload, list): for item in payload: + # Nested list: [[msg, meta], ...] if isinstance(item, list) and len(item) >= 1: - msg = item[0] - if isinstance(msg, dict): - msg_type = msg.get("type", "") - content = msg.get("content", "") - - # Only return AI messages with content - if msg_type == "ai" and content: - return self._create_content_chunk(content) - elif msg_type == "AIMessageChunk" and content: - return self._create_content_chunk(content) + result = self._extract_ai_content(item[0]) + if result is not None: + return result + # Flat list of dicts: [{msg}, {meta}] elif isinstance(item, dict): - msg_type = item.get("type", "") - content = item.get("content", "") - if msg_type in ("ai", "AIMessageChunk") and content: - return self._create_content_chunk(content) + result = self._extract_ai_content(item) + if result is not None: + return result return None + def _extract_ai_content(self, msg: Any) -> Optional[ModelResponseStream]: + """ + Return a content chunk if *msg* is an AI message dict with content. + """ + if not isinstance(msg, dict): + return None + msg_type = msg.get("type", "") + content = msg.get("content", "") + if msg_type in ("ai", "AIMessageChunk") and content: + return self._create_content_chunk(content) + return None + def _process_metadata_event(self, payload) -> Optional[ModelResponseStream]: """ Process a metadata event, which may signal the end of the stream. diff --git a/tests/test_litellm/test_langgraph_sse_parser.py b/tests/test_litellm/test_langgraph_sse_parser.py new file mode 100644 index 00000000000..94ed57ac202 --- /dev/null +++ b/tests/test_litellm/test_langgraph_sse_parser.py @@ -0,0 +1,246 @@ +""" +Tests for LangGraphSSEStreamIterator SSE parsing - Bug #24093. + +Validates that: +1. Standard SSE format (event: + data:) is parsed correctly. +2. Legacy tuple format (data: ["messages", ...]) still works. +3. Mixed / edge-case payloads are handled gracefully. +""" + +import json +from typing import List +from unittest.mock import MagicMock + +import pytest + +from litellm.llms.langgraph.chat.sse_iterator import LangGraphSSEStreamIterator +from litellm.types.utils import ModelResponseStream + +MODEL = "langgraph/my-agent" + + +# Helpers + +def _make_iterator(lines: List[str]) -> LangGraphSSEStreamIterator: + """Create an iterator backed by a mock httpx.Response whose iter_lines + yields the provided *lines* one-by-one.""" + response = MagicMock() + response.iter_lines.return_value = iter(lines) + it = LangGraphSSEStreamIterator(response=response, model=MODEL) + # Trigger __iter__ so line_iterator is populated + iter(it) + return it + + +def _collect_content(it: LangGraphSSEStreamIterator) -> List[str]: + """Exhaust the iterator and return a list of content strings from chunks.""" + contents: List[str] = [] + for chunk in it: + for choice in chunk.choices: + if choice.delta and choice.delta.content: + contents.append(choice.delta.content) + return contents + + +# Tests - Standard SSE format (event: header + data:) + +class TestStandardSSEFormat: + """Standard LangGraph SSE: ``event: messages\\ndata: [...]``.""" + + def test_should_parse_ai_message_chunk(self): + """event: messages with AIMessageChunk type extracts content.""" + lines = [ + "event: messages", + 'data: [{"content": "Hello world", "type": "AIMessageChunk"}, {}]', + "", + ] + contents = _collect_content(_make_iterator(lines)) + assert contents == ["Hello world"] + + def test_should_parse_ai_type_message(self): + """event: messages with 'ai' type extracts content.""" + lines = [ + "event: messages", + 'data: [{"content": "Hi there", "type": "ai"}, {"some": "meta"}]', + "", + ] + contents = _collect_content(_make_iterator(lines)) + assert contents == ["Hi there"] + + def test_should_skip_human_messages(self): + """Human messages should not produce content chunks.""" + lines = [ + "event: messages", + 'data: [{"content": "I am a human", "type": "human"}, {}]', + "", + ] + contents = _collect_content(_make_iterator(lines)) + assert contents == [] + + def test_should_handle_multiple_events(self): + """Multiple SSE event blocks should each produce content.""" + lines = [ + "event: messages", + 'data: [{"content": "First", "type": "AIMessageChunk"}, {}]', + "", + "event: messages", + 'data: [{"content": "Second", "type": "AIMessageChunk"}, {}]', + "", + ] + contents = _collect_content(_make_iterator(lines)) + assert contents == ["First", "Second"] + + def test_should_handle_metadata_event_with_run_id(self): + """Metadata event with run_id should produce a stop chunk.""" + lines = [ + "event: messages", + 'data: [{"content": "answer", "type": "ai"}, {}]', + "", + "event: metadata", + 'data: {"run_id": "abc-123"}', + "", + ] + it = _make_iterator(lines) + chunks = list(it) + # Should have a content chunk + a final stop chunk + assert len(chunks) == 2 + assert chunks[0].choices[0].delta.content == "answer" + assert chunks[1].choices[0].finish_reason == "stop" + + def test_should_handle_messages_partial_event(self): + """``event: messages/partial`` should be treated like messages.""" + lines = [ + "event: messages/partial", + 'data: [{"content": "partial", "type": "AIMessageChunk"}, {}]', + "", + ] + contents = _collect_content(_make_iterator(lines)) + assert contents == ["partial"] + + +# Tests - Legacy tuple format (data: ["messages", ...]) + +class TestLegacyTupleFormat: + """Legacy ``stream_mode="messages-tuple"``: ``data: ["messages", ...]``.""" + + def test_should_parse_nested_tuple_payload(self): + """Nested list payload: ["messages", [[msg, meta]]].""" + payload = [ + "messages", + [[{"content": "legacy hello", "type": "ai"}, {"meta": True}]], + ] + lines = [f"data: {json.dumps(payload)}"] + contents = _collect_content(_make_iterator(lines)) + assert contents == ["legacy hello"] + + def test_should_parse_flat_dict_tuple_payload(self): + """Flat dict payload: ["messages", [msg_dict, ...]].""" + payload = [ + "messages", + [{"content": "flat msg", "type": "AIMessageChunk"}], + ] + lines = [f"data: {json.dumps(payload)}"] + contents = _collect_content(_make_iterator(lines)) + assert contents == ["flat msg"] + + def test_should_handle_metadata_tuple(self): + """Metadata tuple: ["metadata", {run_id: ...}] should signal stop.""" + content_payload = ["messages", [{"content": "ok", "type": "ai"}]] + meta_payload = ["metadata", {"run_id": "xyz"}] + lines = [ + f"data: {json.dumps(content_payload)}", + f"data: {json.dumps(meta_payload)}", + ] + it = _make_iterator(lines) + chunks = list(it) + assert len(chunks) == 2 + assert chunks[0].choices[0].delta.content == "ok" + assert chunks[1].choices[0].finish_reason == "stop" + + +# Tests - Edge cases / backward-compatibility + +class TestEdgeCases: + """Mixed formats, empty data, and malformed lines.""" + + def test_should_skip_empty_lines(self): + """Blank lines and empty data should not crash.""" + lines = ["", "data: ", "", "data: not-json", ""] + it = _make_iterator(lines) + chunks = list(it) + # Only the auto-generated final stop chunk + assert len(chunks) == 1 + assert chunks[0].choices[0].finish_reason == "stop" + + def test_should_handle_dict_with_content_key(self): + """Dict payload with a ``content`` key should produce a chunk even + without an event header (heuristic fallback).""" + lines = ['data: {"content": "fallback", "type": "ai"}'] + # This hits the heuristic path since 'content' key is present + contents = _collect_content(_make_iterator(lines)) + assert contents == ["fallback"] + + def test_should_handle_dict_with_messages_key(self): + """Dict payload with a ``messages`` list (heuristic fallback).""" + lines = [ + 'data: {"messages": [{"content": "inner", "type": "ai"}]}', + ] + contents = _collect_content(_make_iterator(lines)) + assert contents == ["inner"] + + def test_should_reset_event_type_on_blank_line(self): + """After a blank line the event type must be reset so a subsequent + data line without an event header doesn't inherit the old type.""" + lines = [ + "event: messages", + "", # reset + 'data: [{"content": "no-header", "type": "AIMessageChunk"}, {}]', + ] + # Without the event header the parser should still handle the list + # via the headerless-list fallback + contents = _collect_content(_make_iterator(lines)) + assert contents == ["no-header"] + + def test_should_ignore_unknown_event_types(self): + """Unknown event types should be silently skipped.""" + lines = [ + "event: custom_event", + 'data: {"foo": "bar"}', + "", + "event: messages", + 'data: [{"content": "valid", "type": "ai"}, {}]', + "", + ] + contents = _collect_content(_make_iterator(lines)) + assert contents == ["valid"] + + def test_should_produce_final_stop_chunk(self): + """When the stream ends without an explicit metadata event, a final + stop chunk should still be emitted.""" + lines = [ + "event: messages", + 'data: [{"content": "hello", "type": "ai"}, {}]', + "", + ] + it = _make_iterator(lines) + chunks = list(it) + # content chunk + auto final stop + assert len(chunks) == 2 + assert chunks[-1].choices[0].finish_reason == "stop" + + def test_model_response_structure(self): + """Verify that emitted chunks have the correct ModelResponseStream + shape expected by the rest of LiteLLM.""" + lines = [ + "event: messages", + 'data: [{"content": "check", "type": "ai"}, {}]', + "", + ] + it = _make_iterator(lines) + chunk = next(it) + assert isinstance(chunk, ModelResponseStream) + assert chunk.model == MODEL + assert chunk.object == "chat.completion.chunk" + assert chunk.choices[0].delta.role == "assistant" + assert chunk.choices[0].delta.content == "check" + assert chunk.choices[0].finish_reason is None