mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-08 09:25:08 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b06dbaae7a | ||
|
|
d5a835e5fd |
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user