mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-13 13:17:52 +02:00
lib: Add interrupts to stream_mode=updates
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user