mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-23 10:05:08 +02:00
Merge pull request #2378 from langchain-ai/nc/8nov/send-future
Imperative API
This commit is contained in:
@@ -1,5 +1,5 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Generic, Optional, Sequence, Type, TypeVar
|
||||
from typing import Any, Generic, Optional, Sequence, TypeVar
|
||||
|
||||
from typing_extensions import Self
|
||||
|
||||
@@ -13,7 +13,7 @@ C = TypeVar("C")
|
||||
class BaseChannel(Generic[Value, Update, C], ABC):
|
||||
__slots__ = ("key", "typ")
|
||||
|
||||
def __init__(self, typ: Type[Any], key: str = "") -> None:
|
||||
def __init__(self, typ: Any, key: str = "") -> None:
|
||||
self.typ = typ
|
||||
self.key = key
|
||||
|
||||
|
||||
@@ -40,12 +40,16 @@ 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")
|
||||
# holds the `write` function that accepts writes to state/edges/reserved keys
|
||||
CONFIG_KEY_READ = sys.intern("__pregel_read")
|
||||
# holds the `read` function that returns a copy of the current state
|
||||
CONFIG_KEY_CALL = sys.intern("__pregel_call")
|
||||
# holds the `call` function that accepts a node/func, args and returns a future
|
||||
CONFIG_KEY_CHECKPOINTER = sys.intern("__pregel_checkpointer")
|
||||
# holds a `BaseCheckpointSaver` passed from parent graph to child graphs
|
||||
CONFIG_KEY_STREAM = sys.intern("__pregel_stream")
|
||||
|
||||
@@ -0,0 +1,97 @@
|
||||
import asyncio
|
||||
import concurrent
|
||||
import concurrent.futures
|
||||
import types
|
||||
from functools import partial, update_wrapper
|
||||
from typing import (
|
||||
Any,
|
||||
Awaitable,
|
||||
Callable,
|
||||
Optional,
|
||||
TypeVar,
|
||||
Union,
|
||||
overload,
|
||||
)
|
||||
|
||||
from typing_extensions import ParamSpec
|
||||
|
||||
from langgraph.channels.ephemeral_value import EphemeralValue
|
||||
from langgraph.channels.last_value import LastValue
|
||||
from langgraph.checkpoint.base import BaseCheckpointSaver
|
||||
from langgraph.constants import END, START, TAG_HIDDEN
|
||||
from langgraph.pregel import Pregel
|
||||
from langgraph.pregel.call import get_runnable_for_func
|
||||
from langgraph.pregel.read import PregelNode
|
||||
from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry
|
||||
from langgraph.store.base import BaseStore
|
||||
from langgraph.types import RetryPolicy
|
||||
|
||||
P = ParamSpec("P")
|
||||
P1 = TypeVar("P1")
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
def call(
|
||||
func: Callable[[P1], T],
|
||||
input: P1,
|
||||
*,
|
||||
retry: Optional[RetryPolicy] = None,
|
||||
) -> concurrent.futures.Future[T]:
|
||||
from langgraph.constants import CONFIG_KEY_CALL
|
||||
from langgraph.utils.config import get_configurable
|
||||
|
||||
conf = get_configurable()
|
||||
impl = conf[CONFIG_KEY_CALL]
|
||||
fut = impl(func, input, retry=retry)
|
||||
return fut
|
||||
|
||||
|
||||
@overload
|
||||
def task(
|
||||
*, retry: Optional[RetryPolicy] = None
|
||||
) -> Callable[[Callable[P, Awaitable[T]]], Callable[P, asyncio.Future[T]]]: ...
|
||||
|
||||
|
||||
@overload
|
||||
def task( # type: ignore[overload-cannot-match]
|
||||
*, retry: Optional[RetryPolicy] = None
|
||||
) -> Callable[[Callable[P, T]], Callable[P, concurrent.futures.Future[T]]]: ...
|
||||
|
||||
|
||||
def task(
|
||||
*, retry: Optional[RetryPolicy] = None
|
||||
) -> Union[
|
||||
Callable[[Callable[P, Awaitable[T]]], Callable[P, asyncio.Future[T]]],
|
||||
Callable[[Callable[P, T]], Callable[P, concurrent.futures.Future[T]]],
|
||||
]:
|
||||
def _task(func: Callable[P, T]) -> Callable[P, concurrent.futures.Future[T]]:
|
||||
return update_wrapper(partial(call, func, retry=retry), func)
|
||||
|
||||
return _task
|
||||
|
||||
|
||||
def entrypoint(
|
||||
*,
|
||||
checkpointer: Optional[BaseCheckpointSaver] = None,
|
||||
store: Optional[BaseStore] = None,
|
||||
) -> Callable[[types.FunctionType], Pregel]:
|
||||
def _imp(func: types.FunctionType) -> Pregel:
|
||||
return Pregel(
|
||||
nodes={
|
||||
func.__name__: PregelNode(
|
||||
bound=get_runnable_for_func(func),
|
||||
triggers=[START],
|
||||
channels=[START],
|
||||
writers=[ChannelWrite([ChannelWriteEntry(END)], tags=[TAG_HIDDEN])],
|
||||
)
|
||||
},
|
||||
channels={START: EphemeralValue(Any), END: LastValue(Any, END)},
|
||||
input_channels=START,
|
||||
output_channels=END,
|
||||
stream_channels=END,
|
||||
stream_mode="updates",
|
||||
checkpointer=checkpointer,
|
||||
store=store,
|
||||
)
|
||||
|
||||
return _imp
|
||||
@@ -43,6 +43,7 @@ from langgraph.constants import (
|
||||
CONFIG_KEY_TASK_ID,
|
||||
CONFIG_KEY_WRITES,
|
||||
EMPTY_SEQ,
|
||||
ERROR,
|
||||
INTERRUPT,
|
||||
NO_WRITES,
|
||||
NS_END,
|
||||
@@ -52,18 +53,26 @@ from langgraph.constants import (
|
||||
PUSH,
|
||||
RESERVED,
|
||||
RESUME,
|
||||
RETURN,
|
||||
TAG_HIDDEN,
|
||||
TASKS,
|
||||
Send,
|
||||
)
|
||||
from langgraph.errors import EmptyChannelError, InvalidUpdateError
|
||||
from langgraph.managed.base import ManagedValueMapping
|
||||
from langgraph.pregel.call import get_runnable_for_func
|
||||
from langgraph.pregel.io import read_channel, read_channels
|
||||
from langgraph.pregel.log import logger
|
||||
from langgraph.pregel.manager import ChannelsManager
|
||||
from langgraph.pregel.read import PregelNode
|
||||
from langgraph.store.base import BaseStore
|
||||
from langgraph.types import All, LoopProtocol, PregelExecutableTask, PregelTask
|
||||
from langgraph.types import (
|
||||
All,
|
||||
LoopProtocol,
|
||||
PregelExecutableTask,
|
||||
PregelTask,
|
||||
RetryPolicy,
|
||||
)
|
||||
from langgraph.utils.config import merge_configs, patch_config
|
||||
|
||||
GetNextVersion = Callable[[Optional[V], BaseChannel], V]
|
||||
@@ -97,6 +106,21 @@ class PregelTaskWrites(NamedTuple):
|
||||
triggers: Sequence[str]
|
||||
|
||||
|
||||
class Call:
|
||||
__slots__ = ("func", "input", "retry")
|
||||
|
||||
func: Callable
|
||||
input: Any
|
||||
retry: Optional[RetryPolicy]
|
||||
|
||||
def __init__(
|
||||
self, func: Callable, input: Any, *, retry: Optional[RetryPolicy]
|
||||
) -> None:
|
||||
self.func = func
|
||||
self.input = input
|
||||
self.retry = retry
|
||||
|
||||
|
||||
def should_interrupt(
|
||||
checkpoint: Checkpoint,
|
||||
interrupt_nodes: Union[All, Sequence[str]],
|
||||
@@ -179,7 +203,7 @@ def local_write(
|
||||
"""Function injected under CONFIG_KEY_SEND in task config, to write to channels.
|
||||
Validates writes and forwards them to `commit` function."""
|
||||
for chan, value in writes:
|
||||
if chan in (PUSH, TASKS):
|
||||
if chan in (PUSH, TASKS) and value is not None:
|
||||
if not isinstance(value, Send):
|
||||
raise InvalidUpdateError(f"Expected Send, got {value}")
|
||||
if value.node not in process_keys:
|
||||
@@ -247,7 +271,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, ERROR):
|
||||
pass
|
||||
elif chan == TASKS: # TODO: remove branch in 1.0
|
||||
checkpoint["pending_sends"].append(val)
|
||||
@@ -438,7 +462,7 @@ def prepare_next_tasks(
|
||||
|
||||
|
||||
def prepare_single_task(
|
||||
task_path: tuple[Union[str, int, tuple], ...],
|
||||
task_path: tuple[Any, ...],
|
||||
task_id_checksum: Optional[str],
|
||||
*,
|
||||
checkpoint: Checkpoint,
|
||||
@@ -459,7 +483,94 @@ def prepare_single_task(
|
||||
configurable = config.get(CONF, {})
|
||||
parent_ns = configurable.get(CONFIG_KEY_CHECKPOINT_NS, "")
|
||||
|
||||
if task_path[0] == PUSH:
|
||||
if task_path[0] == PUSH and isinstance(task_path[-1], Call):
|
||||
# (PUSH, parent task path, idx of PUSH write, id of parent task, Call)
|
||||
task_path_t = cast(tuple[str, tuple, int, str, Call], task_path)
|
||||
call = task_path_t[-1]
|
||||
proc_ = get_runnable_for_func(call.func)
|
||||
name = proc_.name
|
||||
if name is None:
|
||||
raise ValueError("`call` functions must have a `__name__` attribute")
|
||||
# create task id
|
||||
triggers = [PUSH]
|
||||
checkpoint_ns = f"{parent_ns}{NS_SEP}{name}" if parent_ns else name
|
||||
task_id = _uuid5_str(
|
||||
checkpoint_id,
|
||||
checkpoint_ns,
|
||||
str(step),
|
||||
name,
|
||||
PUSH,
|
||||
_tuple_str(task_path[1]),
|
||||
str(task_path[2]),
|
||||
)
|
||||
task_checkpoint_ns = f"{checkpoint_ns}:{task_id}"
|
||||
metadata = {
|
||||
"langgraph_step": step,
|
||||
"langgraph_node": name,
|
||||
"langgraph_triggers": triggers,
|
||||
"langgraph_path": task_path[:3],
|
||||
"langgraph_checkpoint_ns": task_checkpoint_ns,
|
||||
}
|
||||
if task_id_checksum is not None:
|
||||
assert task_id == task_id_checksum, f"{task_id} != {task_id_checksum}"
|
||||
if for_execution:
|
||||
writes: deque[tuple[str, Any]] = deque()
|
||||
return PregelExecutableTask(
|
||||
name,
|
||||
call.input,
|
||||
proc_,
|
||||
writes,
|
||||
patch_config(
|
||||
merge_configs(config, {"metadata": metadata}),
|
||||
run_name=name,
|
||||
callbacks=(
|
||||
manager.get_child(f"graph:step:{step}") if manager else None
|
||||
),
|
||||
configurable={
|
||||
CONFIG_KEY_TASK_ID: task_id,
|
||||
# deque.extend is thread-safe
|
||||
CONFIG_KEY_SEND: partial(
|
||||
local_write,
|
||||
writes.extend,
|
||||
processes.keys(),
|
||||
),
|
||||
CONFIG_KEY_READ: partial(
|
||||
local_read,
|
||||
step,
|
||||
checkpoint,
|
||||
channels,
|
||||
managed,
|
||||
PregelTaskWrites(task_path[:3], name, writes, triggers),
|
||||
config,
|
||||
),
|
||||
CONFIG_KEY_STORE: (store or configurable.get(CONFIG_KEY_STORE)),
|
||||
CONFIG_KEY_CHECKPOINTER: (
|
||||
checkpointer or configurable.get(CONFIG_KEY_CHECKPOINTER)
|
||||
),
|
||||
CONFIG_KEY_CHECKPOINT_MAP: {
|
||||
**configurable.get(CONFIG_KEY_CHECKPOINT_MAP, {}),
|
||||
parent_ns: checkpoint["id"],
|
||||
},
|
||||
CONFIG_KEY_CHECKPOINT_ID: None,
|
||||
CONFIG_KEY_CHECKPOINT_NS: task_checkpoint_ns,
|
||||
CONFIG_KEY_WRITES: [
|
||||
w
|
||||
for w in pending_writes
|
||||
+ configurable.get(CONFIG_KEY_WRITES, [])
|
||||
if w[0] in (NULL_TASK_ID, task_id)
|
||||
],
|
||||
CONFIG_KEY_SCRATCHPAD: {},
|
||||
},
|
||||
),
|
||||
triggers,
|
||||
call.retry,
|
||||
None,
|
||||
task_id,
|
||||
task_path[:3],
|
||||
)
|
||||
else:
|
||||
return PregelTask(task_id, name, task_path[:3])
|
||||
elif task_path[0] == PUSH:
|
||||
if len(task_path) == 2: # TODO: remove branch in 1.0
|
||||
# legacy SEND tasks, executed in superstep n+1
|
||||
# (PUSH, idx of pending send)
|
||||
@@ -490,17 +601,19 @@ def prepare_single_task(
|
||||
PUSH,
|
||||
str(idx),
|
||||
)
|
||||
elif len(task_path) == 4:
|
||||
elif len(task_path) >= 4:
|
||||
# new PUSH tasks, executed in superstep n
|
||||
# (PUSH, parent task path, idx of PUSH write, id of parent task)
|
||||
task_path_t = cast(tuple[str, tuple, int, str], task_path)
|
||||
writes_for_path = [w for w in pending_writes if w[0] == task_path_t[3]]
|
||||
if task_path_t[2] >= len(writes_for_path):
|
||||
task_path_tt = cast(tuple[str, tuple, int, str], task_path)
|
||||
writes_for_path = [w for w in pending_writes if w[0] == task_path_tt[3]]
|
||||
if task_path_tt[2] >= len(writes_for_path):
|
||||
logger.warning(
|
||||
f"Ignoring invalid write index {task_path[2]} in pending writes"
|
||||
)
|
||||
return
|
||||
packet = writes_for_path[task_path_t[2]][2]
|
||||
packet = writes_for_path[task_path_tt[2]][2]
|
||||
if packet is None:
|
||||
return
|
||||
if not isinstance(packet, Send):
|
||||
logger.warning(
|
||||
f"Ignoring invalid packet type {type(packet)} in pending writes"
|
||||
@@ -533,7 +646,7 @@ def prepare_single_task(
|
||||
"langgraph_step": step,
|
||||
"langgraph_node": packet.node,
|
||||
"langgraph_triggers": triggers,
|
||||
"langgraph_path": task_path,
|
||||
"langgraph_path": task_path[:3],
|
||||
"langgraph_checkpoint_ns": task_checkpoint_ns,
|
||||
}
|
||||
if task_id_checksum is not None:
|
||||
@@ -543,7 +656,7 @@ def prepare_single_task(
|
||||
if node := proc.node:
|
||||
if proc.metadata:
|
||||
metadata.update(proc.metadata)
|
||||
writes: deque[tuple[str, Any]] = deque()
|
||||
writes = deque()
|
||||
return PregelExecutableTask(
|
||||
packet.node,
|
||||
packet.arg,
|
||||
@@ -572,7 +685,7 @@ def prepare_single_task(
|
||||
channels,
|
||||
managed,
|
||||
PregelTaskWrites(
|
||||
task_path, packet.node, writes, triggers
|
||||
task_path[:3], packet.node, writes, triggers
|
||||
),
|
||||
config,
|
||||
),
|
||||
@@ -602,12 +715,11 @@ def prepare_single_task(
|
||||
proc.retry_policy,
|
||||
None,
|
||||
task_id,
|
||||
task_path,
|
||||
task_path[:3],
|
||||
writers=proc.flat_writers,
|
||||
)
|
||||
|
||||
else:
|
||||
return PregelTask(task_id, packet.node, task_path)
|
||||
return PregelTask(task_id, packet.node, task_path[:3])
|
||||
elif task_path[0] == PULL:
|
||||
# (PULL, node name)
|
||||
name = cast(str, task_path[1])
|
||||
@@ -657,7 +769,7 @@ def prepare_single_task(
|
||||
"langgraph_step": step,
|
||||
"langgraph_node": name,
|
||||
"langgraph_triggers": triggers,
|
||||
"langgraph_path": task_path,
|
||||
"langgraph_path": task_path[:3],
|
||||
"langgraph_checkpoint_ns": task_checkpoint_ns,
|
||||
}
|
||||
if task_id_checksum is not None:
|
||||
@@ -696,7 +808,9 @@ def prepare_single_task(
|
||||
checkpoint,
|
||||
channels,
|
||||
managed,
|
||||
PregelTaskWrites(task_path, name, writes, triggers),
|
||||
PregelTaskWrites(
|
||||
task_path[:3], name, writes, triggers
|
||||
),
|
||||
config,
|
||||
),
|
||||
CONFIG_KEY_STORE: (
|
||||
@@ -725,11 +839,11 @@ def prepare_single_task(
|
||||
proc.retry_policy,
|
||||
None,
|
||||
task_id,
|
||||
task_path,
|
||||
task_path[:3],
|
||||
writers=proc.flat_writers,
|
||||
)
|
||||
else:
|
||||
return PregelTask(task_id, name, task_path)
|
||||
return PregelTask(task_id, name, task_path[:3])
|
||||
|
||||
|
||||
def _proc_input(
|
||||
|
||||
@@ -0,0 +1,123 @@
|
||||
import sys
|
||||
import types
|
||||
from typing import Any, Callable, Optional
|
||||
|
||||
from langgraph.constants import RETURN
|
||||
from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry
|
||||
from langgraph.utils.runnable import RunnableSeq, coerce_to_runnable
|
||||
|
||||
"""
|
||||
Utilities borrowed from cloudpickle.
|
||||
https://github.com/cloudpipe/cloudpickle/blob/6220b0ce83ffee5e47e06770a1ee38ca9e47c850/cloudpickle/cloudpickle.py#L265
|
||||
"""
|
||||
|
||||
|
||||
def _getattribute(obj: Any, name: str) -> Any:
|
||||
for subpath in name.split("."):
|
||||
if subpath == "<locals>":
|
||||
raise AttributeError(
|
||||
"Can't get local attribute {!r} on {!r}".format(name, obj)
|
||||
)
|
||||
try:
|
||||
parent = obj
|
||||
obj = getattr(obj, subpath)
|
||||
except AttributeError:
|
||||
raise AttributeError(
|
||||
"Can't get attribute {!r} on {!r}".format(name, obj)
|
||||
) from None
|
||||
return obj, parent
|
||||
|
||||
|
||||
def _whichmodule(obj: Any, name: str) -> Optional[str]:
|
||||
"""Find the module an object belongs to.
|
||||
|
||||
This function differs from ``pickle.whichmodule`` in two ways:
|
||||
- it does not mangle the cases where obj's module is __main__ and obj was
|
||||
not found in any module.
|
||||
- Errors arising during module introspection are ignored, as those errors
|
||||
are considered unwanted side effects.
|
||||
"""
|
||||
module_name = getattr(obj, "__module__", None)
|
||||
|
||||
if module_name is not None:
|
||||
return module_name
|
||||
# Protect the iteration by using a copy of sys.modules against dynamic
|
||||
# modules that trigger imports of other modules upon calls to getattr or
|
||||
# other threads importing at the same time.
|
||||
for module_name, module in sys.modules.copy().items():
|
||||
# Some modules such as coverage can inject non-module objects inside
|
||||
# sys.modules
|
||||
if (
|
||||
module_name == "__main__"
|
||||
or module_name == "__mp_main__"
|
||||
or module is None
|
||||
or not isinstance(module, types.ModuleType)
|
||||
):
|
||||
continue
|
||||
try:
|
||||
if _getattribute(module, name)[0] is obj:
|
||||
return module_name
|
||||
except Exception:
|
||||
pass
|
||||
return None
|
||||
|
||||
|
||||
def _lookup_module_and_qualname(
|
||||
obj: Any, name: Optional[str] = None
|
||||
) -> Optional[tuple[types.ModuleType, str]]:
|
||||
if name is None:
|
||||
name = getattr(obj, "__qualname__", None)
|
||||
if name is None: # pragma: no cover
|
||||
# This used to be needed for Python 2.7 support but is probably not
|
||||
# needed anymore. However we keep the __name__ introspection in case
|
||||
# users of cloudpickle rely on this old behavior for unknown reasons.
|
||||
name = getattr(obj, "__name__", None)
|
||||
if name is None:
|
||||
return None
|
||||
|
||||
module_name = _whichmodule(obj, name)
|
||||
|
||||
if module_name is None:
|
||||
# In this case, obj.__module__ is None AND obj was not found in any
|
||||
# imported module. obj is thus treated as dynamic.
|
||||
return None
|
||||
|
||||
if module_name == "__main__":
|
||||
return None
|
||||
|
||||
# Note: if module_name is in sys.modules, the corresponding module is
|
||||
# assumed importable at unpickling time. See #357
|
||||
module = sys.modules.get(module_name, None)
|
||||
if module is None:
|
||||
# The main reason why obj's module would not be imported is that this
|
||||
# module has been dynamically created, using for example
|
||||
# types.ModuleType. The other possibility is that module was removed
|
||||
# from sys.modules after obj was created/imported. But this case is not
|
||||
# supported, as the standard pickle does not support it either.
|
||||
return None
|
||||
|
||||
try:
|
||||
obj2, parent = _getattribute(module, name)
|
||||
except AttributeError:
|
||||
# obj was not found inside the module it points to
|
||||
return None
|
||||
if obj2 is not obj:
|
||||
return None
|
||||
return module, name
|
||||
|
||||
|
||||
def get_runnable_for_func(func: Callable[..., Any]) -> RunnableSeq:
|
||||
if func in CACHE:
|
||||
return CACHE[func]
|
||||
else:
|
||||
seq = RunnableSeq(
|
||||
coerce_to_runnable(func, name=None, trace=False),
|
||||
ChannelWrite([ChannelWriteEntry(RETURN)]),
|
||||
name=func.__name__,
|
||||
)
|
||||
if not _lookup_module_and_qualname(func):
|
||||
return seq
|
||||
return CACHE.setdefault(func, seq)
|
||||
|
||||
|
||||
CACHE: dict[Callable[..., Any], RunnableSeq] = {}
|
||||
@@ -1,6 +1,7 @@
|
||||
import asyncio
|
||||
import concurrent.futures
|
||||
import sys
|
||||
import time
|
||||
from contextlib import ExitStack
|
||||
from contextvars import copy_context
|
||||
from types import TracebackType
|
||||
@@ -34,6 +35,7 @@ class Submit(Protocol[P, T]):
|
||||
__name__: Optional[str] = None,
|
||||
__cancel_on_exit__: bool = False,
|
||||
__reraise_on_exit__: bool = True,
|
||||
__next_tick__: bool = False,
|
||||
**kwargs: P.kwargs,
|
||||
) -> concurrent.futures.Future[T]: ...
|
||||
|
||||
@@ -58,9 +60,13 @@ class BackgroundExecutor(ContextManager):
|
||||
__name__: Optional[str] = None, # currently not used in sync version
|
||||
__cancel_on_exit__: bool = False, # for sync, can cancel only if not started
|
||||
__reraise_on_exit__: bool = True,
|
||||
__next_tick__: bool = False,
|
||||
**kwargs: P.kwargs,
|
||||
) -> concurrent.futures.Future[T]:
|
||||
task = self.executor.submit(fn, *args, **kwargs)
|
||||
if __next_tick__:
|
||||
task = self.executor.submit(next_tick, fn, *args, **kwargs)
|
||||
else:
|
||||
task = self.executor.submit(fn, *args, **kwargs)
|
||||
self.tasks[task] = (__cancel_on_exit__, __reraise_on_exit__)
|
||||
task.add_done_callback(self.done)
|
||||
return task
|
||||
@@ -137,11 +143,14 @@ class AsyncBackgroundExecutor(AsyncContextManager):
|
||||
__name__: Optional[str] = None,
|
||||
__cancel_on_exit__: bool = False,
|
||||
__reraise_on_exit__: bool = True,
|
||||
__next_tick__: bool = False,
|
||||
**kwargs: P.kwargs,
|
||||
) -> asyncio.Task[T]:
|
||||
coro = cast(Coroutine[None, None, T], fn(*args, **kwargs))
|
||||
if self.semaphore:
|
||||
coro = gated(self.semaphore, coro)
|
||||
if __next_tick__:
|
||||
coro = anext_tick(coro)
|
||||
if self.context_not_supported:
|
||||
task = self.loop.create_task(coro, name=__name__)
|
||||
else:
|
||||
@@ -197,3 +206,15 @@ async def gated(semaphore: asyncio.Semaphore, coro: Coroutine[None, None, T]) ->
|
||||
"""A coroutine that waits for a semaphore before running another coroutine."""
|
||||
async with semaphore:
|
||||
return await coro
|
||||
|
||||
|
||||
def next_tick(fn: Callable[P, T], *args: P.args, **kwargs: P.kwargs) -> T:
|
||||
"""A function that yields control to other threads before running another function."""
|
||||
time.sleep(0)
|
||||
return fn(*args, **kwargs)
|
||||
|
||||
|
||||
async def anext_tick(coro: Coroutine[None, None, T]) -> T:
|
||||
"""A coroutine that yields control to event loop before running another coroutine."""
|
||||
await asyncio.sleep(0)
|
||||
return await coro
|
||||
|
||||
@@ -13,6 +13,7 @@ from langgraph.constants import (
|
||||
NULL_TASK_ID,
|
||||
PUSH,
|
||||
RESUME,
|
||||
RETURN,
|
||||
TAG_HIDDEN,
|
||||
TASKS,
|
||||
)
|
||||
@@ -167,22 +168,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)
|
||||
|
||||
@@ -73,6 +73,7 @@ from langgraph.managed.base import (
|
||||
WritableManagedValue,
|
||||
)
|
||||
from langgraph.pregel.algo import (
|
||||
Call,
|
||||
GetNextVersion,
|
||||
PregelTaskWrites,
|
||||
apply_writes,
|
||||
@@ -289,16 +290,15 @@ class PregelLoop(LoopProtocol):
|
||||
if self.checkpointer_put_writes is not None:
|
||||
self.submit(
|
||||
self.checkpointer_put_writes,
|
||||
{
|
||||
**self.checkpoint_config,
|
||||
CONF: {
|
||||
**self.checkpoint_config[CONF],
|
||||
patch_configurable(
|
||||
self.checkpoint_config,
|
||||
{
|
||||
CONFIG_KEY_CHECKPOINT_NS: self.config[CONF].get(
|
||||
CONFIG_KEY_CHECKPOINT_NS, ""
|
||||
),
|
||||
CONFIG_KEY_CHECKPOINT_ID: self.checkpoint["id"],
|
||||
},
|
||||
},
|
||||
),
|
||||
writes,
|
||||
task_id,
|
||||
)
|
||||
@@ -307,12 +307,9 @@ class PregelLoop(LoopProtocol):
|
||||
self._output_writes(task_id, writes)
|
||||
|
||||
def accept_push(
|
||||
self, task: PregelExecutableTask, write_idx: int
|
||||
self, task: PregelExecutableTask, write_idx: int, call: Optional[Call] = None
|
||||
) -> Optional[PregelExecutableTask]:
|
||||
"""Accept a PUSH from a task, potentially returning a new task to start."""
|
||||
# don't start if an earlier PUSH has already triggered an interrupt
|
||||
if self.to_interrupt:
|
||||
return
|
||||
# don't start if we should interrupt *after* the original task
|
||||
if should_interrupt(self.checkpoint, self.interrupt_after, [task]):
|
||||
self.to_interrupt.append(task)
|
||||
@@ -320,7 +317,7 @@ class PregelLoop(LoopProtocol):
|
||||
if pushed := cast(
|
||||
Optional[PregelExecutableTask],
|
||||
prepare_single_task(
|
||||
(PUSH, task.path, write_idx, task.id),
|
||||
(PUSH, task.path, write_idx, task.id, call),
|
||||
None,
|
||||
checkpoint=self.checkpoint,
|
||||
pending_writes=[(task.id, *w) for w in task.writes],
|
||||
@@ -349,9 +346,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,
|
||||
|
||||
@@ -4,14 +4,12 @@ import random
|
||||
import sys
|
||||
import time
|
||||
from dataclasses import replace
|
||||
from functools import partial
|
||||
from typing import Any, Callable, Optional, Sequence
|
||||
from typing import Any, Optional, Sequence
|
||||
|
||||
from langgraph.constants import (
|
||||
CONF,
|
||||
CONFIG_KEY_CHECKPOINT_NS,
|
||||
CONFIG_KEY_RESUMING,
|
||||
CONFIG_KEY_SEND,
|
||||
NS_SEP,
|
||||
)
|
||||
from langgraph.errors import _SEEN_CHECKPOINT_NS, GraphBubbleUp, ParentCommand
|
||||
@@ -25,25 +23,21 @@ SUPPORTS_EXC_NOTES = sys.version_info >= (3, 11)
|
||||
def run_with_retry(
|
||||
task: PregelExecutableTask,
|
||||
retry_policy: Optional[RetryPolicy],
|
||||
writer: Optional[
|
||||
Callable[[PregelExecutableTask, Sequence[tuple[str, Any]]], None]
|
||||
] = None,
|
||||
configurable: Optional[dict[str, Any]] = None,
|
||||
) -> None:
|
||||
"""Run a task with retries."""
|
||||
retry_policy = task.retry_policy or retry_policy
|
||||
interval = retry_policy.initial_interval if retry_policy else 0
|
||||
attempts = 0
|
||||
config = task.config
|
||||
if writer is not None:
|
||||
config = patch_configurable(config, {CONFIG_KEY_SEND: partial(writer, task)})
|
||||
if configurable is not None:
|
||||
config = patch_configurable(config, configurable)
|
||||
while True:
|
||||
try:
|
||||
# clear any writes from previous attempts
|
||||
task.writes.clear()
|
||||
# run the task
|
||||
task.proc.invoke(task.input, config)
|
||||
# if successful, end
|
||||
break
|
||||
return task.proc.invoke(task.input, config)
|
||||
except ParentCommand as exc:
|
||||
ns: str = config[CONF][CONFIG_KEY_CHECKPOINT_NS]
|
||||
cmd = exc.args[0]
|
||||
@@ -115,17 +109,15 @@ async def arun_with_retry(
|
||||
task: PregelExecutableTask,
|
||||
retry_policy: Optional[RetryPolicy],
|
||||
stream: bool = False,
|
||||
writer: Optional[
|
||||
Callable[[PregelExecutableTask, Sequence[tuple[str, Any]]], None]
|
||||
] = None,
|
||||
configurable: Optional[dict[str, Any]] = None,
|
||||
) -> None:
|
||||
"""Run a task asynchronously with retries."""
|
||||
retry_policy = task.retry_policy or retry_policy
|
||||
interval = retry_policy.initial_interval if retry_policy else 0
|
||||
attempts = 0
|
||||
config = task.config
|
||||
if writer is not None:
|
||||
config = patch_configurable(config, {CONFIG_KEY_SEND: partial(writer, task)})
|
||||
if configurable is not None:
|
||||
config = patch_configurable(config, configurable)
|
||||
while True:
|
||||
try:
|
||||
# clear any writes from previous attempts
|
||||
@@ -134,10 +126,10 @@ async def arun_with_retry(
|
||||
if stream:
|
||||
async for _ in task.proc.astream(task.input, config):
|
||||
pass
|
||||
# if successful, end
|
||||
break
|
||||
else:
|
||||
await task.proc.ainvoke(task.input, config)
|
||||
# if successful, end
|
||||
break
|
||||
return await task.proc.ainvoke(task.input, config)
|
||||
except ParentCommand as exc:
|
||||
ns: str = config[CONF][CONFIG_KEY_CHECKPOINT_NS]
|
||||
cmd = exc.args[0]
|
||||
|
||||
@@ -1,9 +1,12 @@
|
||||
import asyncio
|
||||
import concurrent.futures
|
||||
import threading
|
||||
import time
|
||||
from functools import partial
|
||||
from typing import (
|
||||
Any,
|
||||
AsyncIterator,
|
||||
Awaitable,
|
||||
Callable,
|
||||
Iterable,
|
||||
Iterator,
|
||||
@@ -16,18 +19,22 @@ from typing import (
|
||||
|
||||
from langgraph.constants import (
|
||||
CONF,
|
||||
CONFIG_KEY_CALL,
|
||||
CONFIG_KEY_SEND,
|
||||
ERROR,
|
||||
INTERRUPT,
|
||||
NO_WRITES,
|
||||
PUSH,
|
||||
RESUME,
|
||||
RETURN,
|
||||
TAG_HIDDEN,
|
||||
)
|
||||
from langgraph.errors import GraphBubbleUp, GraphInterrupt
|
||||
from langgraph.pregel.algo import Call
|
||||
from langgraph.pregel.executor import Submit
|
||||
from langgraph.pregel.retry import arun_with_retry, run_with_retry
|
||||
from langgraph.types import PregelExecutableTask, RetryPolicy
|
||||
from langgraph.utils.future import chain_future
|
||||
|
||||
|
||||
class PregelRunner:
|
||||
@@ -41,7 +48,7 @@ class PregelRunner:
|
||||
submit: Submit,
|
||||
put_writes: Callable[[str, Sequence[tuple[str, Any]]], None],
|
||||
schedule_task: Callable[
|
||||
[PregelExecutableTask, int], Optional[PregelExecutableTask]
|
||||
[PregelExecutableTask, int, Optional[Call]], Optional[PregelExecutableTask]
|
||||
],
|
||||
use_astream: bool = False,
|
||||
node_finished: Optional[Callable[[str], None]] = None,
|
||||
@@ -61,73 +68,143 @@ class PregelRunner:
|
||||
retry_policy: Optional[RetryPolicy] = None,
|
||||
get_waiter: Optional[Callable[[], concurrent.futures.Future[None]]] = None,
|
||||
) -> Iterator[None]:
|
||||
locks: dict[str, threading.Lock] = {}
|
||||
|
||||
def writer(
|
||||
task: PregelExecutableTask, writes: Sequence[tuple[str, Any]]
|
||||
) -> None:
|
||||
prev_length = len(task.writes)
|
||||
# delegate to the underlying writer
|
||||
task.config[CONF][CONFIG_KEY_SEND](writes)
|
||||
for idx, w in enumerate(task.writes):
|
||||
# find the index for the newly inserted writes
|
||||
if idx < prev_length:
|
||||
continue
|
||||
assert writes[idx - prev_length] is w
|
||||
task: PregelExecutableTask,
|
||||
writes: Sequence[tuple[str, Any]],
|
||||
*,
|
||||
calls: Optional[Sequence[Call]] = None,
|
||||
) -> Sequence[Optional[concurrent.futures.Future]]:
|
||||
if all(w[0] != PUSH for w in writes):
|
||||
return task.config[CONF][CONFIG_KEY_SEND](writes)
|
||||
|
||||
if task.id not in locks:
|
||||
locks[task.id] = threading.Lock()
|
||||
with locks[task.id]:
|
||||
prev_length = len(task.writes)
|
||||
# delegate to the underlying writer
|
||||
task.config[CONF][CONFIG_KEY_SEND](writes)
|
||||
# confirm no other concurrent writes were added
|
||||
assert len(task.writes) == prev_length + len(writes)
|
||||
# schedule PUSH tasks, collect futures
|
||||
rtn: dict[int, Optional[concurrent.futures.Future]] = {}
|
||||
for idx, w in enumerate(writes, start=prev_length):
|
||||
# bail if not a PUSH write
|
||||
if w[0] != PUSH:
|
||||
continue
|
||||
# schedule the next task, if the callback returns one
|
||||
if next_task := self.schedule_task(task, idx):
|
||||
# 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 next_task := self.schedule_task(
|
||||
task, idx, calls[idx - prev_length] if calls else 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
|
||||
futures[
|
||||
self.submit(
|
||||
# if the parent task was retried,
|
||||
# the next task might already be running
|
||||
rtn[idx - prev_length] = fut
|
||||
elif next_task.writes:
|
||||
# if it already ran, return the result
|
||||
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,
|
||||
writer=writer,
|
||||
configurable={
|
||||
CONFIG_KEY_SEND: partial(writer, next_task),
|
||||
CONFIG_KEY_CALL: partial(call, next_task),
|
||||
},
|
||||
__reraise_on_exit__=reraise,
|
||||
# starting a new task in the next tick ensures
|
||||
# updates from this tick are committed/streamed first
|
||||
__next_tick__=True,
|
||||
)
|
||||
] = next_task
|
||||
fut.add_done_callback(partial(self.commit, next_task))
|
||||
futures[fut] = next_task
|
||||
rtn[idx - prev_length] = fut
|
||||
return [rtn.get(i) for i in range(len(writes))]
|
||||
|
||||
def call(
|
||||
task: PregelExecutableTask,
|
||||
func: Callable[[Any], Union[Awaitable[Any], Any]],
|
||||
input: Any,
|
||||
*,
|
||||
retry: Optional[RetryPolicy] = None,
|
||||
) -> concurrent.futures.Future[Any]:
|
||||
(fut,) = writer(
|
||||
task, [(PUSH, None)], calls=[Call(func, input, retry=retry)]
|
||||
)
|
||||
assert fut is not None, "writer did not return a future for call"
|
||||
return fut
|
||||
|
||||
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
|
||||
if len(tasks) == 1 and timeout is None and get_waiter is None:
|
||||
t = tasks[0]
|
||||
try:
|
||||
run_with_retry(t, retry_policy, writer=writer)
|
||||
run_with_retry(
|
||||
t,
|
||||
retry_policy,
|
||||
configurable={
|
||||
CONFIG_KEY_SEND: partial(writer, t),
|
||||
CONFIG_KEY_CALL: partial(call, t),
|
||||
},
|
||||
)
|
||||
self.commit(t, None)
|
||||
except Exception as exc:
|
||||
self.commit(t, exc)
|
||||
if reraise:
|
||||
self.commit(t, None, exc)
|
||||
if reraise and futures:
|
||||
# will be re-raised after futures are done
|
||||
fut: concurrent.futures.Future = concurrent.futures.Future()
|
||||
fut.set_exception(exc)
|
||||
done_futures.add(fut)
|
||||
elif reraise:
|
||||
raise
|
||||
if not futures: # maybe `t` schuduled another task
|
||||
return
|
||||
# add waiter task if requested
|
||||
if get_waiter is not None:
|
||||
futures[get_waiter()] = None
|
||||
# schedule tasks
|
||||
for t in tasks:
|
||||
if not t.writes:
|
||||
fut = self.submit(
|
||||
run_with_retry,
|
||||
t,
|
||||
retry_policy,
|
||||
configurable={
|
||||
CONFIG_KEY_SEND: partial(writer, t),
|
||||
CONFIG_KEY_CALL: partial(call, t),
|
||||
},
|
||||
__reraise_on_exit__=reraise,
|
||||
)
|
||||
fut.add_done_callback(partial(self.commit, t))
|
||||
futures[fut] = t
|
||||
# execute tasks, and wait for one to fail or all to finish.
|
||||
# each task is independent from all other concurrent tasks
|
||||
# yield updates/debug output as each task finishes
|
||||
for t in tasks:
|
||||
if not t.writes:
|
||||
futures[
|
||||
self.submit(
|
||||
run_with_retry,
|
||||
t,
|
||||
retry_policy,
|
||||
writer=writer,
|
||||
__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(
|
||||
@@ -146,8 +223,6 @@ class PregelRunner:
|
||||
else:
|
||||
# store for panic check
|
||||
done_futures.add(fut)
|
||||
# task finished, commit writes
|
||||
self.commit(task, _exception(fut))
|
||||
else:
|
||||
# remove references to loop vars
|
||||
del fut, task
|
||||
@@ -156,6 +231,10 @@ class PregelRunner:
|
||||
break
|
||||
# give control back to the caller
|
||||
yield
|
||||
# wait for pending done callbacks
|
||||
# if a 2nd future finishes while `wait` is returning, it's possible
|
||||
# that done callbacks for the 2nd future aren't called until next tick
|
||||
time.sleep(0)
|
||||
# panic on failure or timeout
|
||||
_panic_or_proceed(
|
||||
done_futures.union(f for f, t in futures.items() if t is not None),
|
||||
@@ -171,48 +250,109 @@ class PregelRunner:
|
||||
retry_policy: Optional[RetryPolicy] = None,
|
||||
get_waiter: Optional[Callable[[], asyncio.Future[None]]] = None,
|
||||
) -> AsyncIterator[None]:
|
||||
locks: dict[str, threading.Lock] = {}
|
||||
|
||||
def writer(
|
||||
task: PregelExecutableTask, writes: Sequence[tuple[str, Any]]
|
||||
) -> None:
|
||||
prev_length = len(task.writes)
|
||||
# delegate to the underlying writer
|
||||
task.config[CONF][CONFIG_KEY_SEND](writes)
|
||||
for idx, w in enumerate(task.writes):
|
||||
# find the index for the newly inserted writes
|
||||
if idx < prev_length:
|
||||
continue
|
||||
assert writes[idx - prev_length] is w
|
||||
task: PregelExecutableTask,
|
||||
writes: Sequence[tuple[str, Any]],
|
||||
*,
|
||||
calls: Optional[Sequence[Call]] = None,
|
||||
) -> Sequence[Optional[asyncio.Future]]:
|
||||
if all(w[0] != PUSH for w in writes):
|
||||
return task.config[CONF][CONFIG_KEY_SEND](writes)
|
||||
|
||||
if task.id not in locks:
|
||||
locks[task.id] = threading.Lock()
|
||||
with locks[task.id]:
|
||||
prev_length = len(task.writes)
|
||||
# delegate to the underlying writer
|
||||
task.config[CONF][CONFIG_KEY_SEND](writes)
|
||||
# confirm no other concurrent writes were added
|
||||
assert len(task.writes) == prev_length + len(writes)
|
||||
# schedule PUSH tasks, collect futures
|
||||
rtn: dict[int, Optional[asyncio.Future]] = {}
|
||||
for idx, w in enumerate(writes, start=prev_length):
|
||||
# bail if not a PUSH write
|
||||
if w[0] != PUSH:
|
||||
continue
|
||||
# schedule the next task, if the callback returns one
|
||||
if next_task := self.schedule_task(task, idx):
|
||||
wcall = calls[idx - prev_length] if calls is not None else None
|
||||
if next_task := self.schedule_task(task, idx, wcall):
|
||||
# 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
|
||||
futures[
|
||||
cast(
|
||||
# if the parent task was retried,
|
||||
# the next task might already be running
|
||||
rtn[idx - prev_length] = fut
|
||||
elif next_task.writes:
|
||||
# if it already ran, return the result
|
||||
fut = asyncio.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 = cast(
|
||||
asyncio.Future,
|
||||
self.submit(
|
||||
arun_with_retry,
|
||||
next_task,
|
||||
retry_policy,
|
||||
stream=self.use_astream,
|
||||
writer=writer,
|
||||
configurable={
|
||||
CONFIG_KEY_SEND: partial(writer, next_task),
|
||||
CONFIG_KEY_CALL: partial(call, next_task),
|
||||
},
|
||||
__name__=t.name,
|
||||
__cancel_on_exit__=True,
|
||||
__reraise_on_exit__=reraise,
|
||||
# starting a new task in the next tick ensures
|
||||
# updates from this tick are committed/streamed first
|
||||
__next_tick__=True,
|
||||
),
|
||||
)
|
||||
] = next_task
|
||||
fut.add_done_callback(partial(self.commit, next_task))
|
||||
futures[fut] = next_task
|
||||
rtn[idx - prev_length] = fut
|
||||
return [rtn.get(i) for i in range(len(writes))]
|
||||
|
||||
def call(
|
||||
task: PregelExecutableTask,
|
||||
func: Callable[[Any], Union[Awaitable[Any], Any]],
|
||||
input: Any,
|
||||
*,
|
||||
retry: Optional[RetryPolicy] = None,
|
||||
) -> Union[asyncio.Future[Any], concurrent.futures.Future[Any]]:
|
||||
(fut,) = writer(
|
||||
task, [(PUSH, None)], calls=[Call(func, input, retry=retry)]
|
||||
)
|
||||
assert fut is not None, "writer did not return a future for call"
|
||||
if asyncio.iscoroutinefunction(func):
|
||||
return fut
|
||||
# adapted from asyncio.run_coroutine_threadsafe
|
||||
sfut: concurrent.futures.Future = concurrent.futures.Future()
|
||||
loop.call_soon_threadsafe(chain_future, fut, sfut)
|
||||
return sfut
|
||||
|
||||
loop = asyncio.get_event_loop()
|
||||
tasks = tuple(tasks)
|
||||
futures: dict[asyncio.Future, Optional[PregelExecutableTask]] = {}
|
||||
done_futures: set[asyncio.Future] = set()
|
||||
# give control back to the caller
|
||||
yield
|
||||
# fast path if single task with no waiter and no timeout
|
||||
@@ -220,39 +360,53 @@ class PregelRunner:
|
||||
t = tasks[0]
|
||||
try:
|
||||
await arun_with_retry(
|
||||
t, retry_policy, stream=self.use_astream, writer=writer
|
||||
t,
|
||||
retry_policy,
|
||||
stream=self.use_astream,
|
||||
configurable={
|
||||
CONFIG_KEY_SEND: partial(writer, t),
|
||||
CONFIG_KEY_CALL: partial(call, t),
|
||||
},
|
||||
)
|
||||
self.commit(t, None)
|
||||
except Exception as exc:
|
||||
self.commit(t, exc)
|
||||
if reraise:
|
||||
self.commit(t, None, exc)
|
||||
if reraise and futures:
|
||||
# will be re-raised after futures are done
|
||||
fut: asyncio.Future = loop.create_future()
|
||||
fut.set_exception(exc)
|
||||
done_futures.add(fut)
|
||||
elif reraise:
|
||||
raise
|
||||
if not futures: # maybe `t` schuduled another task
|
||||
return
|
||||
# add waiter task if requested
|
||||
if get_waiter is not None:
|
||||
futures[get_waiter()] = None
|
||||
# schedule tasks
|
||||
for t in tasks:
|
||||
if not t.writes:
|
||||
fut = cast(
|
||||
asyncio.Future,
|
||||
self.submit(
|
||||
arun_with_retry,
|
||||
t,
|
||||
retry_policy,
|
||||
stream=self.use_astream,
|
||||
configurable={
|
||||
CONFIG_KEY_SEND: partial(writer, t),
|
||||
CONFIG_KEY_CALL: partial(call, t),
|
||||
},
|
||||
__name__=t.name,
|
||||
__cancel_on_exit__=True,
|
||||
__reraise_on_exit__=reraise,
|
||||
),
|
||||
)
|
||||
fut.add_done_callback(partial(self.commit, t))
|
||||
futures[fut] = t
|
||||
# execute tasks, and wait for one to fail or all to finish.
|
||||
# each task is independent from all other concurrent tasks
|
||||
# yield updates/debug output as each task finishes
|
||||
for t in tasks:
|
||||
if not t.writes:
|
||||
futures[
|
||||
cast(
|
||||
asyncio.Future,
|
||||
self.submit(
|
||||
arun_with_retry,
|
||||
t,
|
||||
retry_policy,
|
||||
stream=self.use_astream,
|
||||
writer=writer,
|
||||
__name__=t.name,
|
||||
__cancel_on_exit__=True,
|
||||
__reraise_on_exit__=reraise,
|
||||
),
|
||||
)
|
||||
] = t
|
||||
done_futures: set[asyncio.Future] = set()
|
||||
end_time = timeout + loop.time() if timeout else None
|
||||
while len(futures) > (1 if get_waiter is not None else 0):
|
||||
done, inflight = await asyncio.wait(
|
||||
@@ -271,8 +425,6 @@ class PregelRunner:
|
||||
else:
|
||||
# store for panic check
|
||||
done_futures.add(fut)
|
||||
# task finished, commit writes
|
||||
self.commit(task, _exception(fut))
|
||||
else:
|
||||
# remove references to loop vars
|
||||
del fut, task
|
||||
@@ -281,6 +433,10 @@ class PregelRunner:
|
||||
break
|
||||
# give control back to the caller
|
||||
yield
|
||||
# wait for pending done callbacks
|
||||
# if a 2nd future finishes while `wait` is returning, it's possible
|
||||
# that done callbacks for the 2nd future aren't called until next tick
|
||||
await asyncio.sleep(0)
|
||||
# cancel waiter task
|
||||
for fut in futures:
|
||||
fut.cancel()
|
||||
@@ -292,9 +448,19 @@ class PregelRunner:
|
||||
)
|
||||
|
||||
def commit(
|
||||
self, task: PregelExecutableTask, exception: Optional[BaseException]
|
||||
self,
|
||||
task: PregelExecutableTask,
|
||||
fut: Union[None, concurrent.futures.Future[Any], asyncio.Future[Any]],
|
||||
exception: Optional[BaseException] = None,
|
||||
) -> None:
|
||||
if exception:
|
||||
if fut is not None:
|
||||
exception = _exception(fut)
|
||||
if isinstance(exception, asyncio.CancelledError):
|
||||
# for cancelled tasks, also save error in task,
|
||||
# so loop can finish super-step
|
||||
task.writes.append((ERROR, exception))
|
||||
self.put_writes(task.id, task.writes)
|
||||
elif exception:
|
||||
if isinstance(exception, GraphInterrupt):
|
||||
# save interrupt to checkpointer
|
||||
if interrupts := [(INTERRUPT, i) for i in exception.args[0]]:
|
||||
@@ -325,11 +491,12 @@ def _should_stop_others(
|
||||
GraphInterrupts are not considered failures."""
|
||||
for fut in done:
|
||||
if fut.cancelled():
|
||||
return True
|
||||
if exc := fut.exception():
|
||||
return not isinstance(exc, GraphBubbleUp)
|
||||
else:
|
||||
return False
|
||||
continue
|
||||
elif exc := fut.exception():
|
||||
if not isinstance(exc, GraphBubbleUp):
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
|
||||
def _exception(
|
||||
@@ -355,7 +522,9 @@ def _panic_or_proceed(
|
||||
done: set[Union[concurrent.futures.Future[Any], asyncio.Future[Any]]] = set()
|
||||
inflight: set[Union[concurrent.futures.Future[Any], asyncio.Future[Any]]] = set()
|
||||
for fut in futs:
|
||||
if fut.done():
|
||||
if fut.cancelled():
|
||||
continue
|
||||
elif fut.done():
|
||||
done.add(fut)
|
||||
else:
|
||||
inflight.add(fut)
|
||||
@@ -368,8 +537,6 @@ def _panic_or_proceed(
|
||||
# raise the exception
|
||||
if panic:
|
||||
raise exc
|
||||
else:
|
||||
return
|
||||
if inflight:
|
||||
# if we got here means we timed out
|
||||
while inflight:
|
||||
|
||||
@@ -0,0 +1,124 @@
|
||||
import asyncio
|
||||
import concurrent.futures
|
||||
from typing import Union
|
||||
|
||||
AnyFuture = Union[asyncio.Future, concurrent.futures.Future]
|
||||
|
||||
|
||||
def _get_loop(fut: asyncio.Future) -> asyncio.AbstractEventLoop:
|
||||
# Tries to call Future.get_loop() if it's available.
|
||||
# Otherwise fallbacks to using the old '_loop' property.
|
||||
try:
|
||||
get_loop = fut.get_loop
|
||||
except AttributeError:
|
||||
pass
|
||||
else:
|
||||
return get_loop()
|
||||
return fut._loop
|
||||
|
||||
|
||||
def _convert_future_exc(exc: BaseException) -> BaseException:
|
||||
exc_class = type(exc)
|
||||
if exc_class is concurrent.futures.CancelledError:
|
||||
return asyncio.CancelledError(*exc.args)
|
||||
elif exc_class is concurrent.futures.TimeoutError:
|
||||
return asyncio.TimeoutError(*exc.args)
|
||||
elif exc_class is concurrent.futures.InvalidStateError:
|
||||
return asyncio.InvalidStateError(*exc.args)
|
||||
else:
|
||||
return exc
|
||||
|
||||
|
||||
def _set_concurrent_future_state(
|
||||
concurrent: concurrent.futures.Future,
|
||||
source: AnyFuture,
|
||||
) -> None:
|
||||
"""Copy state from a future to a concurrent.futures.Future."""
|
||||
assert source.done()
|
||||
if source.cancelled():
|
||||
concurrent.cancel()
|
||||
if not concurrent.set_running_or_notify_cancel():
|
||||
return
|
||||
exception = source.exception()
|
||||
if exception is not None:
|
||||
concurrent.set_exception(_convert_future_exc(exception))
|
||||
else:
|
||||
result = source.result()
|
||||
concurrent.set_result(result)
|
||||
|
||||
|
||||
def _copy_future_state(source: AnyFuture, dest: asyncio.Future) -> None:
|
||||
"""Internal helper to copy state from another Future.
|
||||
|
||||
The other Future may be a concurrent.futures.Future.
|
||||
"""
|
||||
assert source.done()
|
||||
if dest.cancelled():
|
||||
return
|
||||
assert not dest.done()
|
||||
if source.cancelled():
|
||||
dest.cancel()
|
||||
else:
|
||||
exception = source.exception()
|
||||
if exception is not None:
|
||||
dest.set_exception(_convert_future_exc(exception))
|
||||
else:
|
||||
result = source.result()
|
||||
dest.set_result(result)
|
||||
|
||||
|
||||
def _chain_future(source: AnyFuture, destination: AnyFuture) -> None:
|
||||
"""Chain two futures so that when one completes, so does the other.
|
||||
|
||||
The result (or exception) of source will be copied to destination.
|
||||
If destination is cancelled, source gets cancelled too.
|
||||
Compatible with both asyncio.Future and concurrent.futures.Future.
|
||||
"""
|
||||
if not asyncio.isfuture(source) and not isinstance(
|
||||
source, concurrent.futures.Future
|
||||
):
|
||||
raise TypeError("A future is required for source argument")
|
||||
if not asyncio.isfuture(destination) and not isinstance(
|
||||
destination, concurrent.futures.Future
|
||||
):
|
||||
raise TypeError("A future is required for destination argument")
|
||||
source_loop = _get_loop(source) if asyncio.isfuture(source) else None
|
||||
dest_loop = _get_loop(destination) if asyncio.isfuture(destination) else None
|
||||
|
||||
def _set_state(future: AnyFuture, other: AnyFuture) -> None:
|
||||
if asyncio.isfuture(future):
|
||||
_copy_future_state(other, future)
|
||||
else:
|
||||
_set_concurrent_future_state(future, other)
|
||||
|
||||
def _call_check_cancel(destination: AnyFuture) -> None:
|
||||
if destination.cancelled():
|
||||
if source_loop is None or source_loop is dest_loop:
|
||||
source.cancel()
|
||||
else:
|
||||
source_loop.call_soon_threadsafe(source.cancel)
|
||||
|
||||
def _call_set_state(source: AnyFuture) -> None:
|
||||
if destination.cancelled() and dest_loop is not None and dest_loop.is_closed():
|
||||
return
|
||||
if dest_loop is None or dest_loop is source_loop:
|
||||
_set_state(destination, source)
|
||||
else:
|
||||
if dest_loop.is_closed():
|
||||
return
|
||||
dest_loop.call_soon_threadsafe(_set_state, destination, source)
|
||||
|
||||
destination.add_done_callback(_call_check_cancel)
|
||||
source.add_done_callback(_call_set_state)
|
||||
|
||||
|
||||
def chain_future(source: AnyFuture, destination: concurrent.futures.Future) -> None:
|
||||
# adapted from asyncio.run_coroutine_threadsafe
|
||||
try:
|
||||
_chain_future(source, destination)
|
||||
except (SystemExit, KeyboardInterrupt):
|
||||
raise
|
||||
except BaseException as exc:
|
||||
if destination.set_running_or_notify_cancel():
|
||||
destination.set_exception(exc)
|
||||
raise
|
||||
@@ -65,6 +65,7 @@ from langgraph.constants import (
|
||||
START,
|
||||
)
|
||||
from langgraph.errors import InvalidUpdateError, MultipleSubgraphsError, NodeInterrupt
|
||||
from langgraph.func import entrypoint, task
|
||||
from langgraph.graph import END, Graph, StateGraph
|
||||
from langgraph.graph.message import MessageGraph, MessagesState, add_messages
|
||||
from langgraph.managed.shared_value import SharedValue
|
||||
@@ -1957,6 +1958,85 @@ 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: int) -> str:
|
||||
nonlocal mapper_calls
|
||||
mapper_calls += 1
|
||||
time.sleep(input / 100)
|
||||
return str(input) * 2
|
||||
|
||||
@entrypoint(checkpointer=checkpointer)
|
||||
def graph(input: list[int]) -> list[str]:
|
||||
futures = [mapper(i) for i in input]
|
||||
mapped = [f.result() for f in futures]
|
||||
answer = interrupt("question")
|
||||
return [m + answer for m in mapped]
|
||||
|
||||
thread1 = {"configurable": {"thread_id": "1"}}
|
||||
assert [*graph.stream([0, 1], thread1)] == [
|
||||
{"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)
|
||||
def test_imp_stream_order(
|
||||
request: pytest.FixtureRequest, checkpointer_name: str
|
||||
) -> None:
|
||||
checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}")
|
||||
|
||||
@task()
|
||||
def foo(state: dict) -> dict:
|
||||
return {"a": state["a"] + "foo", "b": "bar"}
|
||||
|
||||
@task()
|
||||
def bar(state: dict) -> dict:
|
||||
return {"a": state["a"] + state["b"], "c": "bark"}
|
||||
|
||||
@task()
|
||||
def baz(state: dict) -> dict:
|
||||
return {"a": state["a"] + "baz", "c": "something else"}
|
||||
|
||||
@entrypoint(checkpointer=checkpointer)
|
||||
def graph(state: dict) -> dict:
|
||||
fut_foo = foo(state)
|
||||
fut_bar = bar(fut_foo.result())
|
||||
fut_baz = baz(fut_bar.result())
|
||||
return fut_baz.result()
|
||||
|
||||
thread1 = {"configurable": {"thread_id": "1"}}
|
||||
assert [c for c in graph.stream({"a": "0"}, thread1)] == [
|
||||
{"foo": {"a": "0foo", "b": "bar"}},
|
||||
{"bar": {"a": "0foobar", "c": "bark"}},
|
||||
{"baz": {"a": "0foobarbaz", "c": "something else"}},
|
||||
{"graph": {"a": "0foobarbaz", "c": "something else"}},
|
||||
]
|
||||
|
||||
assert graph.get_state(thread1).values == {"a": "0foobarbaz", "c": "something else"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
|
||||
def test_send_dedupe_on_resume(
|
||||
request: pytest.FixtureRequest, checkpointer_name: str
|
||||
@@ -2484,7 +2564,7 @@ def test_send_react_interrupt(
|
||||
PregelTask(
|
||||
id=AnyStr(),
|
||||
name="foo",
|
||||
path=("__pregel_push", ("__pregel_pull", "agent"), 2, AnyStr()),
|
||||
path=("__pregel_push", ("__pregel_pull", "agent"), 2),
|
||||
error=None,
|
||||
interrupts=(),
|
||||
state=None,
|
||||
@@ -2641,7 +2721,7 @@ def test_send_react_interrupt(
|
||||
PregelTask(
|
||||
id=AnyStr(),
|
||||
name="foo",
|
||||
path=("__pregel_push", ("__pregel_pull", "agent"), 2, AnyStr()),
|
||||
path=("__pregel_push", ("__pregel_pull", "agent"), 2),
|
||||
error=None,
|
||||
interrupts=(),
|
||||
state=None,
|
||||
@@ -2728,7 +2808,7 @@ def test_send_react_interrupt(
|
||||
PregelTask(
|
||||
id=AnyStr(),
|
||||
name="foo",
|
||||
path=("__pregel_push", (), 0, AnyStr()),
|
||||
path=("__pregel_push", (), 0),
|
||||
error=None,
|
||||
interrupts=(),
|
||||
state=None,
|
||||
@@ -2953,7 +3033,7 @@ def test_send_react_interrupt_control(
|
||||
PregelTask(
|
||||
id=AnyStr(),
|
||||
name="foo",
|
||||
path=("__pregel_push", ("__pregel_pull", "agent"), 2, AnyStr()),
|
||||
path=("__pregel_push", ("__pregel_pull", "agent"), 2),
|
||||
error=None,
|
||||
interrupts=(),
|
||||
state=None,
|
||||
@@ -5745,6 +5825,7 @@ def test_state_graph_packets(
|
||||
@tool()
|
||||
def search_api(query: str) -> str:
|
||||
"""Searches the API for the query."""
|
||||
time.sleep(0.1)
|
||||
return f"result for {query}"
|
||||
|
||||
tools = [search_api]
|
||||
@@ -6031,9 +6112,7 @@ def test_state_graph_packets(
|
||||
)
|
||||
},
|
||||
),
|
||||
PregelTask(
|
||||
AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 2, AnyStr())
|
||||
),
|
||||
PregelTask(AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 2)),
|
||||
),
|
||||
next=("tools",),
|
||||
config=(app_w_interrupt.checkpointer.get_tuple(config)).config,
|
||||
@@ -6073,7 +6152,7 @@ def test_state_graph_packets(
|
||||
),
|
||||
]
|
||||
},
|
||||
tasks=(PregelTask(AnyStr(), "tools", (PUSH, (), 0, AnyStr())),),
|
||||
tasks=(PregelTask(AnyStr(), "tools", (PUSH, (), 0)),),
|
||||
next=("tools",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
created_at=(app_w_interrupt.checkpointer.get_tuple(config)).checkpoint["ts"],
|
||||
@@ -6202,12 +6281,8 @@ def test_state_graph_packets(
|
||||
)
|
||||
},
|
||||
),
|
||||
PregelTask(
|
||||
AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 2, AnyStr())
|
||||
),
|
||||
PregelTask(
|
||||
AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 3, AnyStr())
|
||||
),
|
||||
PregelTask(AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 2)),
|
||||
PregelTask(AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 3)),
|
||||
),
|
||||
next=("tools", "tools"),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
@@ -6354,9 +6429,7 @@ def test_state_graph_packets(
|
||||
)
|
||||
},
|
||||
),
|
||||
PregelTask(
|
||||
AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 2, AnyStr())
|
||||
),
|
||||
PregelTask(AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 2)),
|
||||
),
|
||||
next=("tools",),
|
||||
config=(app_w_interrupt.checkpointer.get_tuple(config)).config,
|
||||
@@ -6396,7 +6469,7 @@ def test_state_graph_packets(
|
||||
),
|
||||
]
|
||||
},
|
||||
tasks=(PregelTask(AnyStr(), "tools", (PUSH, (), 0, AnyStr())),),
|
||||
tasks=(PregelTask(AnyStr(), "tools", (PUSH, (), 0)),),
|
||||
next=("tools",),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
created_at=(app_w_interrupt.checkpointer.get_tuple(config)).checkpoint["ts"],
|
||||
@@ -6525,12 +6598,8 @@ def test_state_graph_packets(
|
||||
)
|
||||
},
|
||||
),
|
||||
PregelTask(
|
||||
AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 2, AnyStr())
|
||||
),
|
||||
PregelTask(
|
||||
AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 3, AnyStr())
|
||||
),
|
||||
PregelTask(AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 2)),
|
||||
PregelTask(AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 3)),
|
||||
),
|
||||
next=("tools", "tools"),
|
||||
config=app_w_interrupt.checkpointer.get_tuple(config).config,
|
||||
@@ -12774,7 +12843,7 @@ def test_send_to_nested_graphs(
|
||||
PregelTask(
|
||||
AnyStr(),
|
||||
"generate_joke",
|
||||
(PUSH, ("__pregel_pull", "__start__"), 1, AnyStr()),
|
||||
(PUSH, ("__pregel_pull", "__start__"), 1),
|
||||
state={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
@@ -12785,7 +12854,7 @@ def test_send_to_nested_graphs(
|
||||
PregelTask(
|
||||
AnyStr(),
|
||||
"generate_joke",
|
||||
(PUSH, ("__pregel_pull", "__start__"), 2, AnyStr()),
|
||||
(PUSH, ("__pregel_pull", "__start__"), 2),
|
||||
state={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
@@ -12838,7 +12907,7 @@ def test_send_to_nested_graphs(
|
||||
"checkpoint_ns": AnyStr("generate_joke:"),
|
||||
"langgraph_checkpoint_ns": AnyStr("generate_joke:"),
|
||||
"langgraph_node": "generate_joke",
|
||||
"langgraph_path": [PUSH, ["__pregel_pull", "__start__"], 1, AnyStr()],
|
||||
"langgraph_path": [PUSH, ["__pregel_pull", "__start__"], 1],
|
||||
"langgraph_step": 0,
|
||||
"langgraph_triggers": [PUSH],
|
||||
},
|
||||
@@ -12883,7 +12952,7 @@ def test_send_to_nested_graphs(
|
||||
"checkpoint_ns": AnyStr("generate_joke:"),
|
||||
"langgraph_checkpoint_ns": AnyStr("generate_joke:"),
|
||||
"langgraph_node": "generate_joke",
|
||||
"langgraph_path": [PUSH, ["__pregel_pull", "__start__"], 2, AnyStr()],
|
||||
"langgraph_path": [PUSH, ["__pregel_pull", "__start__"], 2],
|
||||
"langgraph_step": 0,
|
||||
"langgraph_triggers": [PUSH],
|
||||
},
|
||||
@@ -13009,7 +13078,7 @@ def test_send_to_nested_graphs(
|
||||
PregelTask(
|
||||
AnyStr(),
|
||||
"generate_joke",
|
||||
(PUSH, ("__pregel_pull", "__start__"), 1, AnyStr()),
|
||||
(PUSH, ("__pregel_pull", "__start__"), 1),
|
||||
state={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
@@ -13021,7 +13090,7 @@ def test_send_to_nested_graphs(
|
||||
PregelTask(
|
||||
AnyStr(),
|
||||
"generate_joke",
|
||||
(PUSH, ("__pregel_pull", "__start__"), 2, AnyStr()),
|
||||
(PUSH, ("__pregel_pull", "__start__"), 2),
|
||||
state={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
|
||||
@@ -62,6 +62,7 @@ from langgraph.constants import (
|
||||
START,
|
||||
)
|
||||
from langgraph.errors import InvalidUpdateError, MultipleSubgraphsError, NodeInterrupt
|
||||
from langgraph.func import entrypoint, task
|
||||
from langgraph.graph import END, Graph, StateGraph
|
||||
from langgraph.graph.message import MessageGraph, MessagesState, add_messages
|
||||
from langgraph.managed.shared_value import SharedValue
|
||||
@@ -2647,6 +2648,178 @@ async def test_send_sequences(checkpointer_name: str) -> None:
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
sys.version_info < (3, 11),
|
||||
reason="Python 3.11+ is required for async contextvars support",
|
||||
)
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
|
||||
async def test_imp_task(checkpointer_name: str) -> None:
|
||||
async with awith_checkpointer(checkpointer_name) as checkpointer:
|
||||
mapper_calls = 0
|
||||
|
||||
@task()
|
||||
async def mapper(input: int) -> str:
|
||||
nonlocal mapper_calls
|
||||
mapper_calls += 1
|
||||
return str(input) * 2
|
||||
|
||||
@entrypoint(checkpointer=checkpointer)
|
||||
async def graph(input: list[int]) -> list[str]:
|
||||
futures = [mapper(i) for i in input]
|
||||
mapped = await asyncio.gather(*futures)
|
||||
answer = interrupt("question")
|
||||
return [m + answer for m in mapped]
|
||||
|
||||
thread1 = {"configurable": {"thread_id": "1"}}
|
||||
assert [c async for c in graph.astream([0, 1], thread1)] == [
|
||||
{"mapper": "00"},
|
||||
{"mapper": "11"},
|
||||
{
|
||||
"__interrupt__": (
|
||||
Interrupt(
|
||||
value="question",
|
||||
resumable=True,
|
||||
ns=[AnyStr("graph:")],
|
||||
when="during",
|
||||
),
|
||||
)
|
||||
},
|
||||
]
|
||||
assert mapper_calls == 2
|
||||
|
||||
assert await graph.ainvoke(Command(resume="answer"), thread1) == [
|
||||
"00answer",
|
||||
"11answer",
|
||||
]
|
||||
assert mapper_calls == 2
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
sys.version_info < (3, 11),
|
||||
reason="Python 3.11+ is required for async contextvars support",
|
||||
)
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
|
||||
async def test_imp_task_cancel(checkpointer_name: str) -> None:
|
||||
async with awith_checkpointer(checkpointer_name) as checkpointer:
|
||||
mapper_calls = 0
|
||||
mapper_cancels = 0
|
||||
|
||||
@task()
|
||||
async def mapper(input: int) -> str:
|
||||
nonlocal mapper_calls, mapper_cancels
|
||||
mapper_calls += 1
|
||||
try:
|
||||
await asyncio.sleep(1)
|
||||
except asyncio.CancelledError:
|
||||
mapper_cancels += 1
|
||||
raise
|
||||
return str(input) * 2
|
||||
|
||||
@entrypoint(checkpointer=checkpointer)
|
||||
async def graph(input: list[int]) -> list[str]:
|
||||
futures = [mapper(i) for i in input]
|
||||
await asyncio.sleep(0.1)
|
||||
futures.pop().cancel() # cancel one
|
||||
mapped = await asyncio.gather(*futures)
|
||||
answer = interrupt("question")
|
||||
return [m + answer for m in mapped]
|
||||
|
||||
thread1 = {"configurable": {"thread_id": "1"}}
|
||||
assert [c async for c in graph.astream([0, 1], thread1)] == [
|
||||
{"mapper": "00"},
|
||||
{
|
||||
"__interrupt__": (
|
||||
Interrupt(
|
||||
value="question",
|
||||
resumable=True,
|
||||
ns=[AnyStr("graph:")],
|
||||
when="during",
|
||||
),
|
||||
)
|
||||
},
|
||||
]
|
||||
assert mapper_calls == 2
|
||||
assert mapper_cancels == 1
|
||||
|
||||
assert await graph.ainvoke(Command(resume="answer"), thread1) == [
|
||||
"00answer",
|
||||
]
|
||||
assert mapper_calls == 3
|
||||
assert mapper_cancels == 2
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
sys.version_info < (3, 11),
|
||||
reason="Python 3.11+ is required for async contextvars support",
|
||||
)
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
|
||||
async def test_imp_sync_from_async(checkpointer_name: str) -> None:
|
||||
async with awith_checkpointer(checkpointer_name) as checkpointer:
|
||||
|
||||
@task()
|
||||
def foo(state: dict) -> dict:
|
||||
return {"a": state["a"] + "foo", "b": "bar"}
|
||||
|
||||
@task()
|
||||
def bar(state: dict) -> dict:
|
||||
return {"a": state["a"] + state["b"], "c": "bark"}
|
||||
|
||||
@task()
|
||||
def baz(state: dict) -> dict:
|
||||
return {"a": state["a"] + "baz", "c": "something else"}
|
||||
|
||||
@entrypoint(checkpointer=checkpointer)
|
||||
def graph(state: dict) -> dict:
|
||||
fut_foo = foo(state)
|
||||
fut_bar = bar(fut_foo.result())
|
||||
fut_baz = baz(fut_bar.result())
|
||||
return fut_baz.result()
|
||||
|
||||
thread1 = {"configurable": {"thread_id": "1"}}
|
||||
assert [c async for c in graph.astream({"a": "0"}, thread1)] == [
|
||||
{"foo": {"a": "0foo", "b": "bar"}},
|
||||
{"bar": {"a": "0foobar", "c": "bark"}},
|
||||
{"baz": {"a": "0foobarbaz", "c": "something else"}},
|
||||
{"graph": {"a": "0foobarbaz", "c": "something else"}},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.skipif(
|
||||
sys.version_info < (3, 11),
|
||||
reason="Python 3.11+ is required for async contextvars support",
|
||||
)
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
|
||||
async def test_imp_stream_order(checkpointer_name: str) -> None:
|
||||
async with awith_checkpointer(checkpointer_name) as checkpointer:
|
||||
|
||||
@task()
|
||||
async def foo(state: dict) -> dict:
|
||||
return {"a": state["a"] + "foo", "b": "bar"}
|
||||
|
||||
@task()
|
||||
async def bar(state: dict) -> dict:
|
||||
return {"a": state["a"] + state["b"], "c": "bark"}
|
||||
|
||||
@task()
|
||||
async def baz(state: dict) -> dict:
|
||||
return {"a": state["a"] + "baz", "c": "something else"}
|
||||
|
||||
@entrypoint(checkpointer=checkpointer)
|
||||
async def graph(state: dict) -> dict:
|
||||
fut_foo = foo(state)
|
||||
fut_bar = bar(await fut_foo)
|
||||
fut_baz = baz(await fut_bar)
|
||||
return await fut_baz
|
||||
|
||||
thread1 = {"configurable": {"thread_id": "1"}}
|
||||
assert [c async for c in graph.astream({"a": "0"}, thread1)] == [
|
||||
{"foo": {"a": "0foo", "b": "bar"}},
|
||||
{"bar": {"a": "0foobar", "c": "bark"}},
|
||||
{"baz": {"a": "0foobarbaz", "c": "something else"}},
|
||||
{"graph": {"a": "0foobarbaz", "c": "something else"}},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
|
||||
async def test_send_dedupe_on_resume(checkpointer_name: str) -> None:
|
||||
if not FF_SEND_V2:
|
||||
@@ -2864,12 +3037,7 @@ async def test_send_dedupe_on_resume(checkpointer_name: str) -> None:
|
||||
PregelTask(
|
||||
id=AnyStr(),
|
||||
name="2",
|
||||
path=(
|
||||
"__pregel_push",
|
||||
("__pregel_pull", "1"),
|
||||
2,
|
||||
AnyStr(),
|
||||
),
|
||||
path=("__pregel_push", ("__pregel_pull", "1"), 2),
|
||||
error=None,
|
||||
interrupts=(),
|
||||
state=None,
|
||||
@@ -2878,12 +3046,7 @@ async def test_send_dedupe_on_resume(checkpointer_name: str) -> None:
|
||||
PregelTask(
|
||||
id=AnyStr(),
|
||||
name="2",
|
||||
path=(
|
||||
"__pregel_push",
|
||||
("__pregel_pull", "1"),
|
||||
3,
|
||||
AnyStr(),
|
||||
),
|
||||
path=("__pregel_push", ("__pregel_pull", "1"), 3),
|
||||
error=None,
|
||||
interrupts=(),
|
||||
state=None,
|
||||
@@ -2894,14 +3057,8 @@ async def test_send_dedupe_on_resume(checkpointer_name: str) -> None:
|
||||
name="2",
|
||||
path=(
|
||||
"__pregel_push",
|
||||
(
|
||||
"__pregel_push",
|
||||
("__pregel_pull", "1"),
|
||||
2,
|
||||
AnyStr(),
|
||||
),
|
||||
("__pregel_push", ("__pregel_pull", "1"), 2),
|
||||
2,
|
||||
AnyStr(),
|
||||
),
|
||||
error=None,
|
||||
interrupts=(),
|
||||
@@ -2913,14 +3070,8 @@ async def test_send_dedupe_on_resume(checkpointer_name: str) -> None:
|
||||
name="flaky",
|
||||
path=(
|
||||
"__pregel_push",
|
||||
(
|
||||
"__pregel_push",
|
||||
("__pregel_pull", "1"),
|
||||
3,
|
||||
AnyStr(),
|
||||
),
|
||||
("__pregel_push", ("__pregel_pull", "1"), 3),
|
||||
2,
|
||||
AnyStr(),
|
||||
),
|
||||
error=None,
|
||||
interrupts=(Interrupt(value="Bahh", when="during"),),
|
||||
@@ -3157,7 +3308,7 @@ async def test_send_react_interrupt(checkpointer_name: str) -> None:
|
||||
PregelTask(
|
||||
id=AnyStr(),
|
||||
name="foo",
|
||||
path=("__pregel_push", ("__pregel_pull", "agent"), 2, AnyStr()),
|
||||
path=("__pregel_push", ("__pregel_pull", "agent"), 2),
|
||||
error=None,
|
||||
interrupts=(),
|
||||
state=None,
|
||||
@@ -3314,7 +3465,7 @@ async def test_send_react_interrupt(checkpointer_name: str) -> None:
|
||||
PregelTask(
|
||||
id=AnyStr(),
|
||||
name="foo",
|
||||
path=("__pregel_push", ("__pregel_pull", "agent"), 2, AnyStr()),
|
||||
path=("__pregel_push", ("__pregel_pull", "agent"), 2),
|
||||
error=None,
|
||||
interrupts=(),
|
||||
state=None,
|
||||
@@ -3401,7 +3552,7 @@ async def test_send_react_interrupt(checkpointer_name: str) -> None:
|
||||
PregelTask(
|
||||
id=AnyStr(),
|
||||
name="foo",
|
||||
path=("__pregel_push", (), 0, AnyStr()),
|
||||
path=("__pregel_push", (), 0),
|
||||
error=None,
|
||||
interrupts=(),
|
||||
state=None,
|
||||
@@ -3625,7 +3776,7 @@ async def test_send_react_interrupt_control(
|
||||
PregelTask(
|
||||
id=AnyStr(),
|
||||
name="foo",
|
||||
path=("__pregel_push", ("__pregel_pull", "agent"), 2, AnyStr()),
|
||||
path=("__pregel_push", ("__pregel_pull", "agent"), 2),
|
||||
error=None,
|
||||
interrupts=(),
|
||||
state=None,
|
||||
@@ -6420,9 +6571,7 @@ async def test_state_graph_packets(checkpointer_name: str) -> None:
|
||||
)
|
||||
},
|
||||
),
|
||||
PregelTask(
|
||||
AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 2, AnyStr())
|
||||
),
|
||||
PregelTask(AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 2)),
|
||||
),
|
||||
next=("tools",),
|
||||
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
|
||||
@@ -6465,7 +6614,7 @@ async def test_state_graph_packets(checkpointer_name: str) -> None:
|
||||
),
|
||||
]
|
||||
},
|
||||
tasks=(PregelTask(AnyStr(), "tools", (PUSH, (), 0, AnyStr())),),
|
||||
tasks=(PregelTask(AnyStr(), "tools", (PUSH, (), 0)),),
|
||||
next=("tools",),
|
||||
config=tup.config,
|
||||
created_at=tup.checkpoint["ts"],
|
||||
@@ -6596,12 +6745,8 @@ async def test_state_graph_packets(checkpointer_name: str) -> None:
|
||||
)
|
||||
},
|
||||
),
|
||||
PregelTask(
|
||||
AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 2, AnyStr())
|
||||
),
|
||||
PregelTask(
|
||||
AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 3, AnyStr())
|
||||
),
|
||||
PregelTask(AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 2)),
|
||||
PregelTask(AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 3)),
|
||||
),
|
||||
next=("tools", "tools"),
|
||||
config=tup.config,
|
||||
@@ -6751,9 +6896,7 @@ async def test_state_graph_packets(checkpointer_name: str) -> None:
|
||||
)
|
||||
},
|
||||
),
|
||||
PregelTask(
|
||||
AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 2, AnyStr())
|
||||
),
|
||||
PregelTask(AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 2)),
|
||||
),
|
||||
next=("tools",),
|
||||
config=(await app_w_interrupt.checkpointer.aget_tuple(config)).config,
|
||||
@@ -6796,7 +6939,7 @@ async def test_state_graph_packets(checkpointer_name: str) -> None:
|
||||
),
|
||||
]
|
||||
},
|
||||
tasks=(PregelTask(AnyStr(), "tools", (PUSH, (), 0, AnyStr())),),
|
||||
tasks=(PregelTask(AnyStr(), "tools", (PUSH, (), 0)),),
|
||||
next=("tools",),
|
||||
config=tup.config,
|
||||
created_at=tup.checkpoint["ts"],
|
||||
@@ -6929,12 +7072,8 @@ async def test_state_graph_packets(checkpointer_name: str) -> None:
|
||||
)
|
||||
},
|
||||
),
|
||||
PregelTask(
|
||||
AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 2, AnyStr())
|
||||
),
|
||||
PregelTask(
|
||||
AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 3, AnyStr())
|
||||
),
|
||||
PregelTask(AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 2)),
|
||||
PregelTask(AnyStr(), "tools", (PUSH, ("__pregel_pull", "agent"), 3)),
|
||||
),
|
||||
next=("tools", "tools"),
|
||||
config=tup.config,
|
||||
@@ -11612,7 +11751,7 @@ async def test_send_to_nested_graphs(checkpointer_name: str) -> None:
|
||||
PregelTask(
|
||||
AnyStr(),
|
||||
"generate_joke",
|
||||
(PUSH, ("__pregel_pull", "__start__"), 1, AnyStr()),
|
||||
(PUSH, ("__pregel_pull", "__start__"), 1),
|
||||
state={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
@@ -11623,7 +11762,7 @@ async def test_send_to_nested_graphs(checkpointer_name: str) -> None:
|
||||
PregelTask(
|
||||
AnyStr(),
|
||||
"generate_joke",
|
||||
(PUSH, ("__pregel_pull", "__start__"), 2, AnyStr()),
|
||||
(PUSH, ("__pregel_pull", "__start__"), 2),
|
||||
state={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
@@ -11764,7 +11903,7 @@ async def test_send_to_nested_graphs(checkpointer_name: str) -> None:
|
||||
PregelTask(
|
||||
AnyStr(),
|
||||
"generate_joke",
|
||||
(PUSH, ("__pregel_pull", "__start__"), 1, AnyStr()),
|
||||
(PUSH, ("__pregel_pull", "__start__"), 1),
|
||||
state={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
@@ -11776,7 +11915,7 @@ async def test_send_to_nested_graphs(checkpointer_name: str) -> None:
|
||||
PregelTask(
|
||||
AnyStr(),
|
||||
"generate_joke",
|
||||
(PUSH, ("__pregel_pull", "__start__"), 2, AnyStr()),
|
||||
(PUSH, ("__pregel_pull", "__start__"), 2),
|
||||
state={
|
||||
"configurable": {
|
||||
"thread_id": "1",
|
||||
|
||||
@@ -191,6 +191,7 @@ async def test_subgraph_w_interrupt(
|
||||
"__pregel_delegate": False,
|
||||
"__pregel_read": None,
|
||||
"__pregel_send": None,
|
||||
"__pregel_call": None,
|
||||
"__pregel_ensure_latest": True,
|
||||
"__pregel_dedupe_tasks": True,
|
||||
"__pregel_resuming": False,
|
||||
@@ -257,6 +258,7 @@ async def test_subgraph_w_interrupt(
|
||||
"__pregel_delegate": False,
|
||||
"__pregel_read": None,
|
||||
"__pregel_send": None,
|
||||
"__pregel_call": None,
|
||||
"__pregel_ensure_latest": True,
|
||||
"__pregel_dedupe_tasks": True,
|
||||
"__pregel_resuming": False,
|
||||
@@ -353,6 +355,7 @@ async def test_subgraph_w_interrupt(
|
||||
"__pregel_delegate": False,
|
||||
"__pregel_read": None,
|
||||
"__pregel_send": None,
|
||||
"__pregel_call": None,
|
||||
"__pregel_ensure_latest": True,
|
||||
"__pregel_dedupe_tasks": True,
|
||||
"__pregel_resuming": False,
|
||||
@@ -459,6 +462,7 @@ async def test_subgraph_w_interrupt(
|
||||
"__pregel_delegate": False,
|
||||
"__pregel_read": None,
|
||||
"__pregel_send": None,
|
||||
"__pregel_call": None,
|
||||
"__pregel_ensure_latest": True,
|
||||
"__pregel_dedupe_tasks": True,
|
||||
"__pregel_resuming": True,
|
||||
@@ -520,6 +524,7 @@ async def test_subgraph_w_interrupt(
|
||||
"__pregel_delegate": False,
|
||||
"__pregel_read": None,
|
||||
"__pregel_send": None,
|
||||
"__pregel_call": None,
|
||||
"__pregel_ensure_latest": True,
|
||||
"__pregel_dedupe_tasks": True,
|
||||
"__pregel_resuming": True,
|
||||
@@ -637,6 +642,7 @@ async def test_subgraph_w_interrupt(
|
||||
"__pregel_delegate": False,
|
||||
"__pregel_read": None,
|
||||
"__pregel_send": None,
|
||||
"__pregel_call": None,
|
||||
"__pregel_ensure_latest": True,
|
||||
"__pregel_dedupe_tasks": True,
|
||||
"__pregel_resuming": True,
|
||||
|
||||
@@ -190,6 +190,7 @@ def test_subgraph_w_interrupt(
|
||||
"__pregel_delegate": False,
|
||||
"__pregel_read": None,
|
||||
"__pregel_send": None,
|
||||
"__pregel_call": None,
|
||||
"__pregel_ensure_latest": True,
|
||||
"__pregel_dedupe_tasks": True,
|
||||
"__pregel_resuming": False,
|
||||
@@ -256,6 +257,7 @@ def test_subgraph_w_interrupt(
|
||||
"__pregel_delegate": False,
|
||||
"__pregel_read": None,
|
||||
"__pregel_send": None,
|
||||
"__pregel_call": None,
|
||||
"__pregel_ensure_latest": True,
|
||||
"__pregel_store": None,
|
||||
"__pregel_dedupe_tasks": True,
|
||||
@@ -352,6 +354,7 @@ def test_subgraph_w_interrupt(
|
||||
"__pregel_delegate": False,
|
||||
"__pregel_read": None,
|
||||
"__pregel_send": None,
|
||||
"__pregel_call": None,
|
||||
"__pregel_ensure_latest": True,
|
||||
"__pregel_dedupe_tasks": True,
|
||||
"__pregel_store": None,
|
||||
@@ -457,6 +460,7 @@ def test_subgraph_w_interrupt(
|
||||
"__pregel_delegate": False,
|
||||
"__pregel_read": None,
|
||||
"__pregel_send": None,
|
||||
"__pregel_call": None,
|
||||
"__pregel_ensure_latest": True,
|
||||
"__pregel_dedupe_tasks": True,
|
||||
"__pregel_store": None,
|
||||
@@ -518,6 +522,7 @@ def test_subgraph_w_interrupt(
|
||||
"__pregel_delegate": False,
|
||||
"__pregel_read": None,
|
||||
"__pregel_send": None,
|
||||
"__pregel_call": None,
|
||||
"__pregel_ensure_latest": True,
|
||||
"__pregel_dedupe_tasks": True,
|
||||
"__pregel_store": None,
|
||||
@@ -635,6 +640,7 @@ def test_subgraph_w_interrupt(
|
||||
"__pregel_delegate": False,
|
||||
"__pregel_read": None,
|
||||
"__pregel_send": None,
|
||||
"__pregel_call": None,
|
||||
"__pregel_ensure_latest": True,
|
||||
"__pregel_dedupe_tasks": True,
|
||||
"__pregel_resuming": True,
|
||||
|
||||
Reference in New Issue
Block a user