diff --git a/langgraph/managed/base.py b/langgraph/managed/base.py index 15b47070b..b06f6a7c7 100644 --- a/langgraph/managed/base.py +++ b/langgraph/managed/base.py @@ -8,9 +8,10 @@ from typing import ( AsyncGenerator, Generator, Generic, - Sequence, + NamedTuple, Type, TypeVar, + Union, ) from langchain_core.runnables import RunnableConfig @@ -32,10 +33,10 @@ class ManagedValue(ABC, Generic[V]): @classmethod @contextmanager def enter( - cls, config: RunnableConfig, graph: "Pregel" + cls, config: RunnableConfig, graph: "Pregel", **kwargs: Any ) -> Generator[Self, None, None]: try: - value = cls(config, graph) + value = cls(config, graph, **kwargs) yield value finally: # because managed value and Pregel have reference to each other @@ -48,10 +49,10 @@ class ManagedValue(ABC, Generic[V]): @classmethod @asynccontextmanager async def aenter( - cls, config: RunnableConfig, graph: "Pregel" + cls, config: RunnableConfig, graph: "Pregel", **kwargs: Any ) -> AsyncGenerator[Self, None]: try: - value = cls(config, graph) + value = cls(config, graph, **kwargs) yield value finally: # because managed value and Pregel have reference to each other @@ -66,40 +67,64 @@ class ManagedValue(ABC, Generic[V]): ... -def is_managed_value(value: Any) -> TypeGuard[Type[ManagedValue]]: - return isclass(value) and issubclass(value, ManagedValue) +class ConfiguredManagedValue(NamedTuple): + cls: Type[ManagedValue] + kwargs: dict[str, Any] + + +ManagedValueSpec = Union[Type[ManagedValue], ConfiguredManagedValue] + +ManagedValueMapping = dict[str, ManagedValue] + + +def is_managed_value(value: Any) -> TypeGuard[ManagedValueSpec]: + return (isclass(value) and issubclass(value, ManagedValue)) or isinstance( + value, ConfiguredManagedValue + ) @contextmanager def ManagedValuesManager( - values: Sequence[Type[ManagedValue]], + values: dict[str, ManagedValueSpec], config: RunnableConfig, graph: "Pregel", -) -> Generator[Sequence[ManagedValue], None, None]: - with ExitStack() as stack: - unique: list[Type[ManagedValue]] = [] - for value in values: - if value not in unique: - unique.append(value) - - yield [stack.enter_context(value.enter(config, graph)) for value in unique] +) -> Generator[ManagedValueMapping, None, None]: + if values: + with ExitStack() as stack: + yield { + key: stack.enter_context( + value.cls.enter(config, graph, **value.kwargs) + if isinstance(value, ConfiguredManagedValue) + else value.enter(config, graph) + ) + for key, value in values.items() + } + else: + yield {} @asynccontextmanager async def AsyncManagedValuesManager( - values: Sequence[Type[ManagedValue]], + values: dict[str, ManagedValueSpec], config: RunnableConfig, graph: "Pregel", -) -> AsyncGenerator[Sequence[ManagedValue], None]: - async with AsyncExitStack() as stack: - unique: list[Type[ManagedValue]] = [] - for value in values: - if value not in unique: - unique.append(value) - - yield await asyncio.gather( - *( - stack.enter_async_context(value.aenter(config, graph)) - for value in unique - ) - ) +) -> 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, graph, **value.kwargs) + if isinstance(value, ConfiguredManagedValue) + else value.aenter(config, graph) + ) + ): 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 {} diff --git a/langgraph/managed/few_shot.py b/langgraph/managed/few_shot.py index 1e2147006..d7cf399f4 100644 --- a/langgraph/managed/few_shot.py +++ b/langgraph/managed/few_shot.py @@ -1,6 +1,7 @@ from contextlib import asynccontextmanager, contextmanager from typing import ( TYPE_CHECKING, + Any, AsyncGenerator, AsyncIterator, Generator, @@ -13,7 +14,8 @@ from langchain_core.runnables import RunnableConfig from typing_extensions import Self from langgraph.channels.base import AsyncChannelsManager, ChannelsManager -from langgraph.managed.base import ManagedValue, V +from langgraph.managed.base import ConfiguredManagedValue, ManagedValue, V +from langgraph.pregel import Pregel from langgraph.pregel.io import read_channels from langgraph.pregel.types import PregelTaskDescription @@ -24,13 +26,40 @@ if TYPE_CHECKING: class FewShotExamples(ManagedValue[Sequence[V]], Generic[V]): examples: list[V] - def iter(self, score: int = 1, k: int = 5) -> Iterator[V]: - for example in self.graph.checkpointer.search({"score": score}, limit=k): + def __init__( + self, + config: RunnableConfig, + graph: Pregel, + k: int = 5, + metadata_filter: dict[str, Any] = None, + ) -> None: + super().__init__(config, graph) + self.k = k + self.metadata_filter = metadata_filter or {} + + @classmethod + def configure( + cls, k: int = 5, metadata_filter: dict[str, Any] = None + ) -> ConfiguredManagedValue: + return ConfiguredManagedValue( + cls, + { + "k": k, + "metadata_filter": metadata_filter, + }, + ) + + def iter(self, score: int = 1) -> Iterator[V]: + for example in self.graph.checkpointer.search( + {"score": score, **self.metadata_filter}, limit=self.k + ): with ChannelsManager(self.graph.channels, example.checkpoint) as channels: yield read_channels(channels, self.graph.output_channels) - async def aiter(self, score: int = 1, k: int = 5) -> AsyncIterator[V]: - async for example in self.graph.checkpointer.asearch({"score": score}, limit=k): + async def aiter(self, score: int = 1) -> AsyncIterator[V]: + async for example in self.graph.checkpointer.asearch( + {"score": score, **self.metadata_filter}, limit=self.k + ): async with AsyncChannelsManager( self.graph.channels, example.checkpoint ) as channels: @@ -39,18 +68,18 @@ class FewShotExamples(ManagedValue[Sequence[V]], Generic[V]): @classmethod @contextmanager def enter( - cls, config: RunnableConfig, graph: "Pregel" + cls, config: RunnableConfig, graph: "Pregel", **kwargs: Any ) -> Generator[Self, None, None]: - with super().enter(config, graph) as value: + with super().enter(config, graph, **kwargs) as value: value.examples = list(value.iter()) yield value @classmethod @asynccontextmanager async def aenter( - cls, config: RunnableConfig, graph: "Pregel" + cls, config: RunnableConfig, graph: "Pregel", **kwargs: Any ) -> AsyncGenerator[Self, None]: - async with super().aenter(config, graph) as value: + async with super().aenter(config, graph, **kwargs) as value: value.examples = [e async for e in value.aiter()] yield value diff --git a/langgraph/pregel/__init__.py b/langgraph/pregel/__init__.py index 597f617d5..d7039b734 100644 --- a/langgraph/pregel/__init__.py +++ b/langgraph/pregel/__init__.py @@ -4,7 +4,6 @@ import asyncio import concurrent.futures from collections import defaultdict, deque from functools import partial -from inspect import isclass from typing import ( Any, AsyncIterator, @@ -70,8 +69,9 @@ from langgraph.constants import ( from langgraph.errors import GraphRecursionError, InvalidUpdateError from langgraph.managed.base import ( AsyncManagedValuesManager, - ManagedValue, + ManagedValueMapping, ManagedValuesManager, + ManagedValueSpec, is_managed_value, ) from langgraph.pregel.debug import ( @@ -327,14 +327,14 @@ class Pregel( return self.stream_channels or [k for k in self.channels] @property - def managed_values_list(self) -> Sequence[Type[ManagedValue]]: - return [ - v + def managed_values_dict(self) -> dict[str, ManagedValueSpec]: + return { + k: v for node in self.nodes.values() if isinstance(node.channels, dict) - for v in node.channels.values() + 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.""" @@ -347,7 +347,7 @@ class Pregel( with ChannelsManager( self.channels, checkpoint ) as channels, ManagedValuesManager( - self.managed_values_list, ensure_config(config), self + self.managed_values_dict, ensure_config(config), self ) as managed: _, next_tasks = _prepare_next_tasks( checkpoint, @@ -378,7 +378,7 @@ class Pregel( async with AsyncChannelsManager( self.channels, checkpoint ) as channels, AsyncManagedValuesManager( - self.managed_values_list, ensure_config(config), self + self.managed_values_dict, ensure_config(config), self ) as managed: _, next_tasks = _prepare_next_tasks( checkpoint, @@ -414,7 +414,7 @@ class Pregel( with ChannelsManager( self.channels, checkpoint ) as channels, ManagedValuesManager( - self.managed_values_list, ensure_config(config), self + self.managed_values_dict, ensure_config(config), self ) as managed: _, next_tasks = _prepare_next_tasks( checkpoint, @@ -453,7 +453,7 @@ class Pregel( async with AsyncChannelsManager( self.channels, checkpoint ) as channels, AsyncManagedValuesManager( - self.managed_values_list, ensure_config(config), self + self.managed_values_dict, ensure_config(config), self ) as managed: _, next_tasks = _prepare_next_tasks( checkpoint, @@ -725,7 +725,7 @@ class Pregel( ) as channels, get_executor_for_config( config ) as executor, ManagedValuesManager( - self.managed_values_list, config, self + self.managed_values_dict, config, self ) as managed: # map inputs to channel updates if input_writes := deque(map_input(input_keys, input)): @@ -1021,7 +1021,7 @@ class Pregel( async with AsyncChannelsManager( self.channels, checkpoint ) as channels, AsyncManagedValuesManager( - self.managed_values_list, config, self + self.managed_values_dict, config, self ) as managed: # map inputs to channel updates if input_writes := deque(map_input(input_keys, input)): @@ -1478,7 +1478,7 @@ def _prepare_next_tasks( checkpoint: Checkpoint, processes: Mapping[str, PregelNode], channels: Mapping[str, BaseChannel], - managed: Sequence[ManagedValue], + managed: ManagedValueMapping, config: RunnableConfig, step: int, for_execution: Literal[False], @@ -1491,7 +1491,7 @@ def _prepare_next_tasks( checkpoint: Checkpoint, processes: Mapping[str, PregelNode], channels: Mapping[str, BaseChannel], - managed: Sequence[ManagedValue], + managed: ManagedValueMapping, config: RunnableConfig, step: int, for_execution: Literal[True], @@ -1503,7 +1503,7 @@ def _prepare_next_tasks( checkpoint: Checkpoint, processes: Mapping[str, PregelNode], channels: Mapping[str, BaseChannel], - managed: Sequence[ManagedValue], + managed: ManagedValueMapping, config: RunnableConfig, step: int, *, @@ -1536,11 +1536,10 @@ def _prepare_next_tasks( managed_values = {} for key, chan in proc.channels.items(): - for mv in managed: - if isclass(chan) and isinstance(mv, chan): - managed_values[key] = mv( - step, PregelTaskDescription(name, val) - ) + if is_managed_value(chan): + managed_values[key] = managed[key]( + step, PregelTaskDescription(name, val) + ) val.update(managed_values) except EmptyChannelError: diff --git a/langgraph/pregel/read.py b/langgraph/pregel/read.py index e8576e22f..eefed2b47 100644 --- a/langgraph/pregel/read.py +++ b/langgraph/pregel/read.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Any, Callable, Mapping, Optional, Sequence, Type, Union +from typing import Any, Callable, Mapping, Optional, Sequence, Union from langchain_core.pydantic_v1 import Field from langchain_core.runnables import ( @@ -15,7 +15,7 @@ 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 ManagedValue +from langgraph.managed.base import ManagedValueSpec from langgraph.pregel.write import ChannelWrite from langgraph.utils import RunnableCallable @@ -100,7 +100,7 @@ DEFAULT_BOUND: RunnablePassthrough = RunnablePassthrough() class PregelNode(RunnableBindingBase): - channels: Union[list[str], Mapping[str, Union[str, Type[ManagedValue]]]] + channels: Union[list[str], Mapping[str, Union[str, ManagedValueSpec]]] triggers: list[str] = Field(default_factory=list) diff --git a/tests/test_pregel.py b/tests/test_pregel.py index 7ccc1bd7d..01e0dbb0f 100644 --- a/tests/test_pregel.py +++ b/tests/test_pregel.py @@ -2596,7 +2596,9 @@ def test_state_graph_few_shot(snapshot: SnapshotAssertion) -> None: messages: Annotated[list[AnyMessage], add_messages] class AgentState(BaseState): - examples: Annotated[Sequence[BaseState], FewShotExamples[BaseState]] + examples: Annotated[ + Sequence[BaseState], FewShotExamples[BaseState].configure(k=1) + ] # Assemble the tools @tool() @@ -2709,6 +2711,27 @@ Some examples of past conversations: assert len(hiscored) == 1 assert hiscored[0].checkpoint["channel_values"]["messages"] == first_messages + second_messages = [ + HumanMessage(content="what is weather in la", id=AnyStr()), + AIMessage( + content="", + id=AnyStr(), + tool_calls=[ + { + "name": "search_api", + "args": {"query": "query"}, + "id": "tool_call123", + } + ], + ), + ToolMessage( + content="result for query", + name="search_api", + id=AnyStr(), + tool_call_id="tool_call123", + ), + AIMessage(content="answer", id=AnyStr()), + ] assert app.invoke( {"messages": "what is weather in la"}, { @@ -2718,9 +2741,37 @@ Some examples of past conversations: "expected_examples": [{"messages": first_messages}], } }, + ) == {"messages": second_messages} + + # get first checkpoint + chkpnt_tuple_2 = saver.get_tuple({"configurable": {"thread_id": "2"}}) + config = chkpnt_tuple_2.config + checkpoint = chkpnt_tuple_2.checkpoint + metadata = chkpnt_tuple_2.metadata + + # not needed in application code, only for testing + hiscored = list(saver.search({"score": 1})) + assert len(hiscored) == 1 + + # mark as "good" + metadata["score"] = 1 + saver.put(config, checkpoint, metadata) + + hiscored = list(saver.search({"score": 1})) + assert len(hiscored) == 2 + + assert app.invoke( + {"messages": "what is weather in ny"}, + { + "configurable": { + "thread_id": "3", + # below is only for testing purposes, not part of few shot api + "expected_examples": [{"messages": second_messages}], + } + }, ) == { "messages": [ - HumanMessage(content="what is weather in la", id=AnyStr()), + HumanMessage(content="what is weather in ny", id=AnyStr()), AIMessage( content="", id=AnyStr(),