Merge pull request #427 from langchain-ai/nc/9may/checkpoint-metadata-writes

Add writes property in checkpoint metadata
This commit is contained in:
Nuno Campos
2024-05-10 10:01:52 -07:00
committed by GitHub
4 changed files with 811 additions and 84 deletions
+5
View File
@@ -31,6 +31,11 @@ class CheckpointMetadata(TypedDict, total=False):
0 for the first "loop" checkpoint.
... for the nth checkpoint afterwards.
"""
writes: dict[str, Any]
"""The writes that were made between the previous checkpoint and this one.
Mapping from node name to writes emitted by that node.
"""
class Checkpoint(TypedDict):
+34 -4
View File
@@ -486,6 +486,7 @@ class Pregel(
{
"source": "update",
"step": saved.metadata.get("step", 0) + 1 if saved else 0,
"writes": {as_node: values},
},
)
@@ -559,6 +560,7 @@ class Pregel(
{
"source": "update",
"step": saved.metadata.get("step", 0) + 1 if saved else 0,
"writes": {as_node: values},
},
)
@@ -677,7 +679,7 @@ class Pregel(
self.checkpointer.put,
checkpoint_config,
copy_checkpoint(checkpoint),
{"source": "input", "step": start},
{"source": "input", "step": start, "writes": input},
)
)
checkpoint_config = {
@@ -806,7 +808,21 @@ class Pregel(
self.checkpointer.put,
checkpoint_config,
copy_checkpoint(checkpoint),
{"source": "loop", "step": step},
{
"source": "loop",
"step": step,
"writes": next(
map_output_updates(output_keys, next_tasks),
None,
)
if self.stream_mode == "updates"
else next(
map_output_values(
output_keys, pending_writes, channels
),
None,
),
},
)
)
checkpoint_config = {
@@ -943,7 +959,7 @@ class Pregel(
self.checkpointer.aput(
checkpoint_config,
copy_checkpoint(checkpoint),
{"source": "input", "step": start},
{"source": "input", "step": start, "writes": input},
)
)
)
@@ -1084,7 +1100,21 @@ class Pregel(
self.checkpointer.aput(
checkpoint_config,
checkpoint,
{"source": "loop", "step": step},
{
"source": "loop",
"step": step,
"writes": next(
map_output_updates(output_keys, next_tasks),
None,
)
if self.stream_mode == "updates"
else next(
map_output_values(
output_keys, pending_writes, channels
),
None,
),
},
)
)
)
+451 -45
View File
@@ -19,6 +19,7 @@ from langgraph.channels.topic import Topic
from langgraph.checkpoint.sqlite import SqliteSaver
from langgraph.errors import InvalidUpdateError
from langgraph.graph import END, Graph
from langgraph.graph.graph import START
from langgraph.graph.message import MessageGraph
from langgraph.graph.state import StateGraph
from langgraph.prebuilt.chat_agent_executor import (
@@ -1228,7 +1229,20 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None:
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "loop", "step": 0},
metadata={
"source": "loop",
"step": 0,
"writes": {
"agent": {
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:query",
),
},
},
},
)
assert (
app_w_interrupt.checkpointer.get_tuple(config).config["configurable"][
@@ -1262,7 +1276,20 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None:
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "update", "step": 1},
metadata={
"source": "update",
"step": 1,
"writes": {
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"input": "what is weather in sf",
},
},
},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -1346,7 +1373,29 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None:
},
next=(),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "update", "step": 4},
metadata={
"source": "update",
"step": 4,
"writes": {
"agent": {
"input": "what is weather in sf",
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
),
}
},
},
)
# test state get/update methods with interrupt_before
@@ -1382,7 +1431,20 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None:
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "loop", "step": 0},
metadata={
"source": "loop",
"step": 0,
"writes": {
"agent": {
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:query",
),
}
},
},
)
app_w_interrupt.update_state(
@@ -1410,7 +1472,20 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None:
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "update", "step": 1},
metadata={
"source": "update",
"step": 1,
"writes": {
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"input": "what is weather in sf",
}
},
},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -1494,7 +1569,29 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None:
},
next=(),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "update", "step": 4},
metadata={
"source": "update",
"step": 4,
"writes": {
"agent": {
"input": "what is weather in sf",
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
),
}
},
},
)
# test re-invoke to continue with interrupt_before
@@ -1530,7 +1627,20 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None:
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "loop", "step": 0},
metadata={
"source": "loop",
"step": 0,
"writes": {
"agent": {
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:query",
),
}
},
},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -1862,7 +1972,19 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:query",
),
}
},
},
)
app_w_interrupt.update_state(
@@ -1888,7 +2010,19 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "update", "step": 2},
metadata={
"source": "update",
"step": 2,
"writes": {
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
)
},
},
},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -1947,7 +2081,18 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
},
next=(),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "update", "step": 5},
metadata={
"source": "update",
"step": 5,
"writes": {
"agent": {
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
)
}
},
},
)
# test state get/update methods with interrupt_before
@@ -1982,7 +2127,19 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:query",
),
}
},
},
)
app_w_interrupt.update_state(
@@ -2008,7 +2165,19 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "update", "step": 2},
metadata={
"source": "update",
"step": 2,
"writes": {
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
)
}
},
},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -2067,7 +2236,18 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
},
next=(),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "update", "step": 5},
metadata={
"source": "update",
"step": 5,
"writes": {
"agent": {
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
)
}
},
},
)
# test w interrupt before all
@@ -2090,7 +2270,7 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
},
next=("agent",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "loop", "step": 0},
metadata={"source": "loop", "step": 0, "writes": None},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -2113,7 +2293,19 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:query",
),
}
},
},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -2152,7 +2344,24 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
},
next=("agent",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "loop", "step": 2},
metadata={
"source": "loop",
"step": 2,
"writes": {
"tools": {
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:query",
),
"result for query",
)
],
}
},
},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -2197,7 +2406,19 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
},
next=("tools",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:query",
),
}
},
},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -2236,7 +2457,24 @@ def test_conditional_state_graph(snapshot: SnapshotAssertion) -> None:
},
next=("agent",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "loop", "step": 2},
metadata={
"source": "loop",
"step": 2,
"writes": {
"tools": {
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:query",
),
"result for query",
)
],
}
},
},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -3051,7 +3289,23 @@ def test_message_graph(
],
next=("action",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {
"agent": AIMessage(
content="",
tool_calls=[
{
"id": "tool_call123",
"name": "search_api",
"args": {"query": "query"},
}
],
id="ai1",
)
},
},
)
# modify ai message
@@ -3077,7 +3331,23 @@ def test_message_graph(
],
next=("action",),
config=next_config,
metadata={"source": "update", "step": 2},
metadata={
"source": "update",
"step": 2,
"writes": {
"agent": AIMessage(
content="",
tool_calls=[
{
"id": "tool_call123",
"name": "search_api",
"args": {"query": "a different query"},
}
],
id="ai1",
)
},
},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -3143,7 +3413,23 @@ def test_message_graph(
],
next=("action",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "loop", "step": 4},
metadata={
"source": "loop",
"step": 4,
"writes": {
"agent": AIMessage(
content="",
tool_calls=[
{
"id": "tool_call456",
"name": "search_api",
"args": {"query": "another"},
}
],
id="ai2",
)
},
},
)
app_w_interrupt.update_state(
@@ -3179,7 +3465,11 @@ def test_message_graph(
],
next=(),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "update", "step": 5},
metadata={
"source": "update",
"step": 5,
"writes": {"agent": AIMessage(content="answer", id="ai2")},
},
)
app_w_interrupt = workflow.compile(
@@ -3225,7 +3515,23 @@ def test_message_graph(
],
next=("action",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {
"agent": AIMessage(
content="",
tool_calls=[
{
"id": "tool_call123",
"name": "search_api",
"args": {"query": "query"},
}
],
id="ai1",
)
},
},
)
# modify ai message
@@ -3254,7 +3560,23 @@ def test_message_graph(
],
next=("action",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "update", "step": 2},
metadata={
"source": "update",
"step": 2,
"writes": {
"agent": AIMessage(
content="",
tool_calls=[
{
"id": "tool_call123",
"name": "search_api",
"args": {"query": "a different query"},
}
],
id="ai1",
)
},
},
)
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -3320,7 +3642,23 @@ def test_message_graph(
],
next=("action",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "loop", "step": 4},
metadata={
"source": "loop",
"step": 4,
"writes": {
"agent": AIMessage(
content="",
tool_calls=[
{
"id": "tool_call456",
"name": "search_api",
"args": {"query": "another"},
}
],
id="ai2",
)
},
},
)
app_w_interrupt.update_state(
@@ -3356,7 +3694,11 @@ def test_message_graph(
],
next=(),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "update", "step": 5},
metadata={
"source": "update",
"step": 5,
"writes": {"agent": AIMessage(content="answer", id="ai2")},
},
)
# add an extra message as if it came from "action" node
@@ -3392,7 +3734,11 @@ def test_message_graph(
],
next=("agent",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "update", "step": 6},
metadata={
"source": "update",
"step": 6,
"writes": {"action": ("ai", "an extra message")},
},
)
@@ -3515,7 +3861,7 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None:
values={"my_key": "value", "market": "DE"},
next=("tool_two_slow",),
config=tool_two.checkpointer.get_tuple(thread1).config,
metadata={"source": "loop", "step": 0},
metadata={"source": "loop", "step": 0, "writes": None},
parent_config=[*tool_two.checkpointer.list(thread1, limit=2)][-1].config,
)
# resume, for same result as above
@@ -3527,7 +3873,11 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None:
values={"my_key": "value slow", "market": "DE"},
next=(),
config=tool_two.checkpointer.get_tuple(thread1).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {"tool_two_slow": {"my_key": " slow"}},
},
parent_config=[*tool_two.checkpointer.list(thread1, limit=2)][-1].config,
)
@@ -3541,7 +3891,7 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None:
values={"my_key": "value", "market": "US"},
next=("tool_two_fast",),
config=tool_two.checkpointer.get_tuple(thread2).config,
metadata={"source": "loop", "step": 0},
metadata={"source": "loop", "step": 0, "writes": None},
parent_config=[*tool_two.checkpointer.list(thread2, limit=2)][-1].config,
)
# resume, for same result as above
@@ -3553,7 +3903,11 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None:
values={"my_key": "value fast", "market": "US"},
next=(),
config=tool_two.checkpointer.get_tuple(thread2).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {"tool_two_fast": {"my_key": " fast"}},
},
parent_config=[*tool_two.checkpointer.list(thread2, limit=2)][-1].config,
)
@@ -3567,7 +3921,7 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None:
values={"my_key": "value", "market": "US"},
next=("tool_two_fast",),
config=tool_two.checkpointer.get_tuple(thread3).config,
metadata={"source": "loop", "step": 0},
metadata={"source": "loop", "step": 0, "writes": None},
parent_config=[*tool_two.checkpointer.list(thread3, limit=2)][-1].config,
)
# update state
@@ -3576,7 +3930,11 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None:
values={"my_key": "valuekey", "market": "US"},
next=("tool_two_fast",),
config=tool_two.checkpointer.get_tuple(thread3).config,
metadata={"source": "update", "step": 1},
metadata={
"source": "update",
"step": 1,
"writes": {START: {"my_key": "key"}},
},
parent_config=[*tool_two.checkpointer.list(thread3, limit=2)][-1].config,
)
# resume, for same result as above
@@ -3588,7 +3946,11 @@ def test_start_branch_then(snapshot: SnapshotAssertion) -> None:
values={"my_key": "valuekey fast", "market": "US"},
next=(),
config=tool_two.checkpointer.get_tuple(thread3).config,
metadata={"source": "loop", "step": 2},
metadata={
"source": "loop",
"step": 2,
"writes": {"tool_two_fast": {"my_key": " fast"}},
},
parent_config=[*tool_two.checkpointer.list(thread3, limit=2)][-1].config,
)
@@ -3778,7 +4140,11 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
values={"my_key": "value prepared", "market": "DE"},
next=("tool_two_slow",),
config=tool_two.checkpointer.get_tuple(thread1).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {"prepare": {"my_key": " prepared"}},
},
parent_config=[*tool_two.checkpointer.list(thread1, limit=2)][-1].config,
)
# resume, for same result as above
@@ -3790,7 +4156,11 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
values={"my_key": "value prepared slow finished", "market": "DE"},
next=(),
config=tool_two.checkpointer.get_tuple(thread1).config,
metadata={"source": "loop", "step": 3},
metadata={
"source": "loop",
"step": 3,
"writes": {"finish": {"my_key": " finished"}},
},
parent_config=[*tool_two.checkpointer.list(thread1, limit=2)][-1].config,
)
@@ -3804,7 +4174,11 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
values={"my_key": "value prepared", "market": "US"},
next=("tool_two_fast",),
config=tool_two.checkpointer.get_tuple(thread2).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {"prepare": {"my_key": " prepared"}},
},
parent_config=[*tool_two.checkpointer.list(thread2, limit=2)][-1].config,
)
# resume, for same result as above
@@ -3816,7 +4190,11 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
values={"my_key": "value prepared fast finished", "market": "US"},
next=(),
config=tool_two.checkpointer.get_tuple(thread2).config,
metadata={"source": "loop", "step": 3},
metadata={
"source": "loop",
"step": 3,
"writes": {"finish": {"my_key": " finished"}},
},
parent_config=[*tool_two.checkpointer.list(thread2, limit=2)][-1].config,
)
@@ -3839,7 +4217,11 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
values={"my_key": "value prepared", "market": "DE"},
next=("tool_two_slow",),
config=tool_two.checkpointer.get_tuple(thread1).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {"prepare": {"my_key": " prepared"}},
},
parent_config=[*tool_two.checkpointer.list(thread1, limit=2)][-1].config,
)
# resume, for same result as above
@@ -3851,7 +4233,11 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
values={"my_key": "value prepared slow finished", "market": "DE"},
next=(),
config=tool_two.checkpointer.get_tuple(thread1).config,
metadata={"source": "loop", "step": 3},
metadata={
"source": "loop",
"step": 3,
"writes": {"finish": {"my_key": " finished"}},
},
parent_config=[*tool_two.checkpointer.list(thread1, limit=2)][-1].config,
)
@@ -3865,7 +4251,11 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
values={"my_key": "value prepared", "market": "US"},
next=("tool_two_fast",),
config=tool_two.checkpointer.get_tuple(thread2).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {"prepare": {"my_key": " prepared"}},
},
parent_config=[*tool_two.checkpointer.list(thread2, limit=2)][-1].config,
)
# resume, for same result as above
@@ -3877,7 +4267,11 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
values={"my_key": "value prepared fast finished", "market": "US"},
next=(),
config=tool_two.checkpointer.get_tuple(thread2).config,
metadata={"source": "loop", "step": 3},
metadata={
"source": "loop",
"step": 3,
"writes": {"finish": {"my_key": " finished"}},
},
parent_config=[*tool_two.checkpointer.list(thread2, limit=2)][-1].config,
)
@@ -3889,7 +4283,11 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
values={"my_key": "key", "market": "DE"},
next=("prepare",),
config=uconfig,
metadata={"source": "update", "step": 0},
metadata={
"source": "update",
"step": 0,
"writes": {START: {"my_key": "key", "market": "DE"}},
},
)
# run from this point
assert tool_two.invoke(None, thread3) == {
@@ -3901,7 +4299,11 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
values={"my_key": "key prepared", "market": "DE"},
next=("tool_two_slow",),
config=tool_two.checkpointer.get_tuple(thread3).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {"prepare": {"my_key": " prepared"}},
},
parent_config=uconfig,
)
# resume, for same result as above
@@ -3913,7 +4315,11 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
values={"my_key": "key prepared slow finished", "market": "DE"},
next=(),
config=tool_two.checkpointer.get_tuple(thread3).config,
metadata={"source": "loop", "step": 3},
metadata={
"source": "loop",
"step": 3,
"writes": {"finish": {"my_key": " finished"}},
},
parent_config=[*tool_two.checkpointer.list(thread3, limit=2)][-1].config,
)
+321 -35
View File
@@ -27,6 +27,7 @@ from langgraph.channels.topic import Topic
from langgraph.checkpoint.aiosqlite import AsyncSqliteSaver
from langgraph.errors import InvalidUpdateError
from langgraph.graph import END, Graph, StateGraph
from langgraph.graph.graph import START
from langgraph.graph.message import MessageGraph
from langgraph.prebuilt.chat_agent_executor import (
create_function_calling_executor,
@@ -1300,7 +1301,20 @@ async def test_conditional_graph() -> None:
},
next=("tools",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={"source": "loop", "step": 0},
metadata={
"source": "loop",
"step": 0,
"writes": {
"agent": {
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:query",
),
}
},
},
)
await app_w_interrupt.aupdate_state(
@@ -1328,7 +1342,20 @@ async def test_conditional_graph() -> None:
},
next=("tools",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={"source": "update", "step": 1},
metadata={
"source": "update",
"step": 1,
"writes": {
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"input": "what is weather in sf",
}
},
},
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
@@ -1412,7 +1439,29 @@ async def test_conditional_graph() -> None:
},
next=(),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={"source": "update", "step": 4},
metadata={
"source": "update",
"step": 4,
"writes": {
"agent": {
"input": "what is weather in sf",
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
),
}
},
},
)
# test state get/update methods with interrupt_before
@@ -1451,7 +1500,20 @@ async def test_conditional_graph() -> None:
},
next=("tools",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={"source": "loop", "step": 0},
metadata={
"source": "loop",
"step": 0,
"writes": {
"agent": {
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:query",
),
}
},
},
)
await app_w_interrupt.aupdate_state(
@@ -1479,7 +1541,20 @@ async def test_conditional_graph() -> None:
},
next=("tools",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={"source": "update", "step": 1},
metadata={
"source": "update",
"step": 1,
"writes": {
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"input": "what is weather in sf",
}
},
},
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
@@ -1563,7 +1638,29 @@ async def test_conditional_graph() -> None:
},
next=(),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={"source": "update", "step": 4},
metadata={
"source": "update",
"step": 4,
"writes": {
"agent": {
"input": "what is weather in sf",
"intermediate_steps": [
(
AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
),
"result for query",
)
],
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
),
}
},
},
)
# test re-invoke to continue with interrupt_before
@@ -1602,7 +1699,20 @@ async def test_conditional_graph() -> None:
},
next=("tools",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={"source": "loop", "step": 0},
metadata={
"source": "loop",
"step": 0,
"writes": {
"agent": {
"input": "what is weather in sf",
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:query",
),
}
},
},
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
@@ -1922,7 +2032,19 @@ async def test_conditional_graph_state() -> None:
},
next=("tools",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:query",
),
}
},
},
)
await app_w_interrupt.aupdate_state(
@@ -1948,7 +2070,19 @@ async def test_conditional_graph_state() -> None:
},
next=("tools",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={"source": "update", "step": 2},
metadata={
"source": "update",
"step": 2,
"writes": {
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
)
}
},
},
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
@@ -2007,7 +2141,18 @@ async def test_conditional_graph_state() -> None:
},
next=(),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={"source": "update", "step": 5},
metadata={
"source": "update",
"step": 5,
"writes": {
"agent": {
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
)
}
},
},
)
# test state get/update methods with interrupt_before
@@ -2044,7 +2189,19 @@ async def test_conditional_graph_state() -> None:
},
next=("tools",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:query",
),
}
},
},
)
await app_w_interrupt.aupdate_state(
@@ -2070,7 +2227,19 @@ async def test_conditional_graph_state() -> None:
},
next=("tools",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={"source": "update", "step": 2},
metadata={
"source": "update",
"step": 2,
"writes": {
"agent": {
"agent_outcome": AgentAction(
tool="search_api",
tool_input="query",
log="tool:search_api:a different query",
)
}
},
},
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
@@ -2129,7 +2298,18 @@ async def test_conditional_graph_state() -> None:
},
next=(),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={"source": "update", "step": 5},
metadata={
"source": "update",
"step": 5,
"writes": {
"agent": {
"agent_outcome": AgentFinish(
return_values={"answer": "a really nice answer"},
log="finish:a really nice answer",
)
}
},
},
)
@@ -2735,7 +2915,19 @@ async def test_message_graph() -> None:
],
next=("action",),
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {
"agent": AIMessage(
content="",
additional_kwargs={
"function_call": {"name": "search_api", "arguments": '"query"'}
},
id="ai1",
)
},
},
)
# modify ai message
@@ -2763,7 +2955,22 @@ async def test_message_graph() -> None:
],
next=("action",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "update", "step": 2},
metadata={
"source": "update",
"step": 2,
"writes": {
"agent": AIMessage(
content="",
additional_kwargs={
"function_call": {
"name": "search_api",
"arguments": '"a different query"',
}
},
id="ai1",
)
},
},
)
assert [c async for c in app_w_interrupt.astream(None, config)] == [
@@ -2816,7 +3023,22 @@ async def test_message_graph() -> None:
],
next=("action",),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "loop", "step": 4},
metadata={
"source": "loop",
"step": 4,
"writes": {
"agent": AIMessage(
content="",
additional_kwargs={
"function_call": {
"name": "search_api",
"arguments": '"another"',
}
},
id="ai2",
)
},
},
)
await app_w_interrupt.aupdate_state(
@@ -2850,7 +3072,11 @@ async def test_message_graph() -> None:
],
next=(),
config=app_w_interrupt.checkpointer.get_tuple(config).config,
metadata={"source": "update", "step": 5},
metadata={
"source": "update",
"step": 5,
"writes": {"agent": AIMessage(content="answer", id="ai2")},
},
)
@@ -2967,7 +3193,7 @@ async def test_start_branch_then() -> None:
values={"my_key": "value", "market": "DE"},
next=("tool_two_slow",),
config=(await tool_two.checkpointer.aget_tuple(thread1)).config,
metadata={"source": "loop", "step": 0},
metadata={"source": "loop", "step": 0, "writes": None},
parent_config=[
c async for c in tool_two.checkpointer.alist(thread1, limit=2)
][-1].config,
@@ -2981,7 +3207,11 @@ async def test_start_branch_then() -> None:
values={"my_key": "value slow", "market": "DE"},
next=(),
config=(await tool_two.checkpointer.aget_tuple(thread1)).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {"tool_two_slow": {"my_key": " slow"}},
},
parent_config=[
c async for c in tool_two.checkpointer.alist(thread1, limit=2)
][-1].config,
@@ -2997,7 +3227,7 @@ async def test_start_branch_then() -> None:
values={"my_key": "value", "market": "US"},
next=("tool_two_fast",),
config=(await tool_two.checkpointer.aget_tuple(thread2)).config,
metadata={"source": "loop", "step": 0},
metadata={"source": "loop", "step": 0, "writes": None},
parent_config=[
c async for c in tool_two.checkpointer.alist(thread2, limit=2)
][-1].config,
@@ -3011,7 +3241,11 @@ async def test_start_branch_then() -> None:
values={"my_key": "value fast", "market": "US"},
next=(),
config=(await tool_two.checkpointer.aget_tuple(thread2)).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {"tool_two_fast": {"my_key": " fast"}},
},
parent_config=[
c async for c in tool_two.checkpointer.alist(thread2, limit=2)
][-1].config,
@@ -3027,7 +3261,7 @@ async def test_start_branch_then() -> None:
values={"my_key": "value", "market": "US"},
next=("tool_two_fast",),
config=(await tool_two.checkpointer.aget_tuple(thread3)).config,
metadata={"source": "loop", "step": 0},
metadata={"source": "loop", "step": 0, "writes": None},
parent_config=[
c async for c in tool_two.checkpointer.alist(thread3, limit=2)
][-1].config,
@@ -3038,7 +3272,11 @@ async def test_start_branch_then() -> None:
values={"my_key": "valuekey", "market": "US"},
next=("tool_two_fast",),
config=(await tool_two.checkpointer.aget_tuple(thread3)).config,
metadata={"source": "update", "step": 1},
metadata={
"source": "update",
"step": 1,
"writes": {START: {"my_key": "key"}},
},
parent_config=[
c async for c in tool_two.checkpointer.alist(thread3, limit=2)
][-1].config,
@@ -3052,7 +3290,11 @@ async def test_start_branch_then() -> None:
values={"my_key": "valuekey fast", "market": "US"},
next=(),
config=(await tool_two.checkpointer.aget_tuple(thread3)).config,
metadata={"source": "loop", "step": 2},
metadata={
"source": "loop",
"step": 2,
"writes": {"tool_two_fast": {"my_key": " fast"}},
},
parent_config=[
c async for c in tool_two.checkpointer.alist(thread3, limit=2)
][-1].config,
@@ -3229,7 +3471,11 @@ async def test_branch_then() -> None:
values={"my_key": "value prepared", "market": "DE"},
next=("tool_two_slow",),
config=(await tool_two.checkpointer.aget_tuple(thread1)).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {"prepare": {"my_key": " prepared"}},
},
parent_config=[
c async for c in tool_two.checkpointer.alist(thread1, limit=2)
][-1].config,
@@ -3243,7 +3489,11 @@ async def test_branch_then() -> None:
values={"my_key": "value prepared slow finished", "market": "DE"},
next=(),
config=(await tool_two.checkpointer.aget_tuple(thread1)).config,
metadata={"source": "loop", "step": 3},
metadata={
"source": "loop",
"step": 3,
"writes": {"finish": {"my_key": " finished"}},
},
parent_config=[
c async for c in tool_two.checkpointer.alist(thread1, limit=2)
][-1].config,
@@ -3259,7 +3509,11 @@ async def test_branch_then() -> None:
values={"my_key": "value prepared", "market": "US"},
next=("tool_two_fast",),
config=(await tool_two.checkpointer.aget_tuple(thread2)).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {"prepare": {"my_key": " prepared"}},
},
parent_config=[
c async for c in tool_two.checkpointer.alist(thread2, limit=2)
][-1].config,
@@ -3273,7 +3527,11 @@ async def test_branch_then() -> None:
values={"my_key": "value prepared fast finished", "market": "US"},
next=(),
config=(await tool_two.checkpointer.aget_tuple(thread2)).config,
metadata={"source": "loop", "step": 3},
metadata={
"source": "loop",
"step": 3,
"writes": {"finish": {"my_key": " finished"}},
},
parent_config=[
c async for c in tool_two.checkpointer.alist(thread2, limit=2)
][-1].config,
@@ -3298,7 +3556,11 @@ async def test_branch_then() -> None:
values={"my_key": "value prepared", "market": "DE"},
next=("tool_two_slow",),
config=(await tool_two.checkpointer.aget_tuple(thread1)).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {"prepare": {"my_key": " prepared"}},
},
parent_config=[
c async for c in tool_two.checkpointer.alist(thread1, limit=2)
][-1].config,
@@ -3312,7 +3574,11 @@ async def test_branch_then() -> None:
values={"my_key": "value prepared slow finished", "market": "DE"},
next=(),
config=(await tool_two.checkpointer.aget_tuple(thread1)).config,
metadata={"source": "loop", "step": 3},
metadata={
"source": "loop",
"step": 3,
"writes": {"finish": {"my_key": " finished"}},
},
parent_config=[
c async for c in tool_two.checkpointer.alist(thread1, limit=2)
][-1].config,
@@ -3328,7 +3594,11 @@ async def test_branch_then() -> None:
values={"my_key": "value prepared", "market": "US"},
next=("tool_two_fast",),
config=(await tool_two.checkpointer.aget_tuple(thread2)).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {"prepare": {"my_key": " prepared"}},
},
parent_config=[
c async for c in tool_two.checkpointer.alist(thread2, limit=2)
][-1].config,
@@ -3342,7 +3612,11 @@ async def test_branch_then() -> None:
values={"my_key": "value prepared fast finished", "market": "US"},
next=(),
config=(await tool_two.checkpointer.aget_tuple(thread2)).config,
metadata={"source": "loop", "step": 3},
metadata={
"source": "loop",
"step": 3,
"writes": {"finish": {"my_key": " finished"}},
},
parent_config=[
c async for c in tool_two.checkpointer.alist(thread2, limit=2)
][-1].config,
@@ -3358,7 +3632,11 @@ async def test_branch_then() -> None:
values={"my_key": "key", "market": "DE"},
next=("prepare",),
config=uconfig,
metadata={"source": "update", "step": 0},
metadata={
"source": "update",
"step": 0,
"writes": {START: {"my_key": "key", "market": "DE"}},
},
)
# run from this point
assert await tool_two.ainvoke(None, thread3) == {
@@ -3370,7 +3648,11 @@ async def test_branch_then() -> None:
values={"my_key": "key prepared", "market": "DE"},
next=("tool_two_slow",),
config=(await tool_two.checkpointer.aget_tuple(thread3)).config,
metadata={"source": "loop", "step": 1},
metadata={
"source": "loop",
"step": 1,
"writes": {"prepare": {"my_key": " prepared"}},
},
parent_config=uconfig,
)
# resume, for same result as above
@@ -3382,7 +3664,11 @@ async def test_branch_then() -> None:
values={"my_key": "key prepared slow finished", "market": "DE"},
next=(),
config=(await tool_two.checkpointer.aget_tuple(thread3)).config,
metadata={"source": "loop", "step": 3},
metadata={
"source": "loop",
"step": 3,
"writes": {"finish": {"my_key": " finished"}},
},
parent_config=[
c async for c in tool_two.checkpointer.alist(thread3, limit=2)
][-1].config,