mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-01 05:55:14 +02:00
chore(prebuilt): restructure tool node and tool injection logic (#5562)
* Cleaning up the underlying tool injection logic which is happening in multiple locations. * State was being injected into the ToolCall via Send in two places in create react agent and the logic doesn't belong there, the actual injection should be happening inside the ToolNode where there's awareness of what run time parameters the tool accepts. Change is required to unblock: https://github.com/langchain-ai/langgraph/pull/5537
This commit is contained in:
@@ -5,6 +5,7 @@ from functools import partial
|
||||
from typing import (
|
||||
Annotated,
|
||||
List,
|
||||
Literal,
|
||||
Optional,
|
||||
Type,
|
||||
TypeVar,
|
||||
@@ -509,7 +510,7 @@ class CustomStatePydantic(AgentStatePydantic):
|
||||
@pytest.mark.parametrize("state_schema", [CustomState, CustomStatePydantic])
|
||||
def test_react_agent_update_state(
|
||||
sync_checkpointer: BaseCheckpointSaver,
|
||||
version: str,
|
||||
version: Literal["v1", "v2"],
|
||||
state_schema: StateSchemaType,
|
||||
) -> None:
|
||||
@dec_tool
|
||||
@@ -557,7 +558,7 @@ def test_react_agent_update_state(
|
||||
version=version,
|
||||
)
|
||||
config = {"configurable": {"thread_id": "1"}}
|
||||
# run until interrpupted
|
||||
# Run until interrupted
|
||||
agent.invoke({"messages": [("user", "what's my name")]}, config)
|
||||
# supply the value for the interrupt
|
||||
response = agent.invoke(Command(resume="Archibald"), config)
|
||||
@@ -781,8 +782,9 @@ class AgentStateExtraKeyPydantic(AgentStatePydantic):
|
||||
"state_schema", [AgentStateExtraKey, AgentStateExtraKeyPydantic]
|
||||
)
|
||||
def test_create_react_agent_inject_vars(
|
||||
version: str, state_schema: StateSchemaType
|
||||
version: Literal["v1", "v2"], state_schema: StateSchemaType
|
||||
) -> None:
|
||||
"""Test that the agent can inject state and store into tool functions."""
|
||||
store = InMemoryStore()
|
||||
namespace = ("test",)
|
||||
store.put(namespace, "test_key", {"bar": 3})
|
||||
@@ -817,15 +819,14 @@ def test_create_react_agent_inject_vars(
|
||||
model = FakeToolCallingModel(tool_calls=[[tool_call], []])
|
||||
agent = create_react_agent(
|
||||
model,
|
||||
[tool1],
|
||||
ToolNode([tool1], handle_tool_errors=False),
|
||||
state_schema=state_schema,
|
||||
store=store,
|
||||
version=version,
|
||||
)
|
||||
input_message = HumanMessage("hi")
|
||||
result = agent.invoke({"messages": [input_message], "foo": 2})
|
||||
result = agent.invoke({"messages": [{"role": "user", "content": "hi"}], "foo": 2})
|
||||
assert result["messages"] == [
|
||||
input_message,
|
||||
_AnyIdHumanMessage(content="hi"),
|
||||
AIMessage(content="hi", tool_calls=[tool_call], id="0"),
|
||||
_AnyIdToolMessage(content="6", name="tool1", tool_call_id="some 0"),
|
||||
AIMessage("hi-hi-6", id="1"),
|
||||
@@ -1580,13 +1581,14 @@ def test_create_react_agent_inject_vars_with_post_model_hook(
|
||||
"type": "tool_call",
|
||||
}
|
||||
|
||||
def post_model_hook(state: dict) -> None:
|
||||
return
|
||||
def post_model_hook(state: dict) -> dict:
|
||||
"""Post model hook is injecting a new foo key."""
|
||||
return {"foo": 2}
|
||||
|
||||
model = FakeToolCallingModel(tool_calls=[[tool_call], []])
|
||||
agent = create_react_agent(
|
||||
model,
|
||||
[tool1],
|
||||
ToolNode([tool1], handle_tool_errors=False),
|
||||
state_schema=state_schema,
|
||||
store=store,
|
||||
post_model_hook=post_model_hook,
|
||||
|
||||
Reference in New Issue
Block a user