mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-29 04:55:09 +02:00
Get it working with interrupt (sync)
This commit is contained in:
@@ -40,6 +40,8 @@ SCHEDULED = sys.intern("__scheduled__")
|
||||
# marker to signal node was scheduled (in distributed mode)
|
||||
TASKS = sys.intern("__pregel_tasks")
|
||||
# for Send objects returned by nodes/edges, corresponds to PUSH below
|
||||
RETURN = sys.intern("__return__")
|
||||
# for writes of a task where we simply record the return value
|
||||
|
||||
# --- Reserved config.configurable keys ---
|
||||
CONFIG_KEY_SEND = sys.intern("__pregel_send")
|
||||
|
||||
@@ -52,6 +52,7 @@ from langgraph.constants import (
|
||||
PUSH,
|
||||
RESERVED,
|
||||
RESUME,
|
||||
RETURN,
|
||||
TAG_HIDDEN,
|
||||
TASKS,
|
||||
Send,
|
||||
@@ -259,7 +260,7 @@ def apply_writes(
|
||||
pending_writes_by_managed: dict[str, list[Any]] = defaultdict(list)
|
||||
for task in tasks:
|
||||
for chan, val in task.writes:
|
||||
if chan in (NO_WRITES, PUSH, RESUME, INTERRUPT):
|
||||
if chan in (NO_WRITES, PUSH, RESUME, INTERRUPT, RETURN):
|
||||
pass
|
||||
elif chan == TASKS: # TODO: remove branch in 1.0
|
||||
checkpoint["pending_sends"].append(val)
|
||||
@@ -509,7 +510,7 @@ def prepare_single_task(
|
||||
proc,
|
||||
writes,
|
||||
patch_config(
|
||||
merge_configs(config, {"metadata": metadata, "tags": proc.tags}),
|
||||
merge_configs(config, {"metadata": metadata}),
|
||||
run_name=name,
|
||||
callbacks=(
|
||||
manager.get_child(f"graph:step:{step}") if manager else None
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
import sys
|
||||
import types
|
||||
|
||||
from langgraph.utils.runnable import RunnableCallable
|
||||
from langgraph.constants import RETURN
|
||||
from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry
|
||||
from langgraph.utils.runnable import RunnableCallable, RunnableSeq
|
||||
|
||||
"""
|
||||
Utilities borrowed from cloudpickle.
|
||||
@@ -101,13 +103,24 @@ def _lookup_module_and_qualname(obj, name=None):
|
||||
|
||||
def get_runnable_for_func(
|
||||
func: types.FunctionType,
|
||||
) -> RunnableCallable:
|
||||
) -> RunnableSeq:
|
||||
if func in CACHE:
|
||||
return CACHE[func]
|
||||
elif not _lookup_module_and_qualname(func):
|
||||
return RunnableCallable(func)
|
||||
return RunnableSeq(
|
||||
RunnableCallable(func, trace=False),
|
||||
ChannelWrite([ChannelWriteEntry(RETURN)]),
|
||||
name=func.__name__,
|
||||
)
|
||||
else:
|
||||
return CACHE.setdefault(func, RunnableCallable(func))
|
||||
return CACHE.setdefault(
|
||||
func,
|
||||
RunnableSeq(
|
||||
RunnableCallable(func, trace=False),
|
||||
ChannelWrite([ChannelWriteEntry(RETURN)]),
|
||||
name=func.__name__,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
CACHE: dict[types.FunctionType, RunnableCallable] = {}
|
||||
|
||||
@@ -13,6 +13,7 @@ from langgraph.constants import (
|
||||
NULL_TASK_ID,
|
||||
PUSH,
|
||||
RESUME,
|
||||
RETURN,
|
||||
TAG_HIDDEN,
|
||||
TASKS,
|
||||
)
|
||||
@@ -171,22 +172,21 @@ def map_output_updates(
|
||||
]
|
||||
if not output_tasks:
|
||||
return
|
||||
if isinstance(output_channels, str):
|
||||
updated = (
|
||||
(task.name, value)
|
||||
for task, writes in output_tasks
|
||||
for chan, value in writes
|
||||
if chan == output_channels
|
||||
)
|
||||
else:
|
||||
updated = (
|
||||
(
|
||||
task.name,
|
||||
{chan: value for chan, value in writes if chan in output_channels},
|
||||
updated: list[tuple[str, Any]] = []
|
||||
for task, writes in output_tasks:
|
||||
if rtn := next((value for chan, value in writes if chan == RETURN), None):
|
||||
updated.append((task.name, rtn))
|
||||
elif isinstance(output_channels, str):
|
||||
updated.extend(
|
||||
(task.name, value) for chan, value in writes if chan == output_channels
|
||||
)
|
||||
elif any(chan in output_channels for chan, _ in writes):
|
||||
updated.append(
|
||||
(
|
||||
task.name,
|
||||
{chan: value for chan, value in writes if chan in output_channels},
|
||||
)
|
||||
)
|
||||
for task, writes in output_tasks
|
||||
if any(chan in output_channels for chan, _ in writes)
|
||||
)
|
||||
grouped: dict[str, list[Any]] = {t.name: [] for t, _ in output_tasks}
|
||||
for node, value in updated:
|
||||
grouped[node].append(value)
|
||||
|
||||
@@ -347,9 +347,8 @@ class PregelLoop(LoopProtocol):
|
||||
# match any pending writes to the new task
|
||||
if self.skip_done_tasks:
|
||||
self._match_writes({pushed.id: pushed})
|
||||
# return the new task, to be started, if not run before
|
||||
if not pushed.writes:
|
||||
return pushed
|
||||
# return the new task, to be started if not run before
|
||||
return pushed
|
||||
|
||||
def tick(
|
||||
self,
|
||||
|
||||
@@ -25,6 +25,7 @@ from langgraph.constants import (
|
||||
NO_WRITES,
|
||||
PUSH,
|
||||
RESUME,
|
||||
RETURN,
|
||||
TAG_HIDDEN,
|
||||
)
|
||||
from langgraph.errors import GraphBubbleUp, GraphInterrupt
|
||||
@@ -88,23 +89,42 @@ class PregelRunner:
|
||||
):
|
||||
# if the parent task was retried,
|
||||
# the next task might already be running
|
||||
if any(
|
||||
t == next_task.id for t in futures.values() if t is not None
|
||||
if fut := next(
|
||||
(
|
||||
f
|
||||
for f, t in futures.items()
|
||||
if t is not None and t == next_task.id
|
||||
),
|
||||
None,
|
||||
):
|
||||
continue
|
||||
# schedule the next task
|
||||
fut = self.submit(
|
||||
run_with_retry,
|
||||
next_task,
|
||||
retry_policy,
|
||||
configurable={
|
||||
CONFIG_KEY_SEND: partial(writer, next_task),
|
||||
CONFIG_KEY_CALL: partial(call, next_task),
|
||||
},
|
||||
__reraise_on_exit__=reraise,
|
||||
)
|
||||
futures[fut] = next_task
|
||||
rtn[idx - prev_length] = fut
|
||||
rtn[idx - prev_length] = fut
|
||||
elif next_task.writes:
|
||||
fut = concurrent.futures.Future()
|
||||
if val := next(v for c, v in next_task.writes if c == RETURN):
|
||||
fut.set_result(val)
|
||||
elif exc := next(v for c, v in next_task.writes if c == ERROR):
|
||||
fut.set_exception(
|
||||
exc
|
||||
if isinstance(exc, BaseException)
|
||||
else Exception(exc)
|
||||
)
|
||||
else:
|
||||
fut.set_result(None)
|
||||
rtn[idx - prev_length] = fut
|
||||
else:
|
||||
# schedule the next task
|
||||
fut = self.submit(
|
||||
run_with_retry,
|
||||
next_task,
|
||||
retry_policy,
|
||||
configurable={
|
||||
CONFIG_KEY_SEND: partial(writer, next_task),
|
||||
CONFIG_KEY_CALL: partial(call, next_task),
|
||||
},
|
||||
__reraise_on_exit__=reraise,
|
||||
)
|
||||
futures[fut] = next_task
|
||||
rtn[idx - prev_length] = fut
|
||||
return [rtn.get(i) for i in range(len(writes))]
|
||||
|
||||
def call(
|
||||
@@ -116,6 +136,7 @@ class PregelRunner:
|
||||
|
||||
tasks = tuple(tasks)
|
||||
futures: dict[concurrent.futures.Future, Optional[PregelExecutableTask]] = {}
|
||||
done_futures: set[concurrent.futures.Future] = set()
|
||||
# give control back to the caller
|
||||
yield
|
||||
# fast path if single task with no timeout and no waiter
|
||||
@@ -133,7 +154,12 @@ class PregelRunner:
|
||||
self.commit(t, None)
|
||||
except Exception as exc:
|
||||
self.commit(t, exc)
|
||||
if reraise:
|
||||
if reraise and futures:
|
||||
# will be re-raised after futures are done
|
||||
fut = concurrent.futures.Future()
|
||||
fut.set_exception(exc)
|
||||
done_futures.add(fut)
|
||||
elif reraise:
|
||||
raise
|
||||
if not futures: # maybe `t` schuduled another task
|
||||
return
|
||||
@@ -157,7 +183,6 @@ class PregelRunner:
|
||||
__reraise_on_exit__=reraise,
|
||||
)
|
||||
] = t
|
||||
done_futures: set[concurrent.futures.Future] = set()
|
||||
end_time = timeout + time.monotonic() if timeout else None
|
||||
while len(futures) > (1 if get_waiter is not None else 0):
|
||||
done, inflight = concurrent.futures.wait(
|
||||
@@ -183,6 +208,7 @@ class PregelRunner:
|
||||
del fut, task
|
||||
# maybe stop other tasks
|
||||
if _should_stop_others(done):
|
||||
print("stopping others")
|
||||
break
|
||||
# give control back to the caller
|
||||
yield
|
||||
|
||||
@@ -1973,24 +1973,44 @@ def test_send_sequences() -> None:
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
|
||||
def test_imp_task(request: pytest.FixtureRequest, checkpointer_name: str) -> None:
|
||||
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
|
||||
mapper_calls = 0
|
||||
|
||||
@task()
|
||||
def mapper(input: str) -> str:
|
||||
print(f"mapper {input}")
|
||||
nonlocal mapper_calls
|
||||
mapper_calls += 1
|
||||
return input * 2
|
||||
|
||||
@imp(checkpointer=checkpointer)
|
||||
def graph(input: list[str]) -> list[str]:
|
||||
futures = [mapper(i) for i in input]
|
||||
mapped = [f.result() for f in futures]
|
||||
# answer = interrupt("question")
|
||||
# TODO raises NodeInterrupt if no answer provided yet
|
||||
# returns answer (saved in writes?) if provided
|
||||
# what is the API for passing the answer?
|
||||
return mapped
|
||||
answer = interrupt("question")
|
||||
return [m + answer for m in mapped]
|
||||
|
||||
thread1 = {"configurable": {"thread_id": "1"}}
|
||||
assert graph.invoke(["0", "1"], thread1) == ["00", "11"]
|
||||
assert [*graph.stream(["0", "1"], thread1)] == [
|
||||
# TODO make test not depend on order of execution (which is not guaranteed)
|
||||
{"mapper": "00"},
|
||||
{"mapper": "11"},
|
||||
{
|
||||
"__interrupt__": (
|
||||
Interrupt(
|
||||
value="question",
|
||||
resumable=True,
|
||||
ns=[AnyStr("graph:")],
|
||||
when="during",
|
||||
),
|
||||
)
|
||||
},
|
||||
]
|
||||
assert mapper_calls == 2
|
||||
|
||||
assert graph.invoke(Command(resume="answer"), thread1) == [
|
||||
"00answer",
|
||||
"11answer",
|
||||
]
|
||||
assert mapper_calls == 2
|
||||
|
||||
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
|
||||
|
||||
Reference in New Issue
Block a user