diff --git a/libs/langgraph/langgraph/constants.py b/libs/langgraph/langgraph/constants.py index b9d6f37f7..cd847f9be 100644 --- a/libs/langgraph/langgraph/constants.py +++ b/libs/langgraph/langgraph/constants.py @@ -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") diff --git a/libs/langgraph/langgraph/pregel/algo.py b/libs/langgraph/langgraph/pregel/algo.py index c95c699c8..78242a9ac 100644 --- a/libs/langgraph/langgraph/pregel/algo.py +++ b/libs/langgraph/langgraph/pregel/algo.py @@ -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 diff --git a/libs/langgraph/langgraph/pregel/call.py b/libs/langgraph/langgraph/pregel/call.py index 8932697b1..6d218f702 100644 --- a/libs/langgraph/langgraph/pregel/call.py +++ b/libs/langgraph/langgraph/pregel/call.py @@ -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] = {} diff --git a/libs/langgraph/langgraph/pregel/io.py b/libs/langgraph/langgraph/pregel/io.py index b2596d3ad..df841ffd5 100644 --- a/libs/langgraph/langgraph/pregel/io.py +++ b/libs/langgraph/langgraph/pregel/io.py @@ -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) diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index 31397d22b..cb141c4a1 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -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, diff --git a/libs/langgraph/langgraph/pregel/runner.py b/libs/langgraph/langgraph/pregel/runner.py index 6155ce74e..3dec1d84c 100644 --- a/libs/langgraph/langgraph/pregel/runner.py +++ b/libs/langgraph/langgraph/pregel/runner.py @@ -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 diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 356ecbadf..49ee6bec5 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -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)