This commit is contained in:
Nuno Campos
2025-05-27 16:07:54 -07:00
parent 936176eb21
commit bc404d14f6
5 changed files with 298 additions and 27 deletions
+34 -4
View File
@@ -34,6 +34,7 @@ from langgraph.checkpoint.base import (
CheckpointTuple,
copy_checkpoint,
)
from langgraph.config import get_config
from langgraph.constants import (
CACHE_NS_WRITES,
CONF,
@@ -2432,13 +2433,27 @@ class Pregel(PregelProtocol):
# set up messages stream mode
if "messages" in stream_modes:
run_manager.inheritable_handlers.append(
StreamMessagesHandler(stream.put)
StreamMessagesHandler(stream.put, subgraphs)
)
# set up custom stream mode
if "custom" in stream_modes:
config[CONF][CONFIG_KEY_STREAM_WRITER] = lambda c: stream.put(
((), "custom", c)
(
tuple(
get_config()[CONF][CONFIG_KEY_CHECKPOINT_NS].split(NS_SEP)[
:-1
]
),
"custom",
c,
)
)
elif (
CONFIG_KEY_STREAM not in config[CONF]
and CONFIG_KEY_STREAM_WRITER in config[CONF]
):
# remove parent graph stream writer if subgraph streaming not requested
del config[CONF][CONFIG_KEY_STREAM_WRITER]
# set checkpointing mode for subgraphs
if checkpoint_during is not None:
config[CONF][CONFIG_KEY_CHECKPOINT_DURING] = checkpoint_during
@@ -2658,15 +2673,30 @@ class Pregel(PregelProtocol):
# set up messages stream mode
if "messages" in stream_modes:
run_manager.inheritable_handlers.append(
StreamMessagesHandler(stream_put)
StreamMessagesHandler(stream_put, subgraphs)
)
# set up custom stream mode
if "custom" in stream_modes:
config[CONF][CONFIG_KEY_STREAM_WRITER] = (
lambda c: aioloop.call_soon_threadsafe(
stream.put_nowait, ((), "custom", c)
stream.put_nowait,
(
tuple(
get_config()[CONF][CONFIG_KEY_CHECKPOINT_NS].split(
NS_SEP
)[:-1]
),
"custom",
c,
),
)
)
elif (
CONFIG_KEY_STREAM not in config[CONF]
and CONFIG_KEY_STREAM_WRITER in config[CONF]
):
# remove parent graph stream writer if subgraph streaming not requested
del config[CONF][CONFIG_KEY_STREAM_WRITER]
# set checkpointing mode for subgraphs
if checkpoint_during is not None:
config[CONF][CONFIG_KEY_CHECKPOINT_DURING] = checkpoint_during
+23 -13
View File
@@ -11,7 +11,7 @@ from uuid import UUID, uuid4
from langchain_core.callbacks import BaseCallbackHandler
from langchain_core.messages import BaseMessage
from langchain_core.outputs import ChatGenerationChunk, LLMResult
from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, LLMResult
from langgraph.constants import NS_SEP, TAG_HIDDEN, TAG_NOSTREAM, TAG_NOSTREAM_ALT
from langgraph.types import Command, StreamChunk
@@ -32,8 +32,9 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler):
run_inline = True
"""We want this callback to run in the main thread, to avoid order/locking issues."""
def __init__(self, stream: Callable[[StreamChunk], None]):
def __init__(self, stream: Callable[[StreamChunk], None], subgraphs: bool):
self.stream = stream
self.subgraphs = subgraphs
self.metadata: dict[UUID, Meta] = {}
self.seen: set[Union[int, str]] = set()
@@ -96,10 +97,15 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler):
if metadata and (
not tags or (TAG_NOSTREAM not in tags and TAG_NOSTREAM_ALT not in tags)
):
self.metadata[run_id] = (
tuple(cast(str, metadata["langgraph_checkpoint_ns"]).split(NS_SEP)),
metadata,
)
ns = tuple(cast(str, metadata["langgraph_checkpoint_ns"]).split(NS_SEP))[
:-1
]
if not self.subgraphs and len(ns) > 0:
return
if tags:
if filtered_tags := [t for t in tags if not t.startswith("seq:step")]:
metadata["tags"] = filtered_tags
self.metadata[run_id] = (ns, metadata)
def on_llm_new_token(
self,
@@ -114,9 +120,6 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler):
if not isinstance(chunk, ChatGenerationChunk):
return
if meta := self.metadata.get(run_id):
filtered_tags = [t for t in (tags or []) if not t.startswith("seq:step")]
if filtered_tags:
meta[1]["tags"] = filtered_tags
self._emit(meta, chunk.message)
def on_llm_end(
@@ -127,6 +130,11 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler):
parent_run_id: Optional[UUID] = None,
**kwargs: Any,
) -> Any:
if meta := self.metadata.get(run_id):
if response.generations and response.generations[0]:
gen = response.generations[0][0]
if isinstance(gen, ChatGeneration):
self._emit(meta, gen.message, dedupe=True)
self.metadata.pop(run_id, None)
def on_llm_error(
@@ -155,10 +163,12 @@ class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler):
and kwargs.get("name") == metadata.get("langgraph_node")
and (not tags or TAG_HIDDEN not in tags)
):
self.metadata[run_id] = (
tuple(cast(str, metadata["langgraph_checkpoint_ns"]).split(NS_SEP)),
metadata,
)
ns = tuple(cast(str, metadata["langgraph_checkpoint_ns"]).split(NS_SEP))[
:-1
]
if not self.subgraphs and len(ns) > 0:
return
self.metadata[run_id] = (ns, metadata)
if isinstance(inputs, dict):
for key, value in inputs.items():
if isinstance(value, BaseMessage):
+117 -5
View File
@@ -9279,12 +9279,16 @@ def test_weather_subgraph(
inputs = {"messages": [{"role": "user", "content": "what's the weather in sf"}]}
# run with custom output
assert [c for c in graph.stream(inputs, thread2, stream_mode="custom")] == [
"I'm",
" very",
assert [
c for c in graph.stream(inputs, thread2, stream_mode="custom", subgraphs=True)
] == [
((), "I'm"),
((AnyStr("weather_graph:"),), " very"),
]
assert [c for c in graph.stream(None, thread2, stream_mode="custom")] == [
" good",
assert [
c for c in graph.stream(None, thread2, stream_mode="custom", subgraphs=True)
] == [
((AnyStr("weather_graph:"),), " good"),
]
# run until interrupt
@@ -9602,3 +9606,111 @@ def test_weather_subgraph(
},
),
]
# run with custom output, without subgraph streaming, should omit subgraph chunks
assert [
c
for c in graph.stream(
inputs, {"configurable": {"thread_id": "3"}}, stream_mode="custom"
)
] == [
"I'm",
]
# run with messages output, with subgraph streaming, should inc subgraph messages
assert [
c
for c in graph.stream(
inputs,
{"configurable": {"thread_id": "4"}},
stream_mode="messages",
subgraphs=True,
)
] == [
(
(),
(
_AnyIdAIMessage(
content="",
tool_calls=[
ToolCall(
id="tool_call123",
name="router",
args={"dest": "weather"},
)
],
),
{
"thread_id": "4",
"langgraph_step": 1,
"langgraph_node": "router_node",
"langgraph_triggers": ("branch:to:router_node",),
"langgraph_path": ("__pregel_pull", "router_node"),
"langgraph_checkpoint_ns": AnyStr("router_node:"),
"checkpoint_ns": AnyStr("router_node:"),
"ls_provider": "fakemessageslistchatmodel",
"ls_model_type": "chat",
},
),
),
(
(AnyStr("weather_graph:"),),
(
_AnyIdAIMessage(
content="",
tool_calls=[
ToolCall(
id="tool_call123",
name="get_weather",
args={"city": "San Francisco"},
)
],
),
{
"thread_id": "4",
"langgraph_step": 1,
"langgraph_node": "model_node",
"langgraph_triggers": ("branch:to:model_node",),
"langgraph_path": ("__pregel_pull", "model_node"),
"langgraph_checkpoint_ns": AnyStr("weather_graph:"),
"checkpoint_ns": AnyStr("weather_graph:"),
"ls_provider": "fakemessageslistchatmodel",
"ls_model_type": "chat",
},
),
),
]
# run with messages output, without subgraph streaming, should exc subgraph messages
assert [
c
for c in graph.stream(
inputs,
{"configurable": {"thread_id": "5"}},
stream_mode="messages",
)
] == [
(
_AnyIdAIMessage(
content="",
tool_calls=[
ToolCall(
id="tool_call123",
name="router",
args={"dest": "weather"},
)
],
),
{
"thread_id": "5",
"langgraph_step": 1,
"langgraph_node": "router_node",
"langgraph_triggers": ("branch:to:router_node",),
"langgraph_path": ("__pregel_pull", "router_node"),
"langgraph_checkpoint_ns": AnyStr("router_node:"),
"checkpoint_ns": AnyStr("router_node:"),
"ls_provider": "fakemessageslistchatmodel",
"ls_model_type": "chat",
},
),
]
+123 -5
View File
@@ -6452,12 +6452,22 @@ async def test_weather_subgraph(
inputs = {"messages": [{"role": "user", "content": "what's the weather in sf"}]}
# run with custom output
assert [c async for c in graph.astream(inputs, thread2, stream_mode="custom")] == [
"I'm",
" very",
assert [
c
async for c in graph.astream(
inputs, thread2, stream_mode="custom", subgraphs=True
)
] == [
((), "I'm"),
((AnyStr("weather_graph:"),), " very"),
]
assert [c async for c in graph.astream(None, thread2, stream_mode="custom")] == [
" good",
assert [
c
async for c in graph.astream(
None, thread2, stream_mode="custom", subgraphs=True
)
] == [
((AnyStr("weather_graph:"),), " good"),
]
# run until interrupt
@@ -6777,3 +6787,111 @@ async def test_weather_subgraph(
},
),
]
# run with custom output, without subgraph streaming, should omit subgraph chunks
assert [
c
async for c in graph.astream(
inputs, {"configurable": {"thread_id": "3"}}, stream_mode="custom"
)
] == [
"I'm",
]
# run with messages output, with subgraph streaming, should inc subgraph messages
assert [
c
async for c in graph.astream(
inputs,
{"configurable": {"thread_id": "4"}},
stream_mode="messages",
subgraphs=True,
)
] == [
(
(),
(
_AnyIdAIMessage(
content="",
tool_calls=[
ToolCall(
id="tool_call123",
name="router",
args={"dest": "weather"},
)
],
),
{
"thread_id": "4",
"langgraph_step": 1,
"langgraph_node": "router_node",
"langgraph_triggers": ("branch:to:router_node",),
"langgraph_path": ("__pregel_pull", "router_node"),
"langgraph_checkpoint_ns": AnyStr("router_node:"),
"checkpoint_ns": AnyStr("router_node:"),
"ls_provider": "fakemessageslistchatmodel",
"ls_model_type": "chat",
},
),
),
(
(AnyStr("weather_graph:"),),
(
_AnyIdAIMessage(
content="",
tool_calls=[
ToolCall(
id="tool_call123",
name="get_weather",
args={"city": "San Francisco"},
)
],
),
{
"thread_id": "4",
"langgraph_step": 1,
"langgraph_node": "model_node",
"langgraph_triggers": ("branch:to:model_node",),
"langgraph_path": ("__pregel_pull", "model_node"),
"langgraph_checkpoint_ns": AnyStr("weather_graph:"),
"checkpoint_ns": AnyStr("weather_graph:"),
"ls_provider": "fakemessageslistchatmodel",
"ls_model_type": "chat",
},
),
),
]
# run with messages output, without subgraph streaming, should exc subgraph messages
assert [
c
async for c in graph.astream(
inputs,
{"configurable": {"thread_id": "5"}},
stream_mode="messages",
)
] == [
(
_AnyIdAIMessage(
content="",
tool_calls=[
ToolCall(
id="tool_call123",
name="router",
args={"dest": "weather"},
)
],
),
{
"thread_id": "5",
"langgraph_step": 1,
"langgraph_node": "router_node",
"langgraph_triggers": ("branch:to:router_node",),
"langgraph_path": ("__pregel_pull", "router_node"),
"langgraph_checkpoint_ns": AnyStr("router_node:"),
"checkpoint_ns": AnyStr("router_node:"),
"ls_provider": "fakemessageslistchatmodel",
"ls_model_type": "chat",
},
),
]
@@ -5372,6 +5372,7 @@ async def test_stream_subgraphs_during_execution(
]
@NEEDS_CONTEXTVARS
async def test_stream_buffering_single_node(
async_checkpointer: BaseCheckpointSaver,
) -> None: