diff --git a/langgraph/managed/few_shot.py b/langgraph/managed/few_shot.py index d7cf399f4..a96710b55 100644 --- a/langgraph/managed/few_shot.py +++ b/langgraph/managed/few_shot.py @@ -4,10 +4,14 @@ from typing import ( Any, AsyncGenerator, AsyncIterator, + Callable, + Dict, Generator, Generic, Iterator, + Optional, Sequence, + Union, ) from langchain_core.runnables import RunnableConfig @@ -22,6 +26,11 @@ from langgraph.pregel.types import PregelTaskDescription if TYPE_CHECKING: from langgraph.pregel import Pregel +# Metadata filter can be a dict (static) or a function (dynamic) that takes a +# RunnableConfig and returns a dict. Functions are used for filtering on +# metadata values that are only available at runtime. +MetadataFilter = Union[Dict[str, Any], Callable[[RunnableConfig], Dict[str, Any]]] + class FewShotExamples(ManagedValue[Sequence[V]], Generic[V]): examples: list[V] @@ -31,7 +40,7 @@ class FewShotExamples(ManagedValue[Sequence[V]], Generic[V]): config: RunnableConfig, graph: Pregel, k: int = 5, - metadata_filter: dict[str, Any] = None, + metadata_filter: Optional[MetadataFilter] = None, ) -> None: super().__init__(config, graph) self.k = k @@ -39,7 +48,7 @@ class FewShotExamples(ManagedValue[Sequence[V]], Generic[V]): @classmethod def configure( - cls, k: int = 5, metadata_filter: dict[str, Any] = None + cls, k: int = 5, metadata_filter: Optional[MetadataFilter] = None ) -> ConfiguredManagedValue: return ConfiguredManagedValue( cls, @@ -49,16 +58,23 @@ class FewShotExamples(ManagedValue[Sequence[V]], Generic[V]): }, ) + @property + def metadata_filter_dict(self) -> Dict[str, Any]: + if isinstance(self.metadata_filter, Callable): + return self.metadata_filter(self.config) + else: + return self.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 + {"score": score, **self.metadata_filter_dict}, 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) -> AsyncIterator[V]: async for example in self.graph.checkpointer.asearch( - {"score": score, **self.metadata_filter}, limit=self.k + {"score": score, **self.metadata_filter_dict}, limit=self.k ): async with AsyncChannelsManager( self.graph.channels, example.checkpoint diff --git a/tests/test_pregel.py b/tests/test_pregel.py index 0154490f7..dc8e179c5 100644 --- a/tests/test_pregel.py +++ b/tests/test_pregel.py @@ -8,6 +8,7 @@ from contextlib import contextmanager from typing import ( Annotated, Any, + Dict, Generator, Literal, Optional, @@ -2697,18 +2698,25 @@ def test_state_graph_w_config(snapshot: SnapshotAssertion) -> None: assert app.config_schema().schema_json() == snapshot -def test_state_graph_few_shot(snapshot: SnapshotAssertion) -> None: +def test_state_graph_few_shot() -> None: from langchain.chat_models.fake import FakeMessagesListChatModel from langchain_community.tools import tool from langchain_core.messages import AIMessage, AnyMessage, HumanMessage, ToolMessage from langchain_core.prompts import ChatPromptTemplate + def filter_by_source(config: RunnableConfig) -> Dict[str, Any]: + """This function is a trivial example that demonstrates that passing + a Callable to metadata_filter works as expected. + """ + return {"source": "loop"} + class BaseState(TypedDict): messages: Annotated[list[AnyMessage], add_messages] class AgentState(BaseState): examples: Annotated[ - Sequence[BaseState], FewShotExamples[BaseState].configure(k=1) + Sequence[BaseState], + FewShotExamples[BaseState].configure(k=1, metadata_filter=filter_by_source), ] # Assemble the tools @@ -2800,7 +2808,12 @@ Some examples of past conversations: ] assert app.invoke( {"messages": "what is weather in sf"}, - {"configurable": {"thread_id": "1", "expected_examples": []}}, + { + "configurable": { + "thread_id": "1", + "expected_examples": [], + }, + }, ) == {"messages": first_messages} # get first checkpoint diff --git a/tests/test_pregel_async.py b/tests/test_pregel_async.py index 94b209c8e..ed33e12e7 100644 --- a/tests/test_pregel_async.py +++ b/tests/test_pregel_async.py @@ -8,6 +8,7 @@ from typing import ( Any, AsyncGenerator, AsyncIterator, + Dict, Generator, Optional, Sequence, @@ -2435,11 +2436,20 @@ async def test_state_graph_few_shot() -> None: from langchain_core.messages import AIMessage, AnyMessage, HumanMessage, ToolMessage from langchain_core.prompts import ChatPromptTemplate + def filter_by_source(config: RunnableConfig) -> Dict[str, Any]: + """This function is a trivial example that demonstrates that passing + a Callable to metadata_filter works as expected. + """ + return {"source": "loop"} + class BaseState(TypedDict): messages: Annotated[list[AnyMessage], add_messages] class AgentState(BaseState): - examples: Annotated[Sequence[BaseState], FewShotExamples[BaseState]] + examples: Annotated[ + Sequence[BaseState], + FewShotExamples[BaseState].configure(k=1, metadata_filter=filter_by_source), + ] # Assemble the tools @tool()