diff --git a/libs/langgraph/langgraph/channels/dynamic_barrier_value.py b/libs/langgraph/langgraph/channels/dynamic_barrier_value.py deleted file mode 100644 index c25551082..000000000 --- a/libs/langgraph/langgraph/channels/dynamic_barrier_value.py +++ /dev/null @@ -1,206 +0,0 @@ -from collections.abc import Sequence, Set -from typing import Any, Generic, NamedTuple, Optional, Union - -from typing_extensions import Self - -from langgraph.channels.base import BaseChannel, Value -from langgraph.constants import MISSING -from langgraph.errors import EmptyChannelError, InvalidUpdateError - - -class WaitForNames(NamedTuple): - names: Set[Any] - - -class DynamicBarrierValue( - Generic[Value], BaseChannel[Value, Union[Value, WaitForNames], Set[Value]] -): - """A channel that switches between two states - - - in the "priming" state it can't be read from. - - if it receives a WaitForNames update, it switches to the "waiting" state. - - in the "waiting" state it collects named values until all are received. - - once all named values are received, it can be read once, and it switches - back to the "priming" state. - """ - - __slots__ = ("names", "seen") - - names: Optional[Set[Value]] - seen: set[Value] - - def __init__(self, typ: type[Value]) -> None: - super().__init__(typ) - self.names = None - self.seen = set() - - def __eq__(self, value: object) -> bool: - return isinstance(value, DynamicBarrierValue) and value.names == self.names - - @property - def ValueType(self) -> type[Value]: - """The type of the value stored in the channel.""" - return self.typ - - @property - def UpdateType(self) -> type[Value]: - """The type of the update received by the channel.""" - return self.typ - - def copy(self) -> Self: - """Return a copy of the channel.""" - empty = self.__class__(self.typ) - empty.key = self.key - empty.names = self.names - empty.seen = self.seen.copy() - return empty - - def checkpoint(self) -> tuple[Optional[Set[Value]], set[Value]]: - return (self.names, self.seen) - - def from_checkpoint( - self, checkpoint: tuple[Optional[Set[Value]], set[Value]] - ) -> Self: - empty = self.__class__(self.typ) - empty.key = self.key - if checkpoint is not MISSING: - names, seen = checkpoint - empty.names = names if names is not None else None - empty.seen = seen - return empty - - def update(self, values: Sequence[Union[Value, WaitForNames]]) -> bool: - if wait_for_names := [v for v in values if isinstance(v, WaitForNames)]: - if len(wait_for_names) > 1: - raise InvalidUpdateError( - f"At key '{self.key}': Received multiple WaitForNames updates in the same step." - ) - self.names = wait_for_names[0].names - return True - elif self.names is not None: - updated = False - for value in values: - assert not isinstance(value, WaitForNames) - if value in self.names and value not in self.seen: - self.seen.add(value) - updated = True - return updated - - def get(self) -> Value: - if self.seen != self.names: - raise EmptyChannelError() - return None - - def is_available(self) -> bool: - return self.seen == self.names - - def consume(self) -> bool: - if self.seen == self.names: - self.seen = set() - self.names = None - return True - return False - - -class DynamicBarrierValueAfterFinish( - Generic[Value], BaseChannel[Value, Union[Value, WaitForNames], Set[Value]] -): - """A channel that switches between two states - - - in the "priming" state it can't be read from. - - if it receives a WaitForNames update, it switches to the "waiting" state. - - in the "waiting" state it collects named values until all are received. - - once all named values are received, and the finished flag is set, it can be read once, and it switches - back to the "priming" state. - """ - - __slots__ = ("names", "seen", "finished") - - names: Optional[Set[Value]] - seen: set[Value] - finished: bool - - def __init__(self, typ: type[Value]) -> None: - super().__init__(typ) - self.names = None - self.seen = set() - self.finished = False - - def __eq__(self, value: object) -> bool: - return ( - isinstance(value, DynamicBarrierValueAfterFinish) - and value.names == self.names - ) - - @property - def ValueType(self) -> type[Value]: - """The type of the value stored in the channel.""" - return self.typ - - @property - def UpdateType(self) -> type[Value]: - """The type of the update received by the channel.""" - return self.typ - - def copy(self) -> Self: - """Return a copy of the channel.""" - empty = self.__class__(self.typ) - empty.key = self.key - empty.names = self.names - empty.seen = self.seen.copy() - empty.finished = self.finished - return empty - - def checkpoint(self) -> tuple[Optional[Set[Value]], set[Value], bool]: - return (self.names, self.seen, self.finished) - - def from_checkpoint( - self, checkpoint: tuple[Optional[Set[Value]], set[Value], bool] - ) -> Self: - empty = self.__class__(self.typ) - empty.key = self.key - if checkpoint is not MISSING: - names, seen, finished = checkpoint - empty.names = names if names is not None else None - empty.seen = seen - empty.finished = finished - return empty - - def update(self, values: Sequence[Union[Value, WaitForNames]]) -> bool: - if wait_for_names := [v for v in values if isinstance(v, WaitForNames)]: - if len(wait_for_names) > 1: - raise InvalidUpdateError( - f"At key '{self.key}': Received multiple WaitForNames updates in the same step." - ) - self.names = wait_for_names[0].names - return True - elif self.names is not None: - updated = False - for value in values: - assert not isinstance(value, WaitForNames) - if value in self.names and value not in self.seen: - self.seen.add(value) - updated = True - return updated - - def get(self) -> Value: - if not self.finished and self.seen != self.names: - raise EmptyChannelError() - return None - - def is_available(self) -> bool: - return self.seen == self.names and self.finished - - def consume(self) -> bool: - if self.finished and self.seen == self.names: - self.seen = set() - self.names = None - return True - return False - - def finish(self) -> bool: - if not self.finished and self.seen == self.names: - self.finished = True - return True - else: - return False diff --git a/libs/langgraph/langgraph/graph/branch.py b/libs/langgraph/langgraph/graph/branch.py index c2c6c4da1..6df2ab1b5 100644 --- a/libs/langgraph/langgraph/graph/branch.py +++ b/libs/langgraph/langgraph/graph/branch.py @@ -87,7 +87,6 @@ def _get_branch_path_input_schema( class Branch(NamedTuple): path: Runnable[Any, Union[Hashable, list[Hashable]]] ends: Optional[dict[Hashable, str]] - then: Optional[str] = None input_schema: Optional[type[Any]] = None @classmethod @@ -95,7 +94,6 @@ class Branch(NamedTuple): cls, path: Runnable[Any, Union[Hashable, list[Hashable]]], path_map: Optional[Union[dict[Hashable, str], list[str]]], - then: Optional[str] = None, infer_schema: bool = False, ) -> "Branch": # coerce path_map to a dictionary @@ -123,7 +121,7 @@ class Branch(NamedTuple): # infer input schema input_schema = _get_branch_path_input_schema(path) if infer_schema else None # create branch - return cls(path=path, ends=path_map_, then=then, input_schema=input_schema) + return cls(path=path, ends=path_map_, input_schema=input_schema) def run( self, diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 907d158b0..ccd773a58 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -29,11 +29,6 @@ from langgraph._api.deprecation import LangGraphDeprecationWarning from langgraph.cache.base import BaseCache from langgraph.channels.base import BaseChannel from langgraph.channels.binop import BinaryOperatorAggregate -from langgraph.channels.dynamic_barrier_value import ( - DynamicBarrierValue, - DynamicBarrierValueAfterFinish, - WaitForNames, -) from langgraph.channels.ephemeral_value import EphemeralValue from langgraph.channels.last_value import LastValue, LastValueAfterFinish from langgraph.channels.named_barrier_value import ( @@ -506,7 +501,6 @@ class StateGraph: Runnable[Any, Union[Hashable, list[Hashable]]], ], path_map: Optional[Union[dict[Hashable, str], list[str]]] = None, - then: Optional[str] = None, ) -> Self: """Add a conditional edge from the starting node to any number of destination nodes. @@ -518,8 +512,6 @@ class StateGraph: more nodes. If it returns END, the graph will stop execution. path_map: Optional mapping of paths to node names. If omitted the paths returned by `path` should be node names. - then: The name of a node to execute after the nodes - selected by `path`. Returns: Self: The instance of the graph, allowing for method chaining. @@ -543,7 +535,7 @@ class StateGraph: f"Branch with name `{path.name}` already exists for node `{source}`" ) # save it - self.branches[source][name] = Branch.from_path(path, path_map, then, True) + self.branches[source][name] = Branch.from_path(path, path_map, True) if schema := self.branches[source][name].input_schema: self._add_schema(schema) return self @@ -611,7 +603,6 @@ class StateGraph: Runnable[Any, Union[Hashable, list[Hashable]]], ], path_map: Optional[Union[dict[Hashable, str], list[str]]] = None, - then: Optional[str] = None, ) -> Self: """Sets a conditional entry point in the graph. @@ -621,13 +612,11 @@ class StateGraph: more nodes. If it returns END, the graph will stop execution. path_map: Optional mapping of paths to node names. If omitted the paths returned by `path` should be node names. - then: The name of a node to execute after the nodes - selected by `path`. Returns: Self: The instance of the graph, allowing for method chaining. """ - return self.add_conditional_edges(START, path, path_map, then) + return self.add_conditional_edges(START, path, path_map) def set_finish_point(self, key: str) -> Self: """Marks a node as a finish point of the graph. @@ -647,16 +636,6 @@ class StateGraph: all_sources = {src for src, _ in self._all_edges} for start, branches in self.branches.items(): all_sources.add(start) - for cond, branch in branches.items(): - if branch.then is not None: - if branch.ends is not None: - for end in branch.ends.values(): - if end != END: - all_sources.add(end) - else: - for node in self.nodes: - if node != start and node != branch.then: - all_sources.add(node) for name, spec in self.nodes.items(): if spec.ends: all_sources.add(name) @@ -674,8 +653,6 @@ class StateGraph: all_targets = {end for _, end in self._all_edges} for start, branches in self.branches.items(): for cond, branch in branches.items(): - if branch.then is not None: - all_targets.add(branch.then) if branch.ends is not None: for end in branch.ends.values(): if end not in self.nodes and end != END: @@ -686,7 +663,7 @@ class StateGraph: else: all_targets.add(END) for node in self.nodes: - if node != start and node != branch.then: + if node != start: all_targets.add(node) for name, spec in self.nodes.items(): if spec.ends: @@ -999,17 +976,6 @@ class CompiledStateGraph(Pregel): ] if not writes: return [] - if branch.then and branch.then != END: - writes.append( - ChannelWriteEntry( - f"branch:{start}:{name}::then", - WaitForNames( - frozenset( - p.node if isinstance(p, Send) else p for p in packets - ) - ), - ) - ) return writes if with_reader: @@ -1040,25 +1006,6 @@ class CompiledStateGraph(Pregel): # attach branch publisher self.nodes[start].writers.append(branch.run(get_writes, reader)) - # attach then subscriber - if branch.then and branch.then != END: - ends = ( - branch.ends.values() - if branch.ends - else [node for node in self.builder.nodes if node != branch.then] - ) - channel_name = f"branch:{start}:{name}::then" - if self.builder.nodes[branch.then].defer: - self.channels[channel_name] = DynamicBarrierValueAfterFinish(str) - else: - self.channels[channel_name] = DynamicBarrierValue(str) - self.nodes[branch.then].triggers.append(channel_name) - for end in ends: - if end != END: - self.nodes[end].writers.append( - ChannelWrite((ChannelWriteEntry(channel_name, end),)) - ) - def _migrate_checkpoint(self, checkpoint: Checkpoint) -> None: """Migrate a checkpoint to new channel layout.""" diff --git a/libs/langgraph/tests/__snapshots__/test_large_cases.ambr b/libs/langgraph/tests/__snapshots__/test_large_cases.ambr index c200dae98..e750d04b8 100644 --- a/libs/langgraph/tests/__snapshots__/test_large_cases.ambr +++ b/libs/langgraph/tests/__snapshots__/test_large_cases.ambr @@ -1,44 +1,4 @@ # serializer version: 1 -# name: test_branch_then[memory] - ''' - graph TD; - __start__ --> prepare; - prepare -.-> finish; - prepare -.-> tool_two_fast; - prepare -.-> tool_two_slow; - tool_two_fast --> finish; - tool_two_slow --> finish; - finish --> __end__; - - ''' -# --- -# name: test_branch_then[memory].1 - ''' - --- - config: - flowchart: - curve: linear - --- - graph TD; - __start__([
__start__
]):::first - prepare(prepare) - tool_two_slow(tool_two_slow) - tool_two_fast(tool_two_fast) - finish(finish) - __end__([__end__
]):::last - __start__ --> prepare; - prepare -.-> finish; - prepare -.-> tool_two_fast; - prepare -.-> tool_two_slow; - tool_two_fast --> finish; - tool_two_slow --> finish; - finish --> __end__; - classDef default fill:#f2f0ff,line-height:1.2 - classDef first fill-opacity:0 - classDef last fill:#bfb6fc - - ''' -# --- # name: test_conditional_graph[memory] ''' { @@ -422,28 +382,6 @@ ''' # --- -# name: test_start_branch_then[memory-in_memory] - ''' - --- - config: - flowchart: - curve: linear - --- - graph TD; - __start__([__start__
]):::first - tool_two_slow(tool_two_slow) - tool_two_fast(tool_two_fast) - __end__([__end__
]):::last - __start__ -.-> tool_two_fast; - __start__ -.-> tool_two_slow; - tool_two_fast --> __end__; - tool_two_slow --> __end__; - classDef default fill:#f2f0ff,line-height:1.2 - classDef first fill-opacity:0 - classDef last fill:#bfb6fc - - ''' -# --- # name: test_weather_subgraph[memory] ''' --- diff --git a/libs/langgraph/tests/__snapshots__/test_pregel.ambr b/libs/langgraph/tests/__snapshots__/test_pregel.ambr index 4b991b513..a47c52178 100644 --- a/libs/langgraph/tests/__snapshots__/test_pregel.ambr +++ b/libs/langgraph/tests/__snapshots__/test_pregel.ambr @@ -541,21 +541,6 @@ ''' # --- -# name: test_in_one_fan_out_state_graph_then_defer_node[memory-True] - ''' - graph TD; - __start__ --> rewrite_query; - analyzer_one --> qa; - analyzer_one --> retriever_one; - retriever_one -.-> qa; - retriever_two --> qa; - rewrite_query -.-> analyzer_one; - rewrite_query -.-> qa; - rewrite_query -.-> retriever_two; - qa --> __end__; - - ''' -# --- # name: test_in_one_fan_out_state_graph_waiting_edge[memory] ''' graph TD; diff --git a/libs/langgraph/tests/test_large_cases.py b/libs/langgraph/tests/test_large_cases.py index d0a6e657c..262cd71d9 100644 --- a/libs/langgraph/tests/test_large_cases.py +++ b/libs/langgraph/tests/test_large_cases.py @@ -21,8 +21,6 @@ from langgraph.graph import StateGraph from langgraph.graph.message import MessagesState, add_messages from langgraph.prebuilt.chat_agent_executor import create_react_agent from langgraph.pregel import NodeBuilder, Pregel -from langgraph.store.base import BaseStore -from langgraph.store.memory import InMemoryStore from langgraph.types import ( Command, Interrupt, @@ -3139,931 +3137,6 @@ def test_dynamic_interrupt_subgraph(sync_checkpointer: BaseCheckpointSaver) -> N ) -def test_start_branch_then( - snapshot: SnapshotAssertion, - sync_checkpointer: BaseCheckpointSaver, - sync_store: BaseStore, -) -> None: - class State(TypedDict): - my_key: Annotated[str, operator.add] - market: str - - def tool_two_slow(data: State, config: RunnableConfig) -> State: - return {"my_key": " slow"} - - def tool_two_fast(data: State, config: RunnableConfig) -> State: - return {"my_key": " fast"} - - tool_two_graph = StateGraph(State) - tool_two_graph.add_node("tool_two_slow", tool_two_slow) - tool_two_graph.add_node("tool_two_fast", tool_two_fast) - tool_two_graph.set_conditional_entry_point( - lambda s: "tool_two_slow" if s["market"] == "DE" else "tool_two_fast", - then=END, - path_map=["tool_two_slow", "tool_two_fast"], - ) - tool_two = tool_two_graph.compile() - if isinstance(sync_checkpointer, InMemorySaver) and isinstance( - sync_store, InMemoryStore - ): - assert tool_two.get_graph().draw_mermaid() == snapshot - - assert tool_two.invoke({"my_key": "value", "market": "DE"}) == { - "my_key": "value slow", - "market": "DE", - } - assert tool_two.invoke({"my_key": "value", "market": "US"}) == { - "my_key": "value fast", - "market": "US", - } - - tool_two = tool_two_graph.compile( - store=sync_store, - checkpointer=sync_checkpointer, - interrupt_before=["tool_two_fast", "tool_two_slow"], - ) - - # missing thread_id - with pytest.raises(ValueError, match="thread_id"): - tool_two.invoke({"my_key": "value", "market": "DE"}) - - thread1 = {"configurable": {"thread_id": "1", "assistant_id": "a"}} - # stop when about to enter node - assert tool_two.invoke({"my_key": "value ⛰️", "market": "DE"}, thread1) == { - "my_key": "value ⛰️", - "market": "DE", - } - - assert [c.metadata for c in tool_two.checkpointer.list(thread1)] == [ - { - "parents": {}, - "source": "loop", - "step": 0, - "assistant_id": "a", - "thread_id": "1", - }, - ] - - assert tool_two.get_state(thread1) == StateSnapshot( - values={"my_key": "value ⛰️", "market": "DE"}, - tasks=(PregelTask(AnyStr(), "tool_two_slow", (PULL, "tool_two_slow")),), - next=("tool_two_slow",), - config={ - "configurable": { - "thread_id": "1", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 0, - "assistant_id": "a", - "thread_id": "1", - }, - parent_config=None, - interrupts=(), - ) - # resume, for same result as above - assert tool_two.invoke(None, thread1, debug=1) == { - "my_key": "value ⛰️ slow", - "market": "DE", - } - assert tool_two.get_state(thread1) == StateSnapshot( - values={"my_key": "value ⛰️ slow", "market": "DE"}, - tasks=(), - next=(), - config={ - "configurable": { - "thread_id": "1", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 1, - "assistant_id": "a", - "thread_id": "1", - }, - parent_config=(list(tool_two.checkpointer.list(thread1, limit=2))[-1].config), - interrupts=(), - ) - - thread2 = {"configurable": {"thread_id": "2", "assistant_id": "a"}} - # stop when about to enter node - assert tool_two.invoke({"my_key": "value", "market": "US"}, thread2) == { - "my_key": "value", - "market": "US", - } - assert tool_two.get_state(thread2) == StateSnapshot( - values={"my_key": "value", "market": "US"}, - tasks=(PregelTask(AnyStr(), "tool_two_fast", (PULL, "tool_two_fast")),), - next=("tool_two_fast",), - config={ - "configurable": { - "thread_id": "2", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 0, - "assistant_id": "a", - "thread_id": "2", - }, - parent_config=None, - interrupts=(), - ) - # resume, for same result as above - assert tool_two.invoke(None, thread2, debug=1) == { - "my_key": "value fast", - "market": "US", - } - assert tool_two.get_state(thread2) == StateSnapshot( - values={"my_key": "value fast", "market": "US"}, - tasks=(), - next=(), - config={ - "configurable": { - "thread_id": "2", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 1, - "assistant_id": "a", - "thread_id": "2", - }, - parent_config=(list(tool_two.checkpointer.list(thread2, limit=2))[-1].config), - interrupts=(), - ) - - thread3 = {"configurable": {"thread_id": "3", "assistant_id": "b"}} - # stop when about to enter node - assert tool_two.invoke({"my_key": "value", "market": "US"}, thread3) == { - "my_key": "value", - "market": "US", - } - assert tool_two.get_state(thread3) == StateSnapshot( - values={"my_key": "value", "market": "US"}, - tasks=(PregelTask(AnyStr(), "tool_two_fast", (PULL, "tool_two_fast")),), - next=("tool_two_fast",), - config={ - "configurable": { - "thread_id": "3", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 0, - "assistant_id": "b", - "thread_id": "3", - }, - parent_config=None, - interrupts=(), - ) - # update state - tool_two.update_state(thread3, {"my_key": "key"}) # appends to my_key - assert tool_two.get_state(thread3) == StateSnapshot( - values={"my_key": "valuekey", "market": "US"}, - tasks=(PregelTask(AnyStr(), "tool_two_fast", (PULL, "tool_two_fast")),), - next=("tool_two_fast",), - config={ - "configurable": { - "thread_id": "3", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "update", - "step": 1, - "assistant_id": "b", - "thread_id": "3", - }, - parent_config=(list(tool_two.checkpointer.list(thread3, limit=2))[-1].config), - interrupts=(), - ) - # resume, for same result as above - assert tool_two.invoke(None, thread3, debug=1) == { - "my_key": "valuekey fast", - "market": "US", - } - assert tool_two.get_state(thread3) == StateSnapshot( - values={"my_key": "valuekey fast", "market": "US"}, - tasks=(), - next=(), - config={ - "configurable": { - "thread_id": "3", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 2, - "assistant_id": "b", - "thread_id": "3", - }, - parent_config=(list(tool_two.checkpointer.list(thread3, limit=2))[-1].config), - interrupts=(), - ) - - -def test_branch_then( - snapshot: SnapshotAssertion, sync_checkpointer: BaseCheckpointSaver -) -> None: - class State(TypedDict): - my_key: Annotated[str, operator.add] - market: str - - tool_two_graph = StateGraph(State) - tool_two_graph.set_entry_point("prepare") - tool_two_graph.set_finish_point("finish") - tool_two_graph.add_conditional_edges( - source="prepare", - path=lambda s: "tool_two_slow" if s["market"] == "DE" else "tool_two_fast", - path_map=["tool_two_slow", "tool_two_fast"], - then="finish", - ) - tool_two_graph.add_node("prepare", lambda s: {"my_key": " prepared"}) - tool_two_graph.add_node("tool_two_slow", lambda s: {"my_key": " slow"}) - tool_two_graph.add_node("tool_two_fast", lambda s: {"my_key": " fast"}) - tool_two_graph.add_node("finish", lambda s: {"my_key": " finished"}) - tool_two = tool_two_graph.compile() - - if isinstance(sync_checkpointer, InMemorySaver): - assert tool_two.get_graph().draw_mermaid(with_styles=False) == snapshot - assert tool_two.get_graph().draw_mermaid() == snapshot - - assert tool_two.invoke({"my_key": "value", "market": "DE"}, debug=1) == { - "my_key": "value prepared slow finished", - "market": "DE", - } - assert tool_two.invoke({"my_key": "value", "market": "US"}) == { - "my_key": "value prepared fast finished", - "market": "US", - } - - # test stream_mode=debug - tool_two = tool_two_graph.compile(checkpointer=sync_checkpointer) - thread10 = {"configurable": {"thread_id": "10"}} - - res = [ - *tool_two.stream( - {"my_key": "value", "market": "DE"}, - thread10, - stream_mode="debug", - checkpoint_during=True, - ) - ] - - assert res == [ - { - "type": "checkpoint", - "timestamp": AnyStr(), - "step": -1, - "payload": { - "config": { - "configurable": { - "thread_id": "10", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "values": {"my_key": ""}, - "metadata": { - "parents": {}, - "source": "input", - "step": -1, - "thread_id": "10", - }, - "parent_config": None, - "next": ["__start__"], - "tasks": [ - { - "id": AnyStr(), - "name": "__start__", - "interrupts": (), - "state": None, - } - ], - }, - }, - { - "type": "checkpoint", - "timestamp": AnyStr(), - "step": 0, - "payload": { - "config": { - "configurable": { - "thread_id": "10", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "values": { - "my_key": "value", - "market": "DE", - }, - "metadata": { - "parents": {}, - "source": "loop", - "step": 0, - "thread_id": "10", - }, - "parent_config": { - "configurable": { - "thread_id": "10", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "next": ["prepare"], - "tasks": [ - {"id": AnyStr(), "name": "prepare", "interrupts": (), "state": None} - ], - }, - }, - { - "type": "task", - "timestamp": AnyStr(), - "step": 1, - "payload": { - "id": AnyStr(), - "name": "prepare", - "input": {"my_key": "value", "market": "DE"}, - "triggers": ("branch:to:prepare",), - }, - }, - { - "type": "task_result", - "timestamp": AnyStr(), - "step": 1, - "payload": { - "id": AnyStr(), - "name": "prepare", - "result": [("my_key", " prepared")], - "error": None, - "interrupts": [], - }, - }, - { - "type": "checkpoint", - "timestamp": AnyStr(), - "step": 1, - "payload": { - "config": { - "configurable": { - "thread_id": "10", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "values": { - "my_key": "value prepared", - "market": "DE", - }, - "metadata": { - "parents": {}, - "source": "loop", - "step": 1, - "thread_id": "10", - }, - "parent_config": { - "configurable": { - "thread_id": "10", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "next": ["tool_two_slow"], - "tasks": [ - { - "id": AnyStr(), - "name": "tool_two_slow", - "interrupts": (), - "state": None, - } - ], - }, - }, - { - "type": "task", - "timestamp": AnyStr(), - "step": 2, - "payload": { - "id": AnyStr(), - "name": "tool_two_slow", - "input": {"my_key": "value prepared", "market": "DE"}, - "triggers": ("branch:to:tool_two_slow",), - }, - }, - { - "type": "task_result", - "timestamp": AnyStr(), - "step": 2, - "payload": { - "id": AnyStr(), - "name": "tool_two_slow", - "result": [("my_key", " slow")], - "error": None, - "interrupts": [], - }, - }, - { - "type": "checkpoint", - "timestamp": AnyStr(), - "step": 2, - "payload": { - "config": { - "configurable": { - "thread_id": "10", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "values": { - "my_key": "value prepared slow", - "market": "DE", - }, - "metadata": { - "parents": {}, - "source": "loop", - "step": 2, - "thread_id": "10", - }, - "parent_config": { - "configurable": { - "thread_id": "10", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "next": ["finish"], - "tasks": [ - {"id": AnyStr(), "name": "finish", "interrupts": (), "state": None} - ], - }, - }, - { - "type": "task", - "timestamp": AnyStr(), - "step": 3, - "payload": { - "id": AnyStr(), - "name": "finish", - "input": {"my_key": "value prepared slow", "market": "DE"}, - "triggers": ( - "branch:prepare:condition::then", - "branch:to:finish", - ), - }, - }, - { - "type": "task_result", - "timestamp": AnyStr(), - "step": 3, - "payload": { - "id": AnyStr(), - "name": "finish", - "result": [("my_key", " finished")], - "error": None, - "interrupts": [], - }, - }, - { - "type": "checkpoint", - "timestamp": AnyStr(), - "step": 3, - "payload": { - "config": { - "configurable": { - "thread_id": "10", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "values": { - "my_key": "value prepared slow finished", - "market": "DE", - }, - "metadata": { - "parents": {}, - "source": "loop", - "step": 3, - "thread_id": "10", - }, - "parent_config": { - "configurable": { - "thread_id": "10", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "next": [], - "tasks": [], - }, - }, - ] - - tool_two = tool_two_graph.compile( - checkpointer=sync_checkpointer, - interrupt_before=["tool_two_fast", "tool_two_slow"], - ) - - # missing thread_id - with pytest.raises(ValueError, match="thread_id"): - tool_two.invoke({"my_key": "value", "market": "DE"}) - - thread1 = {"configurable": {"thread_id": "1"}} - # stop when about to enter node - assert tool_two.invoke({"my_key": "value", "market": "DE"}, thread1) == { - "my_key": "value prepared", - "market": "DE", - } - assert tool_two.get_state(thread1) == StateSnapshot( - values={"my_key": "value prepared", "market": "DE"}, - tasks=(PregelTask(AnyStr(), "tool_two_slow", (PULL, "tool_two_slow")),), - next=("tool_two_slow",), - config={ - "configurable": { - "thread_id": "1", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 1, - "thread_id": "1", - }, - parent_config=None, - interrupts=(), - ) - # resume, for same result as above - assert tool_two.invoke(None, thread1, debug=1) == { - "my_key": "value prepared slow finished", - "market": "DE", - } - assert tool_two.get_state(thread1) == StateSnapshot( - values={"my_key": "value prepared slow finished", "market": "DE"}, - tasks=(), - next=(), - config={ - "configurable": { - "thread_id": "1", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 3, - "thread_id": "1", - }, - parent_config=(list(tool_two.checkpointer.list(thread1, limit=2))[-1].config), - interrupts=(), - ) - - thread2 = {"configurable": {"thread_id": "2"}} - # stop when about to enter node - assert tool_two.invoke({"my_key": "value", "market": "US"}, thread2) == { - "my_key": "value prepared", - "market": "US", - } - assert tool_two.get_state(thread2) == StateSnapshot( - values={"my_key": "value prepared", "market": "US"}, - tasks=(PregelTask(AnyStr(), "tool_two_fast", (PULL, "tool_two_fast")),), - next=("tool_two_fast",), - config={ - "configurable": { - "thread_id": "2", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 1, - "thread_id": "2", - }, - parent_config=None, - interrupts=(), - ) - # resume, for same result as above - assert tool_two.invoke(None, thread2, debug=1) == { - "my_key": "value prepared fast finished", - "market": "US", - } - assert tool_two.get_state(thread2) == StateSnapshot( - values={"my_key": "value prepared fast finished", "market": "US"}, - tasks=(), - next=(), - config={ - "configurable": { - "thread_id": "2", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 3, - "thread_id": "2", - }, - parent_config=(list(tool_two.checkpointer.list(thread2, limit=2))[-1].config), - interrupts=(), - ) - - tool_two = tool_two_graph.compile( - checkpointer=sync_checkpointer, interrupt_before=["finish"] - ) - - thread1 = {"configurable": {"thread_id": "11"}} - - # stop when about to enter node - assert tool_two.invoke({"my_key": "value", "market": "DE"}, thread1) == { - "my_key": "value prepared slow", - "market": "DE", - } - assert tool_two.get_state(thread1) == StateSnapshot( - values={ - "my_key": "value prepared slow", - "market": "DE", - }, - tasks=(PregelTask(AnyStr(), "finish", (PULL, "finish")),), - next=("finish",), - config={ - "configurable": { - "thread_id": "11", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 2, - "thread_id": "11", - }, - parent_config=None, - interrupts=(), - ) - - # update state - tool_two.update_state(thread1, {"my_key": "er"}) - assert tool_two.get_state(thread1) == StateSnapshot( - values={ - "my_key": "value prepared slower", - "market": "DE", - }, - tasks=(PregelTask(AnyStr(), "finish", (PULL, "finish")),), - next=("finish",), - config={ - "configurable": { - "thread_id": "11", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "update", - "step": 3, - "thread_id": "11", - }, - parent_config=(list(tool_two.checkpointer.list(thread1, limit=2))[-1].config), - interrupts=(), - ) - - tool_two = tool_two_graph.compile( - checkpointer=sync_checkpointer, interrupt_after=["prepare"] - ) - - # missing thread_id - with pytest.raises(ValueError, match="thread_id"): - tool_two.invoke({"my_key": "value", "market": "DE"}) - - thread1 = {"configurable": {"thread_id": "21"}} - # stop when about to enter node - assert tool_two.invoke({"my_key": "value", "market": "DE"}, thread1) == { - "my_key": "value prepared", - "market": "DE", - } - assert tool_two.get_state(thread1) == StateSnapshot( - values={"my_key": "value prepared", "market": "DE"}, - tasks=(PregelTask(AnyStr(), "tool_two_slow", (PULL, "tool_two_slow")),), - next=("tool_two_slow",), - config={ - "configurable": { - "thread_id": "21", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 1, - "thread_id": "21", - }, - parent_config=None, - interrupts=(), - ) - # resume, for same result as above - assert tool_two.invoke(None, thread1, debug=1) == { - "my_key": "value prepared slow finished", - "market": "DE", - } - assert tool_two.get_state(thread1) == StateSnapshot( - values={"my_key": "value prepared slow finished", "market": "DE"}, - tasks=(), - next=(), - config={ - "configurable": { - "thread_id": "21", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 3, - "thread_id": "21", - }, - parent_config=(list(tool_two.checkpointer.list(thread1, limit=2))[-1].config), - interrupts=(), - ) - - thread2 = {"configurable": {"thread_id": "22"}} - # stop when about to enter node - assert tool_two.invoke({"my_key": "value", "market": "US"}, thread2) == { - "my_key": "value prepared", - "market": "US", - } - assert tool_two.get_state(thread2) == StateSnapshot( - values={"my_key": "value prepared", "market": "US"}, - tasks=(PregelTask(AnyStr(), "tool_two_fast", (PULL, "tool_two_fast")),), - next=("tool_two_fast",), - config={ - "configurable": { - "thread_id": "22", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 1, - "thread_id": "22", - }, - parent_config=None, - interrupts=(), - ) - # resume, for same result as above - assert tool_two.invoke(None, thread2, debug=1) == { - "my_key": "value prepared fast finished", - "market": "US", - } - assert tool_two.get_state(thread2) == StateSnapshot( - values={"my_key": "value prepared fast finished", "market": "US"}, - tasks=(), - next=(), - config={ - "configurable": { - "thread_id": "22", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 3, - "thread_id": "22", - }, - parent_config=(list(tool_two.checkpointer.list(thread2, limit=2))[-1].config), - interrupts=(), - ) - - thread3 = {"configurable": {"thread_id": "23"}} - # update an empty thread before first run - tool_two.update_state(thread3, {"my_key": "key", "market": "DE"}) - # check current state - assert tool_two.get_state(thread3) == StateSnapshot( - values={"my_key": "key", "market": "DE"}, - tasks=(PregelTask(AnyStr(), "prepare", (PULL, "prepare")),), - next=("prepare",), - config={ - "configurable": { - "thread_id": "23", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "update", - "step": 0, - "thread_id": "23", - }, - parent_config=None, - interrupts=(), - ) - # run from this point - assert tool_two.invoke(None, thread3) == { - "my_key": "key prepared", - "market": "DE", - } - # get state after first node - assert tool_two.get_state(thread3) == StateSnapshot( - values={"my_key": "key prepared", "market": "DE"}, - tasks=(PregelTask(AnyStr(), "tool_two_slow", (PULL, "tool_two_slow")),), - next=("tool_two_slow",), - config={ - "configurable": { - "thread_id": "23", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 1, - "thread_id": "23", - }, - parent_config=(list(tool_two.checkpointer.list(thread3, limit=2))[-1].config), - interrupts=(), - ) - # resume, for same result as above - assert tool_two.invoke(None, thread3, debug=1) == { - "my_key": "key prepared slow finished", - "market": "DE", - } - assert tool_two.get_state(thread3) == StateSnapshot( - values={"my_key": "key prepared slow finished", "market": "DE"}, - tasks=(), - next=(), - config={ - "configurable": { - "thread_id": "23", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 3, - "thread_id": "23", - }, - parent_config=(list(tool_two.checkpointer.list(thread3, limit=2))[-1].config), - interrupts=(), - ) - - def test_send_dedupe_on_resume( sync_checkpointer: BaseCheckpointSaver, checkpoint_during: bool ) -> None: diff --git a/libs/langgraph/tests/test_large_cases_async.py b/libs/langgraph/tests/test_large_cases_async.py index 0c2c1ff5a..2b3a99c13 100644 --- a/libs/langgraph/tests/test_large_cases_async.py +++ b/libs/langgraph/tests/test_large_cases_async.py @@ -24,7 +24,6 @@ from langgraph.graph.message import add_messages from langgraph.graph.state import StateGraph from langgraph.prebuilt.chat_agent_executor import create_react_agent from langgraph.pregel import NodeBuilder, Pregel -from langgraph.store.base import BaseStore from langgraph.types import PregelTask, Send, StateSnapshot, StreamWriter from tests.any_int import AnyInt from tests.any_str import AnyDict, AnyStr, UnsortedSequence @@ -2306,1014 +2305,6 @@ async def test_in_one_fan_out_out_one_graph_state() -> None: ] -async def test_start_branch_then( - async_checkpointer: BaseCheckpointSaver, async_store: BaseStore -) -> None: - class State(TypedDict): - my_key: Annotated[str, operator.add] - market: str - - def tool_two_slow(data: State, config: RunnableConfig) -> State: - return {"my_key": " slow"} - - def tool_two_fast(data: State, config: RunnableConfig) -> State: - return {"my_key": " fast"} - - tool_two_graph = StateGraph(State) - tool_two_graph.add_node("tool_two_slow", tool_two_slow) - tool_two_graph.add_node("tool_two_fast", tool_two_fast) - tool_two_graph.set_conditional_entry_point( - lambda s: "tool_two_slow" if s["market"] == "DE" else "tool_two_fast", then=END - ) - tool_two = tool_two_graph.compile() - - assert await tool_two.ainvoke({"my_key": "value", "market": "DE"}) == { - "my_key": "value slow", - "market": "DE", - } - assert await tool_two.ainvoke({"my_key": "value", "market": "US"}) == { - "my_key": "value fast", - "market": "US", - } - - tool_two = tool_two_graph.compile( - store=async_store, - checkpointer=async_checkpointer, - interrupt_before=["tool_two_fast", "tool_two_slow"], - ) - - # missing thread_id - with pytest.raises(ValueError, match="thread_id"): - await tool_two.ainvoke({"my_key": "value", "market": "DE"}) - - thread1 = {"configurable": {"thread_id": "1", "assistant_id": "a"}} - # stop when about to enter node - assert await tool_two.ainvoke({"my_key": "value", "market": "DE"}, thread1) == { - "my_key": "value", - "market": "DE", - } - assert [c.metadata async for c in tool_two.checkpointer.alist(thread1)] == [ - { - "parents": {}, - "source": "loop", - "step": 0, - "assistant_id": "a", - "thread_id": "1", - }, - ] - - assert await tool_two.aget_state(thread1) == StateSnapshot( - values={"my_key": "value", "market": "DE"}, - tasks=(PregelTask(AnyStr(), "tool_two_slow", (PULL, "tool_two_slow")),), - next=("tool_two_slow",), - config={ - "configurable": { - "thread_id": "1", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 0, - "assistant_id": "a", - "thread_id": "1", - }, - parent_config=None, - interrupts=(), - ) - # resume, for same result as above - assert await tool_two.ainvoke(None, thread1, debug=1) == { - "my_key": "value slow", - "market": "DE", - } - assert await tool_two.aget_state(thread1) == StateSnapshot( - values={"my_key": "value slow", "market": "DE"}, - tasks=(), - next=(), - config={ - "configurable": { - "thread_id": "1", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 1, - "assistant_id": "a", - "thread_id": "1", - }, - parent_config=[c async for c in tool_two.checkpointer.alist(thread1, limit=2)][ - -1 - ].config, - interrupts=(), - ) - - thread2 = {"configurable": {"thread_id": "2", "assistant_id": "a"}} - # stop when about to enter node - assert await tool_two.ainvoke({"my_key": "value", "market": "US"}, thread2) == { - "my_key": "value", - "market": "US", - } - assert await tool_two.aget_state(thread2) == StateSnapshot( - values={"my_key": "value", "market": "US"}, - tasks=(PregelTask(AnyStr(), "tool_two_fast", (PULL, "tool_two_fast")),), - next=("tool_two_fast",), - config={ - "configurable": { - "thread_id": "2", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 0, - "assistant_id": "a", - "thread_id": "2", - }, - parent_config=None, - interrupts=(), - ) - # resume, for same result as above - assert await tool_two.ainvoke(None, thread2, debug=1) == { - "my_key": "value fast", - "market": "US", - } - assert await tool_two.aget_state(thread2) == StateSnapshot( - values={"my_key": "value fast", "market": "US"}, - tasks=(), - next=(), - config={ - "configurable": { - "thread_id": "2", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 1, - "assistant_id": "a", - "thread_id": "2", - }, - parent_config=( - [c async for c in tool_two.checkpointer.alist(thread2, limit=2)][-1].config - ), - interrupts=(), - ) - - thread3 = {"configurable": {"thread_id": "3", "assistant_id": "b"}} - # stop when about to enter node - assert await tool_two.ainvoke({"my_key": "value", "market": "US"}, thread3) == { - "my_key": "value", - "market": "US", - } - assert await tool_two.aget_state(thread3) == StateSnapshot( - values={"my_key": "value", "market": "US"}, - tasks=(PregelTask(AnyStr(), "tool_two_fast", (PULL, "tool_two_fast")),), - next=("tool_two_fast",), - config={ - "configurable": { - "thread_id": "3", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 0, - "assistant_id": "b", - "thread_id": "3", - }, - parent_config=None, - interrupts=(), - ) - # update state - await tool_two.aupdate_state(thread3, {"my_key": "key"}) # appends to my_key - assert await tool_two.aget_state(thread3) == StateSnapshot( - values={"my_key": "valuekey", "market": "US"}, - tasks=(PregelTask(AnyStr(), "tool_two_fast", (PULL, "tool_two_fast")),), - next=("tool_two_fast",), - config={ - "configurable": { - "thread_id": "3", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "update", - "step": 1, - "assistant_id": "b", - "thread_id": "3", - }, - parent_config=( - [c async for c in tool_two.checkpointer.alist(thread3, limit=2)][-1].config - ), - interrupts=(), - ) - # resume, for same result as above - assert await tool_two.ainvoke(None, thread3, debug=1) == { - "my_key": "valuekey fast", - "market": "US", - } - assert await tool_two.aget_state(thread3) == StateSnapshot( - values={"my_key": "valuekey fast", "market": "US"}, - tasks=(), - next=(), - config={ - "configurable": { - "thread_id": "3", - "checkpoint_ns": AnyStr(), - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 2, - "assistant_id": "b", - "thread_id": "3", - }, - parent_config=( - [c async for c in tool_two.checkpointer.alist(thread3, limit=2)][-1].config - ), - interrupts=(), - ) - - -async def test_branch_then(async_checkpointer: BaseCheckpointSaver) -> None: - class State(TypedDict): - my_key: Annotated[str, operator.add] - market: str - - tool_two_graph = StateGraph(State) - tool_two_graph.set_entry_point("prepare") - tool_two_graph.set_finish_point("finish") - tool_two_graph.add_conditional_edges( - source="prepare", - path=lambda s: "tool_two_slow" if s["market"] == "DE" else "tool_two_fast", - then="finish", - ) - tool_two_graph.add_node("prepare", lambda s: {"my_key": " prepared"}) - tool_two_graph.add_node("tool_two_slow", lambda s: {"my_key": " slow"}) - tool_two_graph.add_node("tool_two_fast", lambda s: {"my_key": " fast"}) - tool_two_graph.add_node("finish", lambda s: {"my_key": " finished"}) - tool_two = tool_two_graph.compile() - - assert await tool_two.ainvoke({"my_key": "value", "market": "DE"}, debug=1) == { - "my_key": "value prepared slow finished", - "market": "DE", - } - assert await tool_two.ainvoke({"my_key": "value", "market": "US"}) == { - "my_key": "value prepared fast finished", - "market": "US", - } - - # test stream_mode=debug - tool_two = tool_two_graph.compile(checkpointer=async_checkpointer) - thread10 = {"configurable": {"thread_id": "10"}} - assert [ - c - async for c in tool_two.astream( - {"my_key": "value", "market": "DE"}, - thread10, - stream_mode="debug", - checkpoint_during=True, - ) - ] == [ - { - "type": "checkpoint", - "timestamp": AnyStr(), - "step": -1, - "payload": { - "config": { - "configurable": { - "thread_id": "10", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "values": {"my_key": ""}, - "metadata": { - "parents": {}, - "source": "input", - "step": -1, - "thread_id": "10", - }, - "parent_config": None, - "next": ["__start__"], - "tasks": [ - { - "id": AnyStr(), - "name": "__start__", - "interrupts": (), - "state": None, - } - ], - }, - }, - { - "type": "checkpoint", - "timestamp": AnyStr(), - "step": 0, - "payload": { - "config": { - "configurable": { - "thread_id": "10", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "values": { - "my_key": "value", - "market": "DE", - }, - "metadata": { - "parents": {}, - "source": "loop", - "step": 0, - "thread_id": "10", - }, - "parent_config": { - "configurable": { - "thread_id": "10", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "next": ["prepare"], - "tasks": [ - { - "id": AnyStr(), - "name": "prepare", - "interrupts": (), - "state": None, - } - ], - }, - }, - { - "type": "task", - "timestamp": AnyStr(), - "step": 1, - "payload": { - "id": AnyStr(), - "name": "prepare", - "input": {"my_key": "value", "market": "DE"}, - "triggers": ("branch:to:prepare",), - }, - }, - { - "type": "task_result", - "timestamp": AnyStr(), - "step": 1, - "payload": { - "id": AnyStr(), - "name": "prepare", - "result": [("my_key", " prepared")], - "error": None, - "interrupts": [], - }, - }, - { - "type": "checkpoint", - "timestamp": AnyStr(), - "step": 1, - "payload": { - "config": { - "configurable": { - "thread_id": "10", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "values": { - "my_key": "value prepared", - "market": "DE", - }, - "metadata": { - "parents": {}, - "source": "loop", - "step": 1, - "thread_id": "10", - }, - "parent_config": { - "configurable": { - "thread_id": "10", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "next": ["tool_two_slow"], - "tasks": [ - { - "id": AnyStr(), - "name": "tool_two_slow", - "interrupts": (), - "state": None, - } - ], - }, - }, - { - "type": "task", - "timestamp": AnyStr(), - "step": 2, - "payload": { - "id": AnyStr(), - "name": "tool_two_slow", - "input": {"my_key": "value prepared", "market": "DE"}, - "triggers": ("branch:to:tool_two_slow",), - }, - }, - { - "type": "task_result", - "timestamp": AnyStr(), - "step": 2, - "payload": { - "id": AnyStr(), - "name": "tool_two_slow", - "result": [("my_key", " slow")], - "error": None, - "interrupts": [], - }, - }, - { - "type": "checkpoint", - "timestamp": AnyStr(), - "step": 2, - "payload": { - "config": { - "configurable": { - "thread_id": "10", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "values": { - "my_key": "value prepared slow", - "market": "DE", - }, - "metadata": { - "parents": {}, - "source": "loop", - "step": 2, - "thread_id": "10", - }, - "parent_config": { - "configurable": { - "thread_id": "10", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "next": ["finish"], - "tasks": [ - { - "id": AnyStr(), - "name": "finish", - "interrupts": (), - "state": None, - } - ], - }, - }, - { - "type": "task", - "timestamp": AnyStr(), - "step": 3, - "payload": { - "id": AnyStr(), - "name": "finish", - "input": {"my_key": "value prepared slow", "market": "DE"}, - "triggers": ( - "branch:prepare:condition::then", - "branch:to:finish", - ), - }, - }, - { - "type": "task_result", - "timestamp": AnyStr(), - "step": 3, - "payload": { - "id": AnyStr(), - "name": "finish", - "result": [("my_key", " finished")], - "error": None, - "interrupts": [], - }, - }, - { - "type": "checkpoint", - "timestamp": AnyStr(), - "step": 3, - "payload": { - "config": { - "configurable": { - "thread_id": "10", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "values": { - "my_key": "value prepared slow finished", - "market": "DE", - }, - "metadata": { - "parents": {}, - "source": "loop", - "step": 3, - "thread_id": "10", - }, - "parent_config": { - "configurable": { - "thread_id": "10", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "next": [], - "tasks": [], - }, - }, - ] - - tool_two = tool_two_graph.compile( - checkpointer=async_checkpointer, - interrupt_before=["tool_two_fast", "tool_two_slow"], - ) - - # missing thread_id - with pytest.raises(ValueError, match="thread_id"): - await tool_two.ainvoke({"my_key": "value", "market": "DE"}) - - thread1 = {"configurable": {"thread_id": "11"}} - # stop when about to enter node - assert [ - c - async for c in tool_two.astream( - {"my_key": "value", "market": "DE"}, - thread1, - stream_mode="debug", - checkpoint_during=True, - ) - ] == [ - { - "type": "checkpoint", - "timestamp": AnyStr(), - "step": -1, - "payload": { - "config": { - "configurable": { - "thread_id": "11", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "values": {"my_key": ""}, - "metadata": { - "parents": {}, - "source": "input", - "step": -1, - "thread_id": "11", - }, - "parent_config": None, - "next": ["__start__"], - "tasks": [ - { - "id": AnyStr(), - "name": "__start__", - "interrupts": (), - "state": None, - } - ], - }, - }, - { - "type": "checkpoint", - "timestamp": AnyStr(), - "step": 0, - "payload": { - "config": { - "configurable": { - "thread_id": "11", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "values": { - "my_key": "value", - "market": "DE", - }, - "metadata": { - "parents": {}, - "source": "loop", - "step": 0, - "thread_id": "11", - }, - "parent_config": { - "configurable": { - "thread_id": "11", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "next": ["prepare"], - "tasks": [ - { - "id": AnyStr(), - "name": "prepare", - "interrupts": (), - "state": None, - } - ], - }, - }, - { - "type": "task", - "timestamp": AnyStr(), - "step": 1, - "payload": { - "id": AnyStr(), - "name": "prepare", - "input": {"my_key": "value", "market": "DE"}, - "triggers": ("branch:to:prepare",), - }, - }, - { - "type": "task_result", - "timestamp": AnyStr(), - "step": 1, - "payload": { - "id": AnyStr(), - "name": "prepare", - "result": [("my_key", " prepared")], - "error": None, - "interrupts": [], - }, - }, - { - "type": "checkpoint", - "timestamp": AnyStr(), - "step": 1, - "payload": { - "config": { - "configurable": { - "thread_id": "11", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "values": { - "my_key": "value prepared", - "market": "DE", - }, - "metadata": { - "parents": {}, - "source": "loop", - "step": 1, - "thread_id": "11", - }, - "parent_config": { - "configurable": { - "thread_id": "11", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - }, - }, - "next": ["tool_two_slow"], - "tasks": [ - { - "id": AnyStr(), - "name": "tool_two_slow", - "interrupts": (), - "state": None, - } - ], - }, - }, - ] - assert await tool_two.aget_state(thread1) == StateSnapshot( - values={"my_key": "value prepared", "market": "DE"}, - tasks=(PregelTask(AnyStr(), "tool_two_slow", (PULL, "tool_two_slow")),), - next=("tool_two_slow",), - config={ - "configurable": { - "thread_id": "11", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 1, - "thread_id": "11", - }, - parent_config=( - [c async for c in tool_two.checkpointer.alist(thread1, limit=2)][-1].config - ), - interrupts=(), - ) - # resume, for same result as above - assert await tool_two.ainvoke(None, thread1, debug=1) == { - "my_key": "value prepared slow finished", - "market": "DE", - } - assert await tool_two.aget_state(thread1) == StateSnapshot( - values={"my_key": "value prepared slow finished", "market": "DE"}, - tasks=(), - next=(), - config={ - "configurable": { - "thread_id": "11", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 3, - "thread_id": "11", - }, - parent_config=( - [c async for c in tool_two.checkpointer.alist(thread1, limit=2)][-1].config - ), - interrupts=(), - ) - - thread2 = {"configurable": {"thread_id": "12"}} - # stop when about to enter node - assert await tool_two.ainvoke({"my_key": "value", "market": "US"}, thread2) == { - "my_key": "value prepared", - "market": "US", - } - assert await tool_two.aget_state(thread2) == StateSnapshot( - values={"my_key": "value prepared", "market": "US"}, - tasks=(PregelTask(AnyStr(), "tool_two_fast", (PULL, "tool_two_fast")),), - next=("tool_two_fast",), - config={ - "configurable": { - "thread_id": "12", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 1, - "thread_id": "12", - }, - parent_config=None, - interrupts=(), - ) - # resume, for same result as above - assert await tool_two.ainvoke(None, thread2, debug=1) == { - "my_key": "value prepared fast finished", - "market": "US", - } - assert await tool_two.aget_state(thread2) == StateSnapshot( - values={"my_key": "value prepared fast finished", "market": "US"}, - tasks=(), - next=(), - config={ - "configurable": { - "thread_id": "12", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 3, - "thread_id": "12", - }, - parent_config=( - [c async for c in tool_two.checkpointer.alist(thread2, limit=2)][-1].config - ), - interrupts=(), - ) - - tool_two = tool_two_graph.compile( - checkpointer=async_checkpointer, interrupt_after=["prepare"] - ) - - # missing thread_id - with pytest.raises(ValueError, match="thread_id"): - await tool_two.ainvoke({"my_key": "value", "market": "DE"}) - - thread1 = {"configurable": {"thread_id": "21"}} - # stop when about to enter node - assert await tool_two.ainvoke({"my_key": "value", "market": "DE"}, thread1) == { - "my_key": "value prepared", - "market": "DE", - } - assert await tool_two.aget_state(thread1) == StateSnapshot( - values={"my_key": "value prepared", "market": "DE"}, - tasks=(PregelTask(AnyStr(), "tool_two_slow", (PULL, "tool_two_slow")),), - next=("tool_two_slow",), - config={ - "configurable": { - "thread_id": "21", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 1, - "thread_id": "21", - }, - parent_config=None, - interrupts=(), - ) - # resume, for same result as above - assert await tool_two.ainvoke(None, thread1, debug=1) == { - "my_key": "value prepared slow finished", - "market": "DE", - } - assert await tool_two.aget_state(thread1) == StateSnapshot( - values={"my_key": "value prepared slow finished", "market": "DE"}, - tasks=(), - next=(), - config={ - "configurable": { - "thread_id": "21", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 3, - "thread_id": "21", - }, - parent_config=( - [c async for c in tool_two.checkpointer.alist(thread1, limit=2)][-1].config - ), - interrupts=(), - ) - - thread2 = {"configurable": {"thread_id": "22"}} - # stop when about to enter node - assert await tool_two.ainvoke({"my_key": "value", "market": "US"}, thread2) == { - "my_key": "value prepared", - "market": "US", - } - assert await tool_two.aget_state(thread2) == StateSnapshot( - values={"my_key": "value prepared", "market": "US"}, - tasks=(PregelTask(AnyStr(), "tool_two_fast", (PULL, "tool_two_fast")),), - next=("tool_two_fast",), - config={ - "configurable": { - "thread_id": "22", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 1, - "thread_id": "22", - }, - parent_config=None, - interrupts=(), - ) - # resume, for same result as above - assert await tool_two.ainvoke(None, thread2, debug=1) == { - "my_key": "value prepared fast finished", - "market": "US", - } - assert await tool_two.aget_state(thread2) == StateSnapshot( - values={"my_key": "value prepared fast finished", "market": "US"}, - tasks=(), - next=(), - config={ - "configurable": { - "thread_id": "22", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 3, - "thread_id": "22", - }, - parent_config=( - [c async for c in tool_two.checkpointer.alist(thread2, limit=2)][-1].config - ), - interrupts=(), - ) - - thread3 = {"configurable": {"thread_id": "23"}} - # update an empty thread before first run - uconfig = await tool_two.aupdate_state(thread3, {"my_key": "key", "market": "DE"}) - # check current state - assert await tool_two.aget_state(thread3) == StateSnapshot( - values={"my_key": "key", "market": "DE"}, - tasks=(PregelTask(AnyStr(), "prepare", (PULL, "prepare")),), - next=("prepare",), - config=uconfig, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "update", - "step": 0, - "thread_id": "23", - }, - parent_config=None, - interrupts=(), - ) - # run from this point - assert await tool_two.ainvoke(None, thread3) == { - "my_key": "key prepared", - "market": "DE", - } - # get state after first node - assert await tool_two.aget_state(thread3) == StateSnapshot( - values={"my_key": "key prepared", "market": "DE"}, - tasks=(PregelTask(AnyStr(), "tool_two_slow", (PULL, "tool_two_slow")),), - next=("tool_two_slow",), - config={ - "configurable": { - "thread_id": "23", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 1, - "thread_id": "23", - }, - parent_config=(uconfig), - interrupts=(), - ) - # resume, for same result as above - assert await tool_two.ainvoke(None, thread3, debug=1) == { - "my_key": "key prepared slow finished", - "market": "DE", - } - assert await tool_two.aget_state(thread3) == StateSnapshot( - values={"my_key": "key prepared slow finished", "market": "DE"}, - tasks=(), - next=(), - config={ - "configurable": { - "thread_id": "23", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "loop", - "step": 3, - "thread_id": "23", - }, - parent_config=( - [c async for c in tool_two.checkpointer.alist(thread3, limit=2)][-1].config - ), - interrupts=(), - ) - - async def test_nested_graph_state(async_checkpointer: BaseCheckpointSaver) -> None: class InnerState(TypedDict): my_key: str diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 0f0cce810..9336e53e4 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -2337,275 +2337,6 @@ def test_in_one_fan_out_state_graph_defer_node( ] -@pytest.mark.parametrize("with_path_map", (True, False)) -def test_in_one_fan_out_state_graph_then_defer_node( - snapshot: SnapshotAssertion, - sync_checkpointer: BaseCheckpointSaver, - with_path_map: bool, -) -> None: - def sorted_add( - x: list[str], y: Union[list[str], list[tuple[str, str]]] - ) -> list[str]: - if isinstance(y[0], tuple): - for rem, _ in y: - x.remove(rem) - y = [t[1] for t in y] - return sorted(operator.add(x, y)) - - class State(TypedDict, total=False): - query: str - answer: str - docs: Annotated[list[str], sorted_add] - - workflow = StateGraph(State) - - @workflow.add_node - def rewrite_query(data: State) -> State: - return {"query": f"query: {data['query']}"} - - def analyzer_one(data: State) -> State: - return {"query": f"analyzed: {data['query']}"} - - def retriever_one(data: State) -> State: - return {"docs": ["doc1", "doc2"]} - - def retriever_two(data: State) -> State: - time.sleep(0.1) # to ensure stream order - return {"docs": ["doc3", "doc4"]} - - def qa(data: State) -> State: - return {"answer": ",".join(data["docs"])} - - workflow.add_node(analyzer_one) - workflow.add_node(retriever_one) - workflow.add_node(retriever_two) - workflow.add_node(qa, defer=True) - - workflow.set_entry_point("rewrite_query") - workflow.add_conditional_edges( - "rewrite_query", - lambda _: ["analyzer_one", "retriever_two"], - ["analyzer_one", "retriever_two"] if with_path_map else None, - then="qa", - ) - workflow.add_edge("analyzer_one", "retriever_one") - - app = workflow.compile() - - if isinstance(sync_checkpointer, InMemorySaver) and with_path_map: - assert app.get_graph().draw_mermaid(with_styles=False) == snapshot - - assert app.invoke({"query": "what is weather in sf"}) == { - "query": "analyzed: query: what is weather in sf", - "docs": ["doc1", "doc2", "doc3", "doc4"], - "answer": "doc1,doc2,doc3,doc4", - } - - assert [*app.stream({"query": "what is weather in sf"})] == [ - {"rewrite_query": {"query": "query: what is weather in sf"}}, - {"analyzer_one": {"query": "analyzed: query: what is weather in sf"}}, - {"retriever_two": {"docs": ["doc3", "doc4"]}}, - {"retriever_one": {"docs": ["doc1", "doc2"]}}, - {"qa": {"answer": "doc1,doc2,doc3,doc4"}}, - ] - - assert [*app.stream({"query": "what is weather in sf"}, stream_mode="debug")] == [ - { - "type": "task", - "timestamp": AnyStr(), - "step": 1, - "payload": { - "id": AnyStr(), - "name": "rewrite_query", - "input": {"query": "what is weather in sf", "docs": []}, - "triggers": ("branch:to:rewrite_query",), - }, - }, - { - "type": "task_result", - "timestamp": AnyStr(), - "step": 1, - "payload": { - "id": AnyStr(), - "name": "rewrite_query", - "error": None, - "result": [("query", "query: what is weather in sf")], - "interrupts": [], - }, - }, - { - "type": "task", - "timestamp": AnyStr(), - "step": 2, - "payload": { - "id": AnyStr(), - "name": "analyzer_one", - "input": { - "query": "query: what is weather in sf", - "docs": [], - }, - "triggers": ("branch:to:analyzer_one",), - }, - }, - { - "type": "task", - "timestamp": AnyStr(), - "step": 2, - "payload": { - "id": AnyStr(), - "name": "retriever_two", - "input": {"query": "query: what is weather in sf", "docs": []}, - "triggers": ("branch:to:retriever_two",), - }, - }, - { - "type": "task_result", - "timestamp": AnyStr(), - "step": 2, - "payload": { - "id": AnyStr(), - "name": "analyzer_one", - "error": None, - "result": [("query", "analyzed: query: what is weather in sf")], - "interrupts": [], - }, - }, - { - "type": "task_result", - "timestamp": AnyStr(), - "step": 2, - "payload": { - "id": AnyStr(), - "name": "retriever_two", - "error": None, - "result": [("docs", ["doc3", "doc4"])], - "interrupts": [], - }, - }, - { - "type": "task", - "timestamp": AnyStr(), - "step": 3, - "payload": { - "id": AnyStr(), - "name": "retriever_one", - "input": { - "query": "analyzed: query: what is weather in sf", - "docs": ["doc3", "doc4"], - }, - "triggers": ("branch:to:retriever_one",), - }, - }, - { - "type": "task_result", - "timestamp": AnyStr(), - "step": 3, - "payload": { - "id": AnyStr(), - "name": "retriever_one", - "error": None, - "result": [("docs", ["doc1", "doc2"])], - "interrupts": [], - }, - }, - { - "type": "task", - "timestamp": AnyStr(), - "step": 4, - "payload": { - "id": AnyStr(), - "name": "qa", - "input": { - "query": "analyzed: query: what is weather in sf", - "docs": ["doc1", "doc2", "doc3", "doc4"], - }, - "triggers": ("branch:rewrite_query:condition::then", "branch:to:qa"), - }, - }, - { - "type": "task_result", - "timestamp": AnyStr(), - "step": 4, - "payload": { - "id": AnyStr(), - "name": "qa", - "error": None, - "result": [("answer", "doc1,doc2,doc3,doc4")], - "interrupts": [], - }, - }, - ] - - app_w_interrupt = workflow.compile( - checkpointer=sync_checkpointer, - interrupt_after=["analyzer_one"], - ) - config = {"configurable": {"thread_id": "1"}} - - assert [ - c for c in app_w_interrupt.stream({"query": "what is weather in sf"}, config) - ] == [ - {"rewrite_query": {"query": "query: what is weather in sf"}}, - {"analyzer_one": {"query": "analyzed: query: what is weather in sf"}}, - {"retriever_two": {"docs": ["doc3", "doc4"]}}, - {"__interrupt__": ()}, - ] - - assert [c for c in app_w_interrupt.stream(None, config)] == [ - {"retriever_one": {"docs": ["doc1", "doc2"]}}, - {"qa": {"answer": "doc1,doc2,doc3,doc4"}}, - ] - - app_w_interrupt = workflow.compile( - checkpointer=sync_checkpointer, - interrupt_before=["qa"], - ) - config = {"configurable": {"thread_id": "2"}} - - assert [ - c for c in app_w_interrupt.stream({"query": "what is weather in sf"}, config) - ] == [ - {"rewrite_query": {"query": "query: what is weather in sf"}}, - {"analyzer_one": {"query": "analyzed: query: what is weather in sf"}}, - {"retriever_two": {"docs": ["doc3", "doc4"]}}, - {"retriever_one": {"docs": ["doc1", "doc2"]}}, - {"__interrupt__": ()}, - ] - - app_w_interrupt.update_state(config, {"docs": ["doc5"]}) - expected_parent_config = list(app_w_interrupt.checkpointer.list(config, limit=2))[ - -1 - ].config - assert app_w_interrupt.get_state(config) == StateSnapshot( - values={ - "query": "analyzed: query: what is weather in sf", - "docs": ["doc1", "doc2", "doc3", "doc4", "doc5"], - }, - tasks=(PregelTask(AnyStr(), "qa", (PULL, "qa")),), - next=("qa",), - config={ - "configurable": { - "thread_id": "2", - "checkpoint_ns": "", - "checkpoint_id": AnyStr(), - } - }, - created_at=AnyStr(), - metadata={ - "parents": {}, - "source": "update", - "step": 4, - "thread_id": "2", - }, - parent_config=expected_parent_config, - interrupts=(), - ) - - assert [c for c in app_w_interrupt.stream(None, config, debug=1)] == [ - {"qa": {"answer": "doc1,doc2,doc3,doc4,doc5"}}, - ] - - def test_in_one_fan_out_state_graph_waiting_edge_via_branch( snapshot: SnapshotAssertion, sync_checkpointer: BaseCheckpointSaver ) -> None: @@ -3389,7 +3120,6 @@ def test_nested_graph_xray(snapshot: SnapshotAssertion) -> None: tool_two_graph.set_conditional_entry_point( lambda s: "tool_two_slow" if s["market"] == "DE" else "tool_two_fast", ["tool_two_slow", "tool_two_fast"], - then=END, ) tool_two = tool_two_graph.compile() @@ -3398,7 +3128,7 @@ def test_nested_graph_xray(snapshot: SnapshotAssertion) -> None: graph.add_node("tool_two", tool_two) graph.add_node("tool_three", logic) graph.set_conditional_entry_point( - lambda s: "tool_one", ["tool_one", "tool_two", "tool_three"], then=END + lambda s: "tool_one", ["tool_one", "tool_two", "tool_three"] ) app = graph.compile() diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index a6de45d1f..f2b7668fa 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -483,7 +483,7 @@ async def test_node_cancellation_on_other_node_exception() -> None: builder = StateGraph(State) builder.add_node("agent", awhile) builder.add_node("bad", iambad) - builder.set_conditional_entry_point(lambda _: ["agent", "bad"], then=END) + builder.set_conditional_entry_point(lambda _: ["agent", "bad"]) graph = builder.compile() @@ -507,7 +507,7 @@ async def test_node_cancellation_on_other_node_exception_two() -> None: builder = StateGraph(State) builder.add_node("agent", awhile) builder.add_node("bad", iambad) - builder.set_conditional_entry_point(lambda _: ["agent", "bad"], then=END) + builder.set_conditional_entry_point(lambda _: ["agent", "bad"]) graph = builder.compile() @@ -1076,7 +1076,7 @@ async def test_node_not_cancelled_on_other_node_interrupted( builder = StateGraph(State) builder.add_node("agent", awhile) builder.add_node("bad", iambad) - builder.set_conditional_entry_point(lambda _: ["agent", "bad"], then=END) + builder.set_conditional_entry_point(lambda _: ["agent", "bad"]) graph = builder.compile(checkpointer=async_checkpointer) thread = {"configurable": {"thread_id": "1"}} @@ -1141,7 +1141,7 @@ async def test_step_timeout_on_stream_hang(stream_hang_s: float) -> None: builder = StateGraph(State) builder.add_node(awhile) builder.add_node(alittlewhile) - builder.set_conditional_entry_point(lambda _: ["awhile", "alittlewhile"], then=END) + builder.set_conditional_entry_point(lambda _: ["awhile", "alittlewhile"]) graph = builder.compile() graph.step_timeout = 1 @@ -4042,97 +4042,6 @@ async def test_in_one_fan_out_state_graph_defer_node( ] -@pytest.mark.parametrize("with_path_map", (True, False)) -async def test_in_one_fan_out_state_graph_then_defer_node( - async_checkpointer: BaseCheckpointSaver, with_path_map: bool -) -> None: - def sorted_add( - x: list[str], y: Union[list[str], list[tuple[str, str]]] - ) -> list[str]: - if isinstance(y[0], tuple): - for rem, _ in y: - x.remove(rem) - y = [t[1] for t in y] - return sorted(operator.add(x, y)) - - class State(TypedDict, total=False): - query: str - answer: str - docs: Annotated[list[str], sorted_add] - - async def rewrite_query(data: State) -> State: - return {"query": f"query: {data['query']}"} - - async def analyzer_one(data: State) -> State: - return {"query": f"analyzed: {data['query']}"} - - async def retriever_one(data: State) -> State: - return {"docs": ["doc1", "doc2"]} - - async def retriever_two(data: State) -> State: - await asyncio.sleep(0.1) - return {"docs": ["doc3", "doc4"]} - - async def qa(data: State) -> State: - return {"answer": ",".join(data["docs"])} - - workflow = StateGraph(State) - - workflow.add_node("rewrite_query", rewrite_query) - workflow.add_node("analyzer_one", analyzer_one) - workflow.add_node("retriever_one", retriever_one) - workflow.add_node("retriever_two", retriever_two) - workflow.add_node("qa", qa, defer=True) - - workflow.set_entry_point("rewrite_query") - workflow.add_conditional_edges( - "rewrite_query", - lambda _: ["analyzer_one", "retriever_two"], - ["analyzer_one", "retriever_two"] if with_path_map else None, - then="qa", - ) - workflow.add_edge("analyzer_one", "retriever_one") - - app = workflow.compile() - - assert await app.ainvoke({"query": "what is weather in sf"}, debug=True) == { - "query": "analyzed: query: what is weather in sf", - "docs": ["doc1", "doc2", "doc3", "doc4"], - "answer": "doc1,doc2,doc3,doc4", - } - - assert [c async for c in app.astream({"query": "what is weather in sf"})] == [ - {"rewrite_query": {"query": "query: what is weather in sf"}}, - {"analyzer_one": {"query": "analyzed: query: what is weather in sf"}}, - {"retriever_two": {"docs": ["doc3", "doc4"]}}, - {"retriever_one": {"docs": ["doc1", "doc2"]}}, - {"qa": {"answer": "doc1,doc2,doc3,doc4"}}, - ] - - app_w_interrupt = workflow.compile( - checkpointer=async_checkpointer, - interrupt_after=["retriever_one"], - ) - config = {"configurable": {"thread_id": "1"}} - - assert [ - c - async for c in app_w_interrupt.astream( - {"query": "what is weather in sf"}, config - ) - ] == [ - {"rewrite_query": {"query": "query: what is weather in sf"}}, - {"analyzer_one": {"query": "analyzed: query: what is weather in sf"}}, - {"retriever_two": {"docs": ["doc3", "doc4"]}}, - {"retriever_one": {"docs": ["doc1", "doc2"]}}, - {"__interrupt__": ()}, - ] - - assert [c async for c in app_w_interrupt.astream(None, config)] == [ - {"qa": {"answer": "doc1,doc2,doc3,doc4"}}, - ] - - async def test_in_one_fan_out_state_graph_waiting_edge_via_branch( async_checkpointer: BaseCheckpointSaver, ) -> None: