diff --git a/libs/langgraph/langgraph/managed/base.py b/libs/langgraph/langgraph/managed/base.py index b06f6a7c7..0e1476895 100644 --- a/libs/langgraph/langgraph/managed/base.py +++ b/libs/langgraph/langgraph/managed/base.py @@ -17,26 +17,23 @@ from typing import ( from langchain_core.runnables import RunnableConfig from typing_extensions import Self, TypeGuard -from langgraph.pregel.types import PregelTaskDescription - if TYPE_CHECKING: - from langgraph.pregel import Pregel + from langgraph.pregel.types import PregelTaskDescription V = TypeVar("V") class ManagedValue(ABC, Generic[V]): - def __init__(self, config: RunnableConfig, graph: "Pregel") -> None: + def __init__(self, config: RunnableConfig) -> None: self.config = config - self.graph = graph @classmethod @contextmanager def enter( - cls, config: RunnableConfig, graph: "Pregel", **kwargs: Any + cls, config: RunnableConfig, **kwargs: Any ) -> Generator[Self, None, None]: try: - value = cls(config, graph, **kwargs) + value = cls(config, **kwargs) yield value finally: # because managed value and Pregel have reference to each other @@ -49,10 +46,10 @@ class ManagedValue(ABC, Generic[V]): @classmethod @asynccontextmanager async def aenter( - cls, config: RunnableConfig, graph: "Pregel", **kwargs: Any + cls, config: RunnableConfig, **kwargs: Any ) -> AsyncGenerator[Self, None]: try: - value = cls(config, graph, **kwargs) + value = cls(config, **kwargs) yield value finally: # because managed value and Pregel have reference to each other @@ -63,7 +60,7 @@ class ManagedValue(ABC, Generic[V]): pass @abstractmethod - def __call__(self, step: int, task: PregelTaskDescription) -> V: + def __call__(self, step: int, task: "PregelTaskDescription") -> V: ... @@ -87,15 +84,14 @@ def is_managed_value(value: Any) -> TypeGuard[ManagedValueSpec]: def ManagedValuesManager( values: dict[str, ManagedValueSpec], config: RunnableConfig, - graph: "Pregel", ) -> Generator[ManagedValueMapping, None, None]: if values: with ExitStack() as stack: yield { key: stack.enter_context( - value.cls.enter(config, graph, **value.kwargs) + value.cls.enter(config, **value.kwargs) if isinstance(value, ConfiguredManagedValue) - else value.enter(config, graph) + else value.enter(config) ) for key, value in values.items() } @@ -107,7 +103,6 @@ def ManagedValuesManager( async def AsyncManagedValuesManager( values: dict[str, ManagedValueSpec], config: RunnableConfig, - graph: "Pregel", ) -> AsyncGenerator[ManagedValueMapping, None]: if values: async with AsyncExitStack() as stack: @@ -115,9 +110,9 @@ async def AsyncManagedValuesManager( tasks = { asyncio.create_task( stack.enter_async_context( - value.cls.aenter(config, graph, **value.kwargs) + value.cls.aenter(config, **value.kwargs) if isinstance(value, ConfiguredManagedValue) - else value.aenter(config, graph) + else value.aenter(config) ) ): key for key, value in values.items() diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 71df7964a..3a247c579 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -361,7 +361,7 @@ class Pregel( with ChannelsManager( self.channels, checkpoint, config ) as channels, ManagedValuesManager( - self.managed_values_dict, ensure_config(config), self + self.managed_values_dict, ensure_config(config) ) as managed: next_tasks = prepare_next_tasks( checkpoint, @@ -393,7 +393,7 @@ class Pregel( async with AsyncChannelsManager( self.channels, checkpoint, config ) as channels, AsyncManagedValuesManager( - self.managed_values_dict, ensure_config(config), self + self.managed_values_dict, ensure_config(config) ) as managed: next_tasks = prepare_next_tasks( checkpoint, @@ -435,7 +435,7 @@ class Pregel( with ChannelsManager( self.channels, checkpoint, config ) as channels, ManagedValuesManager( - self.managed_values_dict, ensure_config(config), self + self.managed_values_dict, ensure_config(config) ) as managed: next_tasks = prepare_next_tasks( checkpoint, @@ -481,7 +481,7 @@ class Pregel( async with AsyncChannelsManager( self.channels, checkpoint, config ) as channels, AsyncManagedValuesManager( - self.managed_values_dict, ensure_config(config), self + self.managed_values_dict, ensure_config(config) ) as managed: next_tasks = prepare_next_tasks( checkpoint, diff --git a/libs/langgraph/langgraph/pregel/algo.py b/libs/langgraph/langgraph/pregel/algo.py index 7a1a2abbb..e30cec4aa 100644 --- a/libs/langgraph/langgraph/pregel/algo.py +++ b/libs/langgraph/langgraph/pregel/algo.py @@ -220,6 +220,7 @@ def prepare_next_tasks( managed: ManagedValueMapping, config: RunnableConfig, step: int, + *, for_execution: Literal[False], is_resuming: bool = False, checkpointer: Literal[None] = None, @@ -236,6 +237,7 @@ def prepare_next_tasks( managed: ManagedValueMapping, config: RunnableConfig, step: int, + *, for_execution: Literal[True], is_resuming: bool, checkpointer: Optional[BaseCheckpointSaver], diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index 37519cb25..eae303d63 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -428,9 +428,7 @@ class SyncPregelLoop(PregelLoop, ContextManager): ChannelsManager(self.graph.channels, self.checkpoint, self.config) ) self.managed = self.stack.enter_context( - ManagedValuesManager( - self.graph.managed_values_dict, self.config, self.graph - ) + ManagedValuesManager(self.graph.managed_values_dict, self.config) ) self.status = "pending" self.step = self.checkpoint_metadata["step"] + 1 @@ -507,9 +505,7 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager): AsyncChannelsManager(self.graph.channels, self.checkpoint, self.config) ) self.managed = await self.stack.enter_async_context( - AsyncManagedValuesManager( - self.graph.managed_values_dict, self.config, self.graph - ) + AsyncManagedValuesManager(self.graph.managed_values_dict, self.config) ) self.status = "pending" self.step = self.checkpoint_metadata["step"] + 1 diff --git a/libs/langgraph/tests/test_algo.py b/libs/langgraph/tests/test_algo.py new file mode 100644 index 000000000..cdf179c24 --- /dev/null +++ b/libs/langgraph/tests/test_algo.py @@ -0,0 +1,28 @@ +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 + + +def test_prepare_next_tasks() -> None: + config = {} + processes = {} + checkpoint = empty_checkpoint() + + with ManagedValuesManager({}, config) as managed, ChannelsManager( + {}, checkpoint, config + ) as channels: + assert ( + prepare_next_tasks( + checkpoint, processes, channels, managed, config, 0, for_execution=False + ) + == [] + ) + assert ( + prepare_next_tasks( + checkpoint, processes, channels, managed, config, 0, for_execution=True + ) + == [] + ) + + # TODO: add more tests