correctly propagate all subgraph attributes

This commit is contained in:
vbarda
2024-08-14 16:15:08 -04:00
parent 6531ec7669
commit 7fa97898aa
3 changed files with 88 additions and 176 deletions
+40 -78
View File
@@ -190,15 +190,12 @@ class Channel:
)
def _get_nodes_and_channels(
nodes: Mapping[str, PregelNode],
channels: Mapping[str, BaseChannel],
checkpoint_ns: str,
) -> tuple[Mapping[str, PregelNode], Mapping[str, BaseChannel]]:
def _get_subgraph(graph: Pregel, checkpoint_ns: str) -> Pregel:
if checkpoint_ns == "":
return nodes, channels
return graph
path = checkpoint_ns.split(CHECKPOINT_NAMESPACE_SEPARATOR)
nodes = graph.nodes
for subgraph_node_name in path:
# if we have this separator it means we have a node that was triggered by Send
if SEND_CHECKPOINT_NAMESPACE_SEPARATOR in subgraph_node_name:
@@ -210,17 +207,17 @@ def _get_nodes_and_channels(
if subgraph_node_name not in nodes:
raise ValueError(f"Couldn't find node '{subgraph_node_name}'.")
subgraph_node = nodes[subgraph_node_name].get_node()
if not isinstance(subgraph_node, RunnableSequence):
break
first_step = subgraph_node.steps[0]
if isinstance(first_step, Pregel):
nodes = first_step.nodes
channels = first_step.channels
return nodes, channels
subgraph_node = nodes[subgraph_node_name]
if isinstance(subgraph_node.bound, Pregel):
nodes = subgraph_node.bound.nodes
elif isinstance(subgraph_node.bound, RunnableSequence):
for runnable in subgraph_node.bound.steps:
if isinstance(runnable, Pregel):
nodes = runnable.nodes
break
else:
continue
return subgraph_node.bound
def _assemble_state_snapshot_hierarchy(
@@ -259,24 +256,21 @@ def _assemble_state_snapshot_hierarchy(
def _prepare_state_snapshot(
saved: CheckpointTuple,
nodes: Mapping[str, PregelNode],
channels: Mapping[str, BaseChannel],
managed_values_dict: dict[str, ManagedValueSpec],
select_channels: str | list[str],
graph: Pregel,
) -> StateSnapshot:
with ChannelsManager(
{
k: LastValue(None) if isinstance(c, Context) else c
for k, c in channels.items()
for k, c in graph.channels.items()
},
saved.checkpoint,
saved.config,
) as channels, ManagedValuesManager(
managed_values_dict, ensure_config(saved.config)
graph.managed_values_dict, ensure_config(saved.config)
) as managed:
next_tasks = prepare_next_tasks(
saved.checkpoint,
nodes,
graph.nodes,
channels,
managed,
saved.config,
@@ -284,7 +278,7 @@ def _prepare_state_snapshot(
for_execution=False,
)
return StateSnapshot(
values=read_channels(channels, select_channels),
values=read_channels(channels, graph.stream_channels_asis),
next=tuple(t.name for t in next_tasks),
config=saved.config,
metadata=saved.metadata,
@@ -294,25 +288,21 @@ def _prepare_state_snapshot(
async def _prepare_state_snapshot_async(
saved: CheckpointTuple,
nodes: Mapping[str, PregelNode],
channels: Mapping[str, BaseChannel],
managed_values_dict: dict[str, ManagedValueSpec],
select_channels: str | list[str],
saved: CheckpointTuple, graph: Pregel
) -> StateSnapshot:
async with AsyncChannelsManager(
{
k: LastValue(None) if isinstance(c, Context) else c
for k, c in channels.items()
for k, c in graph.channels.items()
},
saved.checkpoint,
saved.config,
) as channels, AsyncManagedValuesManager(
managed_values_dict, ensure_config(saved.config)
graph.managed_values_dict, ensure_config(saved.config)
) as managed:
next_tasks = prepare_next_tasks(
saved.checkpoint,
nodes,
graph.nodes,
channels,
managed,
saved.config,
@@ -320,7 +310,7 @@ async def _prepare_state_snapshot_async(
for_execution=False,
)
return StateSnapshot(
values=read_channels(channels, select_channels),
values=read_channels(channels, graph.stream_channels_asis),
next=tuple(t.name for t in next_tasks),
config=saved.config,
metadata=saved.metadata,
@@ -534,9 +524,7 @@ class Pregel(
checkpoint_id = checkpoint_config["configurable"].get("checkpoint_id")
checkpoint_ns_to_checkpoint_id: dict[str, str] = {}
checkpoint_ns_to_state_snapshots: dict[str, StateSnapshot] = {}
checkpoint_ns_to_nodes_and_channels: dict[
str, tuple[Mapping[str, PregelNode], Mapping[str, BaseChannel]]
] = {}
checkpoint_ns_to_graph: dict[str, Pregel] = {}
for checkpoint_tuple in checkpoint_tuples:
saved_checkpoint_ns = checkpoint_tuple.config["configurable"][
"checkpoint_ns"
@@ -555,22 +543,14 @@ class Pregel(
existing_checkpoint_id is None
or saved_checkpoint_id > existing_checkpoint_id
):
if saved_checkpoint_ns not in checkpoint_ns_to_nodes_and_channels:
checkpoint_ns_to_nodes_and_channels[
saved_checkpoint_ns
] = _get_nodes_and_channels(
self.nodes, self.channels, saved_checkpoint_ns
if saved_checkpoint_ns not in checkpoint_ns_to_graph:
checkpoint_ns_to_graph[saved_checkpoint_ns] = _get_subgraph(
self, saved_checkpoint_ns
)
nodes, channels = checkpoint_ns_to_nodes_and_channels[
saved_checkpoint_ns
]
state_snapshot = _prepare_state_snapshot(
checkpoint_tuple,
nodes,
channels,
self.managed_values_dict,
self.stream_channels_asis,
checkpoint_ns_to_graph[saved_checkpoint_ns],
)
checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot
checkpoint_ns_to_checkpoint_id[
@@ -610,9 +590,7 @@ class Pregel(
checkpoint_id = checkpoint_config["configurable"].get("checkpoint_id")
checkpoint_ns_to_checkpoint_id: dict[str, str] = {}
checkpoint_ns_to_state_snapshots: dict[str, StateSnapshot] = {}
checkpoint_ns_to_nodes_and_channels: dict[
str, tuple[Mapping[str, PregelNode], Mapping[str, BaseChannel]]
] = {}
checkpoint_ns_to_graph: dict[str, Pregel] = {}
async for checkpoint_tuple in checkpoint_tuples:
saved_checkpoint_ns = checkpoint_tuple.config["configurable"][
"checkpoint_ns"
@@ -631,22 +609,14 @@ class Pregel(
existing_checkpoint_id is None
or saved_checkpoint_id > existing_checkpoint_id
):
if saved_checkpoint_ns not in checkpoint_ns_to_nodes_and_channels:
checkpoint_ns_to_nodes_and_channels[
saved_checkpoint_ns
] = _get_nodes_and_channels(
self.nodes, self.channels, saved_checkpoint_ns
if saved_checkpoint_ns not in checkpoint_ns_to_graph:
checkpoint_ns_to_graph[saved_checkpoint_ns] = _get_subgraph(
self, saved_checkpoint_ns
)
nodes, channels = checkpoint_ns_to_nodes_and_channels[
saved_checkpoint_ns
]
state_snapshot = await _prepare_state_snapshot_async(
checkpoint_tuple,
nodes,
channels,
self.managed_values_dict,
self.stream_channels_asis,
checkpoint_ns_to_graph[saved_checkpoint_ns],
)
checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot
checkpoint_ns_to_checkpoint_id[
@@ -698,17 +668,13 @@ class Pregel(
)
yield state_snapshot
else:
nodes, channels = _get_nodes_and_channels(
self.nodes,
self.channels,
graph = _get_subgraph(
self,
checkpoint_tuple.config["configurable"]["checkpoint_ns"],
)
yield _prepare_state_snapshot(
checkpoint_tuple,
nodes,
channels,
self.managed_values_dict,
self.stream_channels_asis,
graph,
)
async def aget_state_history(
@@ -746,17 +712,13 @@ class Pregel(
)
yield state_snapshot
else:
nodes, channels = _get_nodes_and_channels(
self.nodes,
self.channels,
graph = _get_subgraph(
self,
checkpoint_tuple.config["configurable"]["checkpoint_ns"],
)
yield await _prepare_state_snapshot_async(
checkpoint_tuple,
nodes,
channels,
self.managed_values_dict,
self.stream_channels_asis,
graph,
)
def update_state(
+23 -49
View File
@@ -8602,7 +8602,7 @@ def test_nested_graph_interrupts(
]
assert child_state_history == [
StateSnapshot(
values={"my_key": "hi my value here"},
values={"my_key": "hi my value here", "my_other_key": "hi my value"},
next=("inner_2",),
config={
"configurable": {
@@ -9140,6 +9140,7 @@ def test_nested_graph_state(
class State(TypedDict):
my_key: str
other_parent_key: str
def outer_1(state: State):
return {"my_key": "hi " + state["my_key"]}
@@ -9214,7 +9215,7 @@ def test_nested_graph_state(
},
subgraph_state_snapshots={
"inner": StateSnapshot(
values={"my_key": "hi my value here"},
values={"my_key": "hi my value here", "my_other_key": "hi my value"},
next=("inner_2",),
config={
"configurable": {
@@ -9271,7 +9272,10 @@ def test_nested_graph_state(
},
subgraph_state_snapshots={
"inner": StateSnapshot(
values={"my_key": "hi my value here"},
values={
"my_key": "hi my value here",
"my_other_key": "hi my value",
},
next=("inner_2",),
config={
"configurable": {
@@ -9376,7 +9380,10 @@ def test_nested_graph_state(
{"configurable": {"thread_id": "1", "checkpoint_ns": "inner"}}
)
assert child_snapshot == StateSnapshot(
values={"my_key": "hi my value here and there"},
values={
"my_key": "hi my value here and there",
"my_other_key": "hi my value here",
},
next=(),
config={
"configurable": {
@@ -9519,7 +9526,10 @@ def test_nested_graph_state(
},
subgraph_state_snapshots={
"inner": StateSnapshot(
values={"my_key": "hi my value here and there"},
values={
"my_key": "hi my value here and there",
"my_other_key": "hi my value here",
},
next=(),
config={
"configurable": {
@@ -9931,6 +9941,13 @@ def test_send_to_nested_graphs(
for subgraph_node in subgraph_nodes:
assert subgraph_node.split(":")[0] == "generate_joke"
subgraph_state_snapshots = {
subgraph_node: graph.get_state(
{"configurable": {"thread_id": "1", "checkpoint_ns": subgraph_node}}
)
for subgraph_node in subgraph_nodes
}
expected_snapshot = StateSnapshot(
values={"subjects": ["cats", "dogs"], "jokes": []},
next=("generate_joke", "generate_joke"),
@@ -9950,50 +9967,7 @@ def test_send_to_nested_graphs(
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots={
subgraph_nodes[0]: StateSnapshot(
values={"jokes": []},
next=("generate",),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": subgraph_nodes[0],
"checkpoint_id": AnyStr(),
}
},
metadata={"source": "loop", "writes": {"edit": None}, "step": 1},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": subgraph_nodes[0],
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots=None,
),
subgraph_nodes[1]: StateSnapshot(
values={"jokes": []},
next=("generate",),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": subgraph_nodes[1],
"checkpoint_id": AnyStr(),
}
},
metadata={"source": "loop", "writes": {"edit": None}, "step": 1},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": subgraph_nodes[1],
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots=None,
),
},
subgraph_state_snapshots=subgraph_state_snapshots,
)
assert actual_snapshot == expected_snapshot
+25 -49
View File
@@ -7104,7 +7104,7 @@ async def test_nested_graph_interrupts(
]
assert child_state_history == [
StateSnapshot(
values={"my_key": "hi my value here"},
values={"my_key": "hi my value here", "my_other_key": "hi my value"},
next=("inner_2",),
config={
"configurable": {
@@ -7647,6 +7647,7 @@ async def test_nested_graph_state(
class State(TypedDict):
my_key: str
other_parent_key: str
async def outer_1(state: State):
return {"my_key": "hi " + state["my_key"]}
@@ -7721,7 +7722,10 @@ async def test_nested_graph_state(
},
subgraph_state_snapshots={
"inner": StateSnapshot(
values={"my_key": "hi my value here"},
values={
"my_key": "hi my value here",
"my_other_key": "hi my value",
},
next=("inner_2",),
config={
"configurable": {
@@ -7780,7 +7784,10 @@ async def test_nested_graph_state(
},
subgraph_state_snapshots={
"inner": StateSnapshot(
values={"my_key": "hi my value here"},
values={
"my_key": "hi my value here",
"my_other_key": "hi my value",
},
next=("inner_2",),
config={
"configurable": {
@@ -7885,7 +7892,10 @@ async def test_nested_graph_state(
{"configurable": {"thread_id": "1", "checkpoint_ns": "inner"}}
)
assert child_snapshot == StateSnapshot(
values={"my_key": "hi my value here and there"},
values={
"my_key": "hi my value here and there",
"my_other_key": "hi my value here",
},
next=(),
config={
"configurable": {
@@ -8030,7 +8040,10 @@ async def test_nested_graph_state(
},
subgraph_state_snapshots={
"inner": StateSnapshot(
values={"my_key": "hi my value here and there"},
values={
"my_key": "hi my value here and there",
"my_other_key": "hi my value here",
},
next=(),
config={
"configurable": {
@@ -8442,6 +8455,12 @@ async def test_send_to_nested_graphs(
for subgraph_node in subgraph_nodes:
assert subgraph_node.split(":")[0] == "generate_joke"
subgraph_state_snapshots = {
subgraph_node: await graph.aget_state(
{"configurable": {"thread_id": "1", "checkpoint_ns": subgraph_node}}
)
for subgraph_node in subgraph_nodes
}
expected_snapshot = StateSnapshot(
values={"subjects": ["cats", "dogs"], "jokes": []},
next=("generate_joke", "generate_joke"),
@@ -8461,50 +8480,7 @@ async def test_send_to_nested_graphs(
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots={
subgraph_nodes[0]: StateSnapshot(
values={"jokes": []},
next=("generate",),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": subgraph_nodes[0],
"checkpoint_id": AnyStr(),
}
},
metadata={"source": "loop", "writes": {"edit": None}, "step": 1},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": subgraph_nodes[0],
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots=None,
),
subgraph_nodes[1]: StateSnapshot(
values={"jokes": []},
next=("generate",),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": subgraph_nodes[1],
"checkpoint_id": AnyStr(),
}
},
metadata={"source": "loop", "writes": {"edit": None}, "step": 1},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": subgraph_nodes[1],
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots=None,
),
},
subgraph_state_snapshots=subgraph_state_snapshots,
)
assert actual_snapshot == expected_snapshot