Merge pull request #1197 from langchain-ai/nc/2aug/managed-rm-graph-arg

Remove graph arg from ManagedValue
This commit is contained in:
Nuno Campos
2024-08-02 09:28:07 -07:00
committed by GitHub
5 changed files with 47 additions and 26 deletions
+11 -16
View File
@@ -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()
+4 -4
View File
@@ -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,
+2
View File
@@ -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],
+2 -6
View File
@@ -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
+28
View File
@@ -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