mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-12 04:37:51 +02:00
Update FewShotExamples class to support setting filter and limit params.
This commit is contained in:
+55
-30
@@ -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 {}
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
+53
-2
@@ -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(),
|
||||
|
||||
Reference in New Issue
Block a user