update logic for latest snapshot's subgraph snapshots

This commit is contained in:
vbarda
2024-08-13 20:13:13 -04:00
parent 392891f5fc
commit d9618880a3
3 changed files with 351 additions and 418 deletions
+10 -8
View File
@@ -506,14 +506,15 @@ class Pregel(
if not self.checkpointer:
raise ValueError("No checkpointer set")
checkpoint_tuple = self.checkpointer.get_tuple(config)
if include_subgraph_state:
checkpoint_tuples = self.checkpointer.list(config)
else:
checkpoint_tuple = self.checkpointer.get_tuple(config)
checkpoint_tuples = iter([checkpoint_tuple] if checkpoint_tuple else [])
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
checkpoint_id = config["configurable"].get("checkpoint_id")
checkpoint_config = checkpoint_tuple.config if checkpoint_tuple else config
checkpoint_ns = checkpoint_config["configurable"].get("checkpoint_ns", "")
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[
@@ -526,7 +527,7 @@ class Pregel(
saved_checkpoint_id = checkpoint_tuple.config["configurable"][
"checkpoint_id"
]
if checkpoint_id and checkpoint_id != saved_checkpoint_id:
if checkpoint_id != saved_checkpoint_id:
continue
existing_checkpoint_id = checkpoint_ns_to_checkpoint_id.get(
@@ -570,19 +571,20 @@ class Pregel(
if not self.checkpointer:
raise ValueError("No checkpointer set")
checkpoint_tuple = await self.checkpointer.aget_tuple(config)
if include_subgraph_state:
checkpoint_tuples = self.checkpointer.alist(config)
else:
async def alist_checkpoints():
checkpoint_tuple = await self.checkpointer.aget_tuple(config)
if checkpoint_tuple:
yield checkpoint_tuple
checkpoint_tuples = alist_checkpoints()
checkpoint_ns = config["configurable"].get("checkpoint_ns", "")
checkpoint_id = config["configurable"].get("checkpoint_id")
checkpoint_config = checkpoint_tuple.config if checkpoint_tuple else config
checkpoint_ns = checkpoint_config["configurable"].get("checkpoint_ns", "")
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[
@@ -595,7 +597,7 @@ class Pregel(
saved_checkpoint_id = checkpoint_tuple.config["configurable"][
"checkpoint_id"
]
if checkpoint_id and checkpoint_id != saved_checkpoint_id:
if checkpoint_id != saved_checkpoint_id:
continue
existing_checkpoint_id = checkpoint_ns_to_checkpoint_id.get(
+171 -205
View File
@@ -9373,39 +9373,77 @@ def test_nested_graph_state(
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots={
"inner": StateSnapshot(
values={"my_key": "hi my value here and there"},
next=(),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "inner",
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {
"inner_2": {
"my_key": "hi my value here and there",
"my_other_key": "hi my value here",
}
},
"step": 2,
},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "inner",
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots=None,
)
},
)
# test loading inner snapshot
child_snapshot = app.get_state(
{"configurable": {"thread_id": "1", "checkpoint_ns": "inner"}}
)
assert child_snapshot == StateSnapshot(
values={"my_key": "hi my value here and there"},
next=(),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "inner",
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {
"inner_2": {
"my_key": "hi my value here and there",
"my_other_key": "hi my value here",
}
},
"step": 2,
},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "inner",
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots=None,
)
# test looking up parent state by checkpoint ID
assert app.get_state(
{
"configurable": {
"thread_id": "1",
"checkpoint_ns": "",
"checkpoint_id": child_snapshot.config["configurable"]["checkpoint_id"],
}
},
include_subgraph_state=True,
) == StateSnapshot(
values={"my_key": "hi my value"},
next=("inner",),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "",
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {"outer_1": {"my_key": "hi my value"}},
"step": 1,
},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "",
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots={"inner": child_snapshot},
)
# test full history at the end
assert list(app.get_state_history(config, include_subgraph_state=True)) == [
StateSnapshot(
values={"my_key": "hi my value here and there and back again"},
@@ -9482,10 +9520,6 @@ def test_nested_graph_state(
"checkpoint_id": AnyStr(),
}
},
# TODO: this is likely very confusing for an end user, and we'll probably need to update this.
# right now this is happening due to us overwriting the
# subgraph snapshot after we finish the graph with while the checkpoint_id
# is the same as when we interrupted
subgraph_state_snapshots={
"inner": StateSnapshot(
values={"my_key": "hi my value here and there"},
@@ -9749,61 +9783,99 @@ def test_doubly_nested_graph_state(
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots={
"child": StateSnapshot(
values={"my_key": "hi my value here and there"},
next=(),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "child",
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {"child_1": {"my_key": "hi my value here and there"}},
"step": 1,
},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "child",
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots={
"child_1": StateSnapshot(
values={"my_key": "hi my value here and there"},
next=(),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "child|child_1",
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {
"grandchild_2": {"my_key": "hi my value here and there"}
},
"step": 2,
},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "child|child_1",
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots=None,
)
},
)
)
# test getting grandchild snapshot
grandchild_snapshot = app.get_state(
{"configurable": {"thread_id": "1", "checkpoint_ns": "child|child_1"}}
)
assert grandchild_snapshot == StateSnapshot(
values={"my_key": "hi my value here and there"},
next=(),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "child|child_1",
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {"grandchild_2": {"my_key": "hi my value here and there"}},
"step": 2,
},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "child|child_1",
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots=None,
)
# test getting child snapshot
child_snapshot = app.get_state(
{"configurable": {"thread_id": "1", "checkpoint_ns": "child"}},
include_subgraph_state=True,
)
assert child_snapshot == StateSnapshot(
values={"my_key": "hi my value here and there"},
next=(),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "child",
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {"child_1": {"my_key": "hi my value here and there"}},
"step": 1,
},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "child",
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots={"child_1": grandchild_snapshot},
)
# test getting parent snapshot for a checkpoint ID
assert app.get_state(
{
"configurable": {
"thread_id": "1",
"checkpoint_id": child_snapshot.config["configurable"]["checkpoint_id"],
}
},
include_subgraph_state=True,
) == StateSnapshot(
values={"my_key": "hi my value"},
next=("child",),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "",
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {"parent_1": {"my_key": "hi my value"}},
"step": 1,
},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "",
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots={"child": child_snapshot},
)
@@ -9939,17 +10011,6 @@ def test_send_to_nested_graphs(
}
actual_snapshot = graph.get_state(config, include_subgraph_state=True)
subgraph_nodes, _ = zip(
*(
sorted(
actual_snapshot.subgraph_state_snapshots.items(),
key=lambda x: x[1].values["jokes"][0],
)
)
)
assert len(subgraph_nodes) == 2
for subgraph_node in subgraph_nodes:
assert subgraph_node.split(":")[0] == "generate_joke"
expected_snapshot = StateSnapshot(
values={
"subjects": ["cats", "dogs"],
@@ -9981,63 +10042,19 @@ def test_send_to_nested_graphs(
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots={
subgraph_nodes[0]: StateSnapshot(
values={"jokes": ["Joke about cats - hohoho"]},
next=(),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": subgraph_nodes[0],
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {"generate": {"jokes": ["Joke about cats - hohoho"]}},
"step": 2,
},
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": ["Joke about dogs - hohoho"]},
next=(),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": subgraph_nodes[1],
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {"generate": {"jokes": ["Joke about dogs - hohoho"]}},
"step": 2,
},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": subgraph_nodes[1],
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots=None,
),
},
)
assert actual_snapshot == expected_snapshot
# test full history
actual_history = list(graph.get_state_history(config, include_subgraph_state=True))
# get subgraph node state for expected history
subgraph_state_snapshots = {
subgraph_node: graph.get_state(
{"configurable": {"thread_id": "1", "checkpoint_ns": subgraph_node}}
)
for subgraph_node in subgraph_nodes
}
expected_history = [
StateSnapshot(
values={
@@ -10091,58 +10108,7 @@ def test_send_to_nested_graphs(
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots={
subgraph_nodes[0]: StateSnapshot(
values={"jokes": ["Joke about cats - hohoho"]},
next=(),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": subgraph_nodes[0],
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {"generate": {"jokes": ["Joke about cats - hohoho"]}},
"step": 2,
},
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": ["Joke about dogs - hohoho"]},
next=(),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": subgraph_nodes[1],
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {"generate": {"jokes": ["Joke about dogs - hohoho"]}},
"step": 2,
},
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,
),
StateSnapshot(
values={"jokes": []},
+170 -205
View File
@@ -7882,39 +7882,77 @@ async def test_nested_graph_state(
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots={
"inner": StateSnapshot(
values={"my_key": "hi my value here and there"},
next=(),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "inner",
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {
"inner_2": {
"my_key": "hi my value here and there",
"my_other_key": "hi my value here",
}
},
"step": 2,
},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "inner",
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots=None,
)
},
)
# test loading inner snapshot
child_snapshot = await app.aget_state(
{"configurable": {"thread_id": "1", "checkpoint_ns": "inner"}}
)
assert child_snapshot == StateSnapshot(
values={"my_key": "hi my value here and there"},
next=(),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "inner",
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {
"inner_2": {
"my_key": "hi my value here and there",
"my_other_key": "hi my value here",
}
},
"step": 2,
},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "inner",
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots=None,
)
# test looking up parent state by checkpoint ID
assert await app.aget_state(
{
"configurable": {
"thread_id": "1",
"checkpoint_ns": "",
"checkpoint_id": child_snapshot.config["configurable"]["checkpoint_id"],
}
},
include_subgraph_state=True,
) == StateSnapshot(
values={"my_key": "hi my value"},
next=("inner",),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "",
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {"outer_1": {"my_key": "hi my value"}},
"step": 1,
},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "",
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots={"inner": child_snapshot},
)
# test full history at the end
assert [
s async for s in app.aget_state_history(config, include_subgraph_state=True)
] == [
@@ -7993,10 +8031,6 @@ async def test_nested_graph_state(
"checkpoint_id": AnyStr(),
}
},
# TODO: this is likely very confusing for an end user, and we'll probably need to update this.
# right now this is happening due to us overwriting the
# subgraph snapshot after we finish the graph with while the checkpoint_id
# is the same as when we interrupted
subgraph_state_snapshots={
"inner": StateSnapshot(
values={"my_key": "hi my value here and there"},
@@ -8260,61 +8294,99 @@ async def test_doubly_nested_graph_state(
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots={
"child": StateSnapshot(
values={"my_key": "hi my value here and there"},
next=(),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "child",
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {"child_1": {"my_key": "hi my value here and there"}},
"step": 1,
},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "child",
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots={
"child_1": StateSnapshot(
values={"my_key": "hi my value here and there"},
next=(),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "child|child_1",
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {
"grandchild_2": {"my_key": "hi my value here and there"}
},
"step": 2,
},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "child|child_1",
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots=None,
)
},
)
)
# test getting grandchild snapshot
grandchild_snapshot = await app.aget_state(
{"configurable": {"thread_id": "1", "checkpoint_ns": "child|child_1"}}
)
assert grandchild_snapshot == StateSnapshot(
values={"my_key": "hi my value here and there"},
next=(),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "child|child_1",
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {"grandchild_2": {"my_key": "hi my value here and there"}},
"step": 2,
},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "child|child_1",
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots=None,
)
# test getting child snapshot
child_snapshot = await app.aget_state(
{"configurable": {"thread_id": "1", "checkpoint_ns": "child"}},
include_subgraph_state=True,
)
assert child_snapshot == StateSnapshot(
values={"my_key": "hi my value here and there"},
next=(),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "child",
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {"child_1": {"my_key": "hi my value here and there"}},
"step": 1,
},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "child",
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots={"child_1": grandchild_snapshot},
)
# test getting parent snapshot for a checkpoint ID
assert await app.aget_state(
{
"configurable": {
"thread_id": "1",
"checkpoint_id": child_snapshot.config["configurable"]["checkpoint_id"],
}
},
include_subgraph_state=True,
) == StateSnapshot(
values={"my_key": "hi my value"},
next=("child",),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "",
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {"parent_1": {"my_key": "hi my value"}},
"step": 1,
},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": "",
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots={"child": child_snapshot},
)
@@ -8450,17 +8522,6 @@ async def test_send_to_nested_graphs(
}
actual_snapshot = await graph.aget_state(config, include_subgraph_state=True)
subgraph_nodes, _ = zip(
*(
sorted(
actual_snapshot.subgraph_state_snapshots.items(),
key=lambda x: x[1].values["jokes"][0],
)
)
)
assert len(subgraph_nodes) == 2
for subgraph_node in subgraph_nodes:
assert subgraph_node.split(":")[0] == "generate_joke"
expected_snapshot = StateSnapshot(
values={
"subjects": ["cats", "dogs"],
@@ -8492,58 +8553,6 @@ async def test_send_to_nested_graphs(
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots={
subgraph_nodes[0]: StateSnapshot(
values={"jokes": ["Joke about cats - hohoho"]},
next=(),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": subgraph_nodes[0],
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {"generate": {"jokes": ["Joke about cats - hohoho"]}},
"step": 2,
},
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": ["Joke about dogs - hohoho"]},
next=(),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": subgraph_nodes[1],
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {"generate": {"jokes": ["Joke about dogs - hohoho"]}},
"step": 2,
},
created_at=AnyStr(),
parent_config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": subgraph_nodes[1],
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots=None,
),
},
)
assert actual_snapshot == expected_snapshot
@@ -8551,6 +8560,13 @@ async def test_send_to_nested_graphs(
actual_history = [
c async for c in graph.aget_state_history(config, include_subgraph_state=True)
]
# get subgraph node state for expected history
subgraph_state_snapshots = {
subgraph_node: await graph.aget_state(
{"configurable": {"thread_id": "1", "checkpoint_ns": subgraph_node}}
)
for subgraph_node in subgraph_nodes
}
expected_history = [
StateSnapshot(
values={
@@ -8604,58 +8620,7 @@ async def test_send_to_nested_graphs(
"checkpoint_id": AnyStr(),
}
},
subgraph_state_snapshots={
subgraph_nodes[0]: StateSnapshot(
values={"jokes": ["Joke about cats - hohoho"]},
next=(),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": subgraph_nodes[0],
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {"generate": {"jokes": ["Joke about cats - hohoho"]}},
"step": 2,
},
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": ["Joke about dogs - hohoho"]},
next=(),
config={
"configurable": {
"thread_id": "1",
"checkpoint_ns": subgraph_nodes[1],
"checkpoint_id": AnyStr(),
}
},
metadata={
"source": "loop",
"writes": {"generate": {"jokes": ["Joke about dogs - hohoho"]}},
"step": 2,
},
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,
),
StateSnapshot(
values={"jokes": []},