Get it working with interrupt (sync)

This commit is contained in:
Nuno Campos
2024-12-04 15:37:56 -08:00
parent 01a3c23a29
commit 872f54adf1
7 changed files with 110 additions and 49 deletions
+2
View File
@@ -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")
+3 -2
View File
@@ -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
+17 -4
View File
@@ -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] = {}
+15 -15
View File
@@ -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)
+2 -3
View File
@@ -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,
+44 -18
View File
@@ -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
+27 -7
View File
@@ -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)