Compare commits

...
Author SHA1 Message Date
Eugene Yurtsev b06dbaae7a x 2025-08-13 11:17:19 -04:00
Eugene Yurtsev d5a835e5fd x 2025-08-13 11:16:28 -04:00
3 changed files with 184 additions and 12 deletions
@@ -55,6 +55,12 @@ StructuredResponseSchema = Union[dict, type[BaseModel]]
F = TypeVar("F", bound=Callable[..., Any])
class StepCountIs(BaseModel):
"""Stop condition that specifies when to halt agent execution based on step count."""
count: int
# We create the AgentState that we will pass around
# This simply involves a list of messages
# We want steps to return messages to append to the list
@@ -66,6 +72,8 @@ class AgentState(TypedDict):
remaining_steps: NotRequired[RemainingSteps]
model_calls: NotRequired[int]
class AgentStatePydantic(BaseModel):
"""The state of the agent."""
@@ -74,6 +82,8 @@ class AgentStatePydantic(BaseModel):
remaining_steps: RemainingSteps = 25
model_calls: int = 0
class AgentStateWithStructuredResponse(AgentState):
"""The state of the agent with a structured response."""
@@ -270,6 +280,7 @@ class _AgentBuilder:
response_format: Optional[
Union[StructuredResponseSchema, tuple[str, StructuredResponseSchema]]
] = None,
stop_when: Optional[Union[StepCountIs, Callable[[StateSchema], bool]]] = None,
pre_model_hook: Optional[RunnableLike] = None,
post_model_hook: Optional[RunnableLike] = None,
state_schema: Optional[StateSchemaType] = None,
@@ -291,6 +302,7 @@ class _AgentBuilder:
self.tools = tools
self.prompt = prompt
self.response_format = response_format
self.stop_when = stop_when
self.pre_model_hook = pre_model_hook
self.post_model_hook = post_model_hook
self.state_schema = state_schema
@@ -328,6 +340,8 @@ class _AgentBuilder:
required_keys = {"messages", "remaining_steps"}
if self.response_format is not None:
required_keys.add("structured_response")
if self.stop_when is not None:
required_keys.add("model_calls")
schema_keys = set(get_type_hints(self.state_schema))
if missing_keys := required_keys - schema_keys:
@@ -459,6 +473,19 @@ class _AgentBuilder:
return True
return False
def _should_stop_execution(state: StateSchema) -> bool:
"""Check if execution should stop based on stop_when condition."""
if self.stop_when is None:
return False
if isinstance(self.stop_when, StepCountIs):
model_calls = _get_state_value(state, "model_calls", 0)
return model_calls >= self.stop_when.count
elif callable(self.stop_when):
return self.stop_when(state)
return False
def call_model(
state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig
) -> StateSchema:
@@ -468,25 +495,40 @@ class _AgentBuilder:
"Use agent.ainvoke() or agent.astream(), or provide a sync model callable."
)
# Check if we should stop before making the model call
if _should_stop_execution(state):
return {}
model_input = _get_model_input_state(state)
model = self._resolve_model(state, runtime)
response = cast(AIMessage, model.invoke(model_input, config)) # type: ignore[arg-type]
response.name = self.name
# Track model calls
result = {"messages": [response]}
if self.stop_when is not None:
current_calls = _get_state_value(state, "model_calls", 0)
result["model_calls"] = current_calls + 1
if _are_more_steps_needed(state, response):
return {
**result,
"messages": [
AIMessage(
id=response.id,
content="Sorry, need more steps to process this request.",
)
]
],
}
return {"messages": [response]}
return result
async def acall_model(
state: StateSchema, runtime: Runtime[ContextT], config: RunnableConfig
) -> StateSchema:
# Check if we should stop before making the model call
if _should_stop_execution(state):
return {}
model_input = _get_model_input_state(state)
model = await self._aresolve_model(state, runtime)
@@ -495,16 +537,24 @@ class _AgentBuilder:
await model.ainvoke(model_input, config), # type: ignore[arg-type]
)
response.name = self.name
# Track model calls
result = {"messages": [response]}
if self.stop_when is not None:
current_calls = _get_state_value(state, "model_calls", 0)
result["model_calls"] = current_calls + 1
if _are_more_steps_needed(state, response):
return {
**result,
"messages": [
AIMessage(
id=response.id,
content="Sorry, need more steps to process this request.",
)
]
],
}
return {"messages": [response]}
return result
return RunnableCallable(call_model, acall_model)
@@ -852,6 +902,7 @@ def create_react_agent(
response_format: Optional[
Union[StructuredResponseSchema, tuple[str, StructuredResponseSchema]]
] = None,
stop_when: Optional[Union[StepCountIs, Callable[[StateSchema], bool]]] = None,
pre_model_hook: Optional[RunnableLike] = None,
post_model_hook: Optional[RunnableLike] = None,
state_schema: Optional[StateSchemaType] = None,
@@ -940,6 +991,16 @@ def create_react_agent(
The graph will make a separate call to the LLM to generate the structured response after the agent loop is finished.
This is not the only strategy to get structured responses, see more options in [this guide](https://langchain-ai.github.io/langgraph/how-tos/react-agent-structured-output/).
stop_when: An optional condition to stop agent execution.
Can be passed in as:
- `StepCountIs(count=N)`: Stop execution after exactly N model calls
- A callable with signature `(state) -> bool`: Custom stop condition that returns True when execution should stop
When a stop condition is met, the agent will halt further execution and return the current state.
If using `StepCountIs`, the state will include a `model_calls` field tracking the number of model invocations.
pre_model_hook: An optional node to add before the `agent` node (i.e., the node that calls the LLM).
Useful for managing long message histories (e.g., message trimming, summarization, etc.).
Pre-model hook must be a callable or a runnable that takes in current graph state and returns a state update in the form of
@@ -1081,6 +1142,7 @@ def create_react_agent(
tools=tools,
prompt=prompt,
response_format=response_format,
stop_when=stop_when,
pre_model_hook=pre_model_hook,
post_model_hook=post_model_hook,
state_schema=state_schema,
@@ -1113,4 +1175,5 @@ __all__ = [
"AgentStatePydantic",
"AgentStateWithStructuredResponse",
"AgentStateWithStructuredResponsePydantic",
"StepCountIs",
]
+49 -8
View File
@@ -340,14 +340,22 @@ class ToolNode(RunnableCallable):
self.tools_by_name: dict[str, BaseTool] = {}
self.tool_to_state_args: dict[str, dict[str, Optional[str]]] = {}
self.tool_to_store_arg: dict[str, Optional[str]] = {}
self.structured_output_tools: list[str] = []
self.handle_tool_errors = handle_tool_errors
self.messages_key = messages_key
for tool_ in tools:
if not isinstance(tool_, BaseTool):
tool_ = create_tool(tool_)
self.tools_by_name[tool_.name] = tool_
self.tool_to_state_args[tool_.name] = _get_state_args(tool_)
self.tool_to_store_arg[tool_.name] = _get_store_arg(tool_)
if inspect.isclass(tool_) and issubclass(tool_, BaseModel):
# Handle Pydantic model classes as structured output tools
self.tools_by_name[tool_.__name__] = tool_
self.tool_to_state_args[tool_.__name__] = {}
self.tool_to_store_arg[tool_.__name__] = None
self.structured_output_tools.append(tool_.__name__)
else:
if not isinstance(tool_, BaseTool):
tool_ = create_tool(tool_)
self.tools_by_name[tool_.name] = tool_
self.tool_to_state_args[tool_.name] = _get_state_args(tool_)
self.tool_to_store_arg[tool_.name] = _get_store_arg(tool_)
def _func(
self,
@@ -390,7 +398,7 @@ class ToolNode(RunnableCallable):
def _combine_tool_outputs(
self,
outputs: list[ToolMessage],
outputs: list[Union[ToolMessage, Command]],
input_type: Literal["list", "dict", "tool_calls"],
) -> list[Union[Command, list[ToolMessage], dict[str, list[ToolMessage]]]]:
# preserve existing behavior for non-command tool outputs for backwards
@@ -437,10 +445,27 @@ class ToolNode(RunnableCallable):
call: ToolCall,
input_type: Literal["list", "dict", "tool_calls"],
config: RunnableConfig,
) -> ToolMessage:
) -> Union[ToolMessage, Command]:
"""Run a single tool call synchronously."""
if invalid_tool_message := self._validate_tool_call(call):
return invalid_tool_message
# Handle structured output tools
if call["name"] in self.structured_output_tools:
response_schema = self.tools_by_name[call["name"]]
return Command(
update={
"messages": [
ToolMessage(
content=msg_content_output(call["args"]),
name=call["name"],
tool_call_id=call["id"],
)
],
"structured_response": response_schema(**call["args"]),
}
)
try:
call_args = {**call, **{"type": "tool_call"}}
response = self.tools_by_name[call["name"]].invoke(call_args, config)
@@ -493,11 +518,27 @@ class ToolNode(RunnableCallable):
call: ToolCall,
input_type: Literal["list", "dict", "tool_calls"],
config: RunnableConfig,
) -> ToolMessage:
) -> Union[ToolMessage, Command]:
"""Run a single tool call asynchronously."""
if invalid_tool_message := self._validate_tool_call(call):
return invalid_tool_message
# Handle structured output tools
if call["name"] in self.structured_output_tools:
response_schema = self.tools_by_name[call["name"]]
return Command(
update={
"messages": [
ToolMessage(
content=msg_content_output(call["args"]),
name=call["name"],
tool_call_id=call["id"],
)
],
"structured_response": response_schema(**call["args"]),
}
)
try:
call_args = {**call, **{"type": "tool_call"}}
response = await self.tools_by_name[call["name"]].ainvoke(call_args, config)
+68
View File
@@ -1481,3 +1481,71 @@ def test_tool_node_stream_writer() -> None:
},
),
]
async def test_structured_output_tools():
"""Test that ToolNode handles Pydantic model classes as structured output tools."""
class OutputSchema(BaseModel):
name: str
age: int
location: str
tool_node = ToolNode([OutputSchema])
# Test that the structured output tool is registered correctly
assert "OutputSchema" in tool_node.tools_by_name
assert "OutputSchema" in tool_node.structured_output_tools
# Create a tool call that matches the schema
tool_call = {
"name": "OutputSchema",
"args": {"name": "Alice", "age": 30, "location": "NYC"},
"id": "call_123",
"type": "tool_call",
}
# Test sync execution
result = tool_node.invoke(
{"messages": [AIMessage(content="", tool_calls=[tool_call])]}
)
# Should return a Command with structured response
assert isinstance(result, list)
assert len(result) == 1
command = result[0]
assert isinstance(command, Command)
# Check the update structure
assert "messages" in command.update
assert "structured_response" in command.update
# Check the tool message
tool_message = command.update["messages"][0]
assert isinstance(tool_message, ToolMessage)
assert tool_message.name == "OutputSchema"
assert tool_message.tool_call_id == "call_123"
# Check the structured response
structured_response = command.update["structured_response"]
assert isinstance(structured_response, OutputSchema)
assert structured_response.name == "Alice"
assert structured_response.age == 30
assert structured_response.location == "NYC"
# Test async execution
result_async = await tool_node.ainvoke(
{"messages": [AIMessage(content="", tool_calls=[tool_call])]}
)
# Should produce the same result
assert isinstance(result_async, list)
assert len(result_async) == 1
command_async = result_async[0]
assert isinstance(command_async, Command)
assert "structured_response" in command_async.update
structured_response_async = command_async.update["structured_response"]
assert isinstance(structured_response_async, OutputSchema)
assert structured_response_async.name == "Alice"