Update FewShotExamples class to support setting filter and limit params.

This commit is contained in:
Andrew Nguonly
2024-05-14 12:41:55 -07:00
parent 3f34b03ba1
commit 28e5d8f699
5 changed files with 169 additions and 65 deletions
+55 -30
View File
@@ -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 {}
+38 -9
View File
@@ -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
+20 -21
View File
@@ -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:
+3 -3
View File
@@ -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
View File
@@ -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(),