diff --git a/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py b/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py index 575c0c6a0..c0e1a414f 100644 --- a/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py +++ b/libs/prebuilt/langgraph/prebuilt/chat_agent_executor.py @@ -23,6 +23,7 @@ from langchain_core.messages import ( AIMessage, BaseMessage, SystemMessage, + ToolCall, ToolMessage, ) from langchain_core.runnables import ( @@ -43,10 +44,12 @@ from langgraph.prebuilt._internal._typing import ( SyncOrAsync, ) from langgraph.prebuilt.responses import ( + MultipleStructuredOutputsError, NativeOutput, NativeOutputBinding, OutputToolBinding, ResponseFormat, + StructuredOutputParsingError, ToolOutput, ) from langgraph.prebuilt.tool_node import ToolNode @@ -60,6 +63,8 @@ StructuredResponseT = TypeVar( default=None, ) +STRUCTURED_OUTPUT_ERROR_TEMPLATE = "Error: {error}\n Please fix your mistakes." + class AgentState(TypedDict, Generic[StructuredResponseT]): """The state of the agent.""" @@ -292,9 +297,10 @@ class _AgentBuilder(Generic[StateT, ContextT, StructuredResponseT]): Command with structured response update if found, None otherwise Raises: - AssertionError: If multiple structured responses are returned + MultipleStructuredOutputsError: If multiple structured responses are returned and error handling is disabled + StructuredOutputParsingError: If parsing fails and error handling is disabled """ - if not response.tool_calls: + if not isinstance(self.response_format, ToolOutput) or not response.tool_calls: return None structured_tool_calls = [ @@ -303,15 +309,53 @@ class _AgentBuilder(Generic[StateT, ContextT, StructuredResponseT]): if tool_call["name"] in self.structured_output_tools ] + if not structured_tool_calls: + return None + if len(structured_tool_calls) > 1: - raise AssertionError( - "Model incorrectly returned multiple structured responses. " - "Behavior has not yet been defined in this case." + return self._handle_multiple_structured_outputs( + response, structured_tool_calls ) - if len(structured_tool_calls) == 1: - tool_call = structured_tool_calls[0] - structured_tool_binding = self.structured_output_tools[tool_call["name"]] + return self._handle_single_structured_output(response, structured_tool_calls[0]) + + def _handle_multiple_structured_outputs( + self, + response: AIMessage, + structured_tool_calls: list[ToolCall], + ) -> Command: + """Handle multiple structured output tool calls.""" + tool_names = [tool_call["name"] for tool_call in structured_tool_calls] + exception = MultipleStructuredOutputsError(tool_names) + + should_retry, error_message = self._handle_structured_output_error(exception) + + if not should_retry: + raise exception + + tool_messages = [ + ToolMessage( + content=error_message, + tool_call_id=tool_call["id"], + name=tool_call["name"], + ) + for tool_call in structured_tool_calls + ] + + return Command( + update={"messages": [response, *tool_messages]}, + goto="model", + ) + + def _handle_single_structured_output( + self, + response: AIMessage, + tool_call: Any, + ) -> Command: + """Handle a single structured output tool call.""" + structured_tool_binding = self.structured_output_tools[tool_call["name"]] + + try: structured_response = structured_tool_binding.parse(tool_call["args"]) if isinstance(structured_response, BaseModel): @@ -341,8 +385,62 @@ class _AgentBuilder(Generic[StateT, ContextT, StructuredResponseT]): "structured_response": structured_response, } ) + except Exception as parse_error: + exception = StructuredOutputParsingError(tool_call["name"], parse_error) - return None + should_retry, error_message = self._handle_structured_output_error( + exception + ) + + if not should_retry: + raise exception + + return Command( + update={ + "messages": [ + response, + ToolMessage( + content=error_message, + tool_call_id=tool_call["id"], + name=tool_call["name"], + ), + ], + }, + goto="model", + ) + + def _handle_structured_output_error( + self, + exception: Exception, + ) -> tuple[bool, str]: + """Handle structured output error. + + Returns (should_retry, retry_tool_message). + """ + assert isinstance(self.response_format, ToolOutput) + handle_errors = self.response_format.handle_errors + + if handle_errors is False: + return False, "" + elif handle_errors is True: + return True, STRUCTURED_OUTPUT_ERROR_TEMPLATE.format(error=str(exception)) + elif isinstance(handle_errors, str): + return True, handle_errors + elif isinstance(handle_errors, type) and issubclass(handle_errors, Exception): + if isinstance(exception, handle_errors): + return True, STRUCTURED_OUTPUT_ERROR_TEMPLATE.format( + error=str(exception) + ) + return False, "" + elif isinstance(handle_errors, tuple): + if any(isinstance(exception, exc_type) for exc_type in handle_errors): + return True, STRUCTURED_OUTPUT_ERROR_TEMPLATE.format( + error=str(exception) + ) + return False, "" + elif callable(handle_errors): + return True, handle_errors(exception) # type: ignore[call-arg] + return False, "" def _apply_native_output_binding( self, model: LanguageModelLike diff --git a/libs/prebuilt/langgraph/prebuilt/responses.py b/libs/prebuilt/langgraph/prebuilt/responses.py index 59a0e1177..72166988d 100644 --- a/libs/prebuilt/langgraph/prebuilt/responses.py +++ b/libs/prebuilt/langgraph/prebuilt/responses.py @@ -5,7 +5,17 @@ from __future__ import annotations import sys import uuid from dataclasses import dataclass, is_dataclass -from typing import Any, Generic, Iterable, Literal, TypeVar, Union, get_args, get_origin +from typing import ( + Any, + Callable, + Generic, + Iterable, + Literal, + TypeVar, + Union, + get_args, + get_origin, +) from langchain_core.messages import AIMessage from langchain_core.tools import BaseTool, StructuredTool @@ -24,6 +34,31 @@ else: SchemaKind = Literal["pydantic", "dataclass", "typeddict", "json_schema"] +class StructuredOutputError(Exception): + """Base class for structured output errors.""" + + +class MultipleStructuredOutputsError(StructuredOutputError): + """Raised when model returns multiple structured output tool calls when only one is expected.""" + + def __init__(self, tool_names: list[str]): + self.tool_names = tool_names + super().__init__( + f"Model incorrectly returned multiple structured responses ({', '.join(tool_names)}) when only one is expected." + ) + + +class StructuredOutputParsingError(StructuredOutputError): + """Raised when structured output tool call arguments fail to parse according to the schema.""" + + def __init__(self, tool_name: str, parse_error: Exception): + self.tool_name = tool_name + self.parse_error = parse_error + super().__init__( + f"Failed to parse structured output for tool '{tool_name}': {parse_error}." + ) + + def _parse_with_schema( schema: Union[type[SchemaT], dict], schema_kind: SchemaKind, data: dict[str, Any] ) -> Any: @@ -54,7 +89,7 @@ def _parse_with_schema( class _SchemaSpec(Generic[SchemaT]): """Describes a structured output schema.""" - schema: Union[type[SchemaT], dict[str, Any]] + schema: type[SchemaT] """The schema for the response, can be a Pydantic model, dataclass, TypedDict, or JSON schema dict.""" name: str @@ -80,7 +115,7 @@ class _SchemaSpec(Generic[SchemaT]): def __init__( self, - schema: Union[type[SchemaT], dict[str, Any]], + schema: type[SchemaT], *, name: str | None = None, description: str | None = None, @@ -89,12 +124,16 @@ class _SchemaSpec(Generic[SchemaT]): """Initialize SchemaSpec with schema and optional parameters.""" self.schema = schema - # Schema names must be unique so we use a shortened UUID suffix - self.name = name or ( - schema.get("title", f"response_format_{str(uuid.uuid4())[:4]}") - if isinstance(schema, dict) - else getattr(schema, "__name__", f"response_format_{str(uuid.uuid4())[:4]}") - ) + if name: + self.name = name + elif isinstance(schema, dict): + self.name = str( + schema.get("title", f"response_format_{str(uuid.uuid4())[:4]}") + ) + else: + self.name = str( + getattr(schema, "__name__", f"response_format_{str(uuid.uuid4())[:4]}") + ) self.description = description or ( schema.get("description", "") @@ -127,7 +166,7 @@ class _SchemaSpec(Generic[SchemaT]): class ToolOutput(Generic[SchemaT]): """Use a tool calling strategy for model responses.""" - schema: Union[type[SchemaT], dict[str, Any]] + schema: type[SchemaT] """Schema for the tool calls.""" schema_specs: list[_SchemaSpec[SchemaT]] @@ -136,14 +175,39 @@ class ToolOutput(Generic[SchemaT]): tool_message_content: str | None """The content of the tool message to be returned when the model calls an artificial structured output tool.""" + handle_errors: Union[ + bool, + str, + type[Exception], + tuple[type[Exception], ...], + Callable[[Exception], str], + ] + """Error handling strategy for structured output via ToolOutput. Default is True. + + - True: Catch all errors with default error template + - str: Catch all errors with this custom message + - type[Exception]: Only catch this exception type with default message + - tuple[type[Exception], ...]: Only catch these exception types with default message + - Callable[[Exception], str]: Custom function that returns error message + - False: No retry, let exceptions propagate + """ + def __init__( self, - schema: Union[type[SchemaT], dict[str, Any]], + schema: type[SchemaT], tool_message_content: str | None = None, + handle_errors: Union[ + bool, + str, + type[Exception], + tuple[type[Exception], ...], + Callable[[Exception], str], + ] = True, ) -> None: - """Initialize ToolOutput with schemas and tool message content.""" + """Initialize ToolOutput with schemas, tool message content, and error handling strategy.""" self.schema = schema self.tool_message_content = tool_message_content + self.handle_errors = handle_errors def _iter_variants(schema: Any) -> Iterable[Any]: """Yield leaf variants from Union and JSON Schema oneOf.""" @@ -167,7 +231,7 @@ class ToolOutput(Generic[SchemaT]): class NativeOutput(Generic[SchemaT]): """Use the model provider's native structured output method.""" - schema: Union[type[SchemaT], dict[str, Any]] + schema: type[SchemaT] """Schema for native mode.""" schema_spec: _SchemaSpec[SchemaT] @@ -175,7 +239,7 @@ class NativeOutput(Generic[SchemaT]): def __init__( self, - schema: Union[type[SchemaT], dict[str, Any]], + schema: type[SchemaT], ) -> None: self.schema = schema self.schema_spec = _SchemaSpec(schema) @@ -202,7 +266,7 @@ class OutputToolBinding(Generic[SchemaT]): and the corresponding tool implementation used by the tools strategy. """ - schema: Union[type[SchemaT], dict[str, Any]] + schema: type[SchemaT] """The original schema provided for structured output (Pydantic model, dataclass, TypedDict, or JSON schema dict).""" schema_kind: SchemaKind @@ -255,7 +319,7 @@ class NativeOutputBinding(Generic[SchemaT]): its type classification, and parsing logic for provider-enforced JSON. """ - schema: Union[type[SchemaT], dict[str, Any]] + schema: type[SchemaT] """The original schema provided for structured output (Pydantic model, dataclass, TypedDict, or JSON schema dict).""" schema_kind: SchemaKind diff --git a/libs/prebuilt/tests/test_response_format.py b/libs/prebuilt/tests/test_response_format.py index c9d1f8357..1bdd986cc 100644 --- a/libs/prebuilt/tests/test_response_format.py +++ b/libs/prebuilt/tests/test_response_format.py @@ -9,7 +9,12 @@ from pydantic import BaseModel, Field from typing_extensions import TypedDict from langgraph.prebuilt import create_agent -from langgraph.prebuilt.responses import NativeOutput, ToolOutput +from langgraph.prebuilt.responses import ( + MultipleStructuredOutputsError, + NativeOutput, + StructuredOutputParsingError, + ToolOutput, +) from tests.model import FakeToolCallingModel try: @@ -387,18 +392,58 @@ class TestResponseFormatAsToolOutput: assert response_location["structured_response"] == EXPECTED_LOCATION assert len(response_location["messages"]) == 5 - def test_multiple_tool_messages(self) -> None: - """Test response_format as ToolOutput with Pydantic model.""" + def test_multiple_structured_outputs_error_without_retry(self) -> None: + """Test that MultipleStructuredOutputsError is raised when model returns multiple structured tool calls without retry.""" tool_calls = [ - [{"args": {}, "id": "1", "name": "get_weather"}], [ { "name": "WeatherBaseModel", - "id": "2", + "id": "1", "args": WEATHER_DATA, }, { - "name": "WeatherDataclass", + "name": "LocationResponse", + "id": "2", + "args": LOCATION_DATA, + }, + ], + ] + + model = FakeToolCallingModel(tool_calls=tool_calls) + + agent = create_agent( + model, + [], + response_format=ToolOutput( + Union[WeatherBaseModel, LocationResponse], + handle_errors=False, + ), + ) + + with pytest.raises( + MultipleStructuredOutputsError, + match=".*WeatherBaseModel.*LocationResponse.*", + ): + agent.invoke({"messages": [HumanMessage("Give me weather and location")]}) + + def test_multiple_structured_outputs_with_retry(self) -> None: + """Test that retry handles multiple structured output tool calls.""" + tool_calls = [ + [ + { + "name": "WeatherBaseModel", + "id": "1", + "args": WEATHER_DATA, + }, + { + "name": "LocationResponse", + "id": "2", + "args": LOCATION_DATA, + }, + ], + [ + { + "name": "WeatherBaseModel", "id": "3", "args": WEATHER_DATA, }, @@ -409,16 +454,177 @@ class TestResponseFormatAsToolOutput: agent = create_agent( model, - [get_weather], - response_format=ToolOutput(Union[WeatherBaseModel, WeatherDataclass]), + [], + response_format=ToolOutput( + Union[WeatherBaseModel, LocationResponse], + handle_errors=True, + ), + ) + + response = agent.invoke({"messages": [HumanMessage("Give me weather")]}) + + # HumanMessage, AIMessage, ToolMessage, ToolMessage, AI, ToolMessage + assert len(response["messages"]) == 6 + assert response["structured_response"] == EXPECTED_WEATHER_PYDANTIC + + def test_structured_output_parsing_error_without_retry(self) -> None: + """Test that StructuredOutputParsingError is raised when tool args fail to parse without retry.""" + tool_calls = [ + [ + { + "name": "WeatherBaseModel", + "id": "1", + "args": {"invalid": "data"}, + }, + ], + ] + + model = FakeToolCallingModel(tool_calls=tool_calls) + + agent = create_agent( + model, + [], + response_format=ToolOutput( + WeatherBaseModel, + handle_errors=False, + ), ) with pytest.raises( - AssertionError, - match="Model incorrectly returned multiple structured responses.", + StructuredOutputParsingError, + match=".*WeatherBaseModel.*", ): agent.invoke({"messages": [HumanMessage("What's the weather?")]}) + def test_structured_output_parsing_error_with_retry(self) -> None: + """Test that retry handles parsing errors for structured output.""" + tool_calls = [ + [ + { + "name": "WeatherBaseModel", + "id": "1", + "args": {"invalid": "data"}, + }, + ], + [ + { + "name": "WeatherBaseModel", + "id": "2", + "args": WEATHER_DATA, + }, + ], + ] + + model = FakeToolCallingModel(tool_calls=tool_calls) + + agent = create_agent( + model, + [], + response_format=ToolOutput( + WeatherBaseModel, + handle_errors=(StructuredOutputParsingError,), + ), + ) + + response = agent.invoke({"messages": [HumanMessage("What's the weather?")]}) + + # HumanMessage, AIMessage, ToolMessage, AIMessage, ToolMessage + assert len(response["messages"]) == 5 + assert response["structured_response"] == EXPECTED_WEATHER_PYDANTIC + + def test_retry_with_custom_function(self) -> None: + """Test retry with custom message generation.""" + tool_calls = [ + [ + { + "name": "WeatherBaseModel", + "id": "1", + "args": WEATHER_DATA, + }, + { + "name": "LocationResponse", + "id": "2", + "args": LOCATION_DATA, + }, + ], + [ + { + "name": "WeatherBaseModel", + "id": "3", + "args": WEATHER_DATA, + }, + ], + ] + + model = FakeToolCallingModel(tool_calls=tool_calls) + + def custom_message(exception: Exception) -> str: + if isinstance(exception, MultipleStructuredOutputsError): + return "Custom error: Multiple outputs not allowed" + return "Custom error" + + agent = create_agent( + model, + [], + response_format=ToolOutput( + Union[WeatherBaseModel, LocationResponse], + handle_errors=custom_message, + ), + ) + + response = agent.invoke({"messages": [HumanMessage("Give me weather")]}) + + # HumanMessage, AIMessage, ToolMessage, ToolMessage, AI, ToolMessage + assert len(response["messages"]) == 6 + assert ( + response["messages"][2].content + == "Custom error: Multiple outputs not allowed" + ) + assert ( + response["messages"][3].content + == "Custom error: Multiple outputs not allowed" + ) + assert response["structured_response"] == EXPECTED_WEATHER_PYDANTIC + + def test_retry_with_custom_string_message(self) -> None: + """Test retry with custom static string message.""" + tool_calls = [ + [ + { + "name": "WeatherBaseModel", + "id": "1", + "args": {"invalid": "data"}, + }, + ], + [ + { + "name": "WeatherBaseModel", + "id": "2", + "args": WEATHER_DATA, + }, + ], + ] + + model = FakeToolCallingModel(tool_calls=tool_calls) + + agent = create_agent( + model, + [], + response_format=ToolOutput( + WeatherBaseModel, + handle_errors="Please provide valid weather data with temperature and condition.", + ), + ) + + response = agent.invoke({"messages": [HumanMessage("What's the weather?")]}) + + assert len(response["messages"]) == 5 + assert ( + response["messages"][2].content + == "Please provide valid weather data with temperature and condition." + ) + assert response["structured_response"] == EXPECTED_WEATHER_PYDANTIC + class TestResponseFormatAsNativeOutput: def test_pydantic_model(self) -> None: