Combine channel and managed values manager

This commit is contained in:
Nuno Campos
2024-08-21 13:30:18 -07:00
parent 14ec51601c
commit e76f4cc434
10 changed files with 213 additions and 227 deletions
@@ -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()
}
+13 -12
View File
@@ -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:
+1 -49
View File
@@ -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()
+37 -71
View File
@@ -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
+8 -15
View File
@@ -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):
+46 -33
View File
@@ -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"
+104
View File
@@ -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},
)
+1 -2
View File
@@ -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)
+1 -1
View File
@@ -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]
+2 -5
View File
@@ -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