diff --git a/libs/langgraph/langgraph/channels/manager.py b/libs/langgraph/langgraph/channels/manager.py deleted file mode 100644 index 9a7c8d87e..000000000 --- a/libs/langgraph/langgraph/channels/manager.py +++ /dev/null @@ -1,39 +0,0 @@ -from contextlib import AsyncExitStack, ExitStack, asynccontextmanager, contextmanager -from typing import AsyncGenerator, Generator, Mapping - -from langchain_core.runnables import RunnableConfig - -from langgraph.channels.base import BaseChannel -from langgraph.checkpoint.base import Checkpoint - - -@contextmanager -def ChannelsManager( - channels: Mapping[str, BaseChannel], - checkpoint: Checkpoint, - config: RunnableConfig, -) -> Generator[Mapping[str, BaseChannel], None, None]: - """Manage channels for the lifetime of a Pregel invocation (multiple steps).""" - with ExitStack() as stack: - yield { - k: stack.enter_context( - v.from_checkpoint(checkpoint["channel_values"].get(k), config) - ) - for k, v in channels.items() - } - - -@asynccontextmanager -async def AsyncChannelsManager( - channels: Mapping[str, BaseChannel], - checkpoint: Checkpoint, - config: RunnableConfig, -) -> AsyncGenerator[Mapping[str, BaseChannel], None]: - """Manage channels for the lifetime of a Pregel invocation (multiple steps).""" - async with AsyncExitStack() as stack: - yield { - k: await stack.enter_async_context( - v.afrom_checkpoint(checkpoint["channel_values"].get(k), config) - ) - for k, v in channels.items() - } diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 8c85b7760..d846bc8d3 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -44,7 +44,7 @@ from langgraph.managed.base import ( ChannelKeyPlaceholder, ChannelTypePlaceholder, ConfiguredManagedValue, - ManagedValue, + ManagedValueSpec, is_managed_value, is_writable_managed_value, ) @@ -129,8 +129,8 @@ class StateGraph(Graph): nodes: dict[str, StateNodeSpec] channels: dict[str, BaseChannel] - managed: dict[str, Type[ManagedValue]] - schemas: dict[Type[Any], dict[str, Union[BaseChannel, Type[ManagedValue]]]] + managed: dict[str, ManagedValueSpec] + schemas: dict[Type[Any], dict[str, Union[BaseChannel, ManagedValueSpec]]] def __init__( self, @@ -442,7 +442,11 @@ class StateGraph(Graph): builder=self, config_type=self.config_schema, nodes={}, - channels={**self.channels, START: EphemeralValue(self.input)}, + channels={ + **self.channels, + **self.managed, + START: EphemeralValue(self.input), + }, input_channels=START, stream_mode="updates", output_channels=output_channels, @@ -497,7 +501,7 @@ class CompiledStateGraph(CompiledGraph): **{ k: (self.channels[k].UpdateType, None) for k in self.builder.schemas[self.builder.input] - if k in self.channels + if isinstance(self.channels[k], BaseChannel) and not isinstance(self.channels[k], Context) }, ) @@ -572,10 +576,7 @@ class CompiledStateGraph(CompiledGraph): ) else: input_schema = node.input if node else self.builder.schema - input_values = { - k: v if is_managed_value(v) else k - for k, v in self.builder.schemas[input_schema].items() - } + input_values = {k: k for k in self.builder.schemas[input_schema]} is_single_input = len(input_values) == 1 and "__root__" in input_values self.channels[key] = EphemeralValue(Any, guard=False) @@ -694,7 +695,7 @@ def _coerce_state(schema: Type[Any], input: dict[str, Any]) -> dict[str, Any]: def _get_channels( schema: Type[dict], -) -> tuple[dict[str, BaseChannel], dict[str, Type[ManagedValue]]]: +) -> tuple[dict[str, BaseChannel], dict[str, ManagedValueSpec]]: if not hasattr(schema, "__annotations__"): return {"__root__": _get_channel("__root__", schema, allow_managed=False)}, {} @@ -711,7 +712,7 @@ def _get_channels( def _get_channel( name: str, annotation: Any, *, allow_managed: bool = True -) -> Union[BaseChannel, Type[ManagedValue]]: +) -> Union[BaseChannel, ManagedValueSpec]: if manager := _is_field_managed_value(name, annotation): if allow_managed: return manager @@ -751,7 +752,7 @@ def _is_field_binop(typ: Type[Any]) -> Optional[BinaryOperatorAggregate]: return None -def _is_field_managed_value(name: str, typ: Type[Any]) -> Optional[Type[ManagedValue]]: +def _is_field_managed_value(name: str, typ: Type[Any]) -> Optional[ManagedValueSpec]: if hasattr(typ, "__metadata__"): meta = typ.__metadata__ if len(meta) >= 1: diff --git a/libs/langgraph/langgraph/managed/base.py b/libs/langgraph/langgraph/managed/base.py index 8dbf5c30c..bebca1be2 100644 --- a/libs/langgraph/langgraph/managed/base.py +++ b/libs/langgraph/langgraph/managed/base.py @@ -1,12 +1,9 @@ -import asyncio from abc import ABC, abstractmethod -from contextlib import AsyncExitStack, ExitStack, asynccontextmanager, contextmanager +from contextlib import asynccontextmanager, contextmanager from inspect import isclass from typing import ( Any, - AsyncGenerator, AsyncIterator, - Generator, Generic, Iterator, NamedTuple, @@ -104,50 +101,5 @@ def is_writable_managed_value(value: Any) -> TypeGuard[Type[WritableManagedValue ) -@contextmanager -def ManagedValuesManager( - values: dict[str, ManagedValueSpec], - config: RunnableConfig, -) -> Generator[ManagedValueMapping, None, None]: - if values: - with ExitStack() as stack: - yield { - key: stack.enter_context( - value.cls.enter(config, **value.kwargs) - if isinstance(value, ConfiguredManagedValue) - else value.enter(config) - ) - for key, value in values.items() - } - else: - yield {} - - -@asynccontextmanager -async def AsyncManagedValuesManager( - values: dict[str, ManagedValueSpec], - config: RunnableConfig, -) -> AsyncGenerator[ManagedValueMapping, None]: - if values: - async with AsyncExitStack() as stack: - # create enter tasks with reference to spec - tasks = { - asyncio.create_task( - stack.enter_async_context( - value.cls.aenter(config, **value.kwargs) - if isinstance(value, ConfiguredManagedValue) - else value.aenter(config) - ) - ): key - for key, value in values.items() - } - # wait for all enter tasks - done, _ = await asyncio.wait(tasks, return_when=asyncio.ALL_COMPLETED) - # build mapping from spec to result - yield {tasks[task]: task.result() for task in done} - else: - yield {} - - ChannelKeyPlaceholder = object() ChannelTypePlaceholder = object() diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 429226214..be04d187d 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -52,11 +52,6 @@ from langgraph.channels.base import ( BaseChannel, ) from langgraph.channels.context import Context -from langgraph.channels.last_value import LastValue -from langgraph.channels.manager import ( - AsyncChannelsManager, - ChannelsManager, -) from langgraph.checkpoint.base import ( BaseCheckpointSaver, copy_checkpoint, @@ -73,12 +68,7 @@ from langgraph.constants import ( Interrupt, ) from langgraph.errors import GraphInterrupt, GraphRecursionError, InvalidUpdateError -from langgraph.managed.base import ( - AsyncManagedValuesManager, - ManagedValuesManager, - ManagedValueSpec, - is_managed_value, -) +from langgraph.managed.base import ManagedValueSpec from langgraph.pregel.algo import ( apply_writes, local_read, @@ -97,6 +87,7 @@ from langgraph.pregel.io import ( read_channels, ) from langgraph.pregel.loop import AsyncPregelLoop, SyncPregelLoop +from langgraph.pregel.manager import AsyncChannelsManager, ChannelsManager from langgraph.pregel.read import PregelNode from langgraph.pregel.retry import RetryPolicy, arun_with_retry, run_with_retry from langgraph.pregel.types import ( @@ -197,7 +188,9 @@ class Pregel( ): nodes: Mapping[str, PregelNode] - channels: Mapping[str, BaseChannel] = Field(default_factory=dict) + channels: Mapping[str, Union[BaseChannel, ManagedValueSpec]] = Field( + default_factory=dict + ) auto_validate: bool = True @@ -350,16 +343,6 @@ class Pregel( k for k in self.channels if not isinstance(self.channels[k], Context) ] - @property - def managed_values_dict(self) -> dict[str, ManagedValueSpec]: - return { - k: v - for node in self.nodes.values() - if isinstance(node.channels, dict) - for k, v in node.channels.items() - if is_managed_value(v) - } - def get_state(self, config: RunnableConfig) -> StateSnapshot: """Get the current state of the graph.""" if not self.checkpointer: @@ -368,16 +351,10 @@ class Pregel( saved = self.checkpointer.get_tuple(config) checkpoint = saved.checkpoint if saved else empty_checkpoint() config = saved.config if saved else config - with ChannelsManager( - { - k: LastValue(None) if isinstance(c, Context) else c - for k, c in self.channels.items() - }, - checkpoint, - config, - ) as channels, ManagedValuesManager( - self.managed_values_dict, ensure_config(config) - ) as managed: + with ChannelsManager(self.channels, checkpoint, config, skip_context=True) as ( + channels, + managed, + ): next_tasks = prepare_next_tasks( checkpoint, self.nodes, @@ -408,15 +385,8 @@ class Pregel( config = saved.config if saved else config async with AsyncChannelsManager( - { - k: LastValue(None) if isinstance(c, Context) else c - for k, c in self.channels.items() - }, - checkpoint, - config, - ) as channels, AsyncManagedValuesManager( - self.managed_values_dict, ensure_config(config) - ) as managed: + self.channels, checkpoint, config, skip_context=True + ) as (channels, managed): next_tasks = prepare_next_tasks( checkpoint, self.nodes, @@ -460,15 +430,8 @@ class Pregel( pending_writes, ) in self.checkpointer.list(config, before=before, limit=limit, filter=filter): with ChannelsManager( - { - k: LastValue(None) if isinstance(c, Context) else c - for k, c in self.channels.items() - }, - checkpoint, - config, - ) as channels, ManagedValuesManager( - self.managed_values_dict, ensure_config(config) - ) as managed: + self.channels, checkpoint, config, skip_context=True + ) as (channels, managed): next_tasks = prepare_next_tasks( checkpoint, self.nodes, @@ -512,15 +475,8 @@ class Pregel( pending_writes, ) in self.checkpointer.alist(config, before=before, limit=limit, filter=filter): async with AsyncChannelsManager( - { - k: LastValue(None) if isinstance(c, Context) else c - for k, c in self.channels.items() - }, - checkpoint, - config, - ) as channels, AsyncManagedValuesManager( - self.managed_values_dict, ensure_config(config) - ) as managed: + self.channels, checkpoint, config, skip_context=True + ) as (channels, managed): next_tasks = prepare_next_tasks( checkpoint, self.nodes, @@ -613,11 +569,10 @@ class Pregel( if as_node not in self.nodes: raise InvalidUpdateError(f"Node {as_node} does not exist") # update channels - with ChannelsManager( - self.channels, checkpoint, config - ) as channels, ManagedValuesManager( - self.managed_values_dict, ensure_config(config) - ) as managed: + with ChannelsManager(self.channels, checkpoint, config) as ( + channels, + managed, + ): # create task to run all writers of the chosen node writers = self.nodes[as_node].get_writers() if not writers: @@ -757,11 +712,10 @@ class Pregel( if as_node not in self.nodes: raise InvalidUpdateError(f"Node {as_node} does not exist") # update channels, acting as the chosen node - async with AsyncChannelsManager( - self.channels, checkpoint, config - ) as channels, AsyncManagedValuesManager( - self.managed_values_dict, ensure_config(config) - ) as managed: + async with AsyncChannelsManager(self.channels, checkpoint, config) as ( + channels, + managed, + ): # create task to run all writers of the chosen node writers = self.nodes[as_node].get_writers() if not writers: @@ -998,7 +952,13 @@ class Pregel( ) with SyncPregelLoop( - input, config=config, checkpointer=checkpointer, graph=self + input, + config=config, + store=self.store, + checkpointer=checkpointer, + graph=self, + nodes=self.nodes, + specs=self.channels, ) as loop: # Similarly to Bulk Synchronous Parallel / Pregel model # computation proceeds in steps, while there are channel updates @@ -1248,7 +1208,13 @@ class Pregel( debug=debug, ) async with AsyncPregelLoop( - input, config=config, checkpointer=checkpointer, graph=self + input, + config=config, + store=self.store, + checkpointer=checkpointer, + graph=self, + nodes=self.nodes, + specs=self.channels, ) as loop: aioloop = asyncio.get_event_loop() # Similarly to Bulk Synchronous Parallel / Pregel model diff --git a/libs/langgraph/langgraph/pregel/algo.py b/libs/langgraph/langgraph/pregel/algo.py index 37f168bf7..28a8dcd5d 100644 --- a/libs/langgraph/langgraph/pregel/algo.py +++ b/libs/langgraph/langgraph/pregel/algo.py @@ -25,7 +25,6 @@ from langchain_core.runnables.config import ( from langgraph.channels.base import BaseChannel from langgraph.channels.context import Context -from langgraph.channels.manager import ChannelsManager from langgraph.checkpoint.base import ( BaseCheckpointSaver, Checkpoint, @@ -46,9 +45,10 @@ from langgraph.constants import ( Send, ) from langgraph.errors import EmptyChannelError, InvalidUpdateError -from langgraph.managed.base import ManagedValueMapping, is_managed_value +from langgraph.managed.base import ManagedValueMapping 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.pregel.types import All, PregelExecutableTask, PregelTask @@ -105,11 +105,10 @@ def local_read( if fresh: new_checkpoint = create_checkpoint(copy_checkpoint(checkpoint), channels, -1) context_channels = {k: v for k, v in channels.items() if isinstance(v, Context)} - with ChannelsManager( - {k: v for k, v in channels.items() if k not in context_channels}, - new_checkpoint, - config, - ) as channels: + with ChannelsManager(channels, new_checkpoint, config, skip_context=True) as ( + channels, + _, + ): all_channels = {**channels, **context_channels} apply_writes(new_checkpoint, all_channels, [task], None) return read_channels(all_channels, select) @@ -470,16 +469,10 @@ def _proc_input( chan, catch=chan not in proc.triggers, ) + if chan in channels + else managed[k](step) for k, chan in proc.channels.items() - if isinstance(chan, str) } - - managed_values = {} - for key, chan in proc.channels.items(): - if is_managed_value(chan): - managed_values[key] = managed[key](step) - - val.update(managed_values) except EmptyChannelError: return elif isinstance(proc.channels, list): diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index 672d42049..bc19b39cf 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -22,14 +22,10 @@ from typing import ( ) from langchain_core.callbacks import AsyncParentRunManager, ParentRunManager -from langchain_core.runnables import RunnableConfig, patch_config +from langchain_core.runnables import RunnableConfig from typing_extensions import Self from langgraph.channels.base import BaseChannel -from langgraph.channels.manager import ( - AsyncChannelsManager, - ChannelsManager, -) from langgraph.checkpoint.base import ( BaseCheckpointSaver, Checkpoint, @@ -43,7 +39,6 @@ from langgraph.checkpoint.base import ( from langgraph.constants import ( CONFIG_KEY_READ, CONFIG_KEY_RESUMING, - CONFIG_KEY_STORE, ERROR, INPUT, INTERRUPT, @@ -51,9 +46,8 @@ from langgraph.constants import ( ) from langgraph.errors import EmptyInputError, GraphInterrupt from langgraph.managed.base import ( - AsyncManagedValuesManager, ManagedValueMapping, - ManagedValuesManager, + ManagedValueSpec, WritableManagedValue, ) from langgraph.pregel.algo import ( @@ -70,6 +64,8 @@ from langgraph.pregel.executor import ( Submit, ) from langgraph.pregel.io import map_input, map_output_updates, map_output_values, single +from langgraph.pregel.manager import AsyncChannelsManager, ChannelsManager +from langgraph.pregel.read import PregelNode from langgraph.pregel.types import PregelExecutableTask from langgraph.pregel.utils import get_new_channel_versions from langgraph.store.base import BaseStore @@ -88,7 +84,12 @@ EMPTY_SEQ = () class PregelLoop: input: Optional[Any] config: RunnableConfig + store: Optional[BaseStore] checkpointer: Optional[BaseCheckpointSaver] + nodes: Mapping[str, PregelNode] + specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]] + is_nested: bool + checkpointer_get_next_version: Callable[[Optional[V]], V] checkpointer_put_writes: Optional[ Callable[[RunnableConfig, Sequence[tuple[str, Any]], str], Any] @@ -123,7 +124,6 @@ class PregelLoop: ] tasks: Sequence[PregelExecutableTask] stream: deque[Tuple[str, Any]] - is_nested: bool # public @@ -132,16 +132,20 @@ class PregelLoop: input: Optional[Any], *, config: RunnableConfig, + store: Optional[BaseStore], checkpointer: Optional[BaseCheckpointSaver], + nodes: Mapping[str, PregelNode], + specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]], graph: "Pregel", ) -> None: self.stream = deque() self.input = input self.config = config + self.store = store self.checkpointer = checkpointer self.graph = graph - # TODO if managed values no longer needs graph we can replace with - # managed_specs, channel_specs + self.nodes = nodes + self.specs = specs self.is_nested = CONFIG_KEY_READ in self.config.get("configurable", {}) def mark_tasks_scheduled(self, tasks: Sequence[PregelExecutableTask]) -> None: @@ -235,7 +239,7 @@ class PregelLoop: # prepare next tasks self.tasks = prepare_next_tasks( self.checkpoint, - self.graph.nodes, + self.nodes, self.channels, self.managed, self.config, @@ -323,7 +327,7 @@ class PregelLoop: # discard any unfinished tasks from previous checkpoint discard_tasks = prepare_next_tasks( self.checkpoint, - self.graph.nodes, + self.nodes, self.channels, self.managed, self.config, @@ -422,11 +426,21 @@ class SyncPregelLoop(PregelLoop, ContextManager): input: Optional[Any], *, config: RunnableConfig, + store: Optional[BaseStore], checkpointer: Optional[BaseCheckpointSaver], + nodes: Mapping[str, PregelNode], + specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]], graph: "Pregel", ) -> None: - super().__init__(input, config=config, checkpointer=checkpointer, graph=graph) - self.store = graph.store + super().__init__( + input, + config=config, + checkpointer=checkpointer, + graph=graph, + store=store, + nodes=nodes, + specs=specs, + ) self.stack = ExitStack() if checkpointer: self.checkpointer_get_next_version = checkpointer.get_next_version @@ -472,14 +486,8 @@ class SyncPregelLoop(PregelLoop, ContextManager): self.checkpoint_pending_writes = saved.pending_writes or [] self.submit = self.stack.enter_context(BackgroundExecutor(self.config)) - self.channels = self.stack.enter_context( - ChannelsManager(self.graph.channels, self.checkpoint, self.config) - ) - self.managed = self.stack.enter_context( - ManagedValuesManager( - self.graph.managed_values_dict, - patch_config(self.config, configurable={CONFIG_KEY_STORE: self.store}), - ) + self.channels, self.managed = self.stack.enter_context( + ChannelsManager(self.specs, self.checkpoint, self.config, self.store) ) self.stack.push(self._suppress_interrupt) self.status = "pending" @@ -506,11 +514,22 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager): input: Optional[Any], *, config: RunnableConfig, + store: Optional[BaseStore], checkpointer: Optional[BaseCheckpointSaver], + nodes: Mapping[str, PregelNode], + specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]], graph: "Pregel", ) -> None: - super().__init__(input, config=config, checkpointer=checkpointer, graph=graph) - self.store = AsyncBatchedStore(graph.store) if graph.store else None + super().__init__( + input, + config=config, + checkpointer=checkpointer, + graph=graph, + store=store, + nodes=nodes, + specs=specs, + ) + self.store = AsyncBatchedStore(self.store) if self.store else None self.stack = AsyncExitStack() if checkpointer: self.checkpointer_get_next_version = checkpointer.get_next_version @@ -560,14 +579,8 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager): self.checkpoint_pending_writes = saved.pending_writes or [] self.submit = await self.stack.enter_async_context(AsyncBackgroundExecutor()) - self.channels = await self.stack.enter_async_context( - AsyncChannelsManager(self.graph.channels, self.checkpoint, self.config) - ) - self.managed = await self.stack.enter_async_context( - AsyncManagedValuesManager( - self.graph.managed_values_dict, - patch_config(self.config, configurable={CONFIG_KEY_STORE: self.store}), - ) + self.channels, self.managed = await self.stack.enter_async_context( + AsyncChannelsManager(self.specs, self.checkpoint, self.config, self.store) ) self.stack.push(self._suppress_interrupt) self.status = "pending" diff --git a/libs/langgraph/langgraph/pregel/manager.py b/libs/langgraph/langgraph/pregel/manager.py new file mode 100644 index 000000000..437019113 --- /dev/null +++ b/libs/langgraph/langgraph/pregel/manager.py @@ -0,0 +1,104 @@ +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 langgraph.channels.base import BaseChannel +from langgraph.channels.context import Context +from langgraph.channels.last_value import LastValue +from langgraph.checkpoint.base import Checkpoint +from langgraph.constants import CONFIG_KEY_STORE +from langgraph.managed.base import ( + ConfiguredManagedValue, + ManagedValueMapping, + ManagedValueSpec, +) +from langgraph.store.base import BaseStore + + +@contextmanager +def ChannelsManager( + specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]], + checkpoint: Checkpoint, + config: RunnableConfig, + store: Optional[BaseStore] = None, + *, + 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}) + channel_specs: Mapping[str, BaseChannel] = {} + managed_specs: Mapping[str, ManagedValueSpec] = {} + for k, v in specs.items(): + if skip_context and isinstance(v, Context): + channel_specs[k] = LastValue(None) + elif isinstance(v, BaseChannel): + channel_specs[k] = v + else: + managed_specs[k] = v + with ExitStack() as stack: + yield ( + { + k: stack.enter_context( + v.from_checkpoint(checkpoint["channel_values"].get(k), config) + ) + for k, v in channel_specs.items() + }, + { + key: stack.enter_context( + value.cls.enter(config_for_managed, **value.kwargs) + if isinstance(value, ConfiguredManagedValue) + else value.enter(config_for_managed) + ) + for key, value in managed_specs.items() + }, + ) + + +@asynccontextmanager +async def AsyncChannelsManager( + specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]], + checkpoint: Checkpoint, + config: RunnableConfig, + store: Optional[BaseStore] = None, + *, + 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}) + channel_specs: Mapping[str, BaseChannel] = {} + managed_specs: Mapping[str, ManagedValueSpec] = {} + for k, v in specs.items(): + if skip_context and isinstance(v, Context): + channel_specs[k] = LastValue(None) + elif isinstance(v, BaseChannel): + channel_specs[k] = v + else: + managed_specs[k] = v + async with AsyncExitStack() as stack: + # managed: create enter tasks with reference to spec, await them + if tasks := { + asyncio.create_task( + stack.enter_async_context( + value.cls.aenter(config_for_managed, **value.kwargs) + if isinstance(value, ConfiguredManagedValue) + else value.aenter(config_for_managed) + ) + ): key + for key, value in managed_specs.items() + }: + done, _ = await asyncio.wait(tasks, return_when=asyncio.ALL_COMPLETED) + else: + done = set() + yield ( + # channels: enter each channel with checkpoint + { + k: await stack.enter_async_context( + v.afrom_checkpoint(checkpoint["channel_values"].get(k), config) + ) + for k, v in channel_specs.items() + }, + # managed: build mapping from spec to result + {tasks[task]: task.result() for task in done}, + ) diff --git a/libs/langgraph/langgraph/pregel/read.py b/libs/langgraph/langgraph/pregel/read.py index 6e492e2ac..b5c971e4a 100644 --- a/libs/langgraph/langgraph/pregel/read.py +++ b/libs/langgraph/langgraph/pregel/read.py @@ -15,7 +15,6 @@ from langchain_core.runnables.config import merge_configs from langchain_core.runnables.utils import ConfigurableFieldSpec from langgraph.constants import CONFIG_KEY_READ -from langgraph.managed.base import ManagedValueSpec from langgraph.pregel.retry import RetryPolicy from langgraph.pregel.write import ChannelWrite from langgraph.utils import RunnableCallable @@ -101,7 +100,7 @@ DEFAULT_BOUND: RunnablePassthrough = RunnablePassthrough() class PregelNode(RunnableBindingBase): - channels: Union[list[str], Mapping[str, Union[str, ManagedValueSpec]]] + channels: Union[list[str], Mapping[str, str]] triggers: list[str] = Field(default_factory=list) diff --git a/libs/langgraph/pyproject.toml b/libs/langgraph/pyproject.toml index 5ba045061..93dc203f2 100644 --- a/libs/langgraph/pyproject.toml +++ b/libs/langgraph/pyproject.toml @@ -65,7 +65,7 @@ omit = ["tests/*"] [tool.pytest-watcher] now = true delay = 0.1 -runner_args = ["--ff", "-v", "-n", "auto", "--dist", "worksteal", "--snapshot-update", "--tb", "short"] +runner_args = ["--ff", "-v", "-x", "-n", "auto", "--dist", "worksteal", "--snapshot-update", "--tb", "short"] patterns = ["*.py"] [build-system] diff --git a/libs/langgraph/tests/test_algo.py b/libs/langgraph/tests/test_algo.py index cdf179c24..2102261ae 100644 --- a/libs/langgraph/tests/test_algo.py +++ b/libs/langgraph/tests/test_algo.py @@ -1,7 +1,6 @@ -from langgraph.channels.manager import ChannelsManager from langgraph.checkpoint.base import empty_checkpoint -from langgraph.managed.base import ManagedValuesManager from langgraph.pregel.algo import prepare_next_tasks +from langgraph.pregel.manager import ChannelsManager def test_prepare_next_tasks() -> None: @@ -9,9 +8,7 @@ def test_prepare_next_tasks() -> None: processes = {} checkpoint = empty_checkpoint() - with ManagedValuesManager({}, config) as managed, ChannelsManager( - {}, checkpoint, config - ) as channels: + with ChannelsManager({}, checkpoint, config) as (channels, managed): assert ( prepare_next_tasks( checkpoint, processes, channels, managed, config, 0, for_execution=False