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:
Eugene Yurtsev
2025-07-18 09:59:14 -04:00
committed by GitHub
parent dc0f0c5944
commit f63bec8578
4 changed files with 91 additions and 34 deletions
+12 -10
View File
@@ -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,