mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-21 23:22:27 +02:00
Combine channel and managed values manager
This commit is contained in:
@@ -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()
|
||||
}
|
||||
@@ -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,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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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},
|
||||
)
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user