pass subgraph nodes/channels

This commit is contained in:
vbarda
2024-08-12 20:29:09 -04:00
parent abe9b7c08e
commit f65d9b2b7d
3 changed files with 85 additions and 20 deletions
+75 -10
View File
@@ -360,11 +360,41 @@ class Pregel(
if is_managed_value(v)
}
def _prepare_state_snapshot(self, saved: CheckpointTuple) -> StateSnapshot:
def _get_nodes_and_channels(
self, checkpoint_ns: str
) -> tuple[Mapping[str, PregelNode], Mapping[str, BaseChannel]]:
if checkpoint_ns == "":
return self.nodes, self.channels
path = checkpoint_ns.split(CHECKPOINT_NAMESPACE_SEPARATOR)
nodes = self.nodes
channels = self.channels
for subgraph_node_name in path:
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
def _prepare_state_snapshot(
self,
saved: CheckpointTuple,
nodes: Mapping[str, PregelNode],
channels: Mapping[str, BaseChannel],
) -> StateSnapshot:
with ChannelsManager(
{
k: LastValue(None) if isinstance(c, Context) else c
for k, c in self.channels.items()
for k, c in channels.items()
},
saved.checkpoint,
saved.config,
@@ -373,7 +403,7 @@ class Pregel(
) as managed:
next_tasks = prepare_next_tasks(
saved.checkpoint,
self.nodes,
nodes,
channels,
managed,
saved.config,
@@ -390,12 +420,15 @@ class Pregel(
)
async def _prepare_state_snapshot_async(
self, saved: CheckpointTuple
self,
saved: CheckpointTuple,
nodes: Mapping[str, PregelNode],
channels: Mapping[str, BaseChannel],
) -> StateSnapshot:
async with AsyncChannelsManager(
{
k: LastValue(None) if isinstance(c, Context) else c
for k, c in self.channels.items()
for k, c in channels.items()
},
saved.checkpoint,
saved.config,
@@ -404,7 +437,7 @@ class Pregel(
) as managed:
next_tasks = prepare_next_tasks(
saved.checkpoint,
self.nodes,
nodes,
channels,
managed,
saved.config,
@@ -472,6 +505,9 @@ class Pregel(
checkpoint_id = 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]]
] = {}
for checkpoint_tuple in checkpoint_tuples:
saved_checkpoint_ns = checkpoint_tuple.config["configurable"][
"checkpoint_ns"
@@ -490,7 +526,17 @@ class Pregel(
existing_checkpoint_id is None
or saved_checkpoint_id > existing_checkpoint_id
):
state_snapshot = self._prepare_state_snapshot(checkpoint_tuple)
if saved_checkpoint_ns not in checkpoint_ns_to_nodes_and_channels:
checkpoint_ns_to_nodes_and_channels[
saved_checkpoint_ns
] = self._get_nodes_and_channels(saved_checkpoint_ns)
nodes, channels = checkpoint_ns_to_nodes_and_channels[
saved_checkpoint_ns
]
state_snapshot = self._prepare_state_snapshot(
checkpoint_tuple, nodes, channels
)
checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot
checkpoint_ns_to_checkpoint_id[
saved_checkpoint_ns
@@ -526,6 +572,9 @@ class Pregel(
checkpoint_id = 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]]
] = {}
async for checkpoint_tuple in checkpoint_tuples:
saved_checkpoint_ns = checkpoint_tuple.config["configurable"][
"checkpoint_ns"
@@ -544,8 +593,16 @@ 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
] = self._get_nodes_and_channels(saved_checkpoint_ns)
nodes, channels = checkpoint_ns_to_nodes_and_channels[
saved_checkpoint_ns
]
state_snapshot = await self._prepare_state_snapshot_async(
checkpoint_tuple
checkpoint_tuple, nodes, channels
)
checkpoint_ns_to_state_snapshots[saved_checkpoint_ns] = state_snapshot
checkpoint_ns_to_checkpoint_id[
@@ -597,7 +654,10 @@ class Pregel(
)
yield state_snapshot
else:
yield self._prepare_state_snapshot(checkpoint_tuple)
nodes, channels = self._get_nodes_and_channels(
checkpoint_tuple.config["configurable"]["checkpoint_ns"]
)
yield self._prepare_state_snapshot(checkpoint_tuple, nodes, channels)
async def aget_state_history(
self,
@@ -634,7 +694,12 @@ class Pregel(
)
yield state_snapshot
else:
yield await self._prepare_state_snapshot_async(checkpoint_tuple)
nodes, channels = self._get_nodes_and_channels(
checkpoint_tuple.config["configurable"]["checkpoint_ns"]
)
yield await self._prepare_state_snapshot_async(
checkpoint_tuple, nodes, channels
)
def update_state(
self,
+5 -5
View File
@@ -8604,7 +8604,7 @@ def test_nested_graph_interrupts(
assert child_state_history == [
StateSnapshot(
values={"my_key": "hi my value here"},
next=(),
next=("inner_2",),
config={
"configurable": {
"thread_id": "6",
@@ -9218,7 +9218,7 @@ def test_nested_graph_state(
subgraph_state_snapshots={
"inner": StateSnapshot(
values={"my_key": "hi my value here"},
next=(),
next=("inner_2",),
config={
"configurable": {
"thread_id": "1",
@@ -9275,7 +9275,7 @@ def test_nested_graph_state(
subgraph_state_snapshots={
"inner": StateSnapshot(
values={"my_key": "hi my value here"},
next=(),
next=("inner_2",),
config={
"configurable": {
"thread_id": "1",
@@ -9676,7 +9676,7 @@ def test_doubly_nested_graph_state(
subgraph_state_snapshots={
"child": StateSnapshot(
values={"my_key": "hi my value"},
next=(),
next=("child_1",),
config={
"configurable": {
"thread_id": "1",
@@ -9696,7 +9696,7 @@ def test_doubly_nested_graph_state(
subgraph_state_snapshots={
"child_1": StateSnapshot(
values={"my_key": "hi my value here"},
next=(),
next=("grandchild_2",),
config={
"configurable": {
"thread_id": "1",
+5 -5
View File
@@ -7106,7 +7106,7 @@ async def test_nested_graph_interrupts(
assert child_state_history == [
StateSnapshot(
values={"my_key": "hi my value here"},
next=(),
next=("inner_2",),
config={
"configurable": {
"thread_id": "6",
@@ -7725,7 +7725,7 @@ async def test_nested_graph_state(
subgraph_state_snapshots={
"inner": StateSnapshot(
values={"my_key": "hi my value here"},
next=(),
next=("inner_2",),
config={
"configurable": {
"thread_id": "1",
@@ -7784,7 +7784,7 @@ async def test_nested_graph_state(
subgraph_state_snapshots={
"inner": StateSnapshot(
values={"my_key": "hi my value here"},
next=(),
next=("inner_2",),
config={
"configurable": {
"thread_id": "1",
@@ -8187,7 +8187,7 @@ async def test_doubly_nested_graph_state(
subgraph_state_snapshots={
"child": StateSnapshot(
values={"my_key": "hi my value"},
next=(),
next=("child_1",),
config={
"configurable": {
"thread_id": "1",
@@ -8207,7 +8207,7 @@ async def test_doubly_nested_graph_state(
subgraph_state_snapshots={
"child_1": StateSnapshot(
values={"my_key": "hi my value here"},
next=(),
next=("grandchild_2",),
config={
"configurable": {
"thread_id": "1",