From f2e0dc104246caa0100900d8f9a7b319020c5d88 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Tue, 3 Sep 2024 18:13:20 -0700 Subject: [PATCH 1/9] Reduce cpu time spent on langchain-core utilities - shaves off 2s of 11s runtime of a simple graph with 1,000 subgraphs --- libs/langgraph/langgraph/graph/graph.py | 2 +- libs/langgraph/langgraph/graph/state.py | 6 +- .../langgraph/prebuilt/tool_executor.py | 2 +- .../langgraph/langgraph/prebuilt/tool_node.py | 2 +- .../langgraph/prebuilt/tool_validator.py | 2 +- libs/langgraph/langgraph/pregel/__init__.py | 11 +- libs/langgraph/langgraph/pregel/algo.py | 7 +- libs/langgraph/langgraph/pregel/config.py | 34 ---- libs/langgraph/langgraph/pregel/loop.py | 2 +- libs/langgraph/langgraph/pregel/manager.py | 7 +- libs/langgraph/langgraph/pregel/read.py | 4 +- libs/langgraph/langgraph/pregel/write.py | 2 +- libs/langgraph/langgraph/utils/__init__.py | 0 libs/langgraph/langgraph/utils/config.py | 150 ++++++++++++++ libs/langgraph/langgraph/utils/fields.py | 101 ++++++++++ .../langgraph/{utils.py => utils/runnable.py} | 184 +++++++----------- libs/langgraph/tests/test_utils.py | 8 +- 17 files changed, 349 insertions(+), 175 deletions(-) delete mode 100644 libs/langgraph/langgraph/pregel/config.py create mode 100644 libs/langgraph/langgraph/utils/__init__.py create mode 100644 libs/langgraph/langgraph/utils/config.py create mode 100644 libs/langgraph/langgraph/utils/fields.py rename libs/langgraph/langgraph/{utils.py => utils/runnable.py} (58%) diff --git a/libs/langgraph/langgraph/graph/graph.py b/libs/langgraph/langgraph/graph/graph.py index 570c06bad..af6700c2a 100644 --- a/libs/langgraph/langgraph/graph/graph.py +++ b/libs/langgraph/langgraph/graph/graph.py @@ -38,7 +38,7 @@ from langgraph.pregel import Channel, Pregel from langgraph.pregel.read import PregelNode from langgraph.pregel.types import All from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry -from langgraph.utils import RunnableCallable, coerce_to_runnable +from langgraph.utils.runnable import RunnableCallable, coerce_to_runnable logger = logging.getLogger(__name__) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 307b672aa..ce25f6c5d 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -44,7 +44,11 @@ from langgraph.pregel.read import ChannelRead, PregelNode from langgraph.pregel.types import All, RetryPolicy from langgraph.pregel.write import SKIP_WRITE, ChannelWrite, ChannelWriteEntry from langgraph.store.base import BaseStore -from langgraph.utils import RunnableCallable, coerce_to_runnable, get_field_default +from langgraph.utils.fields import get_field_default +from langgraph.utils.runnable import ( + RunnableCallable, + coerce_to_runnable, +) logger = logging.getLogger(__name__) diff --git a/libs/langgraph/langgraph/prebuilt/tool_executor.py b/libs/langgraph/langgraph/prebuilt/tool_executor.py index 7590b3836..341f4874d 100644 --- a/libs/langgraph/langgraph/prebuilt/tool_executor.py +++ b/libs/langgraph/langgraph/prebuilt/tool_executor.py @@ -6,7 +6,7 @@ from langchain_core.tools import BaseTool from langchain_core.tools import tool as create_tool from langgraph._api.deprecation import deprecated -from langgraph.utils import RunnableCallable +from langgraph.utils.runnable import RunnableCallable INVALID_TOOL_MSG_TEMPLATE = ( "{requested_tool_name} is not a valid tool, " diff --git a/libs/langgraph/langgraph/prebuilt/tool_node.py b/libs/langgraph/langgraph/prebuilt/tool_node.py index 05ac1d8e0..54fc1c6f3 100644 --- a/libs/langgraph/langgraph/prebuilt/tool_node.py +++ b/libs/langgraph/langgraph/prebuilt/tool_node.py @@ -21,7 +21,7 @@ from langchain_core.tools import BaseTool, InjectedToolArg from langchain_core.tools import tool as create_tool from typing_extensions import get_args -from langgraph.utils import RunnableCallable +from langgraph.utils.runnable import RunnableCallable INVALID_TOOL_NAME_ERROR_TEMPLATE = ( "Error: {requested_tool} is not a valid tool, try one of [{available_tools}]." diff --git a/libs/langgraph/langgraph/prebuilt/tool_validator.py b/libs/langgraph/langgraph/prebuilt/tool_validator.py index 73b2cce2b..95c84a0a5 100644 --- a/libs/langgraph/langgraph/prebuilt/tool_validator.py +++ b/libs/langgraph/langgraph/prebuilt/tool_validator.py @@ -33,7 +33,7 @@ from langchain_core.tools import BaseTool, create_schema_from_function from pydantic import BaseModel as BaseModelV2 from pydantic import ValidationError as ValidationErrorV2 -from langgraph.utils import RunnableCallable +from langgraph.utils.runnable import RunnableCallable def _default_format_error( diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 93eb1a07d..81b298d87 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -34,8 +34,6 @@ from langchain_core.runnables.config import ( ensure_config, get_async_callback_manager_for_config, get_callback_manager_for_config, - merge_configs, - patch_config, ) from langchain_core.runnables.utils import ( ConfigurableFieldSpec, @@ -76,7 +74,6 @@ from langgraph.pregel.algo import ( local_write, prepare_next_tasks, ) -from langgraph.pregel.config import patch_checkpoint_map, patch_configurable from langgraph.pregel.debug import tasks_w_writes from langgraph.pregel.io import read_channels from langgraph.pregel.loop import AsyncPregelLoop, SyncPregelLoop @@ -96,7 +93,13 @@ from langgraph.pregel.utils import ( from langgraph.pregel.validate import validate_graph, validate_keys from langgraph.pregel.write import ChannelWrite, ChannelWriteEntry from langgraph.store.base import BaseStore -from langgraph.utils import RunnableCallable +from langgraph.utils.config import ( + merge_configs, + patch_checkpoint_map, + patch_config, + patch_configurable, +) +from langgraph.utils.runnable import RunnableCallable WriteValue = Union[ Runnable[Input, Output], diff --git a/libs/langgraph/langgraph/pregel/algo.py b/libs/langgraph/langgraph/pregel/algo.py index f6301a379..9c57e707a 100644 --- a/libs/langgraph/langgraph/pregel/algo.py +++ b/libs/langgraph/langgraph/pregel/algo.py @@ -17,11 +17,7 @@ from typing import ( from uuid import UUID, uuid5 from langchain_core.callbacks.manager import AsyncParentRunManager, ParentRunManager -from langchain_core.runnables.config import ( - RunnableConfig, - merge_configs, - patch_config, -) +from langchain_core.runnables.config import RunnableConfig from langgraph.channels.base import BaseChannel from langgraph.checkpoint.base import ( @@ -51,6 +47,7 @@ from langgraph.pregel.log import logger from langgraph.pregel.manager import ChannelsManager from langgraph.pregel.read import PregelNode from langgraph.pregel.types import All, PregelExecutableTask, PregelTask +from langgraph.utils.config import merge_configs, patch_config class WritesProtocol(Protocol): diff --git a/libs/langgraph/langgraph/pregel/config.py b/libs/langgraph/langgraph/pregel/config.py deleted file mode 100644 index f4c64eb39..000000000 --- a/libs/langgraph/langgraph/pregel/config.py +++ /dev/null @@ -1,34 +0,0 @@ -from typing import Any, Optional - -from langchain_core.runnables import RunnableConfig - -from langgraph.checkpoint.base import CheckpointMetadata -from langgraph.constants import CONFIG_KEY_CHECKPOINT_MAP - - -def patch_configurable( - config: Optional[RunnableConfig], patch: dict[str, Any] -) -> RunnableConfig: - if config is None: - return {"configurable": patch} - else: - return {**config, "configurable": {**config["configurable"], **patch}} - - -def patch_checkpoint_map( - config: RunnableConfig, metadata: Optional[CheckpointMetadata] -) -> RunnableConfig: - if parents := (metadata.get("parents") if metadata else None): - return patch_configurable( - config, - { - CONFIG_KEY_CHECKPOINT_MAP: { - **parents, - config["configurable"]["checkpoint_ns"]: config["configurable"][ - "checkpoint_id" - ], - }, - }, - ) - else: - return config diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index c903730dd..d2052b17a 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -59,7 +59,6 @@ from langgraph.pregel.algo import ( prepare_next_tasks, should_interrupt, ) -from langgraph.pregel.config import patch_configurable from langgraph.pregel.debug import ( map_debug_checkpoint, map_debug_task_results, @@ -86,6 +85,7 @@ from langgraph.pregel.types import PregelExecutableTask from langgraph.pregel.utils import get_new_channel_versions from langgraph.store.base import BaseStore from langgraph.store.batch import AsyncBatchedStore +from langgraph.utils.config import patch_configurable V = TypeVar("V") INPUT_DONE = object() diff --git a/libs/langgraph/langgraph/pregel/manager.py b/libs/langgraph/langgraph/pregel/manager.py index ccfc8de8a..f70c86d46 100644 --- a/libs/langgraph/langgraph/pregel/manager.py +++ b/libs/langgraph/langgraph/pregel/manager.py @@ -2,7 +2,7 @@ import asyncio from contextlib import AsyncExitStack, ExitStack, asynccontextmanager, contextmanager from typing import AsyncIterator, Iterator, Mapping, Optional, Union -from langchain_core.runnables import RunnableConfig, patch_config +from langchain_core.runnables import RunnableConfig from langgraph.channels.base import BaseChannel from langgraph.checkpoint.base import Checkpoint @@ -14,6 +14,7 @@ from langgraph.managed.base import ( ) from langgraph.managed.context import Context from langgraph.store.base import BaseStore +from langgraph.utils.config import patch_configurable @contextmanager @@ -26,7 +27,7 @@ def ChannelsManager( skip_context: bool = False, ) -> Iterator[tuple[Mapping[str, BaseChannel], ManagedValueMapping]]: """Manage channels for the lifetime of a Pregel invocation (multiple steps).""" - config_for_managed = patch_config(config, configurable={CONFIG_KEY_STORE: store}) + config_for_managed = patch_configurable(config, {CONFIG_KEY_STORE: store}) channel_specs: Mapping[str, BaseChannel] = {} managed_specs: Mapping[str, ManagedValueSpec] = {} for k, v in specs.items(): @@ -69,7 +70,7 @@ async def AsyncChannelsManager( skip_context: bool = False, ) -> AsyncIterator[Mapping[str, BaseChannel]]: """Manage channels for the lifetime of a Pregel invocation (multiple steps).""" - config_for_managed = patch_config(config, configurable={CONFIG_KEY_STORE: store}) + config_for_managed = patch_configurable(config, {CONFIG_KEY_STORE: store}) channel_specs: Mapping[str, BaseChannel] = {} managed_specs: Mapping[str, ManagedValueSpec] = {} for k, v in specs.items(): diff --git a/libs/langgraph/langgraph/pregel/read.py b/libs/langgraph/langgraph/pregel/read.py index ca0828fda..e111ad820 100644 --- a/libs/langgraph/langgraph/pregel/read.py +++ b/libs/langgraph/langgraph/pregel/read.py @@ -19,13 +19,13 @@ from langchain_core.runnables import ( RunnableSerializable, ) from langchain_core.runnables.base import Input, Other, Output, coerce_to_runnable -from langchain_core.runnables.config import merge_configs from langchain_core.runnables.utils import ConfigurableFieldSpec from langgraph.constants import CONFIG_KEY_READ from langgraph.pregel.retry import RetryPolicy from langgraph.pregel.write import ChannelWrite -from langgraph.utils import RunnableCallable +from langgraph.utils.config import merge_configs +from langgraph.utils.runnable import RunnableCallable READ_TYPE = Callable[[str, bool], Union[Any, dict[str, Any]]] diff --git a/libs/langgraph/langgraph/pregel/write.py b/libs/langgraph/langgraph/pregel/write.py index e85885cc9..9fcb44696 100644 --- a/libs/langgraph/langgraph/pregel/write.py +++ b/libs/langgraph/langgraph/pregel/write.py @@ -18,7 +18,7 @@ from langchain_core.runnables.utils import ConfigurableFieldSpec from langgraph.constants import CONFIG_KEY_SEND, TASKS, Send from langgraph.errors import InvalidUpdateError -from langgraph.utils import RunnableCallable +from langgraph.utils.runnable import RunnableCallable TYPE_SEND = Callable[[Sequence[tuple[str, Any]]], None] R = TypeVar("R", bound=Runnable) diff --git a/libs/langgraph/langgraph/utils/__init__.py b/libs/langgraph/langgraph/utils/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/libs/langgraph/langgraph/utils/config.py b/libs/langgraph/langgraph/utils/config.py new file mode 100644 index 000000000..44fb29d56 --- /dev/null +++ b/libs/langgraph/langgraph/utils/config.py @@ -0,0 +1,150 @@ +from typing import Any, Optional + +from langchain_core.callbacks import Callbacks +from langchain_core.runnables import RunnableConfig +from langchain_core.runnables.config import COPIABLE_KEYS, DEFAULT_RECURSION_LIMIT + +from langgraph.checkpoint.base import CheckpointMetadata +from langgraph.constants import CONFIG_KEY_CHECKPOINT_MAP + + +def patch_configurable( + config: Optional[RunnableConfig], patch: dict[str, Any] +) -> RunnableConfig: + if config is None: + return {"configurable": patch} + else: + return {**config, "configurable": {**config["configurable"], **patch}} + + +def patch_checkpoint_map( + config: RunnableConfig, metadata: Optional[CheckpointMetadata] +) -> RunnableConfig: + if parents := (metadata.get("parents") if metadata else None): + return patch_configurable( + config, + { + CONFIG_KEY_CHECKPOINT_MAP: { + **parents, + config["configurable"]["checkpoint_ns"]: config["configurable"][ + "checkpoint_id" + ], + }, + }, + ) + else: + return config + + +def merge_configs(*configs: Optional[RunnableConfig]) -> RunnableConfig: + """Merge multiple configs into one. + + Args: + *configs (Optional[RunnableConfig]): The configs to merge. + + Returns: + RunnableConfig: The merged config. + """ + base: RunnableConfig = {} + # Even though the keys aren't literals, this is correct + # because both dicts are the same type + for config in configs: + if config is None: + continue + for key in config: + if key == "metadata": + base[key] = { # type: ignore + **base.get(key, {}), # type: ignore + **(config.get(key) or {}), # type: ignore + } + elif key == "tags": + base[key] = sorted( # type: ignore + set(base.get(key, []) + (config.get(key) or [])), # type: ignore + ) + elif key == "configurable": + base[key] = { # type: ignore + **base.get(key, {}), # type: ignore + **(config.get(key) or {}), # type: ignore + } + elif key == "callbacks": + base_callbacks = base.get("callbacks") + these_callbacks = config["callbacks"] + # callbacks can be either None, list[handler] or manager + # so merging two callbacks values has 6 cases + if isinstance(these_callbacks, list): + if base_callbacks is None: + base["callbacks"] = these_callbacks.copy() + elif isinstance(base_callbacks, list): + base["callbacks"] = base_callbacks + these_callbacks + else: + # base_callbacks is a manager + mngr = base_callbacks.copy() + for callback in these_callbacks: + mngr.add_handler(callback, inherit=True) + base["callbacks"] = mngr + elif these_callbacks is not None: + # these_callbacks is a manager + if base_callbacks is None: + base["callbacks"] = these_callbacks.copy() + elif isinstance(base_callbacks, list): + mngr = these_callbacks.copy() + for callback in base_callbacks: + mngr.add_handler(callback, inherit=True) + base["callbacks"] = mngr + else: + # base_callbacks is also a manager + base["callbacks"] = base_callbacks.merge(these_callbacks) + elif key == "recursion_limit": + if config["recursion_limit"] != DEFAULT_RECURSION_LIMIT: + base["recursion_limit"] = config["recursion_limit"] + elif key in COPIABLE_KEYS and config[key] is not None: # type: ignore[literal-required] + base[key] = config[key].copy() # type: ignore[literal-required] + else: + base[key] = config[key] or base.get(key) # type: ignore + return base + + +def patch_config( + config: Optional[RunnableConfig], + *, + callbacks: Optional[Callbacks] = None, + recursion_limit: Optional[int] = None, + max_concurrency: Optional[int] = None, + run_name: Optional[str] = None, + configurable: Optional[dict[str, Any]] = None, +) -> RunnableConfig: + """Patch a config with new values. + + Args: + config (Optional[RunnableConfig]): The config to patch. + callbacks (Optional[BaseCallbackManager], optional): The callbacks to set. + Defaults to None. + recursion_limit (Optional[int], optional): The recursion limit to set. + Defaults to None. + max_concurrency (Optional[int], optional): The max concurrency to set. + Defaults to None. + run_name (Optional[str], optional): The run name to set. Defaults to None. + configurable (Optional[Dict[str, Any]], optional): The configurable to set. + Defaults to None. + + Returns: + RunnableConfig: The patched config. + """ + config = config or {} + if callbacks is not None: + # If we're replacing callbacks, we need to unset run_name + # As that should apply only to the same run as the original callbacks + config["callbacks"] = callbacks + if "run_name" in config: + del config["run_name"] + if "run_id" in config: + del config["run_id"] + if recursion_limit is not None: + config["recursion_limit"] = recursion_limit + if max_concurrency is not None: + config["max_concurrency"] = max_concurrency + if run_name is not None: + config["run_name"] = run_name + if configurable is not None: + config["configurable"] = {**config.get("configurable", {}), **configurable} + return config diff --git a/libs/langgraph/langgraph/utils/fields.py b/libs/langgraph/langgraph/utils/fields.py new file mode 100644 index 000000000..55a8e81be --- /dev/null +++ b/libs/langgraph/langgraph/utils/fields.py @@ -0,0 +1,101 @@ +from typing import Any, Optional, Type, Union + +from typing_extensions import ( + Annotated, + NotRequired, + ReadOnly, + Required, + get_origin, +) + + +def _is_optional_type(type_: Any) -> bool: + """Check if a type is Optional.""" + + if hasattr(type_, "__origin__") and hasattr(type_, "__args__"): + origin = get_origin(type_) + if origin is Optional: + return True + if origin is Union: + return any( + arg is type(None) or _is_optional_type(arg) for arg in type_.__args__ + ) + if origin is Annotated: + return _is_optional_type(type_.__args__[0]) + return origin is None + if hasattr(type_, "__bound__") and type_.__bound__ is not None: + return _is_optional_type(type_.__bound__) + return type_ is None + + +def _is_required_type(type_: Any) -> Optional[bool]: + """Check if an annotation is marked as Required/NotRequired. + + Returns: + - True if required + - False if not required + - None if not annotated with either + """ + origin = get_origin(type_) + if origin is Required: + return True + if origin is NotRequired: + return False + if origin is Annotated or getattr(origin, "__args__", None): + # See https://typing.readthedocs.io/en/latest/spec/typeddict.html#interaction-with-annotated + return _is_required_type(type_.__args__[0]) + return None + + +def _is_readonly_type(type_: Any) -> bool: + """Check if an annotation is marked as ReadOnly. + + Returns: + - True if is read only + - False if not read only + """ + + # See: https://typing.readthedocs.io/en/latest/spec/typeddict.html#typing-readonly-type-qualifier + origin = get_origin(type_) + if origin is Annotated: + return _is_readonly_type(type_.__args__[0]) + if origin is ReadOnly: + return True + return False + + +_DEFAULT_KEYS = frozenset() + + +def get_field_default(name: str, type_: Any, schema: Type[Any]) -> Any: + """Determine the default value for a field in a state schema. + + This is based on: + If TypedDict: + - Required/NotRequired + - total=False -> everything optional + - Type annotation (Optional/Union[None]) + """ + optional_keys = getattr(schema, "__optional_keys__", _DEFAULT_KEYS) + irq = _is_required_type(type_) + if name in optional_keys: + # Either total=False or explicit NotRequired. + # No type annotation trumps this. + if irq: + # Unless it's earlier versions of python & explicit Required + return ... + return None + if irq is not None: + if irq: + # Handle Required[] + # (we already handled NotRequired and total=False) + return ... + # Handle NotRequired[] for earlier versions of python + return None + # Note, we ignore ReadOnly attributes, + # as they don't make much sense. (we don't care if you mutate the state in your node) + # and mutating state in your node has no effect on our graph state. + # Base case is the annotation + if _is_optional_type(type_): + return None + return ... diff --git a/libs/langgraph/langgraph/utils.py b/libs/langgraph/langgraph/utils/runnable.py similarity index 58% rename from libs/langgraph/langgraph/utils.py rename to libs/langgraph/langgraph/utils/runnable.py index cc3fd4760..bbe344ed5 100644 --- a/libs/langgraph/langgraph/utils.py +++ b/libs/langgraph/langgraph/utils/runnable.py @@ -4,8 +4,9 @@ import inspect import sys from contextvars import copy_context from functools import partial, wraps -from typing import Any, AsyncIterator, Awaitable, Callable, Optional, Type, Union +from typing import Any, AsyncIterator, Awaitable, Callable, Optional +from langchain_core.load.serializable import to_json_not_implemented from langchain_core.runnables.base import ( Runnable, RunnableConfig, @@ -14,19 +15,16 @@ from langchain_core.runnables.base import ( RunnableParallel, ) from langchain_core.runnables.config import ( - merge_configs, + ensure_config, + get_async_callback_manager_for_config, + get_callback_manager_for_config, run_in_executor, var_child_runnable_config, ) -from langchain_core.runnables.utils import accepts_config -from typing_extensions import ( - Annotated, - NotRequired, - ReadOnly, - Required, - TypeGuard, - get_origin, -) +from langchain_core.runnables.utils import accepts_config, accepts_run_manager +from typing_extensions import TypeGuard + +from langgraph.utils.config import merge_configs, patch_config try: from langchain_core.runnables.config import _set_config_context @@ -42,6 +40,9 @@ class StrEnum(str, enum.Enum): """A string enum.""" +ASYNCIO_ACCEPTS_CONTEXT = sys.version_info >= (3, 11) + + class RunnableCallable(Runnable): """A much simpler version of RunnableLambda that requires sync and async functions.""" @@ -70,11 +71,18 @@ class RunnableCallable(Runnable): except AttributeError: pass self.func = func + if func is not None: + self.func_accepts_config = accepts_config(func) + self.func_accepts_run_manager = accepts_run_manager(func) self.afunc = afunc + if afunc is not None: + self.afunc_accepts_config = accepts_config(afunc) + self.afunc_accepts_run_manager = accepts_run_manager(afunc) self.config: Optional[RunnableConfig] = {"tags": tags} if tags else None self.kwargs = kwargs self.trace = trace self.recurse = recurse + self.serialized = to_json_not_implemented(self) def __repr__(self) -> str: repr_args = { @@ -94,15 +102,34 @@ class RunnableCallable(Runnable): " via the async API (ainvoke, astream, etc.)" ) kwargs = {**self.kwargs, **kwargs} + config = ensure_config(merge_configs(self.config, config)) + context = copy_context() if self.trace: - ret = self._call_with_config( - self.func, input, merge_configs(self.config, config), **kwargs + config = ensure_config(config) + callback_manager = get_callback_manager_for_config(config) + run_manager = callback_manager.on_chain_start( + self.serialized, + input, + name=config.get("run_name") or self.get_name(), + run_id=config.pop("run_id", None), ) + try: + child_config = patch_config(config, callbacks=run_manager.get_child()) + context = copy_context() + context.run(_set_config_context, child_config) + if self.func_accepts_config: + kwargs["config"] = config + if self.func_accepts_run_manager: + kwargs["run_manager"] = run_manager + ret = context.run(self.func, input, **kwargs) + except BaseException as e: + run_manager.on_chain_error(e) + raise + else: + run_manager.on_chain_end(ret) else: - config = merge_configs(self.config, config) - context = copy_context() context.run(_set_config_context, config) - if accepts_config(self.func): + if self.func_accepts_config: kwargs["config"] = config ret = context.run(self.func, input, **kwargs) if isinstance(ret, Runnable) and self.recurse: @@ -115,17 +142,38 @@ class RunnableCallable(Runnable): if not self.afunc: return self.invoke(input, config) kwargs = {**self.kwargs, **kwargs} + config = ensure_config(merge_configs(self.config, config)) + context = copy_context() if self.trace: - ret = await self._acall_with_config( - self.afunc, input, merge_configs(self.config, config), **kwargs + callback_manager = get_async_callback_manager_for_config(config) + run_manager = await callback_manager.on_chain_start( + self.serialized, + input, + name=config.get("run_name") or self.name, + run_id=config.pop("run_id", None), ) + try: + child_config = patch_config(config, callbacks=run_manager.get_child()) + context.run(_set_config_context, child_config) + if self.afunc_accepts_config: + kwargs["config"] = config + if self.afunc_accepts_run_manager: + kwargs["run_manager"] = run_manager + coro = self.afunc(input, **kwargs) + if ASYNCIO_ACCEPTS_CONTEXT: + ret = await asyncio.create_task(coro, context=context) + else: + ret = await coro + except BaseException as e: + await run_manager.on_chain_error(e) + raise + else: + await run_manager.on_chain_end(ret) else: - config = merge_configs(self.config, config) - context = copy_context() context.run(_set_config_context, config) - if accepts_config(self.afunc): + if self.afunc_accepts_config: kwargs["config"] = config - if sys.version_info >= (3, 11): + if ASYNCIO_ACCEPTS_CONTEXT: ret = await asyncio.create_task( self.afunc(input, **kwargs), context=context ) @@ -188,95 +236,3 @@ def coerce_to_runnable(thing: RunnableLike, *, name: str, trace: bool) -> Runnab f"Expected a Runnable, callable or dict." f"Instead got an unsupported type: {type(thing)}" ) - - -def _is_optional_type(type_: Any) -> bool: - """Check if a type is Optional.""" - - if hasattr(type_, "__origin__") and hasattr(type_, "__args__"): - origin = get_origin(type_) - if origin is Optional: - return True - if origin is Union: - return any( - arg is type(None) or _is_optional_type(arg) for arg in type_.__args__ - ) - if origin is Annotated: - return _is_optional_type(type_.__args__[0]) - return origin is None - if hasattr(type_, "__bound__") and type_.__bound__ is not None: - return _is_optional_type(type_.__bound__) - return type_ is None - - -def _is_required_type(type_: Any) -> Optional[bool]: - """Check if an annotation is marked as Required/NotRequired. - - Returns: - - True if required - - False if not required - - None if not annotated with either - """ - origin = get_origin(type_) - if origin is Required: - return True - if origin is NotRequired: - return False - if origin is Annotated or getattr(origin, "__args__", None): - # See https://typing.readthedocs.io/en/latest/spec/typeddict.html#interaction-with-annotated - return _is_required_type(type_.__args__[0]) - return None - - -def _is_readonly_type(type_: Any) -> bool: - """Check if an annotation is marked as ReadOnly. - - Returns: - - True if is read only - - False if not read only - """ - - # See: https://typing.readthedocs.io/en/latest/spec/typeddict.html#typing-readonly-type-qualifier - origin = get_origin(type_) - if origin is Annotated: - return _is_readonly_type(type_.__args__[0]) - if origin is ReadOnly: - return True - return False - - -_DEFAULT_KEYS = frozenset() - - -def get_field_default(name: str, type_: Any, schema: Type[Any]) -> Any: - """Determine the default value for a field in a state schema. - - This is based on: - If TypedDict: - - Required/NotRequired - - total=False -> everything optional - - Type annotation (Optional/Union[None]) - """ - optional_keys = getattr(schema, "__optional_keys__", _DEFAULT_KEYS) - irq = _is_required_type(type_) - if name in optional_keys: - # Either total=False or explicit NotRequired. - # No type annotation trumps this. - if irq: - # Unless it's earlier versions of python & explicit Required - return ... - return None - if irq is not None: - if irq: - # Handle Required[] - # (we already handled NotRequired and total=False) - return ... - # Handle NotRequired[] for earlier versions of python - return None - # Note, we ignore ReadOnly attributes, - # as they don't make much sense. (we don't care if you mutate the state in your node) - # and mutating state in your node has no effect on our graph state. - # Base case is the annotation - if _is_optional_type(type_): - return None - return ... diff --git a/libs/langgraph/tests/test_utils.py b/libs/langgraph/tests/test_utils.py index bee8fcd9c..e8ea94fff 100644 --- a/libs/langgraph/tests/test_utils.py +++ b/libs/langgraph/tests/test_utils.py @@ -21,12 +21,8 @@ from typing_extensions import Annotated, NotRequired, Required from langgraph.graph import END, StateGraph from langgraph.graph.graph import CompiledGraph -from langgraph.utils import ( - _is_optional_type, - get_field_default, - is_async_callable, - is_async_generator, -) +from langgraph.utils.fields import _is_optional_type, get_field_default +from langgraph.utils.runnable import is_async_callable, is_async_generator pytestmark = pytest.mark.anyio From 535dbd5b06eb7976316badd3c2891f99bc2e02f6 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 4 Sep 2024 10:47:16 -0700 Subject: [PATCH 2/9] Reduce from 9s to 4s on benchmark graph - Use a simpler version of RunnableSequence without tracing serialization - Remove accepts_run_manager check in RunnableCallable - Remove creation of ChannelWrite dynamically every time conditional edge runs --- libs/langgraph/langgraph/graph/graph.py | 22 +- libs/langgraph/langgraph/graph/state.py | 25 +- libs/langgraph/langgraph/pregel/__init__.py | 22 +- libs/langgraph/langgraph/pregel/algo.py | 4 +- libs/langgraph/langgraph/pregel/read.py | 18 +- libs/langgraph/langgraph/pregel/write.py | 94 +++-- libs/langgraph/langgraph/utils/config.py | 4 +- libs/langgraph/langgraph/utils/runnable.py | 321 ++++++++++++++++-- libs/langgraph/poetry.lock | 15 +- .../tests/__snapshots__/test_pregel.ambr | 39 +++ libs/langgraph/tests/test_pregel.py | 26 +- libs/langgraph/tests/test_pregel_async.py | 4 +- 12 files changed, 446 insertions(+), 148 deletions(-) diff --git a/libs/langgraph/langgraph/graph/graph.py b/libs/langgraph/langgraph/graph/graph.py index af6700c2a..41175c170 100644 --- a/libs/langgraph/langgraph/graph/graph.py +++ b/libs/langgraph/langgraph/graph/graph.py @@ -55,7 +55,7 @@ class Branch(NamedTuple): def run( self, - writer: Callable[[list[str]], Optional[Runnable]], + writer: Callable[[list[str], RunnableConfig], None], reader: Optional[Callable[[RunnableConfig], Any]] = None, ) -> None: return ChannelWrite.register_writer( @@ -75,7 +75,7 @@ class Branch(NamedTuple): config: RunnableConfig, *, reader: Optional[Callable[[], Any]], - writer: Callable[[list[str]], Optional[Runnable]], + writer: Callable[[list[str], RunnableConfig], None], ) -> Runnable: if reader: value = reader(config) @@ -86,7 +86,7 @@ class Branch(NamedTuple): else: value = input result = self.path.invoke(value, config) - return self._finish(writer, input, result) + return self._finish(writer, input, result, config) async def _aroute( self, @@ -94,7 +94,7 @@ class Branch(NamedTuple): config: RunnableConfig, *, reader: Optional[Callable[[], Any]], - writer: Callable[[list[str]], Optional[Runnable]], + writer: Callable[[list[str], RunnableConfig], Optional[Runnable]], ) -> Runnable: if reader: value = reader(config) @@ -105,10 +105,14 @@ class Branch(NamedTuple): else: value = input result = await self.path.ainvoke(value, config) - return self._finish(writer, input, result) + return self._finish(writer, input, result, config) def _finish( - self, writer: Callable[[list[str]], Optional[Runnable]], input: Any, result: Any + self, + writer: Callable[[list[str], RunnableConfig], None], + input: Any, + result: Any, + config: RunnableConfig, ): if not isinstance(result, list): result = [result] @@ -120,7 +124,7 @@ class Branch(NamedTuple): raise ValueError("Branch did not return a valid destination") if any(p.node == END for p in destinations if isinstance(p, Send)): raise InvalidUpdateError("Cannot send a packet to the END node") - return writer(destinations) or input + return writer(destinations, config) or input class Graph: @@ -449,7 +453,9 @@ class CompiledGraph(Pregel): self.nodes[end].channels.append(start) def attach_branch(self, start: str, name: str, branch: Branch) -> None: - def branch_writer(packets: list[Union[str, Send]]) -> Optional[ChannelWrite]: + def branch_writer( + packets: list[Union[str, Send]], config: RunnableConfig + ) -> Optional[ChannelWrite]: writes = [ ( ChannelWriteEntry(f"branch:{start}:{name}:{p}" if p != END else END) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index ce25f6c5d..644ae7adb 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -45,10 +45,7 @@ from langgraph.pregel.types import All, RetryPolicy from langgraph.pregel.write import SKIP_WRITE, ChannelWrite, ChannelWriteEntry from langgraph.store.base import BaseStore from langgraph.utils.fields import get_field_default -from langgraph.utils.runnable import ( - RunnableCallable, - coerce_to_runnable, -) +from langgraph.utils.runnable import coerce_to_runnable logger = logging.getLogger(__name__) @@ -534,9 +531,7 @@ class CompiledStateGraph(CompiledGraph): if is_writable_managed_value(v) ] - def _get_state_key( - input: Union[None, dict, Any], config: RunnableConfig, *, key: str - ) -> Any: + def _get_state_key(input: Union[None, dict, Any], *, key: str) -> Any: if input is None: return SKIP_WRITE elif isinstance(input, dict): @@ -552,12 +547,7 @@ class CompiledStateGraph(CompiledGraph): [ChannelWriteEntry("__root__", skip_none=True)] if output_keys == ["__root__"] else [ - ChannelWriteEntry( - key, - mapper=RunnableCallable( - _get_state_key, key=key, trace=False, recurse=False - ), - ) + ChannelWriteEntry(key, mapper=partial(_get_state_key, key=key)) for key in output_keys ] ) @@ -600,7 +590,8 @@ class CompiledStateGraph(CompiledGraph): ], metadata=node.metadata, retry_policy=node.retry_policy, - ).pipe(node.runnable) + bound=node.runnable, + ) def attach_edge(self, starts: Union[str, Sequence[str]], end: str) -> None: if isinstance(starts, str): @@ -630,7 +621,9 @@ class CompiledStateGraph(CompiledGraph): ) def attach_branch(self, start: str, name: str, branch: Branch) -> None: - def branch_writer(packets: list[Union[str, Send]]) -> Optional[ChannelWrite]: + def branch_writer( + packets: list[Union[str, Send]], config: RunnableConfig + ) -> Optional[ChannelWrite]: if filtered := [p for p in packets if p != END]: writes = [ ( @@ -649,7 +642,7 @@ class CompiledStateGraph(CompiledGraph): ), ) ) - return ChannelWrite(writes, tags=[TAG_HIDDEN]) + ChannelWrite.do_write(config, writes) # attach branch publisher schema = ( diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 81b298d87..089b729fa 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -6,7 +6,6 @@ from functools import partial from typing import ( Any, AsyncIterator, - Awaitable, Callable, Dict, Iterator, @@ -28,7 +27,7 @@ from langchain_core.runnables import ( RunnableLambda, RunnableSequence, ) -from langchain_core.runnables.base import Input, Output, coerce_to_runnable +from langchain_core.runnables.base import Input, Output from langchain_core.runnables.config import ( RunnableConfig, ensure_config, @@ -101,12 +100,7 @@ from langgraph.utils.config import ( ) from langgraph.utils.runnable import RunnableCallable -WriteValue = Union[ - Runnable[Input, Output], - Callable[[Input], Output], - Callable[[Input], Awaitable[Output]], - Any, -] +WriteValue = Union[Callable[[Input], Output], Any] class Channel: @@ -171,11 +165,9 @@ class Channel: return ChannelWrite( [ChannelWriteEntry(c) for c in channels] + [ - ( - ChannelWriteEntry(k, skip_none=True, mapper=coerce_to_runnable(v)) - if isinstance(v, Runnable) or callable(v) - else ChannelWriteEntry(k, value=v) - ) + ChannelWriteEntry(k, mapper=v) + if callable(v) + else ChannelWriteEntry(k, value=v) for k, v in kwargs.items() ] ) @@ -812,7 +804,7 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): managed, ): # create task to run all writers of the chosen node - writers = self.nodes[as_node].get_writers() + writers = self.nodes[as_node].flat_writers if not writers: raise InvalidUpdateError(f"Node {as_node} has no writers") task = PregelExecutableTask( @@ -976,7 +968,7 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): managed, ): # create task to run all writers of the chosen node - writers = self.nodes[as_node].get_writers() + writers = self.nodes[as_node].flat_writers if not writers: raise InvalidUpdateError(f"Node {as_node} has no writers") task = PregelExecutableTask( diff --git a/libs/langgraph/langgraph/pregel/algo.py b/libs/langgraph/langgraph/pregel/algo.py index 9c57e707a..3e465f115 100644 --- a/libs/langgraph/langgraph/pregel/algo.py +++ b/libs/langgraph/langgraph/pregel/algo.py @@ -297,7 +297,7 @@ def prepare_next_tasks( ) if for_execution: proc = processes[packet.node] - if node := proc.get_node(): + if node := proc.node: managed.replace_runtime_placeholders(step, packet.arg) writes = deque() task_checkpoint_ns = f"{checkpoint_ns}:{task_id}" @@ -403,7 +403,7 @@ def prepare_next_tasks( ) if for_execution: - if node := proc.get_node(): + if node := proc.node: writes = deque() task_checkpoint_ns = f"{checkpoint_ns}:{task_id}" tasks.append( diff --git a/libs/langgraph/langgraph/pregel/read.py b/libs/langgraph/langgraph/pregel/read.py index e111ad820..4d9944661 100644 --- a/libs/langgraph/langgraph/pregel/read.py +++ b/libs/langgraph/langgraph/pregel/read.py @@ -1,5 +1,6 @@ from __future__ import annotations +from functools import cached_property from typing import ( Any, AsyncIterator, @@ -15,7 +16,6 @@ from langchain_core.runnables import ( Runnable, RunnableConfig, RunnablePassthrough, - RunnableSequence, RunnableSerializable, ) from langchain_core.runnables.base import Input, Other, Output, coerce_to_runnable @@ -25,7 +25,7 @@ from langgraph.constants import CONFIG_KEY_READ from langgraph.pregel.retry import RetryPolicy from langgraph.pregel.write import ChannelWrite from langgraph.utils.config import merge_configs -from langgraph.utils.runnable import RunnableCallable +from langgraph.utils.runnable import RunnableCallable, RunnableSeq READ_TYPE = Callable[[str, bool], Union[Any, dict[str, Any]]] @@ -149,7 +149,8 @@ class PregelNode(Runnable): attrs = {**self.__dict__, **update} return PregelNode(**attrs) - def get_writers(self) -> list[Runnable]: + @cached_property + def flat_writers(self) -> list[Runnable]: """Get writers with optimizations applied.""" writers = self.writers.copy() while ( @@ -167,16 +168,17 @@ class PregelNode(Runnable): writers.pop() return writers - def get_node(self) -> Optional[Runnable[Any, Any]]: - writers = self.get_writers() + @cached_property + def node(self) -> Optional[Runnable[Any, Any]]: + writers = self.flat_writers if self.bound is DEFAULT_BOUND and not writers: return None elif self.bound is DEFAULT_BOUND and len(writers) == 1: return writers[0] elif self.bound is DEFAULT_BOUND: - return RunnableSequence(*writers) + return RunnableSeq(*writers) elif writers: - return RunnableSequence(self.bound, *writers) + return RunnableSeq(self.bound, *writers) else: return self.bound @@ -209,7 +211,7 @@ class PregelNode(Runnable): elif self.bound is DEFAULT_BOUND: return self.copy(update=dict(bound=coerce_to_runnable(other))) else: - return self.copy(update=dict(bound=self.bound | other)) + return self.copy(update=dict(bound=RunnableSeq(self.bound, other))) def pipe( self, diff --git a/libs/langgraph/langgraph/pregel/write.py b/libs/langgraph/langgraph/pregel/write.py index 9fcb44696..fd732966a 100644 --- a/libs/langgraph/langgraph/pregel/write.py +++ b/libs/langgraph/langgraph/pregel/write.py @@ -4,11 +4,9 @@ import asyncio from typing import ( Any, Callable, - List, NamedTuple, Optional, Sequence, - Tuple, TypeVar, Union, ) @@ -32,7 +30,7 @@ class ChannelWriteEntry(NamedTuple): channel: str value: Any = PASSTHROUGH skip_none: bool = False - mapper: Optional[Runnable] = None + mapper: Optional[Callable] = None class ChannelWrite(RunnableCallable): @@ -59,9 +57,6 @@ class ChannelWrite(RunnableCallable): self.writes = writes self.require_at_least_one_of = require_at_least_one_of - def __repr_args__(self) -> Any: - return [("writes", self.writes)] - def get_name( self, suffix: Optional[str] = None, *, name: Optional[str] = None ) -> str: @@ -82,65 +77,29 @@ class ChannelWrite(RunnableCallable): ] def _write(self, input: Any, config: RunnableConfig) -> None: - # split packets and entries - writes = [(TASKS, packet) for packet in self.writes if isinstance(packet, Send)] - entries = [ - write for write in self.writes if isinstance(write, ChannelWriteEntry) + writes = [ + ChannelWriteEntry(write.channel, input, write.skip_none, write.mapper) + if isinstance(write, ChannelWriteEntry) and write.value is PASSTHROUGH + else write + for write in self.writes ] - for entry in entries: - if entry.channel == TASKS: - raise InvalidUpdateError("Cannot write to the reserved channel TASKS") - # process entries into values - values = [ - input if write.value is PASSTHROUGH else write.value for write in entries - ] - values = [ - val if write.mapper is None else write.mapper.invoke(val, config) - for val, write in zip(values, entries) - ] - values = [ - (write.channel, val) - for val, write in zip(values, entries) - if not write.skip_none or val is not None - ] - # write packets and values self.do_write( config, - writes + values, + writes, self.require_at_least_one_of if input is not None else None, ) return input async def _awrite(self, input: Any, config: RunnableConfig) -> None: - # split packets and entries - writes = [(TASKS, packet) for packet in self.writes if isinstance(packet, Send)] - entries = [ - write for write in self.writes if isinstance(write, ChannelWriteEntry) + writes = [ + ChannelWriteEntry(write.channel, input, write.skip_none, write.mapper) + if isinstance(write, ChannelWriteEntry) and write.value is PASSTHROUGH + else write + for write in self.writes ] - for entry in entries: - if entry.channel == TASKS: - raise InvalidUpdateError("Cannot write to the reserved channel TASKS") - # process entries into values - values = [ - input if write.value is PASSTHROUGH else write.value for write in entries - ] - values = await asyncio.gather( - *( - _mk_future(val) - if write.mapper is None - else write.mapper.ainvoke(val, config) - for val, write in zip(values, entries) - ) - ) - values = [ - (write.channel, val) - for val, write in zip(values, entries) - if not write.skip_none or val is not None - ] - # write packets and values self.do_write( config, - writes + values, + writes, self.require_at_least_one_of if input is not None else None, ) return input @@ -148,9 +107,32 @@ class ChannelWrite(RunnableCallable): @staticmethod def do_write( config: RunnableConfig, - values: List[Tuple[str, Any]], + writes: Sequence[Union[ChannelWriteEntry, Send]], require_at_least_one_of: Optional[Sequence[str]] = None, ) -> None: + # validate + for w in writes: + if isinstance(w, ChannelWriteEntry): + if w.channel == TASKS: + raise InvalidUpdateError( + "Cannot write to the reserved channel TASKS" + ) + if w.value is PASSTHROUGH: + raise InvalidUpdateError("PASSTHROUGH value must be replaced") + # split packets and entries + sends = [(TASKS, packet) for packet in writes if isinstance(packet, Send)] + entries = [write for write in writes if isinstance(write, ChannelWriteEntry)] + # process entries into values + values = [ + write.mapper(write.value) if write.mapper is not None else write.value + for write in entries + ] + values = [ + (write.channel, val) + for val, write in zip(values, entries) + if not write.skip_none or val is not None + ] + # filter out SKIP_WRITE values filtered = [(chan, val) for chan, val in values if val is not SKIP_WRITE] if require_at_least_one_of is not None: if not {chan for chan, _ in filtered} & set(require_at_least_one_of): @@ -158,7 +140,7 @@ class ChannelWrite(RunnableCallable): f"Must write to at least one of {require_at_least_one_of}" ) write: TYPE_SEND = config["configurable"][CONFIG_KEY_SEND] - write(filtered) + write(sends + filtered) @staticmethod def is_writer(runnable: Runnable) -> bool: diff --git a/libs/langgraph/langgraph/utils/config.py b/libs/langgraph/langgraph/utils/config.py index 44fb29d56..461c89024 100644 --- a/libs/langgraph/langgraph/utils/config.py +++ b/libs/langgraph/langgraph/utils/config.py @@ -13,6 +13,8 @@ def patch_configurable( ) -> RunnableConfig: if config is None: return {"configurable": patch} + elif "configurable" not in config: + return {**config, "configurable": patch} else: return {**config, "configurable": {**config["configurable"], **patch}} @@ -130,7 +132,7 @@ def patch_config( Returns: RunnableConfig: The patched config. """ - config = config or {} + config = config.copy() or {} if callbacks is not None: # If we're replacing callbacks, we need to unset run_name # As that should apply only to the same run as the original callbacks diff --git a/libs/langgraph/langgraph/utils/runnable.py b/libs/langgraph/langgraph/utils/runnable.py index bbe344ed5..3f3395e25 100644 --- a/libs/langgraph/langgraph/utils/runnable.py +++ b/libs/langgraph/langgraph/utils/runnable.py @@ -2,17 +2,18 @@ import asyncio import enum import inspect import sys +from contextlib import AsyncExitStack from contextvars import copy_context from functools import partial, wraps -from typing import Any, AsyncIterator, Awaitable, Callable, Optional +from typing import Any, AsyncIterator, Awaitable, Callable, Iterator, Optional -from langchain_core.load.serializable import to_json_not_implemented from langchain_core.runnables.base import ( Runnable, RunnableConfig, RunnableLambda, RunnableLike, RunnableParallel, + RunnableSequence, ) from langchain_core.runnables.config import ( ensure_config, @@ -21,7 +22,8 @@ from langchain_core.runnables.config import ( run_in_executor, var_child_runnable_config, ) -from langchain_core.runnables.utils import accepts_config, accepts_run_manager +from langchain_core.runnables.utils import Input, Output, accepts_config +from langchain_core.tracers._streaming import _StreamingCallbackHandler from typing_extensions import TypeGuard from langgraph.utils.config import merge_configs, patch_config @@ -73,16 +75,13 @@ class RunnableCallable(Runnable): self.func = func if func is not None: self.func_accepts_config = accepts_config(func) - self.func_accepts_run_manager = accepts_run_manager(func) self.afunc = afunc if afunc is not None: self.afunc_accepts_config = accepts_config(afunc) - self.afunc_accepts_run_manager = accepts_run_manager(afunc) self.config: Optional[RunnableConfig] = {"tags": tags} if tags else None self.kwargs = kwargs self.trace = trace self.recurse = recurse - self.serialized = to_json_not_implemented(self) def __repr__(self) -> str: repr_args = { @@ -102,13 +101,15 @@ class RunnableCallable(Runnable): " via the async API (ainvoke, astream, etc.)" ) kwargs = {**self.kwargs, **kwargs} + if self.func_accepts_config: + kwargs["config"] = config config = ensure_config(merge_configs(self.config, config)) context = copy_context() if self.trace: config = ensure_config(config) callback_manager = get_callback_manager_for_config(config) run_manager = callback_manager.on_chain_start( - self.serialized, + None, input, name=config.get("run_name") or self.get_name(), run_id=config.pop("run_id", None), @@ -117,10 +118,6 @@ class RunnableCallable(Runnable): child_config = patch_config(config, callbacks=run_manager.get_child()) context = copy_context() context.run(_set_config_context, child_config) - if self.func_accepts_config: - kwargs["config"] = config - if self.func_accepts_run_manager: - kwargs["run_manager"] = run_manager ret = context.run(self.func, input, **kwargs) except BaseException as e: run_manager.on_chain_error(e) @@ -129,8 +126,6 @@ class RunnableCallable(Runnable): run_manager.on_chain_end(ret) else: context.run(_set_config_context, config) - if self.func_accepts_config: - kwargs["config"] = config ret = context.run(self.func, input, **kwargs) if isinstance(ret, Runnable) and self.recurse: return ret.invoke(input, config) @@ -142,12 +137,14 @@ class RunnableCallable(Runnable): if not self.afunc: return self.invoke(input, config) kwargs = {**self.kwargs, **kwargs} + if self.afunc_accepts_config: + kwargs["config"] = config config = ensure_config(merge_configs(self.config, config)) context = copy_context() if self.trace: callback_manager = get_async_callback_manager_for_config(config) run_manager = await callback_manager.on_chain_start( - self.serialized, + None, input, name=config.get("run_name") or self.name, run_id=config.pop("run_id", None), @@ -155,10 +152,6 @@ class RunnableCallable(Runnable): try: child_config = patch_config(config, callbacks=run_manager.get_child()) context.run(_set_config_context, child_config) - if self.afunc_accepts_config: - kwargs["config"] = config - if self.afunc_accepts_run_manager: - kwargs["run_manager"] = run_manager coro = self.afunc(input, **kwargs) if ASYNCIO_ACCEPTS_CONTEXT: ret = await asyncio.create_task(coro, context=context) @@ -171,8 +164,6 @@ class RunnableCallable(Runnable): await run_manager.on_chain_end(ret) else: context.run(_set_config_context, config) - if self.afunc_accepts_config: - kwargs["config"] = config if ASYNCIO_ACCEPTS_CONTEXT: ret = await asyncio.create_task( self.afunc(input, **kwargs), context=context @@ -236,3 +227,293 @@ def coerce_to_runnable(thing: RunnableLike, *, name: str, trace: bool) -> Runnab f"Expected a Runnable, callable or dict." f"Instead got an unsupported type: {type(thing)}" ) + + +class RunnableSeq(Runnable): + """A simpler version of RunnableSequence.""" + + def __init__( + self, + *steps: RunnableLike, + name: Optional[str] = None, + ) -> None: + """Create a new RunnableSequence. + + Args: + steps: The steps to include in the sequence. + name: The name of the Runnable. Defaults to None. + first: The first Runnable in the sequence. Defaults to None. + middle: The middle Runnables in the sequence. Defaults to None. + last: The last Runnable in the sequence. Defaults to None. + + Raises: + ValueError: If the sequence has less than 2 steps. + """ + steps_flat: list[Runnable] = [] + for step in steps: + if isinstance(step, RunnableSequence): + steps_flat.extend(step.steps) + elif isinstance(step, RunnableSeq): + steps_flat.extend(step.steps) + else: + steps_flat.append(coerce_to_runnable(step, name=None, trace=True)) + if len(steps_flat) < 2: + raise ValueError( + f"RunnableSeq must have at least 2 steps, got {len(steps_flat)}" + ) + self.steps = steps_flat + self.name = name + + def __or__( + self, + other: Any, + ) -> Runnable: + if isinstance(other, RunnableSequence): + return RunnableSeq( + *self.steps, + other.first, + *other.middle, + other.last, + name=self.name or other.name, + ) + elif isinstance(other, RunnableSeq): + return RunnableSeq( + *self.steps, + *other.steps, + name=self.name or other.name, + ) + else: + return RunnableSeq( + *self.steps, + coerce_to_runnable(other), + name=self.name, + ) + + def __ror__( + self, + other: Any, + ) -> Runnable: + if isinstance(other, RunnableSequence): + return RunnableSequence( + other.first, + *other.middle, + other.last, + *self.steps, + name=other.name or self.name, + ) + elif isinstance(other, RunnableSeq): + return RunnableSeq( + *other.steps, + *self.steps, + name=other.name or self.name, + ) + else: + return RunnableSequence( + coerce_to_runnable(other), + *self.steps, + name=self.name, + ) + + def invoke( + self, input: Input, config: Optional[RunnableConfig] = None, **kwargs: Any + ) -> Output: + # setup callbacks and context + config = ensure_config(config) + callback_manager = get_callback_manager_for_config(config) + # start the root run + run_manager = callback_manager.on_chain_start( + None, + input, + name=config.get("run_name") or self.get_name(), + run_id=config.pop("run_id", None), + ) + + # invoke all steps in sequence + try: + for i, step in enumerate(self.steps): + # mark each step as a child run + config = patch_config( + config, callbacks=run_manager.get_child(f"seq:step:{i+1}") + ) + context = copy_context() + context.run(_set_config_context, config) + if i == 0: + input = context.run(step.invoke, input, config, **kwargs) + else: + input = context.run(step.invoke, input, config) + # finish the root run + except BaseException as e: + run_manager.on_chain_error(e) + raise + else: + run_manager.on_chain_end(input) + return input + + async def ainvoke( + self, + input: Input, + config: Optional[RunnableConfig] = None, + **kwargs: Optional[Any], + ) -> Output: + # setup callbacks + config = ensure_config(config) + callback_manager = get_async_callback_manager_for_config(config) + # start the root run + run_manager = await callback_manager.on_chain_start( + None, + input, + name=config.get("run_name") or self.get_name(), + run_id=config.pop("run_id", None), + ) + + # invoke all steps in sequence + try: + for i, step in enumerate(self.steps): + # mark each step as a child run + config = patch_config( + config, callbacks=run_manager.get_child(f"seq:step:{i+1}") + ) + context = copy_context() + context.run(_set_config_context, config) + if i == 0: + coro = step.ainvoke(input, config, **kwargs) + else: + coro = step.ainvoke(input, config) + if ASYNCIO_ACCEPTS_CONTEXT: + input = await asyncio.create_task(coro, context=context) + else: + input = await asyncio.create_task(coro) + # finish the root run + except BaseException as e: + await run_manager.on_chain_error(e) + raise + else: + await run_manager.on_chain_end(input) + return input + + def stream( + self, + input: Input, + config: Optional[RunnableConfig] = None, + **kwargs: Optional[Any], + ) -> Iterator[Output]: + # setup callbacks + config = ensure_config(config) + callback_manager = get_callback_manager_for_config(config) + # start the root run + run_manager = callback_manager.on_chain_start( + None, + input, + name=config.get("run_name") or self.get_name(), + run_id=config.pop("run_id", None), + ) + + try: + # stream the last steps + # transform the input stream of each step with the next + # steps that don't natively support transforming an input stream will + # buffer input in memory until all available, and then start emitting output + for idx, step in enumerate(self.steps): + config = patch_config( + config, + callbacks=run_manager.get_child(f"seq:step:{idx+1}"), + ) + if idx == 0: + iterator = step.stream(input, config, **kwargs) + else: + iterator = step.transform(iterator, config) + if stream_handler := next( + ( + h + for h in run_manager.handlers + if isinstance(h, _StreamingCallbackHandler) + ), + None, + ): + # populates streamed_output in astream_log() output if needed + iterator = stream_handler.tap_output_iter(run_manager.run_id, iterator) + output: Output = None + add_supported = False + for chunk in iterator: + yield chunk + # collect final output + if output is None: + output = chunk + elif add_supported: + try: + output = output + chunk + except TypeError: + output = chunk + add_supported = False + else: + output = chunk + except BaseException as e: + run_manager.on_chain_error(e) + raise + else: + run_manager.on_chain_end(output) + + async def astream( + self, + input: Input, + config: Optional[RunnableConfig] = None, + **kwargs: Optional[Any], + ) -> AsyncIterator[Output]: + # setup callbacks + config = ensure_config(config) + callback_manager = get_async_callback_manager_for_config(config) + # start the root run + run_manager = await callback_manager.on_chain_start( + None, + input, + name=config.get("run_name") or self.get_name(), + run_id=config.pop("run_id", None), + ) + + try: + async with AsyncExitStack() as stack: + # stream the last steps + # transform the input stream of each step with the next + # steps that don't natively support transforming an input stream will + # buffer input in memory until all available, and then start emitting output + for idx, step in enumerate(self.steps): + config = patch_config( + config, + callbacks=run_manager.get_child(f"seq:step:{idx+1}"), + ) + if idx == 0: + aiterator = step.astream(input, config, **kwargs) + else: + aiterator = step.atransform(aiterator, config) + if hasattr(aiterator, "aclose"): + stack.push_async_callback(aiterator.aclose) + if stream_handler := next( + ( + h + for h in run_manager.handlers + if isinstance(h, _StreamingCallbackHandler) + ), + None, + ): + # populates streamed_output in astream_log() output if needed + aiterator = stream_handler.tap_output_aiter( + run_manager.run_id, aiterator + ) + output: Output = None + add_supported = False + async for chunk in aiterator: + yield chunk + # collect final output + if add_supported: + try: + output = output + chunk + except TypeError: + output = chunk + add_supported = False + else: + output = chunk + except BaseException as e: + await run_manager.on_chain_error(e) + raise + else: + await run_manager.on_chain_end(output) diff --git a/libs/langgraph/poetry.lock b/libs/langgraph/poetry.lock index 50d6c6f89..2061c06b2 100644 --- a/libs/langgraph/poetry.lock +++ b/libs/langgraph/poetry.lock @@ -1,4 +1,4 @@ -# This file is automatically @generated by Poetry 1.8.3 and should not be changed by hand. +# This file is automatically @generated by Poetry 1.8.2 and should not be changed by hand. [[package]] name = "aiohappyeyeballs" @@ -1884,18 +1884,22 @@ url = "../checkpoint-sqlite" [[package]] name = "langsmith" -version = "0.1.79" +version = "0.1.111" description = "Client library to connect to the LangSmith LLM Tracing and Evaluation Platform." optional = false python-versions = "<4.0,>=3.8.1" files = [ - {file = "langsmith-0.1.79-py3-none-any.whl", hash = "sha256:c7f2c23981917713b5515b773f37c84ff68a7adf803476e2ebb5adcb36a04202"}, - {file = "langsmith-0.1.79.tar.gz", hash = "sha256:d215718cfdcdf4a011126b7a3d4a37eee96d887e59ac1e628a57e24b2bfa3163"}, + {file = "langsmith-0.1.111-py3-none-any.whl", hash = "sha256:e5c702764911193c9812fe55136ae01cd0b9ddf5dff0b068ce6fd60eeddbcb40"}, + {file = "langsmith-0.1.111.tar.gz", hash = "sha256:bab24fd6125685f588d682693c4a3253e163804242829b1ff902e1a3e984a94c"}, ] [package.dependencies] +httpx = ">=0.23.0,<1" orjson = ">=3.9.14,<4.0.0" -pydantic = ">=1,<3" +pydantic = [ + {version = ">=1,<3", markers = "python_full_version < \"3.12.4\""}, + {version = ">=2.7.4,<3.0.0", markers = "python_full_version >= \"3.12.4\""}, +] requests = ">=2,<3" [[package]] @@ -3051,6 +3055,7 @@ files = [ {file = "PyYAML-6.0.1-cp311-cp311-win_amd64.whl", hash = "sha256:bf07ee2fef7014951eeb99f56f39c9bb4af143d8aa3c21b1677805985307da34"}, {file = "PyYAML-6.0.1-cp312-cp312-macosx_10_9_x86_64.whl", hash = "sha256:855fb52b0dc35af121542a76b9a84f8d1cd886ea97c84703eaa6d88e37a2ad28"}, {file = "PyYAML-6.0.1-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:40df9b996c2b73138957fe23a16a4f0ba614f4c0efce1e9406a184b6d07fa3a9"}, + {file = "PyYAML-6.0.1-cp312-cp312-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:a08c6f0fe150303c1c6b71ebcd7213c2858041a7e01975da3a99aed1e7a378ef"}, {file = "PyYAML-6.0.1-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:6c22bec3fbe2524cde73d7ada88f6566758a8f7227bfbf93a408a9d86bcc12a0"}, {file = "PyYAML-6.0.1-cp312-cp312-musllinux_1_1_x86_64.whl", hash = "sha256:8d4e9c88387b0f5c7d5f281e55304de64cf7f9c0021a3525bd3b1c542da3b0e4"}, {file = "PyYAML-6.0.1-cp312-cp312-win32.whl", hash = "sha256:d483d2cdf104e7c9fa60c544d92981f12ad66a457afae824d146093b8c294c54"}, diff --git a/libs/langgraph/tests/__snapshots__/test_pregel.ambr b/libs/langgraph/tests/__snapshots__/test_pregel.ambr index 3937546f5..4699bb217 100644 --- a/libs/langgraph/tests/__snapshots__/test_pregel.ambr +++ b/libs/langgraph/tests/__snapshots__/test_pregel.ambr @@ -225,6 +225,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "left" @@ -237,6 +238,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "right" @@ -306,6 +308,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "left" @@ -318,6 +321,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "right" @@ -387,6 +391,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "get_weather" @@ -800,6 +805,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "tools" @@ -949,6 +955,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "tools" @@ -1081,6 +1088,7 @@ 'id': list([ 'langgraph', 'utils', + 'runnable', 'RunnableCallable', ]), 'name': 'tools', @@ -1151,6 +1159,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "tools" @@ -1300,6 +1309,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "tools" @@ -1432,6 +1442,7 @@ 'id': list([ 'langgraph', 'utils', + 'runnable', 'RunnableCallable', ]), 'name': 'tools', @@ -1502,6 +1513,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "tools" @@ -1651,6 +1663,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "tools" @@ -1783,6 +1796,7 @@ 'id': list([ 'langgraph', 'utils', + 'runnable', 'RunnableCallable', ]), 'name': 'tools', @@ -1853,6 +1867,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "tools" @@ -2002,6 +2017,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "tools" @@ -2134,6 +2150,7 @@ 'id': list([ 'langgraph', 'utils', + 'runnable', 'RunnableCallable', ]), 'name': 'tools', @@ -2204,6 +2221,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "tools" @@ -2353,6 +2371,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "tools" @@ -2485,6 +2504,7 @@ 'id': list([ 'langgraph', 'utils', + 'runnable', 'RunnableCallable', ]), 'name': 'tools', @@ -2642,6 +2662,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "tools" @@ -2723,6 +2744,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "tools" @@ -2804,6 +2826,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "tools" @@ -2885,6 +2908,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "tools" @@ -2966,6 +2990,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "tools" @@ -3028,6 +3053,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "A" @@ -3040,6 +3066,7 @@ "id": [ "langgraph", "utils", + "runnable", "RunnableCallable" ], "name": "B" @@ -4661,6 +4688,7 @@ 'id': list([ 'langgraph', 'utils', + 'runnable', 'RunnableCallable', ]), 'name': 'tool_one', @@ -4678,6 +4706,7 @@ 'id': list([ 'langgraph', 'utils', + 'runnable', 'RunnableCallable', ]), 'name': 'tool_two:tool_two_slow', @@ -4690,6 +4719,7 @@ 'id': list([ 'langgraph', 'utils', + 'runnable', 'RunnableCallable', ]), 'name': 'tool_two:tool_two_fast', @@ -4707,6 +4737,7 @@ 'id': list([ 'langgraph', 'utils', + 'runnable', 'RunnableCallable', ]), 'name': 'tool_three', @@ -5264,6 +5295,7 @@ 'id': list([ 'langgraph', 'utils', + 'runnable', 'RunnableCallable', ]), 'name': 'ask_question', @@ -5276,6 +5308,7 @@ 'id': list([ 'langgraph', 'utils', + 'runnable', 'RunnableCallable', ]), 'name': 'answer_question', @@ -5323,6 +5356,7 @@ 'id': list([ 'langgraph', 'utils', + 'runnable', 'RunnableCallable', ]), 'name': 'generate_analysts', @@ -5348,6 +5382,7 @@ 'id': list([ 'langgraph', 'utils', + 'runnable', 'RunnableCallable', ]), 'name': 'generate_sections', @@ -5413,6 +5448,7 @@ 'id': list([ 'langgraph', 'utils', + 'runnable', 'RunnableCallable', ]), 'name': 'generate_analysts', @@ -5430,6 +5466,7 @@ 'id': list([ 'langgraph', 'utils', + 'runnable', 'RunnableCallable', ]), 'name': 'conduct_interview:ask_question', @@ -5442,6 +5479,7 @@ 'id': list([ 'langgraph', 'utils', + 'runnable', 'RunnableCallable', ]), 'name': 'conduct_interview:answer_question', @@ -5459,6 +5497,7 @@ 'id': list([ 'langgraph', 'utils', + 'runnable', 'RunnableCallable', ]), 'name': 'generate_sections', diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 10018450c..19db1a992 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -1848,9 +1848,7 @@ def test_invoke_two_processes_one_in_two_out(mocker: MockerFixture) -> None: add_one = mocker.Mock(side_effect=lambda x: x + 1) one = ( - Channel.subscribe_to("input") - | add_one - | Channel.write_to(output=RunnablePassthrough(), between=RunnablePassthrough()) + Channel.subscribe_to("input") | add_one | Channel.write_to("output", "between") ) two = Channel.subscribe_to("between") | add_one | Channel.write_to("output") @@ -4824,7 +4822,7 @@ def test_message_graph( content="result for query", name="search_api", tool_call_id="tool_call123", - id="00000000-0000-4000-8000-000000000011", + id="00000000-0000-4000-8000-000000000010", ), AIMessage( content="", @@ -4841,7 +4839,7 @@ def test_message_graph( content="result for another", name="search_api", tool_call_id="tool_call456", - id="00000000-0000-4000-8000-000000000020", + id="00000000-0000-4000-8000-000000000018", ), AIMessage(content="answer", id="ai3"), ] @@ -4866,7 +4864,7 @@ def test_message_graph( content="result for query", name="search_api", tool_call_id="tool_call123", - id="00000000-0000-4000-8000-000000000036", + id="00000000-0000-4000-8000-000000000033", ) ] }, @@ -4889,7 +4887,7 @@ def test_message_graph( content="result for another", name="search_api", tool_call_id="tool_call456", - id="00000000-0000-4000-8000-000000000045", + id="00000000-0000-4000-8000-000000000041", ) ] }, @@ -5558,7 +5556,7 @@ def test_root_graph( content="result for query", name="search_api", tool_call_id="tool_call123", - id="00000000-0000-4000-8000-000000000011", + id="00000000-0000-4000-8000-000000000010", ), AIMessage( content="", @@ -5575,7 +5573,7 @@ def test_root_graph( content="result for another", name="search_api", tool_call_id="tool_call456", - id="00000000-0000-4000-8000-000000000020", + id="00000000-0000-4000-8000-000000000018", ), AIMessage(content="answer", id="ai3"), ] @@ -5600,7 +5598,7 @@ def test_root_graph( content="result for query", name="search_api", tool_call_id="tool_call123", - id="00000000-0000-4000-8000-000000000036", + id="00000000-0000-4000-8000-000000000033", ) ] }, @@ -5623,7 +5621,7 @@ def test_root_graph( content="result for another", name="search_api", tool_call_id="tool_call456", - id="00000000-0000-4000-8000-000000000045", + id="00000000-0000-4000-8000-000000000041", ) ] }, @@ -6223,7 +6221,7 @@ def test_root_graph( "__root__": [ HumanMessage( content="what is weather in sf", - id="00000000-0000-4000-8000-000000000077", + id="00000000-0000-4000-8000-000000000070", ), AIMessage( content="", @@ -6239,12 +6237,12 @@ def test_root_graph( ToolMessage( content="result for a different query", name="search_api", - id="00000000-0000-4000-8000-000000000091", + id="00000000-0000-4000-8000-000000000082", tool_call_id="tool_call123", ), AIMessage(content="answer", id="ai2"), AIMessage( - content="an extra message", id="00000000-0000-4000-8000-000000000101" + content="an extra message", id="00000000-0000-4000-8000-000000000091" ), HumanMessage(content="what is weather in la"), ], diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index 9d9c975d4..73de0b953 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -2087,9 +2087,7 @@ async def test_invoke_two_processes_one_in_two_out(mocker: MockerFixture) -> Non add_one = mocker.Mock(side_effect=lambda x: x + 1) one = ( - Channel.subscribe_to("input") - | add_one - | Channel.write_to(output=RunnablePassthrough(), between=RunnablePassthrough()) + Channel.subscribe_to("input") | add_one | Channel.write_to("output", "between") ) two = Channel.subscribe_to("between") | add_one | Channel.write_to("output") From cf345ef716758906c3cc5d868c4ecc345f6b5565 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 4 Sep 2024 12:07:02 -0700 Subject: [PATCH 3/9] Store loop in async bg exec init --- libs/langgraph/langgraph/pregel/executor.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/executor.py b/libs/langgraph/langgraph/pregel/executor.py index 3a4175cac..981ebbf7d 100644 --- a/libs/langgraph/langgraph/pregel/executor.py +++ b/libs/langgraph/langgraph/pregel/executor.py @@ -99,6 +99,7 @@ class AsyncBackgroundExecutor(AsyncContextManager): self.context_not_supported = sys.version_info < (3, 11) self.tasks: dict[asyncio.Task, bool] = {} self.sentinel = object() + self.loop = asyncio.get_running_loop() def submit( self, @@ -110,9 +111,9 @@ class AsyncBackgroundExecutor(AsyncContextManager): ) -> asyncio.Task[T]: coro = fn(*args, **kwargs) if self.context_not_supported: - task = asyncio.create_task(coro, name=__name__) + task = self.loop.create_task(coro, name=__name__) else: - task = asyncio.create_task(coro, name=__name__, context=copy_context()) + task = self.loop.create_task(coro, name=__name__, context=copy_context()) self.tasks[task] = __cancel_on_exit__ task.add_done_callback(self.done) return task From 4efab93e2fa65b6502b97c22bcb0bd4c1bb4fafc Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 4 Sep 2024 12:07:19 -0700 Subject: [PATCH 4/9] Use pseudo rng for uuid6 generation --- libs/checkpoint/langgraph/checkpoint/base/id.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/libs/checkpoint/langgraph/checkpoint/base/id.py b/libs/checkpoint/langgraph/checkpoint/base/id.py index 459ac7995..807dc8ad0 100644 --- a/libs/checkpoint/langgraph/checkpoint/base/id.py +++ b/libs/checkpoint/langgraph/checkpoint/base/id.py @@ -3,7 +3,7 @@ https://github.com/oittaa/uuid6-python/blob/main/src/uuid6/__init__.py#L95 Bundled in to avoid install issues with uuid6 package """ -import secrets +import random import time import uuid from typing import Optional, Tuple @@ -96,9 +96,9 @@ def uuid6(node: Optional[int] = None, clock_seq: Optional[int] = None) -> UUID: timestamp = _last_v6_timestamp + 1 _last_v6_timestamp = timestamp if clock_seq is None: - clock_seq = secrets.randbits(14) # instead of stable storage + clock_seq = random.getrandbits(14) # instead of stable storage if node is None: - node = secrets.randbits(48) + node = random.getrandbits(48) time_high_and_time_mid = (timestamp >> 12) & 0xFFFFFFFFFFFF time_low_and_version = timestamp & 0x0FFF uuid_int = time_high_and_time_mid << 80 From b68f9211c2ee7de2c21f75a501aba47db4f71995 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 4 Sep 2024 12:13:38 -0700 Subject: [PATCH 5/9] Run branch.reader in bg thread --- libs/langgraph/langgraph/graph/graph.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/libs/langgraph/langgraph/graph/graph.py b/libs/langgraph/langgraph/graph/graph.py index 41175c170..21868bf7c 100644 --- a/libs/langgraph/langgraph/graph/graph.py +++ b/libs/langgraph/langgraph/graph/graph.py @@ -1,3 +1,4 @@ +import asyncio import logging from collections import defaultdict from typing import ( @@ -97,7 +98,7 @@ class Branch(NamedTuple): writer: Callable[[list[str], RunnableConfig], Optional[Runnable]], ) -> Runnable: if reader: - value = reader(config) + value = await asyncio.to_thread(reader, config) # passthrough additional keys from node to branch # only doable when using dict states if isinstance(value, dict) and isinstance(input, dict): From 67a164ca2de8e19b1944b63310057b44a153ae02 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 4 Sep 2024 12:18:42 -0700 Subject: [PATCH 6/9] Run loop tick in bg thread - prepare_next_tasks, apply_writes do some cpu-heavy work --- libs/langgraph/langgraph/pregel/__init__.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 089b729fa..2a132381b 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -1419,7 +1419,8 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): # channel updates from step N are only visible in step N+1 # channels are guaranteed to be immutable for the duration of the step, # with channel updates applied only at the transition between steps - while loop.tick( + while await asyncio.to_thread( + loop.tick, input_keys=self.input_channels, interrupt_before=interrupt_before, interrupt_after=interrupt_after, From 26b4d9139ea26b405c5880fa6454dae8a8da44a0 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 4 Sep 2024 12:19:14 -0700 Subject: [PATCH 7/9] Avoid using json encoding when generating task ids --- libs/langgraph/langgraph/pregel/algo.py | 21 +++++++++++++++------ 1 file changed, 15 insertions(+), 6 deletions(-) diff --git a/libs/langgraph/langgraph/pregel/algo.py b/libs/langgraph/langgraph/pregel/algo.py index 3e465f115..c3dfd1b80 100644 --- a/libs/langgraph/langgraph/pregel/algo.py +++ b/libs/langgraph/langgraph/pregel/algo.py @@ -1,4 +1,3 @@ -import json from collections import defaultdict, deque from functools import partial from typing import ( @@ -270,16 +269,19 @@ def prepare_next_tasks( checkpointer: Optional[BaseCheckpointSaver] = None, manager: Union[None, ParentRunManager, AsyncParentRunManager] = None, ) -> Union[list[PregelTask], list[PregelExecutableTask]]: + checkpoint_id = UUID(checkpoint["id"]) configurable = config.get("configurable", {}) parent_ns = configurable.get("checkpoint_ns", "") tasks: Union[list[PregelTask], list[PregelExecutableTask]] = [] # Consume pending packets for packet in checkpoint["pending_sends"]: if not isinstance(packet, Send): - logger.warn(f"Ignoring invalid packet type {type(packet)} in pending sends") + logger.warning( + f"Ignoring invalid packet type {type(packet)} in pending sends" + ) continue if packet.node not in processes: - logger.warn(f"Ignoring unknown node name {packet.node} in pending sends") + logger.warning(f"Ignoring unknown node name {packet.node} in pending sends") continue # create task id triggers = [TASKS] @@ -293,7 +295,12 @@ def prepare_next_tasks( f"{parent_ns}{NS_SEP}{packet.node}" if parent_ns else packet.node ) task_id = str( - uuid5(UUID(checkpoint["id"]), json.dumps((checkpoint_ns, metadata))) + uuid5( + checkpoint_id, + "".join( + (checkpoint_ns, str(step), packet.node, *triggers, str(len(tasks))) + ), + ) ) if for_execution: proc = processes[packet.node] @@ -397,8 +404,10 @@ def prepare_next_tasks( checkpoint_ns = f"{parent_ns}{NS_SEP}{name}" if parent_ns else name task_id = str( uuid5( - UUID(checkpoint["id"]), - json.dumps((checkpoint_ns, metadata)), + checkpoint_id, + "".join( + (checkpoint_ns, str(step), name, *triggers, str(len(tasks))) + ), ) ) From b9361d69b4eb64694d9078ddc8ead26a08e4afc4 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 4 Sep 2024 12:30:27 -0700 Subject: [PATCH 8/9] Try to fix test --- libs/langgraph/tests/test_pregel_async.py | 1 + 1 file changed, 1 insertion(+) diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index 73de0b953..7afd292ae 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -514,6 +514,7 @@ async def test_cancel_graph_astream_events_v2(checkpointer_name: Optional[str]) if chunk["event"] == "on_chain_stream" and not chunk["parent_ids"]: got_event = True assert chunk["data"]["chunk"] == {"alittlewhile": {"value": 2}} + await asyncio.sleep(0) break # did break From 3a1b90c8821288e21486e92f200b75e0eea81c63 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Wed, 4 Sep 2024 12:34:42 -0700 Subject: [PATCH 9/9] Fix --- libs/langgraph/tests/test_pregel_async.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index 7afd292ae..f0a81ceb0 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -514,7 +514,7 @@ async def test_cancel_graph_astream_events_v2(checkpointer_name: Optional[str]) if chunk["event"] == "on_chain_stream" and not chunk["parent_ids"]: got_event = True assert chunk["data"]["chunk"] == {"alittlewhile": {"value": 2}} - await asyncio.sleep(0) + await asyncio.sleep(0.1) break # did break