diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index be04d187d..6ac6c737a 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -340,7 +340,10 @@ class Pregel( @property def stream_channels_asis(self) -> Union[str, Sequence[str]]: return self.stream_channels or [ - k for k in self.channels if not isinstance(self.channels[k], Context) + k + for k in self.channels + if isinstance(self.channels[k], BaseChannel) + and not isinstance(self.channels[k], Context) ] def get_state(self, config: RunnableConfig) -> StateSnapshot: @@ -956,7 +959,6 @@ class Pregel( config=config, store=self.store, checkpointer=checkpointer, - graph=self, nodes=self.nodes, specs=self.channels, ) as loop: @@ -966,7 +968,9 @@ class Pregel( # channels are guaranteed to be immutable for the duration of the step, # with channel updates applied only at the transition between steps while loop.tick( + input_keys=self.input_channels, output_keys=output_keys, + stream_keys=self.stream_channels_asis, interrupt_before=interrupt_before, interrupt_after=interrupt_after, manager=run_manager, @@ -1212,7 +1216,6 @@ class Pregel( config=config, store=self.store, checkpointer=checkpointer, - graph=self, nodes=self.nodes, specs=self.channels, ) as loop: @@ -1223,7 +1226,9 @@ class Pregel( # channels are guaranteed to be immutable for the duration of the step, # with channel updates applied only at the transition between steps while loop.tick( + input_keys=self.input_channels, output_keys=output_keys, + stream_keys=self.stream_channels_asis, interrupt_before=interrupt_before, interrupt_after=interrupt_after, manager=run_manager, diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index bc19b39cf..db165513b 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -4,7 +4,6 @@ from collections import deque from contextlib import AsyncExitStack, ExitStack from types import TracebackType from typing import ( - TYPE_CHECKING, Any, AsyncContextManager, Callable, @@ -71,10 +70,6 @@ from langgraph.pregel.utils import get_new_channel_versions from langgraph.store.base import BaseStore from langgraph.store.batch import AsyncBatchedStore -if TYPE_CHECKING: - from langgraph.pregel import Pregel - - V = TypeVar("V") INPUT_DONE = object() INPUT_RESUMING = object() @@ -105,8 +100,6 @@ class PregelLoop: Any, ] ] - graph: "Pregel" - store: Optional[BaseStore] submit: Submit channels: Mapping[str, BaseChannel] managed: ManagedValueMapping @@ -136,14 +129,12 @@ class PregelLoop: checkpointer: Optional[BaseCheckpointSaver], nodes: Mapping[str, PregelNode], specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]], - graph: "Pregel", ) -> None: self.stream = deque() self.input = input self.config = config self.store = store self.checkpointer = checkpointer - self.graph = graph self.nodes = nodes self.specs = specs self.is_nested = CONFIG_KEY_READ in self.config.get("configurable", {}) @@ -175,7 +166,9 @@ class PregelLoop: def tick( self, *, + input_keys: Union[str, Sequence[str]], output_keys: Union[str, Sequence[str]] = EMPTY_SEQ, + stream_keys: Union[str, Sequence[str]] = EMPTY_SEQ, interrupt_after: Sequence[str] = EMPTY_SEQ, interrupt_before: Sequence[str] = EMPTY_SEQ, manager: Union[None, AsyncParentRunManager, ParentRunManager] = None, @@ -187,7 +180,7 @@ class PregelLoop: raise RuntimeError("Cannot tick when status is no longer 'pending'") if self.input not in (INPUT_DONE, INPUT_RESUMING): - self._first() + self._first(input_keys=input_keys) elif all(task.writes for task in self.tasks): writes = [w for t in self.tasks for w in t.writes] # all tasks have finished @@ -211,11 +204,7 @@ class PregelLoop: self._put_checkpoint( { "source": "loop", - "writes": single( - map_output_updates(output_keys, self.tasks) - if self.graph.stream_mode == "updates" - else map_output_values(output_keys, writes, self.channels) - ), + "writes": single(map_output_updates(output_keys, self.tasks)), } ) # after execution, check if we should interrupt @@ -258,7 +247,7 @@ class PregelLoop: self.step - 1, # printing checkpoint for previous step self.checkpoint_config, self.channels, - self.graph.stream_channels_asis, + stream_keys, self.checkpoint_metadata, self.checkpoint, self.tasks, @@ -282,6 +271,7 @@ class PregelLoop: # if all tasks have finished, re-tick if all(task.writes for task in self.tasks): return self.tick( + input_keys=input_keys, output_keys=output_keys, interrupt_after=interrupt_after, interrupt_before=interrupt_before, @@ -306,7 +296,7 @@ class PregelLoop: # private - def _first(self) -> None: + def _first(self, *, input_keys: Union[str, Sequence[str]]) -> None: # resuming from previous checkpoint requires # - finding a previous checkpoint # - receiving None input (outer graph) or RESUMING flag (subgraph) @@ -323,7 +313,7 @@ class PregelLoop: version = self.checkpoint["channel_versions"][k] self.checkpoint["versions_seen"][INTERRUPT][k] = version # map inputs to channel updates - elif input_writes := deque(map_input(self.graph.input_channels, self.input)): + elif input_writes := deque(map_input(input_keys, self.input)): # discard any unfinished tasks from previous checkpoint discard_tasks = prepare_next_tasks( self.checkpoint, @@ -345,7 +335,7 @@ class PregelLoop: # save input checkpoint self._put_checkpoint({"source": "input", "writes": self.input}) else: - raise EmptyInputError(f"Received no input for {self.graph.input_channels}") + raise EmptyInputError(f"Received no input for {input_keys}") # done with input self.input = INPUT_RESUMING if is_resuming else INPUT_DONE @@ -430,13 +420,11 @@ class SyncPregelLoop(PregelLoop, ContextManager): checkpointer: Optional[BaseCheckpointSaver], nodes: Mapping[str, PregelNode], specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]], - graph: "Pregel", ) -> None: super().__init__( input, config=config, checkpointer=checkpointer, - graph=graph, store=store, nodes=nodes, specs=specs, @@ -504,7 +492,6 @@ class SyncPregelLoop(PregelLoop, ContextManager): traceback: Optional[TracebackType], ) -> Optional[bool]: # unwind stack - del self.graph return self.stack.__exit__(exc_type, exc_value, traceback) @@ -518,13 +505,11 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager): checkpointer: Optional[BaseCheckpointSaver], nodes: Mapping[str, PregelNode], specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]], - graph: "Pregel", ) -> None: super().__init__( input, config=config, checkpointer=checkpointer, - graph=graph, store=store, nodes=nodes, specs=specs, @@ -598,7 +583,6 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager): traceback: Optional[TracebackType], ) -> Optional[bool]: # unwind stack - del self.graph return await asyncio.shield( self.stack.__aexit__(exc_type, exc_value, traceback) ) diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 070c59476..197dc6c61 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -657,7 +657,7 @@ def test_invoke_two_processes_in_out_interrupt( "checkpoint_id": AnyStr(), } }, - metadata={"source": "loop", "step": 6, "writes": 5}, + metadata={"source": "loop", "step": 6, "writes": {"two": 5}}, created_at=AnyStr(), parent_config=history[1].config, ), @@ -672,7 +672,7 @@ def test_invoke_two_processes_in_out_interrupt( "checkpoint_id": AnyStr(), } }, - metadata={"source": "loop", "step": 5, "writes": None}, + metadata={"source": "loop", "step": 5, "writes": {"one": None}}, created_at=AnyStr(), parent_config=history[2].config, ), @@ -702,7 +702,7 @@ def test_invoke_two_processes_in_out_interrupt( "checkpoint_id": AnyStr(), } }, - metadata={"source": "loop", "step": 3, "writes": None}, + metadata={"source": "loop", "step": 3, "writes": {"one": None}}, created_at=AnyStr(), parent_config=history[4].config, ), @@ -732,7 +732,7 @@ def test_invoke_two_processes_in_out_interrupt( "checkpoint_id": AnyStr(), } }, - metadata={"source": "loop", "step": 1, "writes": 4}, + metadata={"source": "loop", "step": 1, "writes": {"two": 4}}, created_at=AnyStr(), parent_config=history[6].config, ), @@ -747,7 +747,7 @@ def test_invoke_two_processes_in_out_interrupt( "checkpoint_id": AnyStr(), } }, - metadata={"source": "loop", "step": 0, "writes": None}, + metadata={"source": "loop", "step": 0, "writes": {"one": None}}, created_at=AnyStr(), parent_config=history[7].config, ), @@ -1995,12 +1995,14 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None: "step": 0, "writes": { "agent": { - "input": "what is weather in sf", - "agent_outcome": AgentAction( - tool="search_api", - tool_input="query", - log="tool:search_api:query", - ), + "agent": { + "input": "what is weather in sf", + "agent_outcome": AgentAction( + tool="search_api", + tool_input="query", + log="tool:search_api:query", + ), + } }, }, }, @@ -2206,12 +2208,14 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None: "step": 0, "writes": { "agent": { - "input": "what is weather in sf", - "agent_outcome": AgentAction( - tool="search_api", - tool_input="query", - log="tool:search_api:query", - ), + "agent": { + "input": "what is weather in sf", + "agent_outcome": AgentAction( + tool="search_api", + tool_input="query", + log="tool:search_api:query", + ), + } } }, }, @@ -2411,12 +2415,14 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None: "step": 0, "writes": { "agent": { - "input": "what is weather in sf", - "agent_outcome": AgentAction( - tool="search_api", - tool_input="query", - log="tool:search_api:query", - ), + "agent": { + "input": "what is weather in sf", + "agent_outcome": AgentAction( + tool="search_api", + tool_input="query", + log="tool:search_api:query", + ), + } } }, }, diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index ebbfc7b29..e9987c601 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -876,7 +876,7 @@ async def test_invoke_two_processes_in_out_interrupt( "checkpoint_id": AnyStr(), } }, - metadata={"source": "loop", "step": 6, "writes": 5}, + metadata={"source": "loop", "step": 6, "writes": {"two": 5}}, created_at=AnyStr(), parent_config=history[1].config, ), @@ -891,7 +891,7 @@ async def test_invoke_two_processes_in_out_interrupt( "checkpoint_id": AnyStr(), } }, - metadata={"source": "loop", "step": 5, "writes": None}, + metadata={"source": "loop", "step": 5, "writes": {"one": None}}, created_at=AnyStr(), parent_config=history[2].config, ), @@ -921,7 +921,7 @@ async def test_invoke_two_processes_in_out_interrupt( "checkpoint_id": AnyStr(), } }, - metadata={"source": "loop", "step": 3, "writes": None}, + metadata={"source": "loop", "step": 3, "writes": {"one": None}}, created_at=AnyStr(), parent_config=history[4].config, ), @@ -951,7 +951,7 @@ async def test_invoke_two_processes_in_out_interrupt( "checkpoint_id": AnyStr(), } }, - metadata={"source": "loop", "step": 1, "writes": 4}, + metadata={"source": "loop", "step": 1, "writes": {"two": 4}}, created_at=AnyStr(), parent_config=history[6].config, ), @@ -966,7 +966,7 @@ async def test_invoke_two_processes_in_out_interrupt( "checkpoint_id": AnyStr(), } }, - metadata={"source": "loop", "step": 0, "writes": None}, + metadata={"source": "loop", "step": 0, "writes": {"one": None}}, created_at=AnyStr(), parent_config=history[7].config, ), @@ -2290,12 +2290,14 @@ async def test_conditional_graph() -> None: "step": 0, "writes": { "agent": { - "input": "what is weather in sf", - "agent_outcome": AgentAction( - tool="search_api", - tool_input="query", - log="tool:search_api:query", - ), + "agent": { + "input": "what is weather in sf", + "agent_outcome": AgentAction( + tool="search_api", + tool_input="query", + log="tool:search_api:query", + ), + } } }, }, @@ -2516,12 +2518,14 @@ async def test_conditional_graph() -> None: "step": 0, "writes": { "agent": { - "input": "what is weather in sf", - "agent_outcome": AgentAction( - tool="search_api", - tool_input="query", - log="tool:search_api:query", - ), + "agent": { + "input": "what is weather in sf", + "agent_outcome": AgentAction( + tool="search_api", + tool_input="query", + log="tool:search_api:query", + ), + } } }, }, @@ -2748,12 +2752,14 @@ async def test_conditional_graph() -> None: "step": 0, "writes": { "agent": { - "input": "what is weather in sf", - "agent_outcome": AgentAction( - tool="search_api", - tool_input="query", - log="tool:search_api:query", - ), + "agent": { + "input": "what is weather in sf", + "agent_outcome": AgentAction( + tool="search_api", + tool_input="query", + log="tool:search_api:query", + ), + } } }, },