Add support for passing Callable metadata filter to FewShotExamples managed value.

This commit is contained in:
Andrew Nguonly
2024-05-17 18:41:57 -07:00
parent fb8868b710
commit 18addb52aa
3 changed files with 47 additions and 8 deletions
+20 -4
View File
@@ -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
+16 -3
View File
@@ -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
+11 -1
View File
@@ -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()