diff --git a/langgraph/pregel/__init__.py b/langgraph/pregel/__init__.py index 00418f866..bd4fea0c1 100644 --- a/langgraph/pregel/__init__.py +++ b/langgraph/pregel/__init__.py @@ -661,7 +661,7 @@ class Pregel( self, config: Optional[RunnableConfig] = None, *, - stream_mode: Optional[StreamMode] = None, + stream_mode: Optional[Union[StreamMode, list[StreamMode]]] = None, input_keys: Optional[Union[str, Sequence[str]]] = None, output_keys: Optional[Union[str, Sequence[str]]] = None, interrupt_before: Optional[Union[All, Sequence[str]]] = None, @@ -669,7 +669,7 @@ class Pregel( debug: Optional[bool] = None, ) -> tuple[ bool, - StreamMode, + Sequence[StreamMode], Union[str, Sequence[str]], Union[str, Sequence[str]], Optional[Sequence[str]], @@ -687,9 +687,11 @@ class Pregel( interrupt_before = interrupt_before or self.interrupt_before_nodes interrupt_after = interrupt_after or self.interrupt_after_nodes stream_mode = stream_mode if stream_mode is not None else self.stream_mode + if not isinstance(stream_mode, list): + stream_mode = [stream_mode] if config is not None and config.get("configurable", {}).get(CONFIG_KEY_READ): # if being called as a node in another graph, always use values mode - stream_mode = "values" + stream_mode = ["values"] return ( debug, stream_mode, @@ -704,7 +706,7 @@ class Pregel( input: Union[dict[str, Any], Any], config: Optional[RunnableConfig] = None, *, - stream_mode: Optional[StreamMode] = None, + stream_mode: Optional[Union[StreamMode, list[StreamMode]]] = None, output_keys: Optional[Union[str, Sequence[str]]] = None, input_keys: Optional[Union[str, Sequence[str]]] = None, interrupt_before: Optional[Union[All, Sequence[str]]] = None, @@ -731,7 +733,7 @@ class Pregel( # assign defaults ( debug, - stream_mode, + stream_modes, input_keys, output_keys, interrupt_before, @@ -850,9 +852,12 @@ class Pregel( if debug: print_step_tasks(step, next_tasks) - if stream_mode == "debug": - for chunk in map_debug_tasks(step, next_tasks): - yield chunk + if "debug" in stream_modes: + yield from _with_mode( + "debug", + isinstance(stream_mode, list), + map_debug_tasks(step, next_tasks), + ) futures = [ executor.submit(run_with_retry, task, self.retry_policy) @@ -887,16 +892,26 @@ class Pregel( print_step_checkpoint(step, channels, self.stream_channels_list) # yield current value or updates - if stream_mode == "values": - yield from map_output_values( - output_keys, pending_writes, channels + if "updates" in stream_modes: + yield from _with_mode( + "updates", + isinstance(stream_mode, list), + map_output_updates(output_keys, next_tasks), ) - elif stream_mode == "debug": - yield from map_debug_task_results( - step, next_tasks, self.stream_channels_list + if "debug" in stream_modes: + yield from _with_mode( + "debug", + isinstance(stream_mode, list), + map_debug_task_results( + step, next_tasks, self.stream_channels_list + ), + ) + if "values" in stream_modes: + yield from _with_mode( + "values", + isinstance(stream_mode, list), + map_output_values(output_keys, pending_writes, channels), ) - else: - yield from map_output_updates(output_keys, next_tasks) # save end of step checkpoint if self.checkpointer is not None: @@ -929,12 +944,16 @@ class Pregel( }, } # yield debug checkpoint - if stream_mode == "debug": - yield map_debug_checkpoint( - step, - checkpoint_config if self.checkpointer else None, - channels, - self.stream_channels_asis, + if "debug" in stream_modes: + yield from _with_mode( + "debug", + isinstance(stream_mode, list), + map_debug_checkpoint( + step, + checkpoint_config if self.checkpointer else None, + channels, + self.stream_channels_asis, + ), ) # after execution, check if we should interrupt @@ -976,7 +995,7 @@ class Pregel( input: Union[dict[str, Any], Any], config: Optional[RunnableConfig] = None, *, - stream_mode: Optional[StreamMode] = None, + stream_mode: Optional[Union[StreamMode, list[StreamMode]]] = None, output_keys: Optional[Union[str, Sequence[str]]] = None, input_keys: Optional[Union[str, Sequence[str]]] = None, interrupt_before: Optional[Union[All, Sequence[str]]] = None, @@ -1011,7 +1030,7 @@ class Pregel( # assign defaults ( debug, - stream_mode, + stream_modes, input_keys, output_keys, interrupt_before, @@ -1133,8 +1152,12 @@ class Pregel( if debug: print_step_tasks(step, next_tasks) - if stream_mode == "debug": - for chunk in map_debug_tasks(step, next_tasks): + if "debug" in stream_modes: + for chunk in _with_mode( + "debug", + isinstance(stream_mode, list), + map_debug_tasks(step, next_tasks), + ): yield chunk futures = [ @@ -1172,18 +1195,28 @@ class Pregel( print_step_checkpoint(step, channels, self.stream_channels_list) # yield current value or updates - if stream_mode == "values": - for chunk in map_output_values( - output_keys, pending_writes, channels + if "updates" in stream_modes: + for chunk in _with_mode( + "updates", + isinstance(stream_mode, list), + map_output_updates(output_keys, next_tasks), ): yield chunk - elif stream_mode == "debug": - for chunk in map_debug_task_results( - step, next_tasks, self.stream_channels_list + if "debug" in stream_modes: + for chunk in _with_mode( + "debug", + isinstance(stream_mode, list), + map_debug_task_results( + step, next_tasks, self.stream_channels_list + ), ): yield chunk - else: - for chunk in map_output_updates(output_keys, next_tasks): + if "values" in stream_modes: + for chunk in _with_mode( + "values", + isinstance(stream_mode, list), + map_output_values(output_keys, pending_writes, channels), + ): yield chunk # save end of step checkpoint @@ -1218,13 +1251,18 @@ class Pregel( }, } # yield debug checkpoint - if stream_mode == "debug": - yield map_debug_checkpoint( - step, - checkpoint_config if self.checkpointer else None, - channels, - self.stream_channels_asis, - ) + if "debug" in stream_modes: + for chunk in _with_mode( + "debug", + isinstance(stream_mode, list), + map_debug_checkpoint( + step, + checkpoint_config if self.checkpointer else None, + channels, + self.stream_channels_asis, + ), + ): + yield chunk # after execution, check if we should interrupt if _should_interrupt( @@ -1582,7 +1620,16 @@ def _prepare_next_tasks( node, writes, patch_config( - merge_configs(config, proc.config), + merge_configs( + config, + proc.config, + { + "metadata": { + "langgraph_step": step, + "langgraph_node": name, + } + }, + ), run_name=name, callbacks=manager.get_child(f"graph:step:{step}") if manager @@ -1595,9 +1642,17 @@ def _prepare_next_tasks( ), }, ), - triggers, + sorted(triggers), ) ) else: tasks.append(PregelTaskDescription(name, val)) return checkpoint, tasks + + +def _with_mode(mode: StreamMode, on: bool, iter: Iterator[Any]) -> Iterator[Any]: + if on: + for chunk in iter: + yield (mode, chunk) + else: + yield from iter diff --git a/langgraph/pregel/debug.py b/langgraph/pregel/debug.py index 07414fecc..0fd00bdc8 100644 --- a/langgraph/pregel/debug.py +++ b/langgraph/pregel/debug.py @@ -108,9 +108,9 @@ def map_debug_checkpoint( config: RunnableConfig, channels: Mapping[str, BaseChannel], stream_channels: Union[str, Sequence[str]], -) -> DebugOutputCheckpoint: +) -> Iterator[DebugOutputCheckpoint]: ts = datetime.now(timezone.utc).isoformat() - return { + yield { "type": "checkpoint", "timestamp": ts, "step": step, diff --git a/tests/test_pregel.py b/tests/test_pregel.py index e07429fad..7a38f012d 100644 --- a/tests/test_pregel.py +++ b/tests/test_pregel.py @@ -5051,6 +5051,219 @@ def test_in_one_fan_out_out_one_graph_state() -> None: }, ] + assert [ + *app.stream( + {"query": "what is weather in sf"}, + stream_mode=["values", "updates", "debug"], + ) + ] == [ + ("values", {"query": "what is weather in sf", "docs": []}), + ( + "debug", + { + "type": "checkpoint", + "timestamp": AnyStr(), + "step": 0, + "payload": { + "config": None, + "values": {"query": "what is weather in sf", "docs": []}, + }, + }, + ), + ( + "debug", + { + "type": "task", + "timestamp": AnyStr(), + "step": 1, + "payload": { + "id": "03dadab4-fb41-5308-a8a4-6eeb9ef7b9aa", + "name": "rewrite_query", + "input": { + "query": "what is weather in sf", + "answer": None, + "docs": [], + }, + "triggers": ["start:rewrite_query"], + }, + }, + ), + ("updates", {"rewrite_query": {"query": "query: what is weather in sf"}}), + ( + "debug", + { + "type": "task_result", + "timestamp": AnyStr(), + "step": 1, + "payload": { + "id": "03dadab4-fb41-5308-a8a4-6eeb9ef7b9aa", + "name": "rewrite_query", + "result": [("query", "query: what is weather in sf")], + }, + }, + ), + ("values", {"query": "query: what is weather in sf", "docs": []}), + ( + "debug", + { + "type": "checkpoint", + "timestamp": AnyStr(), + "step": 1, + "payload": { + "config": None, + "values": {"query": "query: what is weather in sf", "docs": []}, + }, + }, + ), + ( + "debug", + { + "type": "task", + "timestamp": AnyStr(), + "step": 2, + "payload": { + "id": "96f499e2-e203-5a13-9259-08cb62f4a2e5", + "name": "retriever_one", + "input": { + "query": "query: what is weather in sf", + "answer": None, + "docs": [], + }, + "triggers": ["rewrite_query"], + }, + }, + ), + ( + "debug", + { + "type": "task", + "timestamp": AnyStr(), + "step": 2, + "payload": { + "id": "6b344a90-a061-5f17-8714-51f0cf67cf01", + "name": "retriever_two", + "input": { + "query": "query: what is weather in sf", + "answer": None, + "docs": [], + }, + "triggers": ["rewrite_query"], + }, + }, + ), + ( + "updates", + { + "retriever_one": {"docs": ["doc1", "doc2"]}, + "retriever_two": {"docs": ["doc3", "doc4"]}, + }, + ), + ( + "debug", + { + "type": "task_result", + "timestamp": AnyStr(), + "step": 2, + "payload": { + "id": "96f499e2-e203-5a13-9259-08cb62f4a2e5", + "name": "retriever_one", + "result": [("docs", ["doc1", "doc2"])], + }, + }, + ), + ( + "debug", + { + "type": "task_result", + "timestamp": AnyStr(), + "step": 2, + "payload": { + "id": "6b344a90-a061-5f17-8714-51f0cf67cf01", + "name": "retriever_two", + "result": [("docs", ["doc3", "doc4"])], + }, + }, + ), + ( + "values", + { + "query": "query: what is weather in sf", + "docs": ["doc1", "doc2", "doc3", "doc4"], + }, + ), + ( + "debug", + { + "type": "checkpoint", + "timestamp": AnyStr(), + "step": 2, + "payload": { + "config": None, + "values": { + "query": "query: what is weather in sf", + "docs": ["doc1", "doc2", "doc3", "doc4"], + }, + }, + }, + ), + ( + "debug", + { + "type": "task", + "timestamp": AnyStr(), + "step": 3, + "payload": { + "id": "0dda6269-4ce3-5b98-9cea-d40737a68500", + "name": "qa", + "input": { + "query": "query: what is weather in sf", + "answer": None, + "docs": ["doc1", "doc2", "doc3", "doc4"], + }, + "triggers": ["retriever_one", "retriever_two"], + }, + }, + ), + ("updates", {"qa": {"answer": "doc1,doc2,doc3,doc4"}}), + ( + "debug", + { + "type": "task_result", + "timestamp": AnyStr(), + "step": 3, + "payload": { + "id": "0dda6269-4ce3-5b98-9cea-d40737a68500", + "name": "qa", + "result": [("answer", "doc1,doc2,doc3,doc4")], + }, + }, + ), + ( + "values", + { + "query": "query: what is weather in sf", + "answer": "doc1,doc2,doc3,doc4", + "docs": ["doc1", "doc2", "doc3", "doc4"], + }, + ), + ( + "debug", + { + "type": "checkpoint", + "timestamp": AnyStr(), + "step": 3, + "payload": { + "config": None, + "values": { + "query": "query: what is weather in sf", + "answer": "doc1,doc2,doc3,doc4", + "docs": ["doc1", "doc2", "doc3", "doc4"], + }, + }, + }, + ), + ] + def test_start_branch_then(snapshot: SnapshotAssertion) -> None: class State(TypedDict): diff --git a/tests/test_pregel_async.py b/tests/test_pregel_async.py index 1fbda625d..e410ea8d1 100644 --- a/tests/test_pregel_async.py +++ b/tests/test_pregel_async.py @@ -3503,6 +3503,220 @@ async def test_in_one_fan_out_out_one_graph_state() -> None: }, ] + assert [ + c + async for c in app.astream( + {"query": "what is weather in sf"}, + stream_mode=["values", "updates", "debug"], + ) + ] == [ + ("values", {"query": "what is weather in sf", "docs": []}), + ( + "debug", + { + "type": "checkpoint", + "timestamp": AnyStr(), + "step": 0, + "payload": { + "config": None, + "values": {"query": "what is weather in sf", "docs": []}, + }, + }, + ), + ( + "debug", + { + "type": "task", + "timestamp": AnyStr(), + "step": 1, + "payload": { + "id": "03dadab4-fb41-5308-a8a4-6eeb9ef7b9aa", + "name": "rewrite_query", + "input": { + "query": "what is weather in sf", + "answer": None, + "docs": [], + }, + "triggers": ["start:rewrite_query"], + }, + }, + ), + ("updates", {"rewrite_query": {"query": "query: what is weather in sf"}}), + ( + "debug", + { + "type": "task_result", + "timestamp": AnyStr(), + "step": 1, + "payload": { + "id": "03dadab4-fb41-5308-a8a4-6eeb9ef7b9aa", + "name": "rewrite_query", + "result": [("query", "query: what is weather in sf")], + }, + }, + ), + ("values", {"query": "query: what is weather in sf", "docs": []}), + ( + "debug", + { + "type": "checkpoint", + "timestamp": AnyStr(), + "step": 1, + "payload": { + "config": None, + "values": {"query": "query: what is weather in sf", "docs": []}, + }, + }, + ), + ( + "debug", + { + "type": "task", + "timestamp": AnyStr(), + "step": 2, + "payload": { + "id": "96f499e2-e203-5a13-9259-08cb62f4a2e5", + "name": "retriever_one", + "input": { + "query": "query: what is weather in sf", + "answer": None, + "docs": [], + }, + "triggers": ["rewrite_query"], + }, + }, + ), + ( + "debug", + { + "type": "task", + "timestamp": AnyStr(), + "step": 2, + "payload": { + "id": "6b344a90-a061-5f17-8714-51f0cf67cf01", + "name": "retriever_two", + "input": { + "query": "query: what is weather in sf", + "answer": None, + "docs": [], + }, + "triggers": ["rewrite_query"], + }, + }, + ), + ( + "updates", + { + "retriever_one": {"docs": ["doc1", "doc2"]}, + "retriever_two": {"docs": ["doc3", "doc4"]}, + }, + ), + ( + "debug", + { + "type": "task_result", + "timestamp": AnyStr(), + "step": 2, + "payload": { + "id": "96f499e2-e203-5a13-9259-08cb62f4a2e5", + "name": "retriever_one", + "result": [("docs", ["doc1", "doc2"])], + }, + }, + ), + ( + "debug", + { + "type": "task_result", + "timestamp": AnyStr(), + "step": 2, + "payload": { + "id": "6b344a90-a061-5f17-8714-51f0cf67cf01", + "name": "retriever_two", + "result": [("docs", ["doc3", "doc4"])], + }, + }, + ), + ( + "values", + { + "query": "query: what is weather in sf", + "docs": ["doc1", "doc2", "doc3", "doc4"], + }, + ), + ( + "debug", + { + "type": "checkpoint", + "timestamp": AnyStr(), + "step": 2, + "payload": { + "config": None, + "values": { + "query": "query: what is weather in sf", + "docs": ["doc1", "doc2", "doc3", "doc4"], + }, + }, + }, + ), + ( + "debug", + { + "type": "task", + "timestamp": AnyStr(), + "step": 3, + "payload": { + "id": "0dda6269-4ce3-5b98-9cea-d40737a68500", + "name": "qa", + "input": { + "query": "query: what is weather in sf", + "answer": None, + "docs": ["doc1", "doc2", "doc3", "doc4"], + }, + "triggers": ["retriever_one", "retriever_two"], + }, + }, + ), + ("updates", {"qa": {"answer": "doc1,doc2,doc3,doc4"}}), + ( + "debug", + { + "type": "task_result", + "timestamp": AnyStr(), + "step": 3, + "payload": { + "id": "0dda6269-4ce3-5b98-9cea-d40737a68500", + "name": "qa", + "result": [("answer", "doc1,doc2,doc3,doc4")], + }, + }, + ), + ( + "values", + { + "query": "query: what is weather in sf", + "answer": "doc1,doc2,doc3,doc4", + "docs": ["doc1", "doc2", "doc3", "doc4"], + }, + ), + ( + "debug", + { + "type": "checkpoint", + "timestamp": AnyStr(), + "step": 3, + "payload": { + "config": None, + "values": { + "query": "query: what is weather in sf", + "answer": "doc1,doc2,doc3,doc4", + "docs": ["doc1", "doc2", "doc3", "doc4"], + }, + }, + }, + ), + ] + async def test_start_branch_then() -> None: class State(TypedDict):