mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-07 02:07:52 +02:00
Add support for passing Callable metadata filter to FewShotExamples managed value.
This commit is contained in:
@@ -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
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user