mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-22 17:45:09 +02:00
chore(prebuilt): revert unnecessary breaking changes (#6020)
* revert change of `agent` node to `model` * revert removal of pydantic agent states + agent state consolidation * revert removal of `version` arg * revert removal of backwards compat `config_schema`
This commit is contained in:
@@ -2,6 +2,8 @@ import dataclasses
|
||||
import inspect
|
||||
from typing import (
|
||||
Annotated,
|
||||
Literal,
|
||||
Optional,
|
||||
Union,
|
||||
)
|
||||
|
||||
@@ -31,6 +33,8 @@ from langgraph.prebuilt import (
|
||||
)
|
||||
from langgraph.prebuilt.chat_agent_executor import (
|
||||
AgentState,
|
||||
AgentStatePydantic,
|
||||
StateT,
|
||||
_validate_chat_history,
|
||||
)
|
||||
from langgraph.prebuilt.tool_node import (
|
||||
@@ -49,14 +53,18 @@ from tests.model import FakeToolCallingModel
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
REACT_TOOL_CALL_VERSIONS = ["v1", "v2"]
|
||||
|
||||
def test_no_prompt(sync_checkpointer: BaseCheckpointSaver) -> None:
|
||||
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
def test_no_prompt(sync_checkpointer: BaseCheckpointSaver, version: str) -> None:
|
||||
model = FakeToolCallingModel()
|
||||
|
||||
agent = create_react_agent(
|
||||
model,
|
||||
[],
|
||||
checkpointer=sync_checkpointer,
|
||||
version=version,
|
||||
)
|
||||
inputs = [HumanMessage("hi?")]
|
||||
thread = {"configurable": {"thread_id": "123"}}
|
||||
@@ -164,7 +172,8 @@ def test_runnable_prompt():
|
||||
assert response == expected_response
|
||||
|
||||
|
||||
def test_prompt_with_store():
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
def test_prompt_with_store(version: Literal["v1", "v2"]):
|
||||
def add(a: int, b: int):
|
||||
"""Adds a and b"""
|
||||
return a + b
|
||||
@@ -189,6 +198,7 @@ def test_prompt_with_store():
|
||||
[add],
|
||||
prompt=prompt,
|
||||
store=in_memory_store,
|
||||
version=version,
|
||||
)
|
||||
response = agent.invoke(
|
||||
{"messages": [("user", "hi")]}, {"configurable": {"user_id": "1"}}
|
||||
@@ -201,6 +211,7 @@ def test_prompt_with_store():
|
||||
[add],
|
||||
prompt=prompt_no_store,
|
||||
store=in_memory_store,
|
||||
version=version,
|
||||
)
|
||||
response = agent.invoke(
|
||||
{"messages": [("user", "hi")]}, {"configurable": {"user_id": "2"}}
|
||||
@@ -241,7 +252,9 @@ async def test_prompt_with_store_async():
|
||||
assert response["messages"][-1].content == "User name is Alice-hi"
|
||||
|
||||
# test state modifier that doesn't use store works
|
||||
agent = create_react_agent(model, [add], prompt=prompt_no_store, store=in_memory_store)
|
||||
agent = create_react_agent(
|
||||
model, [add], prompt=prompt_no_store, store=in_memory_store
|
||||
)
|
||||
response = await agent.ainvoke(
|
||||
{"messages": [("user", "hi")]}, {"configurable": {"user_id": "2"}}
|
||||
)
|
||||
@@ -249,8 +262,9 @@ async def test_prompt_with_store_async():
|
||||
|
||||
|
||||
@pytest.mark.parametrize("tool_style", ["openai", "anthropic"])
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
@pytest.mark.parametrize("include_builtin", [True, False])
|
||||
def test_model_with_tools(tool_style: str, include_builtin: bool) -> None:
|
||||
def test_model_with_tools(tool_style: str, version: str, include_builtin: bool) -> None:
|
||||
model = FakeToolCallingModel(tool_style=tool_style)
|
||||
|
||||
@dec_tool
|
||||
@@ -285,6 +299,7 @@ def test_model_with_tools(tool_style: str, include_builtin: bool) -> None:
|
||||
create_react_agent(
|
||||
model.bind_tools(tools),
|
||||
tools,
|
||||
version=version,
|
||||
)
|
||||
|
||||
|
||||
@@ -414,7 +429,8 @@ def test__infer_handled_types() -> None:
|
||||
_infer_handled_types(handler)
|
||||
|
||||
|
||||
def test_react_agent_with_structured_response() -> None:
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
def test_react_agent_with_structured_response(version: str) -> None:
|
||||
class WeatherResponse(BaseModel):
|
||||
temperature: float = Field(description="The temperature in fahrenheit")
|
||||
|
||||
@@ -435,6 +451,7 @@ def test_react_agent_with_structured_response() -> None:
|
||||
model,
|
||||
[get_weather],
|
||||
response_format=WeatherResponse,
|
||||
version=version,
|
||||
)
|
||||
response = agent.invoke({"messages": [HumanMessage("What's the weather?")]})
|
||||
assert response["structured_response"] == expected_structured_response
|
||||
@@ -463,8 +480,16 @@ class CustomState(AgentState):
|
||||
user_name: str
|
||||
|
||||
|
||||
class CustomStatePydantic(AgentStatePydantic):
|
||||
user_name: Optional[str] = None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
@pytest.mark.parametrize("state_schema", [CustomState, CustomStatePydantic])
|
||||
def test_react_agent_update_state(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
version: Literal["v1", "v2"],
|
||||
state_schema: StateT,
|
||||
) -> None:
|
||||
@dec_tool
|
||||
def get_user_name(tool_call_id: Annotated[str, InjectedToolCallId]):
|
||||
@@ -481,22 +506,34 @@ def test_react_agent_update_state(
|
||||
}
|
||||
)
|
||||
|
||||
def prompt(state: CustomState):
|
||||
user_name = state.get("user_name")
|
||||
if user_name is None:
|
||||
return state["messages"]
|
||||
if issubclass(state_schema, AgentStatePydantic):
|
||||
|
||||
system_msg = f"User name is {user_name}"
|
||||
return [{"role": "system", "content": system_msg}] + state["messages"]
|
||||
def prompt(state: CustomStatePydantic):
|
||||
user_name = state.user_name
|
||||
if user_name is None:
|
||||
return state.messages
|
||||
|
||||
system_msg = f"User name is {user_name}"
|
||||
return [{"role": "system", "content": system_msg}] + state.messages
|
||||
else:
|
||||
|
||||
def prompt(state: CustomState):
|
||||
user_name = state.get("user_name")
|
||||
if user_name is None:
|
||||
return state["messages"]
|
||||
|
||||
system_msg = f"User name is {user_name}"
|
||||
return [{"role": "system", "content": system_msg}] + state["messages"]
|
||||
|
||||
tool_calls = [[{"args": {}, "id": "1", "name": "get_user_name"}]]
|
||||
model = FakeToolCallingModel(tool_calls=tool_calls)
|
||||
agent = create_react_agent(
|
||||
model,
|
||||
[get_user_name],
|
||||
state_schema=CustomState,
|
||||
state_schema=state_schema,
|
||||
prompt=prompt,
|
||||
checkpointer=sync_checkpointer,
|
||||
version=version,
|
||||
)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
# Run until interrupted
|
||||
@@ -512,8 +549,9 @@ def test_react_agent_update_state(
|
||||
assert tool_message.name == "get_user_name"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
def test_react_agent_parallel_tool_calls(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
sync_checkpointer: BaseCheckpointSaver, version: str
|
||||
) -> None:
|
||||
human_assistance_execution_count = 0
|
||||
|
||||
@@ -546,6 +584,7 @@ def test_react_agent_parallel_tool_calls(
|
||||
model,
|
||||
[human_assistance, get_weather],
|
||||
checkpointer=sync_checkpointer,
|
||||
version=version,
|
||||
)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
query = "Get user assistance and also check the weather"
|
||||
@@ -556,11 +595,17 @@ def test_react_agent_parallel_tool_calls(
|
||||
if messages := event.get("messages"):
|
||||
message_types.append([m.type for m in messages])
|
||||
|
||||
assert message_types == [
|
||||
["human"],
|
||||
["human", "ai"],
|
||||
["human", "ai", "tool"],
|
||||
]
|
||||
if version == "v1":
|
||||
assert message_types == [
|
||||
["human"],
|
||||
["human", "ai"],
|
||||
]
|
||||
elif version == "v2":
|
||||
assert message_types == [
|
||||
["human"],
|
||||
["human", "ai"],
|
||||
["human", "ai", "tool"],
|
||||
]
|
||||
|
||||
# Resume
|
||||
message_types = []
|
||||
@@ -576,28 +621,54 @@ def test_react_agent_parallel_tool_calls(
|
||||
["human", "ai", "tool", "tool", "ai"],
|
||||
]
|
||||
|
||||
assert human_assistance_execution_count == 1
|
||||
assert get_weather_execution_count == 1
|
||||
if version == "v1":
|
||||
assert human_assistance_execution_count == 1
|
||||
assert get_weather_execution_count == 2
|
||||
elif version == "v2":
|
||||
assert human_assistance_execution_count == 1
|
||||
assert get_weather_execution_count == 1
|
||||
|
||||
|
||||
class AgentStateExtraKey(AgentState):
|
||||
foo: int
|
||||
|
||||
|
||||
def test_create_react_agent_inject_vars() -> None:
|
||||
class AgentStateExtraKeyPydantic(AgentStatePydantic):
|
||||
foo: int
|
||||
|
||||
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
@pytest.mark.parametrize(
|
||||
"state_schema", [AgentStateExtraKey, AgentStateExtraKeyPydantic]
|
||||
)
|
||||
def test_create_react_agent_inject_vars(
|
||||
version: Literal["v1", "v2"], state_schema: StateT
|
||||
) -> None:
|
||||
"""Test that the agent can inject state and store into tool functions."""
|
||||
store = InMemoryStore()
|
||||
namespace = ("test",)
|
||||
store.put(namespace, "test_key", {"bar": 3})
|
||||
|
||||
def tool1(
|
||||
some_val: int,
|
||||
state: Annotated[dict, InjectedState],
|
||||
store: Annotated[BaseStore, InjectedStore()],
|
||||
) -> str:
|
||||
"""Tool 1 docstring."""
|
||||
store_val = store.get(namespace, "test_key").value["bar"]
|
||||
return some_val + state["foo"] + store_val
|
||||
if issubclass(state_schema, AgentStatePydantic):
|
||||
|
||||
def tool1(
|
||||
some_val: int,
|
||||
state: Annotated[AgentStateExtraKeyPydantic, InjectedState],
|
||||
store: Annotated[BaseStore, InjectedStore()],
|
||||
) -> str:
|
||||
"""Tool 1 docstring."""
|
||||
store_val = store.get(namespace, "test_key").value["bar"]
|
||||
return some_val + state.foo + store_val
|
||||
else:
|
||||
|
||||
def tool1(
|
||||
some_val: int,
|
||||
state: Annotated[dict, InjectedState],
|
||||
store: Annotated[BaseStore, InjectedStore()],
|
||||
) -> str:
|
||||
"""Tool 1 docstring."""
|
||||
store_val = store.get(namespace, "test_key").value["bar"]
|
||||
return some_val + state["foo"] + store_val
|
||||
|
||||
tool_call = {
|
||||
"name": "tool1",
|
||||
@@ -609,8 +680,9 @@ def test_create_react_agent_inject_vars() -> None:
|
||||
agent = create_react_agent(
|
||||
model,
|
||||
ToolNode([tool1], handle_tool_errors=False),
|
||||
state_schema=AgentStateExtraKey,
|
||||
state_schema=state_schema,
|
||||
store=store,
|
||||
version=version,
|
||||
)
|
||||
result = agent.invoke({"messages": [{"role": "user", "content": "hi"}], "foo": 2})
|
||||
assert result["messages"] == [
|
||||
@@ -622,7 +694,8 @@ def test_create_react_agent_inject_vars() -> None:
|
||||
assert result["foo"] == 2
|
||||
|
||||
|
||||
async def test_return_direct() -> None:
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
async def test_return_direct(version: str) -> None:
|
||||
@dec_tool(return_direct=True)
|
||||
def tool_return_direct(input: str) -> str:
|
||||
"""A tool that returns directly."""
|
||||
@@ -649,6 +722,7 @@ async def test_return_direct() -> None:
|
||||
agent = create_react_agent(
|
||||
model,
|
||||
[tool_return_direct, tool_normal],
|
||||
version=version,
|
||||
)
|
||||
|
||||
# Test direct return for tool_return_direct
|
||||
@@ -673,7 +747,9 @@ async def test_return_direct() -> None:
|
||||
),
|
||||
]
|
||||
model = FakeToolCallingModel(tool_calls=[second_tool_call, []])
|
||||
agent = create_react_agent(model, [tool_return_direct, tool_normal])
|
||||
agent = create_react_agent(
|
||||
model, [tool_return_direct, tool_normal], version=version
|
||||
)
|
||||
result = agent.invoke(
|
||||
{"messages": [HumanMessage(content="Test normal", id="hum1")]}
|
||||
)
|
||||
@@ -702,7 +778,9 @@ async def test_return_direct() -> None:
|
||||
),
|
||||
]
|
||||
model = FakeToolCallingModel(tool_calls=[both_tool_calls, []])
|
||||
agent = create_react_agent(model, [tool_return_direct, tool_normal])
|
||||
agent = create_react_agent(
|
||||
model, [tool_return_direct, tool_normal], version=version
|
||||
)
|
||||
result = agent.invoke({"messages": [HumanMessage(content="Test both", id="hum2")]})
|
||||
assert result["messages"] == [
|
||||
HumanMessage(content="Test both", id="hum2"),
|
||||
@@ -740,11 +818,12 @@ def test__get_state_args() -> None:
|
||||
def test_inspect_react() -> None:
|
||||
model = FakeToolCallingModel(tool_calls=[])
|
||||
agent = create_react_agent(model, [])
|
||||
inspect.getclosurevars(agent.nodes["model"].bound.func)
|
||||
inspect.getclosurevars(agent.nodes["agent"].bound.func)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
def test_react_with_subgraph_tools(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
sync_checkpointer: BaseCheckpointSaver, version: Literal["v1", "v2"]
|
||||
) -> None:
|
||||
class State(TypedDict):
|
||||
a: int
|
||||
@@ -800,6 +879,7 @@ def test_react_with_subgraph_tools(
|
||||
model,
|
||||
tool_node,
|
||||
checkpointer=sync_checkpointer,
|
||||
version=version,
|
||||
)
|
||||
result = agent.invoke(
|
||||
{"messages": [HumanMessage(content="What's 2 + 3 and 2 * 3?")]},
|
||||
@@ -830,7 +910,8 @@ def test_react_with_subgraph_tools(
|
||||
]
|
||||
|
||||
|
||||
def test_react_agent_subgraph_streaming_sync() -> None:
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
def test_react_agent_subgraph_streaming_sync(version: Literal["v1", "v2"]) -> None:
|
||||
"""Test React agent streaming when used as a subgraph node sync version"""
|
||||
|
||||
@dec_tool
|
||||
@@ -850,6 +931,7 @@ def test_react_agent_subgraph_streaming_sync() -> None:
|
||||
model,
|
||||
tools=[get_weather],
|
||||
prompt="You are a helpful travel assistant.",
|
||||
version=version,
|
||||
)
|
||||
|
||||
# Create a subgraph that uses the React agent as a node
|
||||
@@ -919,7 +1001,8 @@ def test_react_agent_subgraph_streaming_sync() -> None:
|
||||
assert msg.content.startswith("The weather of Tokyo is sunny.")
|
||||
|
||||
|
||||
async def test_react_agent_subgraph_streaming() -> None:
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
async def test_react_agent_subgraph_streaming(version: Literal["v1", "v2"]) -> None:
|
||||
"""Test React agent streaming when used as a subgraph node."""
|
||||
|
||||
@dec_tool
|
||||
@@ -939,6 +1022,7 @@ async def test_react_agent_subgraph_streaming() -> None:
|
||||
model,
|
||||
tools=[get_weather],
|
||||
prompt="You are a helpful travel assistant.",
|
||||
version=version,
|
||||
)
|
||||
|
||||
# Create a subgraph that uses the React agent as a node
|
||||
@@ -1010,8 +1094,9 @@ async def test_react_agent_subgraph_streaming() -> None:
|
||||
assert msg.content.startswith("The weather of Tokyo is sunny.")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
def test_tool_node_node_interrupt(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
sync_checkpointer: BaseCheckpointSaver, version: str
|
||||
) -> None:
|
||||
def tool_normal(some_val: int) -> str:
|
||||
"""Tool docstring."""
|
||||
@@ -1037,6 +1122,7 @@ def test_tool_node_node_interrupt(
|
||||
model,
|
||||
[tool_interrupt, tool_normal],
|
||||
checkpointer=sync_checkpointer,
|
||||
version=version,
|
||||
)
|
||||
result = agent.invoke({"messages": [HumanMessage("hi?")]}, config)
|
||||
expected_messages = [
|
||||
@@ -1061,7 +1147,11 @@ def test_tool_node_node_interrupt(
|
||||
),
|
||||
_AnyIdToolMessage(content="normal", name="tool_normal", tool_call_id="2"),
|
||||
]
|
||||
assert result["messages"] == expected_messages
|
||||
if version == "v1":
|
||||
# Interrupt blocks second tool result
|
||||
assert result["messages"] == expected_messages[:-1]
|
||||
elif version == "v2":
|
||||
assert result["messages"] == expected_messages
|
||||
|
||||
state = agent.get_state(config)
|
||||
assert state.next == ("tools",)
|
||||
@@ -1075,7 +1165,8 @@ def test_tool_node_node_interrupt(
|
||||
)
|
||||
|
||||
|
||||
def test_dynamic_model_basic() -> None:
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
def test_dynamic_model_basic(version: str) -> None:
|
||||
"""Test basic dynamic model functionality."""
|
||||
|
||||
def dynamic_model(state, runtime: Runtime):
|
||||
@@ -1085,7 +1176,7 @@ def test_dynamic_model_basic() -> None:
|
||||
else:
|
||||
return FakeToolCallingModel(tool_calls=[])
|
||||
|
||||
agent = create_react_agent(dynamic_model, [])
|
||||
agent = create_react_agent(dynamic_model, [], version=version)
|
||||
|
||||
result = agent.invoke({"messages": [HumanMessage("hello")]})
|
||||
assert len(result["messages"]) == 2
|
||||
@@ -1096,7 +1187,8 @@ def test_dynamic_model_basic() -> None:
|
||||
assert result["messages"][-1].content == "urgent help"
|
||||
|
||||
|
||||
def test_dynamic_model_with_tools() -> None:
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
def test_dynamic_model_with_tools(version: Literal["v1", "v2"]) -> None:
|
||||
"""Test dynamic model with tool calling."""
|
||||
|
||||
@dec_tool
|
||||
@@ -1123,7 +1215,9 @@ def test_dynamic_model_with_tools() -> None:
|
||||
tool_calls=[[{"args": {"x": 1}, "id": "1", "name": "basic_tool"}], []]
|
||||
)
|
||||
|
||||
agent = create_react_agent(dynamic_model, [basic_tool, advanced_tool])
|
||||
agent = create_react_agent(
|
||||
dynamic_model, [basic_tool, advanced_tool], version=version
|
||||
)
|
||||
|
||||
# Test basic tool usage
|
||||
result = agent.invoke({"messages": [HumanMessage("basic request")]})
|
||||
@@ -1145,7 +1239,8 @@ class Context:
|
||||
user_id: str
|
||||
|
||||
|
||||
def test_dynamic_model_with_context() -> None:
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
def test_dynamic_model_with_context(version: str) -> None:
|
||||
"""Test dynamic model using config parameters."""
|
||||
|
||||
def dynamic_model(state, runtime: Runtime[Context]):
|
||||
@@ -1156,7 +1251,9 @@ def test_dynamic_model_with_context() -> None:
|
||||
else:
|
||||
return FakeToolCallingModel(tool_calls=[])
|
||||
|
||||
agent = create_react_agent(dynamic_model, [], context_schema=Context)
|
||||
agent = create_react_agent(
|
||||
dynamic_model, [], context_schema=Context, version=version
|
||||
)
|
||||
|
||||
# Test with basic user
|
||||
result = agent.invoke(
|
||||
@@ -1173,7 +1270,8 @@ def test_dynamic_model_with_context() -> None:
|
||||
assert len(result["messages"]) == 2
|
||||
|
||||
|
||||
def test_dynamic_model_with_state_schema() -> None:
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
def test_dynamic_model_with_state_schema(version: Literal["v1", "v2"]) -> None:
|
||||
"""Test dynamic model with custom state schema."""
|
||||
|
||||
class CustomDynamicState(AgentState):
|
||||
@@ -1186,7 +1284,9 @@ def test_dynamic_model_with_state_schema() -> None:
|
||||
else:
|
||||
return FakeToolCallingModel(tool_calls=[])
|
||||
|
||||
agent = create_react_agent(dynamic_model, [], state_schema=CustomDynamicState)
|
||||
agent = create_react_agent(
|
||||
dynamic_model, [], state_schema=CustomDynamicState, version=version
|
||||
)
|
||||
|
||||
result = agent.invoke(
|
||||
{"messages": [HumanMessage("hello")], "model_preference": "advanced"}
|
||||
@@ -1195,14 +1295,15 @@ def test_dynamic_model_with_state_schema() -> None:
|
||||
assert result["model_preference"] == "advanced"
|
||||
|
||||
|
||||
def test_dynamic_model_with_prompt() -> None:
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
def test_dynamic_model_with_prompt(version: Literal["v1", "v2"]) -> None:
|
||||
"""Test dynamic model with different prompt types."""
|
||||
|
||||
def dynamic_model(state: AgentState, runtime: Runtime) -> BaseChatModel:
|
||||
return FakeToolCallingModel(tool_calls=[])
|
||||
|
||||
# Test with string prompt
|
||||
agent = create_react_agent(dynamic_model, [], prompt="system_msg")
|
||||
agent = create_react_agent(dynamic_model, [], prompt="system_msg", version=version)
|
||||
result = agent.invoke({"messages": [HumanMessage("human_msg")]})
|
||||
assert result["messages"][-1].content == "system_msg-human_msg"
|
||||
|
||||
@@ -1211,7 +1312,9 @@ def test_dynamic_model_with_prompt() -> None:
|
||||
"""Generate a dynamic system message based on state."""
|
||||
return [{"role": "system", "content": "system_msg"}] + list(state["messages"])
|
||||
|
||||
agent = create_react_agent(dynamic_model, [], prompt=dynamic_prompt)
|
||||
agent = create_react_agent(
|
||||
dynamic_model, [], prompt=dynamic_prompt, version=version
|
||||
)
|
||||
result = agent.invoke({"messages": [HumanMessage("human_msg")]})
|
||||
assert result["messages"][-1].content == "system_msg-human_msg"
|
||||
|
||||
@@ -1229,7 +1332,8 @@ async def test_dynamic_model_async() -> None:
|
||||
assert result["messages"][-1].content == "hello async"
|
||||
|
||||
|
||||
def test_dynamic_model_with_structured_response() -> None:
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
def test_dynamic_model_with_structured_response(version: str) -> None:
|
||||
"""Test dynamic model with structured response format."""
|
||||
|
||||
class TestResponse(BaseModel):
|
||||
@@ -1250,7 +1354,9 @@ def test_dynamic_model_with_structured_response() -> None:
|
||||
],
|
||||
)
|
||||
|
||||
agent = create_react_agent(dynamic_model, [], response_format=TestResponse)
|
||||
agent = create_react_agent(
|
||||
dynamic_model, [], response_format=TestResponse, version=version
|
||||
)
|
||||
|
||||
result = agent.invoke({"messages": [HumanMessage("hello")]})
|
||||
assert "structured_response" in result
|
||||
@@ -1289,7 +1395,8 @@ def test_dynamic_model_with_checkpointer(sync_checkpointer):
|
||||
assert call_count >= 2
|
||||
|
||||
|
||||
def test_dynamic_model_state_dependent_tools() -> None:
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
def test_dynamic_model_state_dependent_tools(version: Literal["v1", "v2"]) -> None:
|
||||
"""Test dynamic model that changes available tools based on state."""
|
||||
|
||||
@dec_tool
|
||||
@@ -1313,7 +1420,7 @@ def test_dynamic_model_state_dependent_tools() -> None:
|
||||
tool_calls=[[{"args": {"x": 1}, "id": "1", "name": "tool_a"}], []]
|
||||
)
|
||||
|
||||
agent = create_react_agent(dynamic_model, [tool_a, tool_b])
|
||||
agent = create_react_agent(dynamic_model, [tool_a, tool_b], version=version)
|
||||
|
||||
# Ask to use tool B
|
||||
result = agent.invoke({"messages": [HumanMessage("use_b please")]})
|
||||
@@ -1328,7 +1435,8 @@ def test_dynamic_model_state_dependent_tools() -> None:
|
||||
assert last_message.content == "A: 1"
|
||||
|
||||
|
||||
def test_dynamic_model_error_handling() -> None:
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
def test_dynamic_model_error_handling(version: Literal["v1", "v2"]) -> None:
|
||||
"""Test error handling in dynamic model."""
|
||||
|
||||
def failing_dynamic_model(state, runtime: Runtime):
|
||||
@@ -1336,7 +1444,7 @@ def test_dynamic_model_error_handling() -> None:
|
||||
raise ValueError("Dynamic model failed")
|
||||
return FakeToolCallingModel(tool_calls=[])
|
||||
|
||||
agent = create_react_agent(failing_dynamic_model, [])
|
||||
agent = create_react_agent(failing_dynamic_model, [], version=version)
|
||||
|
||||
# Normal operation should work
|
||||
result = agent.invoke({"messages": [HumanMessage("hello")]})
|
||||
@@ -1435,7 +1543,6 @@ async def test_dynamic_model_receives_correct_state_async():
|
||||
assert received_state["messages"][0].content == "hello async"
|
||||
|
||||
|
||||
@pytest.mark.skip(reason="TODO: support with prepare call")
|
||||
def test_pre_model_hook() -> None:
|
||||
model = FakeToolCallingModel(tool_calls=[])
|
||||
|
||||
@@ -1490,7 +1597,7 @@ def test_post_model_hook() -> None:
|
||||
events = list(pmh_agent.stream({"messages": [HumanMessage("hi?")], "flag": False}))
|
||||
assert events == [
|
||||
{
|
||||
"model": {
|
||||
"agent": {
|
||||
"messages": [
|
||||
AIMessage(
|
||||
content="hi?",
|
||||
@@ -1559,7 +1666,7 @@ def test_post_model_hook_with_structured_output() -> None:
|
||||
)
|
||||
assert events == [
|
||||
{
|
||||
"model": {
|
||||
"agent": {
|
||||
"messages": [
|
||||
AIMessage(
|
||||
content="What's the weather?",
|
||||
@@ -1591,7 +1698,7 @@ def test_post_model_hook_with_structured_output() -> None:
|
||||
}
|
||||
},
|
||||
{
|
||||
"model": {
|
||||
"agent": {
|
||||
"messages": [
|
||||
AIMessage(
|
||||
content="What's the weather?-What's the weather?-The weather is sunny and 75°F.",
|
||||
@@ -1620,19 +1727,36 @@ def test_post_model_hook_with_structured_output() -> None:
|
||||
]
|
||||
|
||||
|
||||
def test_create_react_agent_inject_vars_with_post_model_hook() -> None:
|
||||
@pytest.mark.parametrize(
|
||||
"state_schema", [AgentStateExtraKey, AgentStateExtraKeyPydantic]
|
||||
)
|
||||
def test_create_react_agent_inject_vars_with_post_model_hook(
|
||||
state_schema: StateT,
|
||||
) -> None:
|
||||
store = InMemoryStore()
|
||||
namespace = ("test",)
|
||||
store.put(namespace, "test_key", {"bar": 3})
|
||||
|
||||
def tool1(
|
||||
some_val: int,
|
||||
state: Annotated[dict, InjectedState],
|
||||
store: Annotated[BaseStore, InjectedStore()],
|
||||
) -> str:
|
||||
"""Tool 1 docstring."""
|
||||
store_val = store.get(namespace, "test_key").value["bar"]
|
||||
return some_val + state["foo"] + store_val
|
||||
if issubclass(state_schema, AgentStatePydantic):
|
||||
|
||||
def tool1(
|
||||
some_val: int,
|
||||
state: Annotated[AgentStateExtraKeyPydantic, InjectedState],
|
||||
store: Annotated[BaseStore, InjectedStore()],
|
||||
) -> str:
|
||||
"""Tool 1 docstring."""
|
||||
store_val = store.get(namespace, "test_key").value["bar"]
|
||||
return some_val + state.foo + store_val
|
||||
else:
|
||||
|
||||
def tool1(
|
||||
some_val: int,
|
||||
state: Annotated[dict, InjectedState],
|
||||
store: Annotated[BaseStore, InjectedStore()],
|
||||
) -> str:
|
||||
"""Tool 1 docstring."""
|
||||
store_val = store.get(namespace, "test_key").value["bar"]
|
||||
return some_val + state["foo"] + store_val
|
||||
|
||||
tool_call = {
|
||||
"name": "tool1",
|
||||
@@ -1649,7 +1773,7 @@ def test_create_react_agent_inject_vars_with_post_model_hook() -> None:
|
||||
agent = create_react_agent(
|
||||
model,
|
||||
ToolNode([tool1], handle_tool_errors=False),
|
||||
state_schema=AgentStateExtraKey,
|
||||
state_schema=state_schema,
|
||||
store=store,
|
||||
post_model_hook=post_model_hook,
|
||||
)
|
||||
@@ -1664,7 +1788,8 @@ def test_create_react_agent_inject_vars_with_post_model_hook() -> None:
|
||||
assert result["foo"] == 2
|
||||
|
||||
|
||||
def test_response_format_using_tool_choice() -> None:
|
||||
@pytest.mark.parametrize("version", REACT_TOOL_CALL_VERSIONS)
|
||||
def test_response_format_using_tool_choice(version: Literal["v1", "v2"]) -> None:
|
||||
"""Test response format using tool choice."""
|
||||
|
||||
class WeatherResponse(BaseModel):
|
||||
@@ -1685,6 +1810,7 @@ def test_response_format_using_tool_choice() -> None:
|
||||
model,
|
||||
[get_weather],
|
||||
response_format=WeatherResponse,
|
||||
version=version,
|
||||
)
|
||||
response = agent.invoke(
|
||||
{
|
||||
|
||||
Reference in New Issue
Block a user