fix(langgraph): correct SSE event parsing for LangGraph streaming (#24093)

This commit is contained in:
ashish-509 2026-03-19 23:20:00 +05:45
parent e5baa2232f
commit 5b1bbadcba
2 changed files with 353 additions and 49 deletions

View file

@ -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: <type>\\n
data: <payload>\\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.

View file

@ -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