mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-17 21:25:46 +02:00
Merge pull request #1197 from langchain-ai/nc/2aug/managed-rm-graph-arg
Remove graph arg from ManagedValue
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user