mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-07 02:07:52 +02:00
langgraph: use 'prompt' param for model input preprocessing in create_react_agent (#3173)
This commit is contained in:
@@ -1,5 +1,7 @@
|
||||
import functools
|
||||
import inspect
|
||||
from typing import (
|
||||
Any,
|
||||
Callable,
|
||||
Literal,
|
||||
Optional,
|
||||
@@ -21,7 +23,6 @@ from langchain_core.tools import BaseTool
|
||||
from pydantic import BaseModel
|
||||
from typing_extensions import Annotated, TypedDict
|
||||
|
||||
from langgraph._api.deprecation import deprecated_parameter
|
||||
from langgraph.errors import ErrorCode, create_error_message
|
||||
from langgraph.graph import END, StateGraph
|
||||
from langgraph.graph.graph import CompiledGraph
|
||||
@@ -35,6 +36,7 @@ from langgraph.utils.runnable import RunnableCallable
|
||||
|
||||
StructuredResponse = Union[dict, BaseModel]
|
||||
StructuredResponseSchema = Union[dict, type[BaseModel]]
|
||||
F = TypeVar("F", bound=Callable[..., Any])
|
||||
|
||||
|
||||
# We create the AgentState that we will pass around
|
||||
@@ -56,7 +58,7 @@ class AgentState(TypedDict):
|
||||
StateSchema = TypeVar("StateSchema", bound=AgentState)
|
||||
StateSchemaType = Type[StateSchema]
|
||||
|
||||
STATE_MODIFIER_RUNNABLE_NAME = "StateModifier"
|
||||
PROMPT_RUNNABLE_NAME = "Prompt"
|
||||
|
||||
MessagesModifier = Union[
|
||||
SystemMessage,
|
||||
@@ -65,7 +67,7 @@ MessagesModifier = Union[
|
||||
Runnable[Sequence[BaseMessage], Sequence[BaseMessage]],
|
||||
]
|
||||
|
||||
StateModifier = Union[
|
||||
Prompt = Union[
|
||||
SystemMessage,
|
||||
str,
|
||||
Callable[[StateSchema], Sequence[BaseMessage]],
|
||||
@@ -73,81 +75,84 @@ StateModifier = Union[
|
||||
]
|
||||
|
||||
|
||||
def _get_state_modifier_runnable(
|
||||
state_modifier: Optional[StateModifier], store: Optional[BaseStore] = None
|
||||
) -> Runnable:
|
||||
state_modifier_runnable: Runnable
|
||||
if state_modifier is None:
|
||||
state_modifier_runnable = RunnableCallable(
|
||||
lambda state: state["messages"], name=STATE_MODIFIER_RUNNABLE_NAME
|
||||
def _get_prompt_runnable(prompt: Optional[Prompt]) -> Runnable:
|
||||
prompt_runnable: Runnable
|
||||
if prompt is None:
|
||||
prompt_runnable = RunnableCallable(
|
||||
lambda state: state["messages"], name=PROMPT_RUNNABLE_NAME
|
||||
)
|
||||
elif isinstance(state_modifier, str):
|
||||
_system_message: BaseMessage = SystemMessage(content=state_modifier)
|
||||
state_modifier_runnable = RunnableCallable(
|
||||
elif isinstance(prompt, str):
|
||||
_system_message: BaseMessage = SystemMessage(content=prompt)
|
||||
prompt_runnable = RunnableCallable(
|
||||
lambda state: [_system_message] + state["messages"],
|
||||
name=STATE_MODIFIER_RUNNABLE_NAME,
|
||||
name=PROMPT_RUNNABLE_NAME,
|
||||
)
|
||||
elif isinstance(state_modifier, SystemMessage):
|
||||
state_modifier_runnable = RunnableCallable(
|
||||
lambda state: [state_modifier] + state["messages"],
|
||||
name=STATE_MODIFIER_RUNNABLE_NAME,
|
||||
elif isinstance(prompt, SystemMessage):
|
||||
prompt_runnable = RunnableCallable(
|
||||
lambda state: [prompt] + state["messages"],
|
||||
name=PROMPT_RUNNABLE_NAME,
|
||||
)
|
||||
elif inspect.iscoroutinefunction(state_modifier):
|
||||
state_modifier_runnable = RunnableCallable(
|
||||
elif inspect.iscoroutinefunction(prompt):
|
||||
prompt_runnable = RunnableCallable(
|
||||
None,
|
||||
state_modifier,
|
||||
name=STATE_MODIFIER_RUNNABLE_NAME,
|
||||
prompt,
|
||||
name=PROMPT_RUNNABLE_NAME,
|
||||
)
|
||||
elif callable(state_modifier):
|
||||
state_modifier_runnable = RunnableCallable(
|
||||
state_modifier,
|
||||
name=STATE_MODIFIER_RUNNABLE_NAME,
|
||||
elif callable(prompt):
|
||||
prompt_runnable = RunnableCallable(
|
||||
prompt,
|
||||
name=PROMPT_RUNNABLE_NAME,
|
||||
)
|
||||
elif isinstance(state_modifier, Runnable):
|
||||
state_modifier_runnable = state_modifier
|
||||
elif isinstance(prompt, Runnable):
|
||||
prompt_runnable = prompt
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Got unexpected type for `state_modifier`: {type(state_modifier)}"
|
||||
)
|
||||
raise ValueError(f"Got unexpected type for `prompt`: {type(prompt)}")
|
||||
|
||||
return state_modifier_runnable
|
||||
return prompt_runnable
|
||||
|
||||
|
||||
def _convert_messages_modifier_to_state_modifier(
|
||||
def _convert_messages_modifier_to_prompt(
|
||||
messages_modifier: MessagesModifier,
|
||||
) -> StateModifier:
|
||||
state_modifier: StateModifier
|
||||
) -> Prompt:
|
||||
prompt: Prompt
|
||||
if isinstance(messages_modifier, (str, SystemMessage)):
|
||||
return messages_modifier
|
||||
elif callable(messages_modifier):
|
||||
|
||||
def state_modifier(state: AgentState) -> Sequence[BaseMessage]:
|
||||
def prompt(state: AgentState) -> Sequence[BaseMessage]:
|
||||
return messages_modifier(state["messages"])
|
||||
|
||||
return state_modifier
|
||||
return prompt
|
||||
elif isinstance(messages_modifier, Runnable):
|
||||
state_modifier = (lambda state: state["messages"]) | messages_modifier
|
||||
return state_modifier
|
||||
prompt = (lambda state: state["messages"]) | messages_modifier
|
||||
return prompt
|
||||
raise ValueError(
|
||||
f"Got unexpected type for `messages_modifier`: {type(messages_modifier)}"
|
||||
)
|
||||
|
||||
|
||||
def _get_model_preprocessing_runnable(
|
||||
state_modifier: Optional[StateModifier],
|
||||
messages_modifier: Optional[MessagesModifier],
|
||||
store: Optional[BaseStore],
|
||||
) -> Runnable:
|
||||
# Add the state or message modifier, if exists
|
||||
if state_modifier is not None and messages_modifier is not None:
|
||||
raise ValueError(
|
||||
"Expected value for either state_modifier or messages_modifier, got values for both"
|
||||
)
|
||||
def _convert_modifier_to_prompt(func: F) -> F:
|
||||
"""Decorator that converts state_modifier/messages_modifier kwargs to prompt kwarg."""
|
||||
|
||||
if state_modifier is None and messages_modifier is not None:
|
||||
state_modifier = _convert_messages_modifier_to_state_modifier(messages_modifier)
|
||||
@functools.wraps(func)
|
||||
def wrapper(*args: Any, **kwargs: Any) -> Any:
|
||||
prompt = kwargs.get("prompt")
|
||||
state_modifier = kwargs.pop("state_modifier", None)
|
||||
messages_modifier = kwargs.pop("messages_modifier", None)
|
||||
if sum(p is not None for p in (prompt, state_modifier, messages_modifier)) > 1:
|
||||
raise ValueError(
|
||||
"Expected only one of prompt, state_modifier, or messages_modifier, got multiple values"
|
||||
)
|
||||
|
||||
return _get_state_modifier_runnable(state_modifier, store)
|
||||
if state_modifier is not None:
|
||||
prompt = state_modifier
|
||||
elif messages_modifier is not None:
|
||||
prompt = _convert_messages_modifier_to_prompt(messages_modifier)
|
||||
|
||||
kwargs["prompt"] = prompt
|
||||
return func(*args, **kwargs)
|
||||
|
||||
return cast(F, wrapper)
|
||||
|
||||
|
||||
def _should_bind_tools(model: LanguageModelLike, tools: Sequence[BaseTool]) -> bool:
|
||||
@@ -228,14 +233,13 @@ def _validate_chat_history(
|
||||
raise ValueError(error_message)
|
||||
|
||||
|
||||
@deprecated_parameter("messages_modifier", "0.1.9", "state_modifier", removal="0.3.0")
|
||||
@_convert_modifier_to_prompt
|
||||
def create_react_agent(
|
||||
model: Union[str, LanguageModelLike],
|
||||
tools: Union[ToolExecutor, Sequence[BaseTool], ToolNode],
|
||||
*,
|
||||
state_schema: Optional[StateSchemaType] = None,
|
||||
messages_modifier: Optional[MessagesModifier] = None,
|
||||
state_modifier: Optional[StateModifier] = None,
|
||||
prompt: Optional[Prompt] = None,
|
||||
response_format: Optional[
|
||||
Union[StructuredResponseSchema, tuple[str, StructuredResponseSchema]]
|
||||
] = None,
|
||||
@@ -254,26 +258,15 @@ def create_react_agent(
|
||||
state_schema: An optional state schema that defines graph state.
|
||||
Must have `messages` and `is_last_step` keys.
|
||||
Defaults to `AgentState` that defines those two keys.
|
||||
messages_modifier: An optional
|
||||
messages modifier. This applies to messages BEFORE they are passed into the LLM.
|
||||
prompt: An optional prompt for the LLM. Can take a few different forms:
|
||||
|
||||
Can take a few different forms:
|
||||
|
||||
- SystemMessage: this is added to the beginning of the list of messages.
|
||||
- str: This is converted to a SystemMessage and added to the beginning of the list of messages.
|
||||
- Callable: This function should take in a list of messages and the output is then passed to the language model.
|
||||
- Runnable: This runnable should take in a list of messages and the output is then passed to the language model.
|
||||
!!! Warning
|
||||
`messages_modifier` parameter is deprecated as of version 0.1.9 and will be removed in 0.2.0
|
||||
state_modifier: An optional
|
||||
state modifier. This takes full graph state BEFORE the LLM is called and prepares the input to LLM.
|
||||
|
||||
Can take a few different forms:
|
||||
|
||||
- SystemMessage: this is added to the beginning of the list of messages in state["messages"].
|
||||
- str: This is converted to a SystemMessage and added to the beginning of the list of messages in state["messages"].
|
||||
- SystemMessage: this is added to the beginning of the list of messages in state["messages"].
|
||||
- Callable: This function should take in full graph state and the output is then passed to the language model.
|
||||
- Runnable: This runnable should take in full graph state and the output is then passed to the language model.
|
||||
|
||||
!!! Note
|
||||
Prior to `v0.2.68`, the prompt was set using `state_modifier` / `messages_modifier` parameters.
|
||||
response_format: An optional schema for the final agent output.
|
||||
|
||||
If provided, output will be formatted to match the given schema and returned in the 'structured_response' state key.
|
||||
@@ -389,7 +382,7 @@ def create_react_agent(
|
||||
|
||||
```pycon
|
||||
>>> system_prompt = "You are a helpful bot named Fred."
|
||||
>>> graph = create_react_agent(model, tools, state_modifier=system_prompt)
|
||||
>>> graph = create_react_agent(model, tools, prompt=system_prompt)
|
||||
>>> inputs = {"messages": [("user", "What's your name? And what's the weather in SF?")]}
|
||||
>>> for s in graph.stream(inputs, stream_mode="values"):
|
||||
... message = s["messages"][-1]
|
||||
@@ -421,11 +414,8 @@ def create_react_agent(
|
||||
... ("placeholder", "{messages}"),
|
||||
... ("user", "Remember, always be polite!"),
|
||||
... ])
|
||||
>>> def format_for_model(state: AgentState):
|
||||
... # You can do more complex modifications here
|
||||
... return prompt.invoke({"messages": state["messages"]})
|
||||
>>>
|
||||
>>> graph = create_react_agent(model, tools, state_modifier=format_for_model)
|
||||
>>> graph = create_react_agent(model, tools, prompt=prompt)
|
||||
>>> inputs = {"messages": [("user", "What's your name? And what's the weather in SF?")]}
|
||||
>>> for s in graph.stream(inputs, stream_mode="values"):
|
||||
... message = s["messages"][-1]
|
||||
@@ -453,7 +443,7 @@ def create_react_agent(
|
||||
... messages: Annotated[list[BaseMessage], add_messages]
|
||||
... is_last_step: IsLastStep
|
||||
>>>
|
||||
>>> graph = create_react_agent(model, tools, state_schema=CustomState, state_modifier=prompt)
|
||||
>>> graph = create_react_agent(model, tools, state_schema=CustomState, prompt=prompt)
|
||||
>>> inputs = {"messages": [("user", "What's today's date? And what's the weather in SF?")], "today": "July 16, 2004"}
|
||||
>>> for s in graph.stream(inputs, stream_mode="values"):
|
||||
... message = s["messages"][-1]
|
||||
@@ -539,7 +529,7 @@ def create_react_agent(
|
||||
>>> from langgraph.checkpoint.memory import MemorySaver
|
||||
>>> from langgraph.store.memory import InMemoryStore
|
||||
>>> store = InMemoryStore()
|
||||
>>> graph = create_react_agent(model, [save_memory], state_modifier=prepare_model_inputs, store=store, checkpointer=MemorySaver())
|
||||
>>> graph = create_react_agent(model, [save_memory], prompt=prepare_model_inputs, store=store, checkpointer=MemorySaver())
|
||||
>>> config = {"configurable": {"thread_id": "thread-1", "user_id": "1"}}
|
||||
|
||||
>>> inputs = {"messages": [("user", "Hey I'm Will, how's it going?")]}
|
||||
@@ -619,11 +609,7 @@ def create_react_agent(
|
||||
if _should_bind_tools(model, tool_classes) and tool_calling_enabled:
|
||||
model = cast(BaseChatModel, model).bind_tools(tool_classes)
|
||||
|
||||
# we're passing store here for validation
|
||||
preprocessor = _get_model_preprocessing_runnable(
|
||||
state_modifier, messages_modifier, store
|
||||
)
|
||||
model_runnable = preprocessor | model
|
||||
model_runnable = _get_prompt_runnable(prompt) | model
|
||||
|
||||
# If any of the tools are configured to return_directly after running,
|
||||
# our graph needs to check if these were called
|
||||
|
||||
@@ -147,7 +147,7 @@ class FakeToolCallingModel(BaseChatModel):
|
||||
|
||||
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC)
|
||||
def test_no_modifier(request: pytest.FixtureRequest, checkpointer_name: str) -> None:
|
||||
def test_no_prompt(request: pytest.FixtureRequest, checkpointer_name: str) -> None:
|
||||
checkpointer: BaseCheckpointSaver = request.getfixturevalue(
|
||||
"checkpointer_" + checkpointer_name
|
||||
)
|
||||
@@ -181,7 +181,7 @@ def test_no_modifier(request: pytest.FixtureRequest, checkpointer_name: str) ->
|
||||
|
||||
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
|
||||
async def test_no_modifier_async(checkpointer_name: str) -> None:
|
||||
async def test_no_prompt_async(checkpointer_name: str) -> None:
|
||||
async with awith_checkpointer(checkpointer_name) as checkpointer:
|
||||
model = FakeToolCallingModel()
|
||||
|
||||
@@ -217,16 +217,17 @@ def test_passing_two_modifiers():
|
||||
with pytest.raises(ValueError):
|
||||
create_react_agent(model, [], messages_modifier="Foo", state_modifier="Bar")
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
create_react_agent(model, [], messages_modifier="Foo", prompt="Bar")
|
||||
|
||||
def test_system_message_modifier():
|
||||
messages_modifier = SystemMessage(content="Foo")
|
||||
agent_1 = create_react_agent(
|
||||
FakeToolCallingModel(), [], messages_modifier=messages_modifier
|
||||
)
|
||||
agent_2 = create_react_agent(
|
||||
FakeToolCallingModel(), [], state_modifier=messages_modifier
|
||||
)
|
||||
for agent in [agent_1, agent_2]:
|
||||
|
||||
def test_system_message_prompt():
|
||||
prompt = SystemMessage(content="Foo")
|
||||
for agent in (
|
||||
create_react_agent(FakeToolCallingModel(), [], prompt=prompt),
|
||||
create_react_agent(FakeToolCallingModel(), [], messages_modifier=prompt),
|
||||
create_react_agent(FakeToolCallingModel(), [], state_modifier=prompt),
|
||||
):
|
||||
inputs = [HumanMessage("hi?")]
|
||||
response = agent.invoke({"messages": inputs})
|
||||
expected_response = {
|
||||
@@ -235,15 +236,13 @@ def test_system_message_modifier():
|
||||
assert response == expected_response
|
||||
|
||||
|
||||
def test_system_message_string_modifier():
|
||||
messages_modifier = "Foo"
|
||||
agent_1 = create_react_agent(
|
||||
FakeToolCallingModel(), [], messages_modifier=messages_modifier
|
||||
)
|
||||
agent_2 = create_react_agent(
|
||||
FakeToolCallingModel(), [], state_modifier=messages_modifier
|
||||
)
|
||||
for agent in [agent_1, agent_2]:
|
||||
def test_string_prompt():
|
||||
prompt = "Foo"
|
||||
for agent in (
|
||||
create_react_agent(FakeToolCallingModel(), [], prompt=prompt),
|
||||
create_react_agent(FakeToolCallingModel(), [], messages_modifier=prompt),
|
||||
create_react_agent(FakeToolCallingModel(), [], state_modifier=prompt),
|
||||
):
|
||||
inputs = [HumanMessage("hi?")]
|
||||
response = agent.invoke({"messages": inputs})
|
||||
expected_response = {
|
||||
@@ -252,63 +251,71 @@ def test_system_message_string_modifier():
|
||||
assert response == expected_response
|
||||
|
||||
|
||||
def test_callable_messages_modifier():
|
||||
model = FakeToolCallingModel()
|
||||
|
||||
def test_callable_prompt():
|
||||
def messages_modifier(messages):
|
||||
modified_message = f"Bar {messages[-1].content}"
|
||||
return [HumanMessage(content=modified_message)]
|
||||
|
||||
agent = create_react_agent(model, [], messages_modifier=messages_modifier)
|
||||
inputs = [HumanMessage("hi?")]
|
||||
response = agent.invoke({"messages": inputs})
|
||||
expected_response = {"messages": inputs + [AIMessage(content="Bar hi?", id="0")]}
|
||||
assert response == expected_response
|
||||
|
||||
|
||||
def test_callable_state_modifier():
|
||||
model = FakeToolCallingModel()
|
||||
|
||||
def state_modifier(state):
|
||||
def prompt(state):
|
||||
modified_message = f"Bar {state['messages'][-1].content}"
|
||||
return [HumanMessage(content=modified_message)]
|
||||
|
||||
agent = create_react_agent(model, [], state_modifier=state_modifier)
|
||||
inputs = [HumanMessage("hi?")]
|
||||
response = agent.invoke({"messages": inputs})
|
||||
expected_response = {"messages": inputs + [AIMessage(content="Bar hi?", id="0")]}
|
||||
assert response == expected_response
|
||||
for agent in (
|
||||
create_react_agent(FakeToolCallingModel(), [], prompt=prompt),
|
||||
create_react_agent(FakeToolCallingModel(), [], state_modifier=prompt),
|
||||
create_react_agent(
|
||||
FakeToolCallingModel(), [], messages_modifier=messages_modifier
|
||||
),
|
||||
):
|
||||
inputs = [HumanMessage("hi?")]
|
||||
response = agent.invoke({"messages": inputs})
|
||||
expected_response = {
|
||||
"messages": inputs + [AIMessage(content="Bar hi?", id="0")]
|
||||
}
|
||||
assert response == expected_response
|
||||
|
||||
|
||||
def test_runnable_messages_modifier():
|
||||
model = FakeToolCallingModel()
|
||||
async def test_callable_prompt_async():
|
||||
async def prompt(state):
|
||||
modified_message = f"Bar {state['messages'][-1].content}"
|
||||
return [HumanMessage(content=modified_message)]
|
||||
|
||||
for agent in (
|
||||
create_react_agent(FakeToolCallingModel(), [], prompt=prompt),
|
||||
create_react_agent(FakeToolCallingModel(), [], state_modifier=prompt),
|
||||
):
|
||||
inputs = [HumanMessage("hi?")]
|
||||
response = await agent.ainvoke({"messages": inputs})
|
||||
expected_response = {
|
||||
"messages": inputs + [AIMessage(content="Bar hi?", id="0")]
|
||||
}
|
||||
assert response == expected_response
|
||||
|
||||
|
||||
def test_runnable_prompt():
|
||||
messages_modifier = RunnableLambda(
|
||||
lambda messages: [HumanMessage(content=f"Baz {messages[-1].content}")]
|
||||
)
|
||||
|
||||
agent = create_react_agent(model, [], messages_modifier=messages_modifier)
|
||||
inputs = [HumanMessage("hi?")]
|
||||
response = agent.invoke({"messages": inputs})
|
||||
expected_response = {"messages": inputs + [AIMessage(content="Baz hi?", id="0")]}
|
||||
assert response == expected_response
|
||||
|
||||
|
||||
def test_runnable_state_modifier():
|
||||
model = FakeToolCallingModel()
|
||||
|
||||
state_modifier = RunnableLambda(
|
||||
prompt = RunnableLambda(
|
||||
lambda state: [HumanMessage(content=f"Baz {state['messages'][-1].content}")]
|
||||
)
|
||||
|
||||
agent = create_react_agent(model, [], state_modifier=state_modifier)
|
||||
inputs = [HumanMessage("hi?")]
|
||||
response = agent.invoke({"messages": inputs})
|
||||
expected_response = {"messages": inputs + [AIMessage(content="Baz hi?", id="0")]}
|
||||
assert response == expected_response
|
||||
for agent in (
|
||||
create_react_agent(FakeToolCallingModel(), [], prompt=prompt),
|
||||
create_react_agent(FakeToolCallingModel(), [], state_modifier=prompt),
|
||||
create_react_agent(
|
||||
FakeToolCallingModel(), [], messages_modifier=messages_modifier
|
||||
),
|
||||
):
|
||||
inputs = [HumanMessage("hi?")]
|
||||
response = agent.invoke({"messages": inputs})
|
||||
expected_response = {
|
||||
"messages": inputs + [AIMessage(content="Baz hi?", id="0")]
|
||||
}
|
||||
assert response == expected_response
|
||||
|
||||
|
||||
def test_state_modifier_with_store():
|
||||
def test_prompt_with_store():
|
||||
def add(a: int, b: int):
|
||||
"""Adds a and b"""
|
||||
return a + b
|
||||
@@ -317,20 +324,18 @@ def test_state_modifier_with_store():
|
||||
in_memory_store.put(("memories", "1"), "user_name", {"data": "User name is Alice"})
|
||||
in_memory_store.put(("memories", "2"), "user_name", {"data": "User name is Bob"})
|
||||
|
||||
def modify(state, config, *, store):
|
||||
def prompt(state, config, *, store):
|
||||
user_id = config["configurable"]["user_id"]
|
||||
system_str = store.get(("memories", user_id), "user_name").value["data"]
|
||||
return [SystemMessage(system_str)] + state["messages"]
|
||||
|
||||
def modify_no_store(state, config):
|
||||
def prompt_no_store(state, config):
|
||||
return SystemMessage("foo") + state["messages"]
|
||||
|
||||
model = FakeToolCallingModel()
|
||||
|
||||
# test state modifier that uses store works
|
||||
agent = create_react_agent(
|
||||
model, [add], state_modifier=modify, store=in_memory_store
|
||||
)
|
||||
agent = create_react_agent(model, [add], prompt=prompt, store=in_memory_store)
|
||||
response = agent.invoke(
|
||||
{"messages": [("user", "hi")]}, {"configurable": {"user_id": "1"}}
|
||||
)
|
||||
@@ -338,7 +343,7 @@ def test_state_modifier_with_store():
|
||||
|
||||
# test state modifier that doesn't use store works
|
||||
agent = create_react_agent(
|
||||
model, [add], state_modifier=modify_no_store, store=in_memory_store
|
||||
model, [add], prompt=prompt_no_store, store=in_memory_store
|
||||
)
|
||||
response = agent.invoke(
|
||||
{"messages": [("user", "hi")]}, {"configurable": {"user_id": "2"}}
|
||||
@@ -346,7 +351,7 @@ def test_state_modifier_with_store():
|
||||
assert response["messages"][-1].content == "foo-hi"
|
||||
|
||||
|
||||
async def test_state_modifier_with_store_async():
|
||||
async def test_prompt_with_store_async():
|
||||
async def add(a: int, b: int):
|
||||
"""Adds a and b"""
|
||||
return a + b
|
||||
@@ -359,22 +364,20 @@ async def test_state_modifier_with_store_async():
|
||||
("memories", "2"), "user_name", {"data": "User name is Bob"}
|
||||
)
|
||||
|
||||
async def modify(state, config, *, store):
|
||||
async def prompt(state, config, *, store):
|
||||
user_id = config["configurable"]["user_id"]
|
||||
system_str = (await store.aget(("memories", user_id), "user_name")).value[
|
||||
"data"
|
||||
]
|
||||
return [SystemMessage(system_str)] + state["messages"]
|
||||
|
||||
async def modify_no_store(state, config):
|
||||
async def prompt_no_store(state, config):
|
||||
return SystemMessage("foo") + state["messages"]
|
||||
|
||||
model = FakeToolCallingModel()
|
||||
|
||||
# test state modifier that uses store works
|
||||
agent = create_react_agent(
|
||||
model, [add], state_modifier=modify, store=in_memory_store
|
||||
)
|
||||
agent = create_react_agent(model, [add], prompt=prompt, store=in_memory_store)
|
||||
response = await agent.ainvoke(
|
||||
{"messages": [("user", "hi")]}, {"configurable": {"user_id": "1"}}
|
||||
)
|
||||
@@ -382,7 +385,7 @@ async def test_state_modifier_with_store_async():
|
||||
|
||||
# test state modifier that doesn't use store works
|
||||
agent = create_react_agent(
|
||||
model, [add], state_modifier=modify_no_store, store=in_memory_store
|
||||
model, [add], prompt=prompt_no_store, store=in_memory_store
|
||||
)
|
||||
response = await agent.ainvoke(
|
||||
{"messages": [("user", "hi")]}, {"configurable": {"user_id": "2"}}
|
||||
@@ -1663,7 +1666,7 @@ def test_react_agent_update_state():
|
||||
}
|
||||
)
|
||||
|
||||
def state_modifier(state: State):
|
||||
def prompt(state: State):
|
||||
user_name = state.get("user_name")
|
||||
if user_name is None:
|
||||
return state["messages"]
|
||||
@@ -1678,7 +1681,7 @@ def test_react_agent_update_state():
|
||||
model,
|
||||
[get_user_name],
|
||||
state_schema=State,
|
||||
state_modifier=state_modifier,
|
||||
prompt=prompt,
|
||||
checkpointer=checkpointer,
|
||||
)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
|
||||
Reference in New Issue
Block a user