lib: Add interrupts to stream_mode=updates

This commit is contained in:
Nuno Campos
2024-10-11 14:28:16 -07:00
parent 0557fb03a4
commit c6a450b857
4 changed files with 73 additions and 18 deletions
+3 -8
View File
@@ -346,10 +346,7 @@ class PregelLoop:
# after execution, check if we should interrupt
if should_interrupt(self.checkpoint, interrupt_after, self.tasks.values()):
self.status = "interrupt_after"
if self.is_nested:
raise GraphInterrupt()
else:
return False
raise GraphInterrupt()
else:
return False
@@ -441,10 +438,7 @@ class PregelLoop:
# before execution, check if we should interrupt
if should_interrupt(self.checkpoint, interrupt_before, self.tasks.values()):
self.status = "interrupt_before"
if self.is_nested:
raise GraphInterrupt()
else:
return False
raise GraphInterrupt()
# produce debug output
self._emit("debug", map_debug_tasks, self.step, self.tasks.values())
@@ -598,6 +592,7 @@ class PregelLoop:
self.output = read_channels(self.channels, self.output_keys)
if suppress:
# suppress interrupt
self._emit("updates", lambda: iter([{INTERRUPT: exc_value.args[0]}]))
return True
def _emit(
+37 -2
View File
@@ -896,6 +896,7 @@ def test_invoke_two_processes_in_out_interrupt(
]
assert [c for c in app.stream(None, history[2].config, stream_mode="updates")] == [
{"one": {"inbox": 4}},
{"__interrupt__": ()},
]
@@ -3198,6 +3199,7 @@ def test_conditional_state_graph(
),
}
},
{"__interrupt__": ()},
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
@@ -3295,6 +3297,7 @@ def test_conditional_state_graph(
),
}
},
{"__interrupt__": ()},
]
with assert_ctx_once():
@@ -3365,6 +3368,7 @@ def test_conditional_state_graph(
),
}
},
{"__interrupt__": ()},
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
@@ -3460,6 +3464,7 @@ def test_conditional_state_graph(
),
}
},
{"__interrupt__": ()},
]
app_w_interrupt.update_state(
@@ -3520,7 +3525,9 @@ def test_conditional_state_graph(
assert [
c for c in app_w_interrupt.stream({"input": "what is weather in sf"}, config)
] == []
] == [
{"__interrupt__": ()},
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
values={
@@ -3542,6 +3549,7 @@ def test_conditional_state_graph(
),
}
},
{"__interrupt__": ()},
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
@@ -3587,6 +3595,7 @@ def test_conditional_state_graph(
],
}
},
{"__interrupt__": ()},
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
@@ -3641,6 +3650,7 @@ def test_conditional_state_graph(
),
}
},
{"__interrupt__": ()},
]
# test w interrupt after all
@@ -3661,6 +3671,7 @@ def test_conditional_state_graph(
),
}
},
{"__interrupt__": ()},
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
@@ -3706,6 +3717,7 @@ def test_conditional_state_graph(
],
}
},
{"__interrupt__": ()},
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
@@ -3760,6 +3772,7 @@ def test_conditional_state_graph(
),
}
},
{"__interrupt__": ()},
]
@@ -4630,6 +4643,7 @@ def test_state_graph_packets(
)
}
},
{"__interrupt__": ()},
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
@@ -4759,6 +4773,7 @@ def test_state_graph_packets(
)
},
},
{"__interrupt__": ()},
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
@@ -5130,6 +5145,7 @@ def test_message_graph(
id="ai1",
)
},
{"__interrupt__": ()},
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
@@ -5241,6 +5257,7 @@ def test_message_graph(
id="ai2",
)
},
{"__interrupt__": ()},
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
@@ -5360,6 +5377,7 @@ def test_message_graph(
id="ai1",
)
},
{"__interrupt__": ()},
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
@@ -5471,6 +5489,7 @@ def test_message_graph(
id="ai2",
)
},
{"__interrupt__": ()},
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
@@ -5856,6 +5875,7 @@ def test_root_graph(
id="ai1",
)
},
{"__interrupt__": ()},
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
@@ -5967,6 +5987,7 @@ def test_root_graph(
id="ai2",
)
},
{"__interrupt__": ()},
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
@@ -6088,6 +6109,7 @@ def test_root_graph(
id="ai1",
)
},
{"__interrupt__": ()},
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
@@ -6199,6 +6221,7 @@ def test_root_graph(
id="ai2",
)
},
{"__interrupt__": ()},
]
assert app_w_interrupt.get_state(config) == StateSnapshot(
@@ -7653,6 +7676,7 @@ def test_in_one_fan_out_state_graph_waiting_edge(
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
{"retriever_two": {"docs": ["doc3", "doc4"]}},
{"retriever_one": {"docs": ["doc1", "doc2"]}},
{"__interrupt__": ()},
]
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -7672,6 +7696,7 @@ def test_in_one_fan_out_state_graph_waiting_edge(
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
{"retriever_two": {"docs": ["doc3", "doc4"]}},
{"retriever_one": {"docs": ["doc1", "doc2"]}},
{"__interrupt__": ()},
]
app_w_interrupt.update_state(config, {"docs": ["doc5"]})
@@ -7785,6 +7810,7 @@ def test_in_one_fan_out_state_graph_waiting_edge_via_branch(
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
{"retriever_two": {"docs": ["doc3", "doc4"]}},
{"retriever_one": {"docs": ["doc1", "doc2"]}},
{"__interrupt__": ()},
]
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -7941,6 +7967,7 @@ def test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic1(
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
{"retriever_two": {"docs": ["doc3", "doc4"]}},
{"retriever_one": {"docs": ["doc1", "doc2"]}},
{"__interrupt__": ()},
]
with assert_ctx_once():
@@ -8109,6 +8136,7 @@ def test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic2(
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
{"retriever_two": {"docs": ["doc3", "doc4"]}},
{"retriever_one": {"docs": ["doc1", "doc2"]}},
{"__interrupt__": ()},
]
with assert_ctx_once():
@@ -8217,6 +8245,7 @@ def test_in_one_fan_out_state_graph_waiting_edge_plus_regular(
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
{"retriever_two": {"docs": ["doc3", "doc4"]}},
{"retriever_one": {"docs": ["doc1", "doc2"]}},
{"__interrupt__": ()},
]
assert [c for c in app_w_interrupt.stream(None, config)] == [
@@ -8779,6 +8808,7 @@ def test_nested_graph_interrupts_parallel(
# we got to parallel node first
((), {"outer_1": {"my_key": " and parallel"}}),
((AnyStr("inner:"),), {"inner_1": {"my_key": "got here", "my_other_key": ""}}),
((), {"__interrupt__": ()}),
]
assert [*app.stream(None, config)] == [
{"outer_1": {"my_key": " and parallel"}, "__metadata__": {"cached": True}},
@@ -8898,6 +8928,7 @@ def test_doubly_nested_graph_interrupts(
config = {"configurable": {"thread_id": "2"}}
assert [*app.stream({"my_key": "my value"}, config)] == [
{"parent_1": {"my_key": "hi my value"}},
{"__interrupt__": ()},
]
assert [*app.stream(None, config)] == [
{"child": {"my_key": "hi my value here and there"}},
@@ -9493,6 +9524,7 @@ def test_doubly_nested_graph_state(
(AnyStr("child:"), AnyStr("child_1:")),
{"grandchild_1": {"my_key": "hi my value here"}},
),
((), {"__interrupt__": ()}),
]
# get state without subgraphs
outer_state = app.get_state(config)
@@ -10192,7 +10224,8 @@ def test_doubly_nested_graph_state(
(
(AnyStr("child:"), AnyStr("child_1:")),
{"grandchild_1": {"my_key": "hi my value here"}},
)
),
((), {"__interrupt__": ()}),
]
@@ -10644,6 +10677,7 @@ def test_weather_subgraph(
] == [
((), {"router_node": {"route": "weather"}}),
((AnyStr("weather_graph:"),), {"model_node": {"city": "San Francisco"}}),
((), {"__interrupt__": ()}),
]
# check current state
@@ -10732,6 +10766,7 @@ def test_weather_subgraph(
] == [
((), {"router_node": {"route": "weather"}}),
((AnyStr("weather_graph:"),), {"model_node": {"city": "San Francisco"}}),
((), {"__interrupt__": ()}),
]
state = graph.get_state(config, subgraphs=True)
assert state == StateSnapshot(
+29 -8
View File
@@ -291,12 +291,14 @@ async def test_dynamic_interrupt(checkpointer_name: str) -> None:
thread1 = {"configurable": {"thread_id": "1"}}
# stop when about to enter node
assert await tool_two.ainvoke(
{"my_key": "value ⛰️", "market": "DE"}, thread1
) == {
"my_key": "value ⛰️",
"market": "DE",
}
assert [
c
async for c in tool_two.astream(
{"my_key": "value ⛰️", "market": "DE"}, thread1
)
] == [
{"__interrupt__": [Interrupt(value="Just because...", when="during")]},
]
assert [c.metadata async for c in tool_two.checkpointer.alist(thread1)] == [
{
"parents": {},
@@ -330,7 +332,6 @@ async def test_dynamic_interrupt(checkpointer_name: str) -> None:
c async for c in tool_two.checkpointer.alist(thread1, limit=2)
][-1].config,
)
# TODO use aget_state_history
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
@@ -1103,6 +1104,7 @@ async def test_invoke_two_processes_in_out_interrupt(
c async for c in app.astream(None, history[2].config, stream_mode="updates")
] == [
{"one": {"inbox": 4}},
{"__interrupt__": ()},
]
@@ -3452,6 +3454,7 @@ async def test_conditional_graph_state(
),
}
},
{"__interrupt__": ()},
]
assert await app_w_interrupt.aget_state(config) == StateSnapshot(
@@ -3559,6 +3562,7 @@ async def test_conditional_graph_state(
),
}
},
{"__interrupt__": ()},
]
async with assert_ctx_once():
@@ -3637,6 +3641,7 @@ async def test_conditional_graph_state(
),
}
},
{"__interrupt__": ()},
]
assert await app_w_interrupt.aget_state(config) == StateSnapshot(
@@ -3740,6 +3745,7 @@ async def test_conditional_graph_state(
),
}
},
{"__interrupt__": ()},
]
await app_w_interrupt.aupdate_state(
@@ -4448,6 +4454,7 @@ async def test_state_graph_packets(checkpointer_name: str) -> None:
)
}
},
{"__interrupt__": ()},
]
assert await app_w_interrupt.aget_state(config) == StateSnapshot(
@@ -4581,6 +4588,7 @@ async def test_state_graph_packets(checkpointer_name: str) -> None:
)
},
},
{"__interrupt__": ()},
]
tup = await app_w_interrupt.checkpointer.aget_tuple(config)
@@ -4922,6 +4930,7 @@ async def test_message_graph(checkpointer_name: str) -> None:
id="ai1",
)
},
{"__interrupt__": ()},
]
tup = await app_w_interrupt.checkpointer.aget_tuple(config)
@@ -5039,6 +5048,7 @@ async def test_message_graph(checkpointer_name: str) -> None:
id="ai2",
)
},
{"__interrupt__": ()},
]
tup = await app_w_interrupt.checkpointer.aget_tuple(config)
@@ -6434,6 +6444,7 @@ async def test_in_one_fan_out_state_graph_waiting_edge(checkpointer_name: str) -
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
{"retriever_two": {"docs": ["doc3", "doc4"]}},
{"retriever_one": {"docs": ["doc1", "doc2"]}},
{"__interrupt__": ()},
]
assert [c async for c in app_w_interrupt.astream(None, config)] == [
@@ -6525,6 +6536,7 @@ async def test_in_one_fan_out_state_graph_waiting_edge_via_branch(
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
{"retriever_two": {"docs": ["doc3", "doc4"]}},
{"retriever_one": {"docs": ["doc1", "doc2"]}},
{"__interrupt__": ()},
]
assert [c async for c in app_w_interrupt.astream(None, config)] == [
@@ -6668,6 +6680,7 @@ async def test_in_one_fan_out_state_graph_waiting_edge_custom_state_class(
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
{"retriever_two": {"docs": ["doc3", "doc4"]}},
{"retriever_one": {"docs": ["doc1", "doc2"]}},
{"__interrupt__": ()},
]
async with assert_ctx_once():
@@ -6833,6 +6846,7 @@ async def test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydant
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
{"retriever_two": {"docs": ["doc3", "doc4"]}},
{"retriever_one": {"docs": ["doc1", "doc2"]}},
{"__interrupt__": ()},
]
assert [c async for c in app_w_interrupt.astream(None, config)] == [
@@ -6939,6 +6953,7 @@ async def test_in_one_fan_out_state_graph_waiting_edge_plus_regular(
{"analyzer_one": {"query": "analyzed: query: what is weather in sf"}},
{"retriever_two": {"docs": ["doc3", "doc4"]}},
{"retriever_one": {"docs": ["doc1", "doc2"]}},
{"__interrupt__": ()},
]
assert [c async for c in app_w_interrupt.astream(None, config)] == [
@@ -7419,6 +7434,7 @@ async def test_nested_graph_interrupts_parallel(checkpointer_name: str) -> None:
(AnyStr("inner:"),),
{"inner_1": {"my_key": "got here", "my_other_key": ""}},
),
((), {"__interrupt__": ()}),
]
assert [c async for c in app.astream(None, config)] == [
{"outer_1": {"my_key": " and parallel"}, "__metadata__": {"cached": True}},
@@ -7541,6 +7557,7 @@ async def test_doubly_nested_graph_interrupts(checkpointer_name: str) -> None:
config = {"configurable": {"thread_id": "2"}}
assert [c async for c in app.astream({"my_key": "my value"}, config)] == [
{"parent_1": {"my_key": "hi my value"}},
{"__interrupt__": ()},
]
assert [c async for c in app.astream(None, config)] == [
{"child": {"my_key": "hi my value here and there"}},
@@ -8163,6 +8180,7 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
(AnyStr("child:"), AnyStr("child_1:")),
{"grandchild_1": {"my_key": "hi my value here"}},
),
((), {"__interrupt__": ()}),
]
# get state without subgraphs
outer_state = await app.aget_state(config)
@@ -8895,7 +8913,8 @@ async def test_doubly_nested_graph_state(checkpointer_name: str) -> None:
(
(AnyStr("child:"), AnyStr("child_1:")),
{"grandchild_1": {"my_key": "hi my value here"}},
)
),
((), {"__interrupt__": ()}),
]
@@ -9297,6 +9316,7 @@ async def test_weather_subgraph(
] == [
((), {"router_node": {"route": "weather"}}),
((AnyStr("weather_graph:"),), {"model_node": {"city": "San Francisco"}}),
((), {"__interrupt__": ()}),
]
# check current state
@@ -9389,6 +9409,7 @@ async def test_weather_subgraph(
] == [
((), {"router_node": {"route": "weather"}}),
((AnyStr("weather_graph:"),), {"model_node": {"city": "San Francisco"}}),
((), {"__interrupt__": ()}),
]
state = await graph.aget_state(config, subgraphs=True)
assert state == StateSnapshot(
+4
View File
@@ -2,10 +2,14 @@ from __future__ import annotations
from typing import Any
import pytest
from langgraph.store.base import BaseStore
from langgraph.types import StreamWriter
from langgraph.utils.runnable import RunnableCallable
pytestmark = pytest.mark.anyio
def test_runnable_callable_func_accepts():
def sync_func(x: Any) -> str: