mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-17 21:25:46 +02:00
Merge pull request #1172 from langchain-ai/nc/30jul/update-no-values
Allow call to update_state without values
This commit is contained in:
@@ -44,19 +44,22 @@ async def AsyncChannelsManager(
|
||||
|
||||
def create_checkpoint(
|
||||
checkpoint: Checkpoint,
|
||||
channels: Mapping[str, BaseChannel],
|
||||
channels: Optional[Mapping[str, BaseChannel]],
|
||||
step: int,
|
||||
*,
|
||||
id: Optional[str] = None,
|
||||
) -> Checkpoint:
|
||||
"""Create a checkpoint for the given channels."""
|
||||
ts = datetime.now(timezone.utc).isoformat()
|
||||
values: dict[str, Any] = {}
|
||||
for k, v in channels.items():
|
||||
try:
|
||||
values[k] = v.checkpoint()
|
||||
except EmptyChannelError:
|
||||
pass
|
||||
if channels is None:
|
||||
values = checkpoint["channel_values"]
|
||||
else:
|
||||
values: dict[str, Any] = {}
|
||||
for k, v in channels.items():
|
||||
try:
|
||||
values[k] = v.checkpoint()
|
||||
except EmptyChannelError:
|
||||
pass
|
||||
return Checkpoint(
|
||||
v=1,
|
||||
ts=ts,
|
||||
|
||||
@@ -504,7 +504,7 @@ class Pregel(
|
||||
def update_state(
|
||||
self,
|
||||
config: RunnableConfig,
|
||||
values: dict[str, Any] | Any,
|
||||
values: Optional[Union[dict[str, Any], Any]],
|
||||
as_node: Optional[str] = None,
|
||||
) -> RunnableConfig:
|
||||
"""Update the state of the graph with the given values, as if they came from
|
||||
@@ -517,8 +517,28 @@ class Pregel(
|
||||
# get last checkpoint
|
||||
saved = self.checkpointer.get_tuple(config)
|
||||
checkpoint = copy_checkpoint(saved.checkpoint) if saved else empty_checkpoint()
|
||||
step = saved.metadata.get("step", -1) if saved else -1
|
||||
# merge configurable fields with previous checkpoint config
|
||||
checkpoint_config = config
|
||||
if saved:
|
||||
checkpoint_config = {
|
||||
"configurable": {
|
||||
**config.get("configurable", {}),
|
||||
**saved.config["configurable"],
|
||||
}
|
||||
}
|
||||
# find last node that updated the state, if not provided
|
||||
if as_node is None and not any(
|
||||
if values is None and as_node is None:
|
||||
return self.checkpointer.put(
|
||||
checkpoint_config,
|
||||
create_checkpoint(checkpoint, None, step),
|
||||
{
|
||||
"source": "update",
|
||||
"step": step,
|
||||
"writes": {},
|
||||
},
|
||||
)
|
||||
elif as_node is None and not any(
|
||||
v for vv in checkpoint["versions_seen"].values() for v in vv.values()
|
||||
):
|
||||
if (
|
||||
@@ -577,24 +597,13 @@ class Pregel(
|
||||
apply_writes(
|
||||
checkpoint, channels, [task], self.checkpointer.get_next_version
|
||||
)
|
||||
step = saved.metadata.get("step", -2) + 1 if saved else -1
|
||||
|
||||
# merge configurable fields with previous checkpoint config
|
||||
checkpoint_config = config
|
||||
if saved:
|
||||
checkpoint_config = {
|
||||
"configurable": {
|
||||
**config.get("configurable", {}),
|
||||
**saved.config["configurable"],
|
||||
}
|
||||
}
|
||||
|
||||
return self.checkpointer.put(
|
||||
checkpoint_config,
|
||||
create_checkpoint(checkpoint, channels, step),
|
||||
create_checkpoint(checkpoint, channels, step + 1),
|
||||
{
|
||||
"source": "update",
|
||||
"step": step,
|
||||
"step": step + 1,
|
||||
"writes": {as_node: values},
|
||||
},
|
||||
)
|
||||
|
||||
+375
-167
@@ -549,172 +549,380 @@ def test_invoke_two_processes_in_out(mocker: MockerFixture) -> None:
|
||||
assert step == 2
|
||||
|
||||
|
||||
def test_invoke_two_processes_in_out_interrupt(mocker: MockerFixture) -> None:
|
||||
add_one = mocker.Mock(side_effect=lambda x: x + 1)
|
||||
one = Channel.subscribe_to("input") | add_one | Channel.write_to("inbox")
|
||||
two = Channel.subscribe_to("inbox") | add_one | Channel.write_to("output")
|
||||
@pytest.mark.parametrize(
|
||||
"checkpointer",
|
||||
[
|
||||
MemorySaverAssertImmutable(),
|
||||
SqliteSaver.from_conn_string(":memory:"),
|
||||
],
|
||||
ids=[
|
||||
"memory",
|
||||
"sqlite",
|
||||
],
|
||||
)
|
||||
def test_invoke_two_processes_in_out_interrupt(
|
||||
checkpointer: BaseCheckpointSaver, mocker: MockerFixture
|
||||
) -> None:
|
||||
try:
|
||||
add_one = mocker.Mock(side_effect=lambda x: x + 1)
|
||||
one = Channel.subscribe_to("input") | add_one | Channel.write_to("inbox")
|
||||
two = Channel.subscribe_to("inbox") | add_one | Channel.write_to("output")
|
||||
|
||||
memory = MemorySaverAssertImmutable()
|
||||
app = Pregel(
|
||||
nodes={"one": one, "two": two},
|
||||
channels={
|
||||
"inbox": LastValue(int),
|
||||
"output": LastValue(int),
|
||||
"input": LastValue(int),
|
||||
},
|
||||
input_channels="input",
|
||||
output_channels="output",
|
||||
checkpointer=memory,
|
||||
interrupt_after_nodes=["one"],
|
||||
)
|
||||
|
||||
# start execution, stop at inbox
|
||||
assert app.invoke(2, {"configurable": {"thread_id": 1}}) is None
|
||||
|
||||
# inbox == 3
|
||||
checkpoint = memory.get({"configurable": {"thread_id": 1}})
|
||||
assert checkpoint is not None
|
||||
assert checkpoint["channel_values"]["inbox"] == 3
|
||||
|
||||
# resume execution, finish
|
||||
assert app.invoke(None, {"configurable": {"thread_id": 1}}) == 4
|
||||
|
||||
# start execution again, stop at inbox
|
||||
assert app.invoke(20, {"configurable": {"thread_id": 1}}) is None
|
||||
|
||||
# inbox == 21
|
||||
checkpoint = memory.get({"configurable": {"thread_id": 1}})
|
||||
assert checkpoint is not None
|
||||
assert checkpoint["channel_values"]["inbox"] == 21
|
||||
|
||||
# send a new value in, interrupting the previous execution
|
||||
assert app.invoke(3, {"configurable": {"thread_id": 1}}) is None
|
||||
assert app.invoke(None, {"configurable": {"thread_id": 1}}) == 5
|
||||
|
||||
# start execution again, stopping at inbox
|
||||
assert app.invoke(20, {"configurable": {"thread_id": 2}}) is None
|
||||
|
||||
# inbox == 21
|
||||
snapshot = app.get_state({"configurable": {"thread_id": 2}})
|
||||
assert snapshot.values["inbox"] == 21
|
||||
assert snapshot.next == ("two",)
|
||||
|
||||
# update the state, resume
|
||||
app.update_state({"configurable": {"thread_id": 2}}, 25, as_node="one")
|
||||
assert app.invoke(None, {"configurable": {"thread_id": 2}}) == 26
|
||||
|
||||
# no pending tasks
|
||||
snapshot = app.get_state({"configurable": {"thread_id": 2}})
|
||||
assert snapshot.next == ()
|
||||
|
||||
# list history
|
||||
thread1 = {"configurable": {"thread_id": 1}}
|
||||
assert [c for c in app.get_state_history(thread1)] == [
|
||||
StateSnapshot(
|
||||
values={"inbox": 4, "output": 5, "input": 3},
|
||||
next=(),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": 1,
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
app = Pregel(
|
||||
nodes={"one": one, "two": two},
|
||||
channels={
|
||||
"inbox": LastValue(int),
|
||||
"output": LastValue(int),
|
||||
"input": LastValue(int),
|
||||
},
|
||||
metadata={"source": "loop", "step": 6, "writes": 5},
|
||||
created_at=AnyStr(),
|
||||
parent_config=[*app.checkpointer.list(thread1)][1].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 4, "output": 4, "input": 3},
|
||||
next=("two",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": 1,
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 5, "writes": None},
|
||||
created_at=AnyStr(),
|
||||
parent_config=[*app.checkpointer.list(thread1)][2].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 21, "output": 4, "input": 3},
|
||||
next=("one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": 1,
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "input", "step": 4, "writes": 3},
|
||||
created_at=AnyStr(),
|
||||
parent_config=[*app.checkpointer.list(thread1)][3].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 21, "output": 4, "input": 20},
|
||||
next=("two",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": 1,
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 3, "writes": None},
|
||||
created_at=AnyStr(),
|
||||
parent_config=[*app.checkpointer.list(thread1)][4].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 3, "output": 4, "input": 20},
|
||||
next=("one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": 1,
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "input", "step": 2, "writes": 20},
|
||||
created_at=AnyStr(),
|
||||
parent_config=[*app.checkpointer.list(thread1)][5].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 3, "output": 4, "input": 2},
|
||||
next=(),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": 1,
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 1, "writes": 4},
|
||||
created_at=AnyStr(),
|
||||
parent_config=[*app.checkpointer.list(thread1)][6].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 3, "input": 2},
|
||||
next=("two",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": 1,
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 0, "writes": None},
|
||||
created_at=AnyStr(),
|
||||
parent_config=[*app.checkpointer.list(thread1)][7].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"input": 2},
|
||||
next=("one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": 1,
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "input", "step": -1, "writes": 2},
|
||||
created_at=AnyStr(),
|
||||
parent_config=None,
|
||||
),
|
||||
]
|
||||
input_channels="input",
|
||||
output_channels="output",
|
||||
checkpointer=checkpointer,
|
||||
interrupt_after_nodes=["one"],
|
||||
)
|
||||
thread1 = {"configurable": {"thread_id": "1"}}
|
||||
thread2 = {"configurable": {"thread_id": "2"}}
|
||||
|
||||
# start execution, stop at inbox
|
||||
assert app.invoke(2, thread1) is None
|
||||
|
||||
# inbox == 3
|
||||
checkpoint = checkpointer.get(thread1)
|
||||
assert checkpoint is not None
|
||||
assert checkpoint["channel_values"]["inbox"] == 3
|
||||
|
||||
# resume execution, finish
|
||||
assert app.invoke(None, thread1) == 4
|
||||
|
||||
# start execution again, stop at inbox
|
||||
assert app.invoke(20, thread1) is None
|
||||
|
||||
# inbox == 21
|
||||
checkpoint = checkpointer.get(thread1)
|
||||
assert checkpoint is not None
|
||||
assert checkpoint["channel_values"]["inbox"] == 21
|
||||
|
||||
# send a new value in, interrupting the previous execution
|
||||
assert app.invoke(3, thread1) is None
|
||||
assert app.invoke(None, thread1) == 5
|
||||
|
||||
# start execution again, stopping at inbox
|
||||
assert app.invoke(20, thread2) is None
|
||||
|
||||
# inbox == 21
|
||||
snapshot = app.get_state(thread2)
|
||||
assert snapshot.values["inbox"] == 21
|
||||
assert snapshot.next == ("two",)
|
||||
|
||||
# update the state, resume
|
||||
app.update_state(thread2, 25, as_node="one")
|
||||
assert app.invoke(None, thread2) == 26
|
||||
|
||||
# no pending tasks
|
||||
snapshot = app.get_state(thread2)
|
||||
assert snapshot.next == ()
|
||||
|
||||
# list history
|
||||
history = [c for c in app.get_state_history(thread1)]
|
||||
assert history == [
|
||||
StateSnapshot(
|
||||
values={"inbox": 4, "output": 5, "input": 3},
|
||||
next=(),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 6, "writes": 5},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[1].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 4, "output": 4, "input": 3},
|
||||
next=("two",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 5, "writes": None},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[2].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 21, "output": 4, "input": 3},
|
||||
next=("one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "input", "step": 4, "writes": 3},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[3].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 21, "output": 4, "input": 20},
|
||||
next=("two",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 3, "writes": None},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[4].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 3, "output": 4, "input": 20},
|
||||
next=("one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "input", "step": 2, "writes": 20},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[5].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 3, "output": 4, "input": 2},
|
||||
next=(),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 1, "writes": 4},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[6].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 3, "input": 2},
|
||||
next=("two",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 0, "writes": None},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[7].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"input": 2},
|
||||
next=("one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "input", "step": -1, "writes": 2},
|
||||
created_at=AnyStr(),
|
||||
parent_config=None,
|
||||
),
|
||||
]
|
||||
|
||||
# forking from any previous checkpoint w/out forking should do nothing
|
||||
assert [
|
||||
c for c in app.stream(None, history[0].config, stream_mode="updates")
|
||||
] == []
|
||||
assert [
|
||||
c for c in app.stream(None, history[1].config, stream_mode="updates")
|
||||
] == []
|
||||
assert [
|
||||
c for c in app.stream(None, history[2].config, stream_mode="updates")
|
||||
] == []
|
||||
|
||||
# forking and re-running from any prev checkpoint should re-run nodes
|
||||
fork_config = app.update_state(history[0].config, None)
|
||||
assert [c for c in app.stream(None, fork_config, stream_mode="updates")] == []
|
||||
|
||||
fork_config = app.update_state(history[1].config, None)
|
||||
assert [c for c in app.stream(None, fork_config, stream_mode="updates")] == [
|
||||
{"two": {"output": 5}}
|
||||
]
|
||||
|
||||
fork_config = app.update_state(history[2].config, None)
|
||||
assert [c for c in app.stream(None, fork_config, stream_mode="updates")] == [
|
||||
{"one": {"inbox": 4}}
|
||||
]
|
||||
finally:
|
||||
if hasattr(checkpointer, "__exit__"):
|
||||
checkpointer.__exit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"checkpointer",
|
||||
[
|
||||
MemorySaverAssertImmutable(),
|
||||
SqliteSaver.from_conn_string(":memory:"),
|
||||
],
|
||||
ids=[
|
||||
"memory",
|
||||
"sqlite",
|
||||
],
|
||||
)
|
||||
def test_fork_always_re_runs_nodes(
|
||||
checkpointer: BaseCheckpointSaver, mocker: MockerFixture
|
||||
) -> None:
|
||||
try:
|
||||
add_one = mocker.Mock(side_effect=lambda _: 1)
|
||||
|
||||
builder = StateGraph(Annotated[int, operator.add])
|
||||
builder.add_node("add_one", add_one)
|
||||
builder.add_edge(START, "add_one")
|
||||
builder.add_conditional_edges(
|
||||
"add_one", lambda cnt: "add_one" if cnt < 6 else END
|
||||
)
|
||||
graph = builder.compile(checkpointer=checkpointer)
|
||||
|
||||
thread1 = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# start execution, stop at inbox
|
||||
assert [*graph.stream(1, thread1, stream_mode=["values", "updates"])] == [
|
||||
("values", 1),
|
||||
("updates", {"add_one": 1}),
|
||||
("values", 2),
|
||||
("updates", {"add_one": 1}),
|
||||
("values", 3),
|
||||
("updates", {"add_one": 1}),
|
||||
("values", 4),
|
||||
("updates", {"add_one": 1}),
|
||||
("values", 5),
|
||||
("updates", {"add_one": 1}),
|
||||
("values", 6),
|
||||
]
|
||||
|
||||
# list history
|
||||
history = [c for c in graph.get_state_history(thread1)]
|
||||
assert history == [
|
||||
StateSnapshot(
|
||||
values=6,
|
||||
next=(),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 5, "writes": {"add_one": 1}},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[1].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values=5,
|
||||
next=("add_one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 4, "writes": {"add_one": 1}},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[2].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values=4,
|
||||
next=("add_one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 3, "writes": {"add_one": 1}},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[3].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values=3,
|
||||
next=("add_one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 2, "writes": {"add_one": 1}},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[4].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values=2,
|
||||
next=("add_one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 1, "writes": {"add_one": 1}},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[5].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values=1,
|
||||
next=("add_one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 0, "writes": None},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[6].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values=0,
|
||||
next=("__start__",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "input", "step": -1, "writes": 1},
|
||||
created_at=AnyStr(),
|
||||
parent_config=None,
|
||||
),
|
||||
]
|
||||
|
||||
# forking from any previous checkpoint w/out forking should do nothing
|
||||
assert [
|
||||
c for c in graph.stream(None, history[0].config, stream_mode="updates")
|
||||
] == []
|
||||
assert [
|
||||
c for c in graph.stream(None, history[1].config, stream_mode="updates")
|
||||
] == []
|
||||
|
||||
# forking and re-running from any prev checkpoint should re-run nodes
|
||||
fork_config = graph.update_state(history[0].config, None)
|
||||
assert [c for c in graph.stream(None, fork_config, stream_mode="updates")] == []
|
||||
|
||||
fork_config = graph.update_state(history[1].config, None)
|
||||
assert [c for c in graph.stream(None, fork_config, stream_mode="updates")] == [
|
||||
{"add_one": 1}
|
||||
]
|
||||
|
||||
fork_config = graph.update_state(history[2].config, None)
|
||||
assert [c for c in graph.stream(None, fork_config, stream_mode="updates")] == [
|
||||
{"add_one": 1},
|
||||
{"add_one": 1},
|
||||
]
|
||||
finally:
|
||||
if hasattr(checkpointer, "__exit__"):
|
||||
checkpointer.__exit__(None, None, None)
|
||||
|
||||
|
||||
def test_invoke_two_processes_in_dict_out(mocker: MockerFixture) -> None:
|
||||
@@ -6798,7 +7006,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
|
||||
created_at=AnyStr(),
|
||||
metadata={
|
||||
"source": "update",
|
||||
"step": -1,
|
||||
"step": 0,
|
||||
"writes": {START: {"my_key": "key", "market": "DE"}},
|
||||
},
|
||||
)
|
||||
@@ -6815,7 +7023,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
|
||||
created_at=tool_two.checkpointer.get_tuple(thread3).checkpoint["ts"],
|
||||
metadata={
|
||||
"source": "loop",
|
||||
"step": 0,
|
||||
"step": 1,
|
||||
"writes": {"prepare": {"my_key": " prepared"}},
|
||||
},
|
||||
parent_config=uconfig,
|
||||
@@ -6832,7 +7040,7 @@ def test_branch_then(snapshot: SnapshotAssertion) -> None:
|
||||
created_at=tool_two.checkpointer.get_tuple(thread3).checkpoint["ts"],
|
||||
metadata={
|
||||
"source": "loop",
|
||||
"step": 2,
|
||||
"step": 3,
|
||||
"writes": {"finish": {"my_key": " finished"}},
|
||||
},
|
||||
parent_config=[*tool_two.checkpointer.list(thread3, limit=2)][-1].config,
|
||||
|
||||
@@ -665,172 +665,364 @@ async def test_invoke_two_processes_in_out(mocker: MockerFixture) -> None:
|
||||
assert step == 2
|
||||
|
||||
|
||||
async def test_invoke_two_processes_in_out_interrupt(mocker: MockerFixture) -> None:
|
||||
add_one = mocker.Mock(side_effect=lambda x: x + 1)
|
||||
one = Channel.subscribe_to("input") | add_one | Channel.write_to("inbox")
|
||||
two = Channel.subscribe_to("inbox") | add_one | Channel.write_to("output")
|
||||
@pytest.mark.parametrize(
|
||||
"checkpointer",
|
||||
[
|
||||
MemorySaverAssertImmutable(),
|
||||
AsyncSqliteSaver.from_conn_string(":memory:"),
|
||||
],
|
||||
ids=[
|
||||
"memory",
|
||||
"sqlite",
|
||||
],
|
||||
)
|
||||
async def test_invoke_two_processes_in_out_interrupt(
|
||||
checkpointer: BaseCheckpointSaver, mocker: MockerFixture
|
||||
) -> None:
|
||||
try:
|
||||
add_one = mocker.Mock(side_effect=lambda x: x + 1)
|
||||
one = Channel.subscribe_to("input") | add_one | Channel.write_to("inbox")
|
||||
two = Channel.subscribe_to("inbox") | add_one | Channel.write_to("output")
|
||||
|
||||
memory = MemorySaverAssertImmutable()
|
||||
app = Pregel(
|
||||
nodes={"one": one, "two": two},
|
||||
channels={
|
||||
"inbox": LastValue(int),
|
||||
"output": LastValue(int),
|
||||
"input": LastValue(int),
|
||||
},
|
||||
input_channels="input",
|
||||
output_channels="output",
|
||||
checkpointer=memory,
|
||||
interrupt_after_nodes=["one"],
|
||||
)
|
||||
|
||||
# start execution, stop at inbox
|
||||
assert await app.ainvoke(2, {"configurable": {"thread_id": 1}}) is None
|
||||
|
||||
# inbox == 3
|
||||
checkpoint = await memory.aget({"configurable": {"thread_id": 1}})
|
||||
assert checkpoint is not None
|
||||
assert checkpoint["channel_values"]["inbox"] == 3
|
||||
|
||||
# resume execution, finish
|
||||
assert await app.ainvoke(None, {"configurable": {"thread_id": 1}}) == 4
|
||||
|
||||
# start execution again, stop at inbox
|
||||
assert await app.ainvoke(20, {"configurable": {"thread_id": 1}}) is None
|
||||
|
||||
# inbox == 21
|
||||
checkpoint = await memory.aget({"configurable": {"thread_id": 1}})
|
||||
assert checkpoint is not None
|
||||
assert checkpoint["channel_values"]["inbox"] == 21
|
||||
|
||||
# send a new value in, interrupting the previous execution
|
||||
assert await app.ainvoke(3, {"configurable": {"thread_id": 1}}) is None
|
||||
assert await app.ainvoke(None, {"configurable": {"thread_id": 1}}) == 5
|
||||
|
||||
# start execution again, stopping at inbox
|
||||
assert await app.ainvoke(20, {"configurable": {"thread_id": 2}}) is None
|
||||
|
||||
# inbox == 21
|
||||
snapshot = await app.aget_state({"configurable": {"thread_id": 2}})
|
||||
assert snapshot.values["inbox"] == 21
|
||||
assert snapshot.next == ("two",)
|
||||
|
||||
# update the state, resume
|
||||
await app.aupdate_state({"configurable": {"thread_id": 2}}, 25, as_node="one")
|
||||
assert await app.ainvoke(None, {"configurable": {"thread_id": 2}}) == 26
|
||||
|
||||
# no pending tasks
|
||||
snapshot = await app.aget_state({"configurable": {"thread_id": 2}})
|
||||
assert snapshot.next == ()
|
||||
|
||||
# list history
|
||||
thread1 = {"configurable": {"thread_id": 1}}
|
||||
assert [c async for c in app.aget_state_history(thread1)] == [
|
||||
StateSnapshot(
|
||||
values={"inbox": 4, "output": 5, "input": 3},
|
||||
next=(),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": 1,
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
app = Pregel(
|
||||
nodes={"one": one, "two": two},
|
||||
channels={
|
||||
"inbox": LastValue(int),
|
||||
"output": LastValue(int),
|
||||
"input": LastValue(int),
|
||||
},
|
||||
metadata={"source": "loop", "step": 6, "writes": 5},
|
||||
created_at=AnyStr(),
|
||||
parent_config=[c async for c in app.checkpointer.alist(thread1)][1].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 4, "output": 4, "input": 3},
|
||||
next=("two",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": 1,
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 5, "writes": None},
|
||||
created_at=AnyStr(),
|
||||
parent_config=[c async for c in app.checkpointer.alist(thread1)][2].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 21, "output": 4, "input": 3},
|
||||
next=("one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": 1,
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "input", "step": 4, "writes": 3},
|
||||
created_at=AnyStr(),
|
||||
parent_config=[c async for c in app.checkpointer.alist(thread1)][3].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 21, "output": 4, "input": 20},
|
||||
next=("two",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": 1,
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 3, "writes": None},
|
||||
created_at=AnyStr(),
|
||||
parent_config=[c async for c in app.checkpointer.alist(thread1)][4].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 3, "output": 4, "input": 20},
|
||||
next=("one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": 1,
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "input", "step": 2, "writes": 20},
|
||||
created_at=AnyStr(),
|
||||
parent_config=[c async for c in app.checkpointer.alist(thread1)][5].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 3, "output": 4, "input": 2},
|
||||
next=(),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": 1,
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 1, "writes": 4},
|
||||
created_at=AnyStr(),
|
||||
parent_config=[c async for c in app.checkpointer.alist(thread1)][6].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 3, "input": 2},
|
||||
next=("two",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": 1,
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 0, "writes": None},
|
||||
created_at=AnyStr(),
|
||||
parent_config=[c async for c in app.checkpointer.alist(thread1)][7].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"input": 2},
|
||||
next=("one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": 1,
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "input", "step": -1, "writes": 2},
|
||||
created_at=AnyStr(),
|
||||
parent_config=None,
|
||||
),
|
||||
]
|
||||
input_channels="input",
|
||||
output_channels="output",
|
||||
checkpointer=checkpointer,
|
||||
interrupt_after_nodes=["one"],
|
||||
)
|
||||
thread1 = {"configurable": {"thread_id": "1"}}
|
||||
thread2 = {"configurable": {"thread_id": "2"}}
|
||||
|
||||
# start execution, stop at inbox
|
||||
assert await app.ainvoke(2, thread1) is None
|
||||
|
||||
# inbox == 3
|
||||
checkpoint = await checkpointer.aget(thread1)
|
||||
assert checkpoint is not None
|
||||
assert checkpoint["channel_values"]["inbox"] == 3
|
||||
|
||||
# resume execution, finish
|
||||
assert await app.ainvoke(None, thread1) == 4
|
||||
|
||||
# start execution again, stop at inbox
|
||||
assert await app.ainvoke(20, thread1) is None
|
||||
|
||||
# inbox == 21
|
||||
checkpoint = await checkpointer.aget(thread1)
|
||||
assert checkpoint is not None
|
||||
assert checkpoint["channel_values"]["inbox"] == 21
|
||||
|
||||
# send a new value in, interrupting the previous execution
|
||||
assert await app.ainvoke(3, thread1) is None
|
||||
assert await app.ainvoke(None, thread1) == 5
|
||||
|
||||
# start execution again, stopping at inbox
|
||||
assert await app.ainvoke(20, thread2) is None
|
||||
|
||||
# inbox == 21
|
||||
snapshot = await app.aget_state(thread2)
|
||||
assert snapshot.values["inbox"] == 21
|
||||
assert snapshot.next == ("two",)
|
||||
|
||||
# update the state, resume
|
||||
await app.aupdate_state(thread2, 25, as_node="one")
|
||||
assert await app.ainvoke(None, thread2) == 26
|
||||
|
||||
# no pending tasks
|
||||
snapshot = await app.aget_state(thread2)
|
||||
assert snapshot.next == ()
|
||||
|
||||
# list history
|
||||
history = [c async for c in app.aget_state_history(thread1)]
|
||||
assert history == [
|
||||
StateSnapshot(
|
||||
values={"inbox": 4, "output": 5, "input": 3},
|
||||
next=(),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 6, "writes": 5},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[1].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 4, "output": 4, "input": 3},
|
||||
next=("two",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 5, "writes": None},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[2].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 21, "output": 4, "input": 3},
|
||||
next=("one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "input", "step": 4, "writes": 3},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[3].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 21, "output": 4, "input": 20},
|
||||
next=("two",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 3, "writes": None},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[4].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 3, "output": 4, "input": 20},
|
||||
next=("one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "input", "step": 2, "writes": 20},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[5].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 3, "output": 4, "input": 2},
|
||||
next=(),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 1, "writes": 4},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[6].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"inbox": 3, "input": 2},
|
||||
next=("two",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 0, "writes": None},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[7].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values={"input": 2},
|
||||
next=("one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "input", "step": -1, "writes": 2},
|
||||
created_at=AnyStr(),
|
||||
parent_config=None,
|
||||
),
|
||||
]
|
||||
finally:
|
||||
if hasattr(checkpointer, "__aexit__"):
|
||||
await checkpointer.__aexit__(None, None, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"checkpointer",
|
||||
[
|
||||
MemorySaverAssertImmutable(),
|
||||
AsyncSqliteSaver.from_conn_string(":memory:"),
|
||||
],
|
||||
ids=[
|
||||
"memory",
|
||||
"sqlite",
|
||||
],
|
||||
)
|
||||
async def test_fork_always_re_runs_nodes(
|
||||
checkpointer: BaseCheckpointSaver, mocker: MockerFixture
|
||||
) -> None:
|
||||
try:
|
||||
add_one = mocker.Mock(side_effect=lambda _: 1)
|
||||
|
||||
builder = StateGraph(Annotated[int, operator.add])
|
||||
builder.add_node("add_one", add_one)
|
||||
builder.add_edge(START, "add_one")
|
||||
builder.add_conditional_edges(
|
||||
"add_one", lambda cnt: "add_one" if cnt < 6 else END
|
||||
)
|
||||
graph = builder.compile(checkpointer=checkpointer)
|
||||
|
||||
thread1 = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
# start execution, stop at inbox
|
||||
assert [
|
||||
c
|
||||
async for c in graph.astream(1, thread1, stream_mode=["values", "updates"])
|
||||
] == [
|
||||
("values", 1),
|
||||
("updates", {"add_one": 1}),
|
||||
("values", 2),
|
||||
("updates", {"add_one": 1}),
|
||||
("values", 3),
|
||||
("updates", {"add_one": 1}),
|
||||
("values", 4),
|
||||
("updates", {"add_one": 1}),
|
||||
("values", 5),
|
||||
("updates", {"add_one": 1}),
|
||||
("values", 6),
|
||||
]
|
||||
|
||||
# list history
|
||||
history = [c async for c in graph.aget_state_history(thread1)]
|
||||
assert history == [
|
||||
StateSnapshot(
|
||||
values=6,
|
||||
next=(),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 5, "writes": {"add_one": 1}},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[1].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values=5,
|
||||
next=("add_one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 4, "writes": {"add_one": 1}},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[2].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values=4,
|
||||
next=("add_one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 3, "writes": {"add_one": 1}},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[3].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values=3,
|
||||
next=("add_one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 2, "writes": {"add_one": 1}},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[4].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values=2,
|
||||
next=("add_one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 1, "writes": {"add_one": 1}},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[5].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values=1,
|
||||
next=("add_one",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 0, "writes": None},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[6].config,
|
||||
),
|
||||
StateSnapshot(
|
||||
values=0,
|
||||
next=("__start__",),
|
||||
config={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
"thread_ts": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "input", "step": -1, "writes": 1},
|
||||
created_at=AnyStr(),
|
||||
parent_config=None,
|
||||
),
|
||||
]
|
||||
|
||||
# forking from any previous checkpoint w/out forking should do nothing
|
||||
assert [
|
||||
c
|
||||
async for c in graph.astream(None, history[0].config, stream_mode="updates")
|
||||
] == []
|
||||
assert [
|
||||
c
|
||||
async for c in graph.astream(None, history[1].config, stream_mode="updates")
|
||||
] == []
|
||||
|
||||
# forking and re-running from any prev checkpoint should re-run nodes
|
||||
fork_config = await graph.aupdate_state(history[0].config, None)
|
||||
assert [
|
||||
c async for c in graph.astream(None, fork_config, stream_mode="updates")
|
||||
] == []
|
||||
|
||||
fork_config = await graph.aupdate_state(history[1].config, None)
|
||||
assert [
|
||||
c async for c in graph.astream(None, fork_config, stream_mode="updates")
|
||||
] == [{"add_one": 1}]
|
||||
|
||||
fork_config = await graph.aupdate_state(history[2].config, None)
|
||||
assert [
|
||||
c async for c in graph.astream(None, fork_config, stream_mode="updates")
|
||||
] == [
|
||||
{"add_one": 1},
|
||||
{"add_one": 1},
|
||||
]
|
||||
finally:
|
||||
if hasattr(checkpointer, "__aexit__"):
|
||||
await checkpointer.__aexit__(None, None, None)
|
||||
|
||||
|
||||
async def test_invoke_two_processes_in_dict_out(mocker: MockerFixture) -> None:
|
||||
|
||||
Reference in New Issue
Block a user