mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(langgraph): correct SSE event parsing for LangGraph streaming (#24093)
This commit is contained in:
parent
e5baa2232f
commit
5b1bbadcba
2 changed files with 353 additions and 49 deletions
|
|
@ -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.
|
||||
|
|
|
|||
246
tests/test_litellm/test_langgraph_sse_parser.py
Normal file
246
tests/test_litellm/test_langgraph_sse_parser.py
Normal 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
|
||||
Loading…
Add table
Reference in a new issue