diff --git a/libs/langchain_v1/langchain/agents/middleware/tool_call_limit.py b/libs/langchain_v1/langchain/agents/middleware/tool_call_limit.py index 1dfc5e4e6e..620b7bc7b1 100644 --- a/libs/langchain_v1/langchain/agents/middleware/tool_call_limit.py +++ b/libs/langchain_v1/langchain/agents/middleware/tool_call_limit.py @@ -26,9 +26,10 @@ ExitBehavior = Literal["continue", "error", "end"] - `'continue'`: Block exceeded tools with error messages, let other tools continue (default) - `'error'`: Raise a `ToolCallLimitExceededError` exception -- `'end'`: Stop execution immediately, injecting a `ToolMessage` and an `AIMessage` for - the single tool call that exceeded the limit. Raises `NotImplementedError` if there - are other pending tool calls (due to parallel tool calling). +- `'end'`: Stop execution immediately, injecting a `ToolMessage` for each tool call + that exceeded the limit and an `AIMessage` explaining why. Any other pending tool + call on the same `AIMessage` (e.g. from parallel tool calling) also gets an + explanatory `ToolMessage` instead of being executed. """ @@ -148,9 +149,9 @@ class ToolCallLimitMiddleware(AgentMiddleware[ToolCallLimitState[ResponseT], Con - `exit_behavior`: How to handle when limits are exceeded - `'continue'`: Block exceeded tools, let execution continue (default) - `'error'`: Raise an exception - - `'end'`: Stop immediately with a `ToolMessage` + AI message for the single - tool call that exceeded the limit (raises `NotImplementedError` if there - are other pending tool calls (due to parallel tool calling). + - `'end'`: Stop immediately with a `ToolMessage` for each exceeded tool + call, a `ToolMessage` explaining any other pending calls were skipped, + and an AI message explaining why execution stopped Examples: !!! example "Continue execution with blocked tools (default)" @@ -220,10 +221,11 @@ class ToolCallLimitMiddleware(AgentMiddleware[ToolCallLimitState[ResponseT], Con - `'continue'`: Block exceeded tools with error messages, let other tools continue. Model decides when to end. - `'error'`: Raise a `ToolCallLimitExceededError` exception - - `'end'`: Stop execution immediately with a `ToolMessage` + AI message - for the single tool call that exceeded the limit. Raises - `NotImplementedError` if there are multiple parallel tool - calls to other tools or multiple pending tool calls. + - `'end'`: Stop execution immediately with a `ToolMessage` for each + tool call that exceeded the limit and an AI message explaining + why. Any other pending tool call on the same `AIMessage` (e.g. + from parallel tool calling) is given an explanatory + `ToolMessage` instead of being executed. Raises: ValueError: If both limits are `None`, if `exit_behavior` is invalid, @@ -336,13 +338,12 @@ class ToolCallLimitMiddleware(AgentMiddleware[ToolCallLimitState[ResponseT], Con Returns: State updates with incremented tool call counts. If limits are exceeded and exit_behavior is `'end'`, also includes a jump to end with a - `ToolMessage` and AI message for the single exceeded tool call. + `ToolMessage` for each exceeded tool call, an explanatory + `ToolMessage` for any other pending tool calls, and an AI message. Raises: ToolCallLimitExceededError: If limits are exceeded and `exit_behavior` is `'error'`. - NotImplementedError: If limits are exceeded, `exit_behavior` is `'end'`, - and there are multiple tool calls. """ # Get the last AIMessage to check for tool calls messages = state.get("messages", []) @@ -419,20 +420,24 @@ class ToolCallLimitMiddleware(AgentMiddleware[ToolCallLimitState[ResponseT], Con ] if self.exit_behavior == "end": - # Check if there are tool calls to other tools that would continue executing - other_tools = [ - tc - for tc in last_ai_message.tool_calls - if self.tool_name is not None and tc["name"] != self.tool_name + # Since jumping to `"end"` skips tool execution, add `ToolMessage`s for all + # pending calls so none are executed and the message history stays valid. + blocked_ids = {tool_call["id"] for tool_call in blocked_calls} + pending_tool_calls = [ + tc for tc in last_ai_message.tool_calls if tc["id"] not in blocked_ids ] - - if other_tools: - tool_names = ", ".join({tc["name"] for tc in other_tools}) - msg = ( - f"Cannot end execution with other tool calls pending. " - f"Found calls to: {tool_names}. Use 'continue' or 'error' behavior instead." + artificial_messages.extend( + ToolMessage( + content=( + "Execution stopped before this tool call could run because " + "another tool call in the same batch exceeded the limit." + ), + tool_call_id=tool_call["id"], + name=tool_call.get("name"), + status="error", ) - raise NotImplementedError(msg) + for tool_call in pending_tool_calls + ) # Build final AI message content (displayed to user - includes thread/run details) # Use hypothetical thread count (what it would have been if call wasn't blocked) @@ -447,6 +452,10 @@ class ToolCallLimitMiddleware(AgentMiddleware[ToolCallLimitState[ResponseT], Con ) artificial_messages.append(AIMessage(content=final_msg_content)) + # Since no calls execute after jumping to `"end"`, don't include this batch in + # the thread-level count. Only the run-level count tracks attempted calls. + thread_counts[count_key] = current_thread_count + return { "thread_tool_call_count": thread_counts, "run_tool_call_count": run_counts, @@ -476,12 +485,11 @@ class ToolCallLimitMiddleware(AgentMiddleware[ToolCallLimitState[ResponseT], Con Returns: State updates with incremented tool call counts. If limits are exceeded and exit_behavior is `'end'`, also includes a jump to end with a - `ToolMessage` and AI message for the single exceeded tool call. + `ToolMessage` for each exceeded tool call, an explanatory + `ToolMessage` for any other pending tool calls, and an AI message. Raises: ToolCallLimitExceededError: If limits are exceeded and `exit_behavior` is `'error'`. - NotImplementedError: If limits are exceeded, `exit_behavior` is `'end'`, - and there are multiple tool calls. """ return self.after_model(state, runtime) diff --git a/libs/langchain_v1/langchain/agents/middleware/tool_selection.py b/libs/langchain_v1/langchain/agents/middleware/tool_selection.py index 98cf01e4b5..e55e267f10 100644 --- a/libs/langchain_v1/langchain/agents/middleware/tool_selection.py +++ b/libs/langchain_v1/langchain/agents/middleware/tool_selection.py @@ -3,11 +3,12 @@ from __future__ import annotations import logging +from collections.abc import Callable from dataclasses import dataclass -from typing import TYPE_CHECKING, Annotated, Any, Literal, Union +from typing import TYPE_CHECKING, Annotated, Any, Literal, TypeGuard, Union from langchain_core.language_models.chat_models import BaseChatModel -from langchain_core.messages import AIMessage, HumanMessage +from langchain_core.messages import AIMessage, BaseMessage, HumanMessage from pydantic import Field, TypeAdapter from typing_extensions import TypedDict @@ -22,7 +23,9 @@ from langchain.agents.middleware.types import ( from langchain.chat_models.base import init_chat_model if TYPE_CHECKING: - from collections.abc import Awaitable, Callable + from collections.abc import Awaitable + + from langchain_core.runnables import RunnableConfig from langchain.tools import BaseTool @@ -32,6 +35,17 @@ DEFAULT_SYSTEM_PROMPT = ( "Your goal is to select the most relevant tools for answering the user's query." ) +OnParsingFailure = Literal["error", "none", "all"] | list[str] | Callable[[Any], list[str]] +"""Behavior when the selection model keeps returning a malformed response. + +Can be either: +- `'error'`: Raise a `ValueError` (the default). +- `'none'`: Select no tools. +- `'all'`: Select every available tool. +- A `list[str]` of tool names to fall back to. +- A callable that takes the last (malformed) response and returns tool names to use. +""" + @dataclass class _SelectionRequest: @@ -78,6 +92,18 @@ def _create_tool_selection_response(tools: list[BaseTool]) -> TypeAdapter[Any]: return TypeAdapter(ToolSelectionResponse) +def _is_valid_selection_response(response: Any) -> TypeGuard[dict[str, Any]]: + """Check whether a structured-output response is a well-formed tool selection. + + Args: + response: Raw response from the structured-output model call. + + Returns: + `True` if `response` is a dict with a `tools` list, `False` otherwise. + """ + return isinstance(response, dict) and isinstance(response.get("tools"), list) + + def _render_tool_list(tools: list[BaseTool]) -> str: """Format tools as markdown list. @@ -126,6 +152,8 @@ class LLMToolSelectorMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, system_prompt: str = DEFAULT_SYSTEM_PROMPT, max_tools: int | None = None, always_include: list[str] | None = None, + max_retries: int = 1, + on_parsing_failure: OnParsingFailure = "error", ) -> None: """Initialize the tool selector. @@ -144,11 +172,39 @@ class LLMToolSelectorMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, always_include: Tool names to always include regardless of selection. These do not count against the `max_tools` limit. + max_retries: Maximum number of retry attempts after the initial call if + the selection model returns a malformed response (not a dict with a + `tools` list). + + Must be `>= 0`. + on_parsing_failure: Behavior once `max_retries` is exhausted and the + response is still malformed. + + Options: + + - `'error'` (default): Raise a `ValueError`. + - `'none'`: Select no tools. + - `'all'`: Select every available tool. + - A `list[str]` of tool names to fall back to. + - A callable that takes the last (malformed) response and returns + the tool names to use. + + Unlike a normal model selection, the fallback tools are not capped + by `max_tools` -- it's an already-deliberate choice, not raw model + output that needs bounding. + + Raises: + ValueError: If `max_retries < 0`. """ super().__init__() + if max_retries < 0: + msg = "max_retries must be >= 0" + raise ValueError(msg) self.system_prompt = system_prompt self.max_tools = max_tools self.always_include = always_include or [] + self.max_retries = max_retries + self.on_parsing_failure = on_parsing_failure if isinstance(model, (BaseChatModel, type(None))): self.model: BaseChatModel | None = model @@ -234,19 +290,34 @@ class LLMToolSelectorMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, available_tools: list[BaseTool], valid_tool_names: list[str], request: ModelRequest[ContextT], + *, + apply_max_tools: bool = True, ) -> ModelRequest[ContextT]: - """Process the selection response and return filtered `ModelRequest`.""" + """Process the selection response and return filtered `ModelRequest`. + + Args: + response: Selection response, expected to have a `tools` list. + available_tools: Tools eligible for selection. + valid_tool_names: Names of `available_tools`. + request: Original model request to override. + apply_max_tools: Whether to cap the selection at `max_tools`. + + Set to `False` for an already-deliberate fallback selection (e.g. + `on_parsing_failure="all"`), where truncating would contradict the + fallback's own semantics. + """ selected_tool_names: list[str] = [] invalid_tool_selections = [] + max_tools = self.max_tools if apply_max_tools else None - for tool_name in response["tools"]: + for tool_name in response.get("tools", []): if tool_name not in valid_tool_names: invalid_tool_selections.append(tool_name) continue # Only add if not already selected and within max_tools limit if tool_name not in selected_tool_names and ( - self.max_tools is None or len(selected_tool_names) < self.max_tools + max_tools is None or len(selected_tool_names) < max_tools ): selected_tool_names.append(tool_name) @@ -270,6 +341,34 @@ class LLMToolSelectorMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, return request.override(tools=[*selected_tools, *provider_tools]) + def _resolve_parsing_failure(self, response: Any, valid_tool_names: list[str]) -> list[str]: + """Determine which tool names to use once `max_retries` is exhausted. + + Args: + response: The last (still malformed) response from the selection model. + valid_tool_names: Tool names available for selection. + + Returns: + Tool names to select, per `on_parsing_failure`. + + Raises: + ValueError: If `on_parsing_failure == "error"` (the default). + """ + if self.on_parsing_failure == "error": + msg = ( + "LLMToolSelectorMiddleware: selection model returned a malformed " + f"response after {self.max_retries} retries (expected a dict with a " + f"'tools' list): {response!r}" + ) + raise ValueError(msg) + if self.on_parsing_failure == "none": + return [] + if self.on_parsing_failure == "all": + return valid_tool_names + if callable(self.on_parsing_failure): + return self.on_parsing_failure(response) + return list(self.on_parsing_failure) + def wrap_model_call( self, request: ModelRequest[ContextT], @@ -285,7 +384,9 @@ class LLMToolSelectorMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, The model call result. Raises: - AssertionError: If the selection model response is not a dict. + ValueError: If `on_parsing_failure == "error"` (the default) and the + selection model returns a malformed response after `max_retries` + retries. """ selection_request = self._prepare_selection_request(request) if selection_request is None: @@ -296,20 +397,41 @@ class LLMToolSelectorMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, schema = type_adapter.json_schema() structured_model = selection_request.model.with_structured_output(schema) - response = structured_model.invoke( - [ - {"role": "system", "content": selection_request.system_message}, - selection_request.last_user_message, - ] - ) + messages: list[BaseMessage | dict[str, Any]] = [ + {"role": "system", "content": selection_request.system_message}, + selection_request.last_user_message, + ] + config: RunnableConfig = {"metadata": {"lc_source": "tool_selection"}} - # Response should be a dict since we're passing a schema (not a Pydantic model class) - if not isinstance(response, dict): - msg = f"Expected dict response, got {type(response)}" - raise AssertionError(msg) # noqa: TRY004 - modified_request = self._process_selection_response( - response, selection_request.available_tools, selection_request.valid_tool_names, request - ) + response = structured_model.invoke(messages, config=config) + attempts = 0 + while not _is_valid_selection_response(response) and attempts < self.max_retries: + # Malformed structured output is usually transient, retry before giving + # up, rather than silently degrading to no/all tools selected. + response = structured_model.invoke(messages, config=config) + attempts += 1 + + if _is_valid_selection_response(response): + modified_request = self._process_selection_response( + response, + selection_request.available_tools, + selection_request.valid_tool_names, + request, + ) + else: + fallback_tool_names = self._resolve_parsing_failure( + response, selection_request.valid_tool_names + ) + # Bypass max_tools: the fallback is already a deliberate selection + # (e.g. `on_parsing_failure="all"`), not raw model output that needs + # to be capped. + modified_request = self._process_selection_response( + {"tools": fallback_tool_names}, + selection_request.available_tools, + selection_request.valid_tool_names, + request, + apply_max_tools=False, + ) return handler(modified_request) async def awrap_model_call( @@ -327,7 +449,9 @@ class LLMToolSelectorMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, The model call result. Raises: - AssertionError: If the selection model response is not a dict. + ValueError: If `on_parsing_failure == "error"` (the default) and the + selection model returns a malformed response after `max_retries` + retries. """ selection_request = self._prepare_selection_request(request) if selection_request is None: @@ -338,18 +462,39 @@ class LLMToolSelectorMiddleware(AgentMiddleware[AgentState[ResponseT], ContextT, schema = type_adapter.json_schema() structured_model = selection_request.model.with_structured_output(schema) - response = await structured_model.ainvoke( - [ - {"role": "system", "content": selection_request.system_message}, - selection_request.last_user_message, - ] - ) + messages: list[BaseMessage | dict[str, Any]] = [ + {"role": "system", "content": selection_request.system_message}, + selection_request.last_user_message, + ] + config: RunnableConfig = {"metadata": {"lc_source": "tool_selection"}} - # Response should be a dict since we're passing a schema (not a Pydantic model class) - if not isinstance(response, dict): - msg = f"Expected dict response, got {type(response)}" - raise AssertionError(msg) # noqa: TRY004 - modified_request = self._process_selection_response( - response, selection_request.available_tools, selection_request.valid_tool_names, request - ) + response = await structured_model.ainvoke(messages, config=config) + attempts = 0 + while not _is_valid_selection_response(response) and attempts < self.max_retries: + # Malformed structured output is usually transient, retry before giving + # up, rather than silently degrading to no/all tools selected. + response = await structured_model.ainvoke(messages, config=config) + attempts += 1 + + if _is_valid_selection_response(response): + modified_request = self._process_selection_response( + response, + selection_request.available_tools, + selection_request.valid_tool_names, + request, + ) + else: + fallback_tool_names = self._resolve_parsing_failure( + response, selection_request.valid_tool_names + ) + # Bypass max_tools: the fallback is already a deliberate selection + # (e.g. `on_parsing_failure="all"`), not raw model output that needs + # to be capped. + modified_request = self._process_selection_response( + {"tools": fallback_tool_names}, + selection_request.available_tools, + selection_request.valid_tool_names, + request, + apply_max_tools=False, + ) return await handler(modified_request) diff --git a/libs/langchain_v1/tests/unit_tests/agents/middleware/implementations/test_shell_tool.py b/libs/langchain_v1/tests/unit_tests/agents/middleware/implementations/test_shell_tool.py index e1cb50fbb4..52023be0a0 100644 --- a/libs/langchain_v1/tests/unit_tests/agents/middleware/implementations/test_shell_tool.py +++ b/libs/langchain_v1/tests/unit_tests/agents/middleware/implementations/test_shell_tool.py @@ -11,10 +11,12 @@ from typing import Any, cast from unittest.mock import Mock import pytest -from langchain_core.messages import ToolMessage +from langchain_core.messages import HumanMessage, ToolCall, ToolMessage from langchain_core.tools.base import ToolException +from langgraph.checkpoint.memory import InMemorySaver from langgraph.runtime import Runtime +from langchain.agents.factory import create_agent from langchain.agents.middleware.shell_tool import ( HostExecutionPolicy, RedactionRule, @@ -24,6 +26,7 @@ from langchain.agents.middleware.shell_tool import ( _SessionResources, _ShellToolInput, ) +from tests.unit_tests.agents.model import FakeToolCallingModel def _empty_state() -> ShellToolState: @@ -694,3 +697,36 @@ def test_kill_process_noop_without_active_process(tmp_path: Path) -> None: # No process has been started; this must not raise. session._kill_process() + + +def test_shell_tool_with_checkpointer_does_not_raise_msgpack_error(tmp_path: Path) -> None: + """`ShellToolMiddleware` must work with a checkpointer end to end. + + Regression test for #34490: `ShellToolState.shell_session_resources` holds + a live, non-serializable `ShellSession` (subprocess + threads). Before + `create_agent` stopped inlining the full agent state into tool-dispatch + `Send`s, that resource ended up in the `Send` payload and the + checkpointer's msgpack writer raised `TypeError: Type is not msgpack + serializable: Send` the moment a checkpointer was configured. + """ + model = FakeToolCallingModel( + tool_calls=[ + [ToolCall(name="shell", args={"command": "echo hi"}, id="call_1")], + [], + ] + ) + middleware = ShellToolMiddleware(workspace_root=tmp_path / "workspace") + agent = create_agent( + model=model, + middleware=[middleware], + checkpointer=InMemorySaver(), + ) + + result = agent.invoke( + {"messages": [HumanMessage("run echo hi")]}, + {"configurable": {"thread_id": "shell-checkpointer-test"}}, + ) + + tool_messages = [m for m in result["messages"] if isinstance(m, ToolMessage)] + assert tool_messages + assert "hi" in str(tool_messages[0].content) diff --git a/libs/langchain_v1/tests/unit_tests/agents/middleware/implementations/test_tool_call_limit.py b/libs/langchain_v1/tests/unit_tests/agents/middleware/implementations/test_tool_call_limit.py index 9050eb6751..d5207a6ce5 100644 --- a/libs/langchain_v1/tests/unit_tests/agents/middleware/implementations/test_tool_call_limit.py +++ b/libs/langchain_v1/tests/unit_tests/agents/middleware/implementations/test_tool_call_limit.py @@ -163,12 +163,10 @@ def test_middleware_unit_functionality() -> None: def test_middleware_end_behavior_with_unrelated_parallel_tool_calls() -> None: """Test middleware 'end' behavior with unrelated parallel tool calls. - Test that 'end' behavior raises NotImplementedError when there are parallel calls - to unrelated tools. - - When limiting a specific tool with "end" behavior and the model proposes parallel calls - to BOTH the limited tool AND other tools, we can't handle this scenario (we'd be stopping - execution while other tools should run). + When limiting a specific tool with "end" behavior and the model proposes parallel + calls to BOTH the limited tool AND other tools, execution still stops immediately, + but the unrelated call gets an explanatory `ToolMessage` instead of being executed, + so it isn't left orphaned in the message history. """ # Limit search tool specifically middleware = ToolCallLimitMiddleware(tool_name="search", thread_limit=1, exit_behavior="end") @@ -189,10 +187,22 @@ def test_middleware_end_behavior_with_unrelated_parallel_tool_calls() -> None: run_tool_call_count={"search": 1}, ) - with pytest.raises( - NotImplementedError, match="Cannot end execution with other tool calls pending" - ): - middleware.after_model(state, runtime) # type: ignore[arg-type] + result = middleware.after_model(state, runtime) # type: ignore[arg-type] + assert result is not None + assert result["jump_to"] == "end" + + tool_messages = [msg for msg in result["messages"] if isinstance(msg, ToolMessage)] + assert len(tool_messages) == 2, "Both the exceeded call and the unrelated call get a message" + + search_msg = next(msg for msg in tool_messages if msg.tool_call_id == "1") + calculator_msg = next(msg for msg in tool_messages if msg.tool_call_id == "2") + assert search_msg.status == "error" + assert "Tool call limit exceeded" in search_msg.content + assert calculator_msg.status == "error" + assert "exceeded the limit" in calculator_msg.content + assert result["thread_tool_call_count"] == {"search": 1}, ( + "Thread count shouldn't advance since the blocked call never executed" + ) def test_middleware_with_specific_tool() -> None: @@ -686,71 +696,95 @@ def test_parallel_tool_calls_with_limit_continue_mode() -> None: def test_parallel_tool_calls_with_limit_end_mode() -> None: """Test parallel tool calls with a limit of 1 in 'end' mode. - When the model proposes 3 tool calls with a limit of 1: - - The first call would be allowed (within limit) - - The 2nd and 3rd calls exceed the limit and get blocked with error ToolMessages - - Execution stops immediately (jump_to: end) so NO tools actually execute - - An AI message explains why execution stopped + When the model proposes 3 tool calls with a limit of 1, the first call would be + allowed (within limit) while the 2nd and 3rd exceed it. Jumping to "end" skips the + tool-execution node entirely, so `ToolCallLimitMiddleware` gives the allowed call + an explanatory `ToolMessage` too (instead of executing it, or leaving it + "orphaned" with no response at all). """ + middleware = ToolCallLimitMiddleware(thread_limit=1, exit_behavior="end") + runtime = None - @tool - def search(query: str) -> str: - """Search for information.""" - return f"Results: {query}" - - # Model proposes 3 parallel search calls - model = FakeToolCallingModel( - tool_calls=[ - [ - ToolCall(name="search", args={"query": "q1"}, id="1"), - ToolCall(name="search", args={"query": "q2"}, id="2"), - ToolCall(name="search", args={"query": "q3"}, id="3"), - ], - [], - ] + state = ToolCallLimitState( + messages=[ + AIMessage( + "Response", + tool_calls=[ + {"name": "search", "args": {"query": "q1"}, "id": "1"}, + {"name": "search", "args": {"query": "q2"}, "id": "2"}, + {"name": "search", "args": {"query": "q3"}, "id": "3"}, + ], + ), + ], + thread_tool_call_count={}, + run_tool_call_count={}, ) - limiter = ToolCallLimitMiddleware(thread_limit=1, exit_behavior="end") - agent = create_agent( - model=model, tools=[search], middleware=[limiter], checkpointer=InMemorySaver() + result = middleware.after_model(state, runtime) # type: ignore[arg-type] + assert result is not None + assert result["jump_to"] == "end" + + tool_messages = [msg for msg in result["messages"] if isinstance(msg, ToolMessage)] + assert len(tool_messages) == 3, "Every tool call gets a matching ToolMessage" + assert {msg.tool_call_id for msg in tool_messages} == {"1", "2", "3"} + assert all(msg.status == "error" for msg in tool_messages) + + exceeded_msgs = [msg for msg in tool_messages if msg.tool_call_id in {"2", "3"}] + allowed_msg = next(msg for msg in tool_messages if msg.tool_call_id == "1") + assert all("Tool call limit exceeded" in msg.content for msg in exceeded_msgs) + assert "exceeded the limit" in allowed_msg.content + + # None of the 3 calls actually ran (jump_to="end" skips the tool node), including + # call "1" which `_separate_tool_calls` classified as "allowed" — so the thread + # count must stay at its pre-batch value, not be incremented for that call. + assert result["thread_tool_call_count"] == {"__all__": 0}, ( + "Thread count shouldn't advance for calls that never actually executed" ) - result = agent.invoke( - {"messages": [HumanMessage("Test")]}, {"configurable": {"thread_id": "test"}} + +def test_middleware_end_behavior_with_allowed_and_blocked_parallel_calls() -> None: + """Regression test for orphaned tool_calls when tool_name is unset (langchain#34159). + + When `exit_behavior="end"` and the middleware limits all tools (`tool_name=None`), + a single `AIMessage` with parallel tool calls can have some calls allowed (under + the limit) and some blocked (over the limit). Previously, only the blocked call + got a `ToolMessage`, so the allowed call was left without a matching response — an + invalid message history that raises a 400 error on the next turn. The allowed call + must now get an explanatory `ToolMessage` too. + """ + middleware = ToolCallLimitMiddleware(run_limit=2, exit_behavior="end") + runtime = None + + state = ToolCallLimitState( + messages=[ + AIMessage( + "Response", + tool_calls=[ + {"name": "get_summarized_eda_data", "args": {}, "id": "call_1"}, + {"name": "get_org_psych_analysis", "args": {}, "id": "call_2"}, + ], + ), + ], + thread_tool_call_count={"__all__": 1}, + run_tool_call_count={"__all__": 1}, ) - messages = result["messages"] - # Verify tool message counts - # With "end" behavior, when we jump to end, NO tools execute (not even allowed ones) - # We only get error ToolMessages for the 2 blocked calls - tool_messages = [msg for msg in messages if isinstance(msg, ToolMessage)] - successful_tool_messages = [msg for msg in tool_messages if msg.status != "error"] - error_tool_messages = [msg for msg in tool_messages if msg.status == "error"] + result = middleware.after_model(state, runtime) # type: ignore[arg-type] + assert result is not None + assert result["jump_to"] == "end" - assert len(successful_tool_messages) == 0, "No tools execute when we jump to end" - assert len(error_tool_messages) == 2, "Should have 2 blocked tool messages (q2, q3)" + tool_messages = [msg for msg in result["messages"] if isinstance(msg, ToolMessage)] + assert {msg.tool_call_id for msg in tool_messages} == {"call_1", "call_2"}, ( + "Every tool call on the AIMessage must get a matching ToolMessage" + ) + assert all(msg.status == "error" for msg in tool_messages) - # Verify error tool messages (sent to model - include "Do not" instruction) - for error_msg in error_tool_messages: - assert "Tool call limit exceeded" in error_msg.content - assert "Do not" in error_msg.content - - # Verify AI message explaining why execution stopped - # (displayed to user - includes thread/run details) - ai_limit_messages = [] - for msg in messages: - if not isinstance(msg, AIMessage): - continue - assert isinstance(msg.content, str) - if "limit" in msg.content.lower() and not msg.tool_calls: - ai_limit_messages.append(msg) - assert len(ai_limit_messages) == 1, "Should have exactly one AI message explaining the limit" - - ai_msg_content = ai_limit_messages[0].content - assert isinstance(ai_msg_content, str) - assert "thread limit exceeded" in ai_msg_content.lower() or ( - "run limit exceeded" in ai_msg_content.lower() - ), "AI message should include thread/run limit details for the user" + # `call_1` was "allowed" by `_separate_tool_calls`, but it never actually + # executes once we jump to "end", so it must not permanently consume the + # thread's quota — the count should stay at its pre-batch value (1). + assert result["thread_tool_call_count"] == {"__all__": 1}, ( + "Thread count shouldn't advance for a call that never actually executed" + ) def test_parallel_mixed_tool_calls_with_specific_tool_limit() -> None: diff --git a/libs/langchain_v1/tests/unit_tests/agents/middleware/implementations/test_tool_selection.py b/libs/langchain_v1/tests/unit_tests/agents/middleware/implementations/test_tool_selection.py index 231583ce2b..a2eecf82e2 100644 --- a/libs/langchain_v1/tests/unit_tests/agents/middleware/implementations/test_tool_selection.py +++ b/libs/langchain_v1/tests/unit_tests/agents/middleware/implementations/test_tool_selection.py @@ -641,3 +641,224 @@ class TestEdgeCases: """Test that empty tools list raises an error in schema creation.""" with pytest.raises(AssertionError, match="tools must be non-empty"): _create_tool_selection_response([]) + + +def _malformed_response_model(times: int) -> FakeModel: + """A selection model that returns a malformed response `times` times in a row.""" + malformed = AIMessage( + content="", tool_calls=[{"name": "ToolSelectionResponse", "id": "1", "args": {}}] + ) + valid = AIMessage( + content="", + tool_calls=[ + {"name": "ToolSelectionResponse", "id": "2", "args": {"tools": ["get_weather"]}} + ], + ) + return FakeModel(messages=iter([malformed] * times + [valid])) + + +class TestMalformedSelectionResponse: + """Test retry and fallback behavior when the selection model returns malformed output.""" + + def test_recovers_within_max_retries(self) -> None: + """Test that a malformed response is retried and recovers within `max_retries`.""" + tool_selection_model = _malformed_response_model(times=1) + model = FakeModel(messages=iter([AIMessage(content="Done")])) + tool_selector = LLMToolSelectorMiddleware(max_tools=1, model=tool_selection_model) + + agent = create_agent( + model=model, tools=[get_weather, search_web], middleware=[tool_selector] + ) + + response = agent.invoke({"messages": [HumanMessage("test")]}) + + assert isinstance(response["messages"][-1], AIMessage) + + def test_default_raises_after_max_retries_exhausted(self) -> None: + """Test that the default `on_parsing_failure='error'` raises a clear `ValueError`.""" + tool_selection_model = _malformed_response_model(times=100) + model = FakeModel(messages=iter([AIMessage(content="Done")])) + tool_selector = LLMToolSelectorMiddleware(model=tool_selection_model) + + agent = create_agent( + model=model, tools=[get_weather, search_web], middleware=[tool_selector] + ) + + with pytest.raises(ValueError, match="malformed"): + agent.invoke({"messages": [HumanMessage("test")]}) + + async def test_async_default_raises_after_max_retries_exhausted(self) -> None: + """Async counterpart: the default behavior also raises via `awrap_model_call`.""" + tool_selection_model = _malformed_response_model(times=100) + model = FakeModel(messages=iter([AIMessage(content="Done")])) + tool_selector = LLMToolSelectorMiddleware(model=tool_selection_model) + + agent = create_agent( + model=model, tools=[get_weather, search_web], middleware=[tool_selector] + ) + + with pytest.raises(ValueError, match="malformed"): + await agent.ainvoke({"messages": [HumanMessage("test")]}) + + def test_max_retries_zero_gives_up_immediately(self) -> None: + """Test that `max_retries=0` doesn't retry before applying `on_parsing_failure`.""" + tool_selection_model = _malformed_response_model(times=1) + model = FakeModel(messages=iter([AIMessage(content="Done")])) + tool_selector = LLMToolSelectorMiddleware(max_retries=0, model=tool_selection_model) + + agent = create_agent( + model=model, tools=[get_weather, search_web], middleware=[tool_selector] + ) + + with pytest.raises(ValueError, match="malformed"): + agent.invoke({"messages": [HumanMessage("test")]}) + + def test_max_retries_negative_rejected_at_construction(self) -> None: + """Test that a negative `max_retries` is rejected eagerly.""" + with pytest.raises(ValueError, match="max_retries must be >= 0"): + LLMToolSelectorMiddleware(max_retries=-1) + + def test_on_parsing_failure_none_selects_no_tools(self) -> None: + """Test that `on_parsing_failure='none'` selects no tools once retries are exhausted.""" + model_requests: list[ModelRequest] = [] + + @wrap_model_call + def trace_model_requests( + request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse] + ) -> ModelResponse: + model_requests.append(request) + return handler(request) + + tool_selection_model = _malformed_response_model(times=100) + model = FakeModel(messages=iter([AIMessage(content="Done")])) + tool_selector = LLMToolSelectorMiddleware( + on_parsing_failure="none", model=tool_selection_model + ) + + agent = create_agent( + model=model, + tools=[get_weather, search_web], + middleware=[tool_selector, trace_model_requests], + ) + agent.invoke({"messages": [HumanMessage("test")]}) + + assert len(model_requests) > 0 + for request in model_requests: + assert request.tools == [] + + def test_on_parsing_failure_all_ignores_max_tools(self) -> None: + """Test that `on_parsing_failure='all'` selects every tool, bypassing `max_tools`.""" + model_requests: list[ModelRequest] = [] + + @wrap_model_call + def trace_model_requests( + request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse] + ) -> ModelResponse: + model_requests.append(request) + return handler(request) + + tool_selection_model = _malformed_response_model(times=100) + model = FakeModel(messages=iter([AIMessage(content="Done")])) + # max_tools=1 would normally cap a real selection at one tool; the "all" + # fallback must not be truncated by it. + tool_selector = LLMToolSelectorMiddleware( + max_tools=1, on_parsing_failure="all", model=tool_selection_model + ) + + agent = create_agent( + model=model, + tools=[get_weather, search_web, calculate], + middleware=[tool_selector, trace_model_requests], + ) + agent.invoke({"messages": [HumanMessage("test")]}) + + assert len(model_requests) > 0 + for request in model_requests: + tool_names = set() + for tool_ in request.tools: + assert isinstance(tool_, BaseTool) + tool_names.add(tool_.name) + assert tool_names == {"get_weather", "search_web", "calculate"} + + def test_on_parsing_failure_list_falls_back_to_subset(self) -> None: + """Test that a `list[str]` fallback selects exactly that subset.""" + model_requests: list[ModelRequest] = [] + + @wrap_model_call + def trace_model_requests( + request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse] + ) -> ModelResponse: + model_requests.append(request) + return handler(request) + + tool_selection_model = _malformed_response_model(times=100) + model = FakeModel(messages=iter([AIMessage(content="Done")])) + tool_selector = LLMToolSelectorMiddleware( + on_parsing_failure=["calculate"], model=tool_selection_model + ) + + agent = create_agent( + model=model, + tools=[get_weather, search_web, calculate], + middleware=[tool_selector, trace_model_requests], + ) + agent.invoke({"messages": [HumanMessage("test")]}) + + assert len(model_requests) > 0 + for request in model_requests: + tool_names = [] + for tool_ in request.tools: + assert isinstance(tool_, BaseTool) + tool_names.append(tool_.name) + assert tool_names == ["calculate"] + + def test_on_parsing_failure_callable_receives_response(self) -> None: + """Test that a callable fallback receives the last malformed response.""" + received: list[Any] = [] + + def fallback(response: Any) -> list[str]: + received.append(response) + return ["search_web"] + + model_requests: list[ModelRequest] = [] + + @wrap_model_call + def trace_model_requests( + request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse] + ) -> ModelResponse: + model_requests.append(request) + return handler(request) + + tool_selection_model = FakeModel( + messages=cycle( + [ + AIMessage( + content="", + tool_calls=[ + {"name": "ToolSelectionResponse", "id": "1", "args": {"oops": True}} + ], + ) + ] + ) + ) + model = FakeModel(messages=iter([AIMessage(content="Done")])) + tool_selector = LLMToolSelectorMiddleware( + on_parsing_failure=fallback, model=tool_selection_model + ) + + agent = create_agent( + model=model, + tools=[get_weather, search_web], + middleware=[tool_selector, trace_model_requests], + ) + agent.invoke({"messages": [HumanMessage("test")]}) + + assert len(model_requests) > 0 + for request in model_requests: + tool_names = [] + for tool_ in request.tools: + assert isinstance(tool_, BaseTool) + tool_names.append(tool_.name) + assert tool_names == ["search_web"] + assert len(received) == 1 + assert received[0] == {"oops": True} diff --git a/libs/partners/anthropic/langchain_anthropic/_version.py b/libs/partners/anthropic/langchain_anthropic/_version.py index 952af29571..9aa0a74cc5 100644 --- a/libs/partners/anthropic/langchain_anthropic/_version.py +++ b/libs/partners/anthropic/langchain_anthropic/_version.py @@ -1,3 +1,3 @@ """Version information for `langchain-anthropic`.""" -__version__ = "1.5.3" +__version__ = "1.5.4" diff --git a/libs/partners/anthropic/langchain_anthropic/chat_models.py b/libs/partners/anthropic/langchain_anthropic/chat_models.py index 9f68a7e214..aa1b67acc3 100644 --- a/libs/partners/anthropic/langchain_anthropic/chat_models.py +++ b/libs/partners/anthropic/langchain_anthropic/chat_models.py @@ -11,7 +11,7 @@ import warnings from collections.abc import AsyncIterator, Callable, Iterator, Mapping, Sequence from functools import cached_property from operator import itemgetter -from typing import Any, Final, Literal, cast +from typing import Any, Final, Literal, TypeGuard, cast import anthropic from langchain_core.callbacks import ( @@ -175,7 +175,7 @@ _ANTHROPIC_EXTRA_FIELDS: set[str] = { """Valid Anthropic-specific extra fields""" -def _is_builtin_tool(tool: Any) -> bool: +def _is_builtin_tool(tool: Any) -> TypeGuard[dict[str, Any]]: """Check if a tool is a built-in (server-side) Anthropic tool. `tool` must be a `dict` and have a `type` key starting with one of the known @@ -2044,6 +2044,15 @@ class ChatAnthropic(BaseChatModel): See the [docs](https://docs.langchain.com/oss/python/integrations/chat/anthropic#strict-tool-use) for more info. kwargs: Any additional parameters are passed directly to `bind`. + Raises: + ValueError: If every tool in `tools` was dropped for using top-level + schema composition, leaving the model with no callable tool. + ValueError: If `tool_choice` forces tool use (a specific tool, or + `'any'`) while any tool was dropped, since the forced tool may be + unreachable. Does not apply when `thinking` is enabled, as the + forced choice is discarded before the request is sent. + ValueError: If `tool_choice` is neither a `dict`, a `str`, nor `None`. + Example: ```python from langchain_anthropic import ChatAnthropic @@ -2080,12 +2089,84 @@ class ChatAnthropic(BaseChatModel): # Allows built-in tools either by their: # - Raw `dict` format # - Extracting extras["provider_tool_definition"] if provided on a BaseTool - formatted_tools = [ + formatted_tools: list[Mapping[str, Any]] = [ tool if _is_builtin_tool(tool) else convert_to_anthropic_tool(tool, strict=strict) for tool in tools ] + formatted_tools, dropped_tool_names = _drop_unsupported_root_composition_tools( + formatted_tools + ) + + # Dropping salvages a request when usable tools remain. If every tool was + # dropped there is nothing to salvage: the caller asked for a model with + # tools and would get one that cannot call any, so fail loudly instead of + # letting it surface later as a tool call that never happens. + if tools and not formatted_tools: + msg = ( + f"All {len(tools)} bound tool(s) use a top-level " + f"{'/'.join(_TOP_LEVEL_SCHEMA_COMPOSITION_KEYS)} in their " + f"input_schema, which the Anthropic API rejects: " + f"{sorted(dropped_tool_names)}. No tool is left for the model to " + "call. If you control the schema, move the combinator under " + "`properties`; for structured output, `with_structured_output(..., " + "method='json_schema')` accepts these schemas directly. Otherwise " + "these tools cannot be used with Anthropic -- bind a subset that " + "excludes them, or raise it with the tool's author." + ) + raise ValueError(msg) + + # Reconcile tool_choice with the filtered list: forcing a tool that was + # dropped, or forcing tool use at all when the set the caller depends on + # has silently shrunk, still produces a 400 or a stuck agent loop. + if tool_choice and dropped_tool_names: + choice_type: str | None = None + choice_name: str | None = None + if isinstance(tool_choice, dict): + choice_type = tool_choice.get("type") + choice_name = tool_choice.get("name") + elif isinstance(tool_choice, str): + if tool_choice in ("any", "auto"): + choice_type = tool_choice + else: + choice_type, choice_name = "tool", tool_choice + # Thinking discards forced choices before the request is sent, so + # they need not refer to a tool that remains after filtering. + thinking_discards_forced_choice = ( + self.thinking is not None + and self.thinking.get("type") in ("enabled", "adaptive") + and choice_type in ("any", "tool") + ) + if not thinking_discards_forced_choice: + if choice_type == "tool" and choice_name in dropped_tool_names: + msg = ( + f"tool_choice forces {choice_name!r}, but that tool was " + "dropped because its input_schema uses a top-level " + f"{'/'.join(_TOP_LEVEL_SCHEMA_COMPOSITION_KEYS)}, which " + "the Anthropic API rejects. Stop forcing it, or -- if you " + "control the schema -- move the combinator under " + "`properties`. For structured output, " + "`with_structured_output(..., method='json_schema')` " + "accepts these schemas directly." + ) + raise ValueError(msg) + if choice_type == "any": + # Forcing tool use means the caller depends on a specific + # reachable tool set, so losing any member of it is fatal -- + # not just losing all of them. + msg = ( + "tool_choice='any' forces the model to call a tool, but " + f"{sorted(dropped_tool_names)} were dropped because their " + "input_schema uses a top-level " + f"{'/'.join(_TOP_LEVEL_SCHEMA_COMPOSITION_KEYS)}, which " + "the Anthropic API rejects, so the model can no longer " + "call them. Use tool_choice='auto' to proceed with the " + "remaining tools, or -- if you control the schema -- move " + "the combinator under `properties`." + ) + raise ValueError(msg) + if not tool_choice: pass elif isinstance(tool_choice, dict): @@ -2368,7 +2449,11 @@ class ChatAnthropic(BaseChatModel): if isinstance(formatted_system, str): kwargs["system"] = formatted_system if tools: - kwargs["tools"] = [convert_to_anthropic_tool(tool) for tool in tools] + # Filter the same schemas `bind_tools` drops, so counting tokens and + # sending a request agree on which tools the API will accept. + kwargs["tools"], _ = _drop_unsupported_root_composition_tools( + [convert_to_anthropic_tool(tool) for tool in tools] + ) if self.context_management is not None: kwargs["context_management"] = self.context_management @@ -2388,6 +2473,56 @@ class ChatAnthropic(BaseChatModel): return response.input_tokens +_TOP_LEVEL_SCHEMA_COMPOSITION_KEYS = ("oneOf", "anyOf") + + +def _drop_unsupported_root_composition_tools( + tools: Sequence[Mapping[str, Any]], +) -> tuple[list[Mapping[str, Any]], set[str]]: + """Drop tools whose root `input_schema` uses `oneOf`/`anyOf`. + + The Anthropic API rejects these at request validation, failing the entire + request. A tool is dropped only if its `input_schema` is a mapping carrying + a root combinator, so built-in (server-side) tools -- which have no + `input_schema` -- and tools whose combinators are nested under `properties` + are passed through as the same objects, unmodified. + + A `UserWarning` is emitted per dropped tool. + + Args: + tools: Already-formatted tool definitions, as built by `bind_tools`. + + Returns: + The retained tools, and the names of the dropped tools. A dropped tool + with no string `name` contributes no entry to the name set. + """ + kept: list[Mapping[str, Any]] = [] + dropped_tool_names: set[str] = set() + for tool in tools: + input_schema = tool.get("input_schema") + offending_keys = ( + [k for k in _TOP_LEVEL_SCHEMA_COMPOSITION_KEYS if k in input_schema] + if isinstance(input_schema, Mapping) + else [] + ) + if not offending_keys: + kept.append(tool) + continue + tool_name = tool.get("name") + if isinstance(tool_name, str): + dropped_tool_names.add(tool_name) + described = repr(tool_name) + else: + described = "with no name" + warnings.warn( + f"Dropping tool {described}: its input_schema has a " + f"top-level {'/'.join(offending_keys)}, which the Anthropic API does " + "not support. The tool will not be available to the model.", + stacklevel=3, + ) + return kept, dropped_tool_names + + def convert_to_anthropic_tool( tool: Mapping[str, Any] | type | Callable | BaseTool, *, diff --git a/libs/partners/anthropic/pyproject.toml b/libs/partners/anthropic/pyproject.toml index be8899139f..4823077261 100644 --- a/libs/partners/anthropic/pyproject.toml +++ b/libs/partners/anthropic/pyproject.toml @@ -20,7 +20,7 @@ classifiers = [ "Topic :: Scientific/Engineering :: Artificial Intelligence", ] -version = "1.5.3" +version = "1.5.4" requires-python = ">=3.10.0,<4.0.0" dependencies = [ "anthropic>=0.120.0,<1.0.0", diff --git a/libs/partners/anthropic/tests/unit_tests/test_chat_models.py b/libs/partners/anthropic/tests/unit_tests/test_chat_models.py index 0f27b893a6..c7b77b89a3 100644 --- a/libs/partners/anthropic/tests/unit_tests/test_chat_models.py +++ b/libs/partners/anthropic/tests/unit_tests/test_chat_models.py @@ -6,6 +6,7 @@ import copy import os import warnings from collections.abc import Callable +from types import SimpleNamespace from typing import Any, Literal, cast from unittest.mock import MagicMock, patch @@ -26,7 +27,7 @@ from langchain_core.runnables import RunnableBinding from langchain_core.tools import BaseTool, tool from langchain_core.tracers.base import BaseTracer from langchain_core.tracers.schemas import Run -from pydantic import BaseModel, Field, SecretStr, ValidationError +from pydantic import BaseModel, Field, RootModel, SecretStr, ValidationError from pytest import CaptureFixture, MonkeyPatch from langchain_anthropic import ChatAnthropic @@ -34,6 +35,7 @@ from langchain_anthropic._version import __version__ from langchain_anthropic.chat_models import ( _TOOL_CALL_ID_PATTERN, _create_usage_metadata, + _drop_unsupported_root_composition_tools, _format_image, _format_messages, _is_builtin_tool, @@ -1622,6 +1624,341 @@ def test_anthropic_bind_tools_does_not_mutate_tool_choice() -> None: } +def test_bind_tools_drops_top_level_composition() -> None: + """Tools with a root `oneOf`/`anyOf` are dropped with a warning. + + The Anthropic API rejects tool schemas carrying these keywords at the top + level, failing the whole request. MCP servers can emit them. See + https://github.com/langchain-ai/langchain/issues/39271. + """ + chat_model = ChatAnthropic( # type: ignore[call-arg, call-arg] + model=MODEL_NAME, + anthropic_api_key="secret-api-key", + ) + valid_tool = { + "name": "search", + "description": "Search", + "input_schema": { + "type": "object", + "properties": {"query": {"type": "string"}}, + "required": ["query"], + }, + } + invalid_tool = { + "name": "notion_create_attachment", + "description": "Create an attachment", + "input_schema": { + "type": "object", + "anyOf": [ + { + "type": "object", + "properties": {"content": {"type": "string"}}, + "required": ["content"], + }, + { + "type": "object", + "properties": {"source_url": {"type": "string"}}, + "required": ["source_url"], + }, + ], + }, + } + with pytest.warns(UserWarning, match="notion_create_attachment"): + chat_model_with_tools = chat_model.bind_tools([valid_tool, invalid_tool]) + + bound = cast("RunnableBinding", chat_model_with_tools).kwargs["tools"] + assert [t["name"] for t in bound] == ["search"] + + +def test_bind_tools_keeps_nested_composition_without_warning() -> None: + """Combinators nested under `properties` are valid and left untouched.""" + chat_model = ChatAnthropic( # type: ignore[call-arg, call-arg] + model=MODEL_NAME, + anthropic_api_key="secret-api-key", + ) + tool = { + "name": "search", + "description": "Search", + "input_schema": { + "type": "object", + "properties": { + "value": {"anyOf": [{"type": "string"}, {"type": "integer"}]}, + }, + "required": ["value"], + }, + } + with warnings.catch_warnings(): + warnings.simplefilter("error") # no warning expected + chat_model_with_tools = chat_model.bind_tools([tool]) + + bound = cast("RunnableBinding", chat_model_with_tools).kwargs["tools"] + assert [t["name"] for t in bound] == ["search"] + assert bound[0]["input_schema"] == tool["input_schema"] + + +def _composition_tool(name: str, keyword: str = "anyOf") -> dict: + """A tool whose root `input_schema` uses a top-level combinator.""" + return { + "name": name, + "description": "Root schema composition.", + "input_schema": { + "type": "object", + keyword: [ + { + "type": "object", + "properties": {"content": {"type": "string"}}, + "required": ["content"], + } + ], + }, + } + + +def _plain_tool(name: str) -> dict: + """A tool with a supported root `input_schema`.""" + return { + "name": name, + "description": "Supported.", + "input_schema": {"type": "object", "properties": {}}, + } + + +@pytest.mark.parametrize("keyword", ["oneOf", "anyOf"]) +def test_bind_tools_drops_each_root_combinator(keyword: str) -> None: + """Every combinator in the unsupported set is filtered and named in the warning.""" + chat_model = ChatAnthropic( # type: ignore[call-arg, call-arg] + model=MODEL_NAME, + anthropic_api_key="secret-api-key", + ) + with pytest.warns(UserWarning, match=f"top-level {keyword}") as record: + bound = chat_model.bind_tools( + [_plain_tool("search"), _composition_tool("attach", keyword)] + ) + + assert [t["name"] for t in cast("RunnableBinding", bound).kwargs["tools"]] == [ + "search" + ] + assert "attach" in str(record[0].message) + + +def test_bind_tools_keeps_root_all_of_without_warning() -> None: + """A root `allOf` schema is supported and remains available to the model.""" + chat_model = ChatAnthropic( # type: ignore[call-arg, call-arg] + model=MODEL_NAME, + anthropic_api_key="secret-api-key", + ) + tool = _composition_tool("attach", "allOf") + with warnings.catch_warnings(): + warnings.simplefilter("error") + bound = chat_model.bind_tools([tool]) + + bound_tools = cast("RunnableBinding", bound).kwargs["tools"] + assert [tool["name"] for tool in bound_tools] == ["attach"] + assert bound_tools[0]["input_schema"] == tool["input_schema"] + + +def test_bind_tools_warning_names_every_offending_combinator() -> None: + """A schema with several root combinators reports all of them.""" + chat_model = ChatAnthropic( # type: ignore[call-arg, call-arg] + model=MODEL_NAME, + anthropic_api_key="secret-api-key", + ) + tool = _composition_tool("attach") + tool["input_schema"]["oneOf"] = [{"type": "object", "properties": {}}] + with pytest.warns(UserWarning, match="top-level oneOf/anyOf"): + chat_model.bind_tools([_plain_tool("search"), tool]) + + +def test_bind_tools_passes_builtin_tools_through_unfiltered() -> None: + """Built-in server-side tools have no `input_schema` and are never dropped.""" + chat_model = ChatAnthropic( # type: ignore[call-arg, call-arg] + model=MODEL_NAME, + anthropic_api_key="secret-api-key", + ) + builtin = {"type": "mcp_toolset", "mcp_server_name": "notion"} + with warnings.catch_warnings(): + warnings.simplefilter("error") # no warning expected + bound = chat_model.bind_tools([builtin]) + + assert cast("RunnableBinding", bound).kwargs["tools"] == [builtin] + + +def test_drop_unsupported_tools_describes_unnamed_tool() -> None: + """A dropped tool with no `name` is described, not rendered as `None`. + + Exercised on the helper directly: `convert_to_anthropic_tool` rejects a + nameless tool before `bind_tools` could ever reach this branch. + """ + unnamed = _composition_tool("attach") + del unnamed["name"] + with pytest.warns(UserWarning, match="Dropping tool with no name") as record: + kept, dropped_names = _drop_unsupported_root_composition_tools([unnamed]) + + assert kept == [] + assert dropped_names == set() + assert "None" not in str(record[0].message) + + +def test_bind_tools_all_tools_dropped_raises() -> None: + """No usable tool remains, so there is nothing to salvage by dropping.""" + chat_model = ChatAnthropic( # type: ignore[call-arg, call-arg] + model=MODEL_NAME, + anthropic_api_key="secret-api-key", + ) + with ( + pytest.warns(UserWarning, match="Dropping tool"), + pytest.raises(ValueError, match="All 1 bound tool"), + ): + chat_model.bind_tools([_composition_tool("attach")]) + + +def test_bind_tools_no_tools_does_not_claim_tools_were_dropped() -> None: + """An empty tool list is not a dropped tool list, and must not say so.""" + chat_model = ChatAnthropic( # type: ignore[call-arg, call-arg] + model=MODEL_NAME, + anthropic_api_key="secret-api-key", + ) + bound = chat_model.bind_tools([], tool_choice="any") + assert cast("RunnableBinding", bound).kwargs["tools"] == [] + + +@pytest.mark.parametrize("tool_choice", ["attach", {"type": "tool", "name": "attach"}]) +def test_bind_tools_dropped_tool_forced_by_tool_choice_raises( + tool_choice: str | dict, +) -> None: + """A dropped forced tool raises locally even when valid tools remain.""" + chat_model = ChatAnthropic( # type: ignore[call-arg, call-arg] + model=MODEL_NAME, + anthropic_api_key="secret-api-key", + ) + with ( + pytest.warns(UserWarning, match="Dropping tool"), + pytest.raises(ValueError, match="tool_choice forces 'attach'"), + ): + chat_model.bind_tools( + [_plain_tool("search"), _composition_tool("attach")], + tool_choice=tool_choice, + ) + + +@pytest.mark.parametrize("tool_choice", ["any", {"type": "any"}]) +def test_bind_tools_partial_drop_under_forced_any_raises( + tool_choice: str | dict, +) -> None: + """Forcing tool use depends on the whole tool set, so losing any member is fatal. + + Under `create_agent`'s `ToolStrategy`, `tool_choice='any'` plus a dropped + structured-output tool would otherwise loop until `GraphRecursionError`. + """ + chat_model = ChatAnthropic( # type: ignore[call-arg, call-arg] + model=MODEL_NAME, + anthropic_api_key="secret-api-key", + ) + with ( + pytest.warns(UserWarning, match="Dropping tool"), + pytest.raises(ValueError, match="tool_choice='any' forces the model"), + ): + chat_model.bind_tools( + [_plain_tool("search"), _composition_tool("attach")], + tool_choice=tool_choice, + ) + + +def test_bind_tools_unknown_forced_tool_choice_is_left_to_the_api() -> None: + """Names the client can't see are not rejected locally. + + `mcp_toolset` tools expose no per-tool `name` -- the names live on the MCP + server -- so a forced choice naming one is only resolvable server-side. + """ + chat_model = ChatAnthropic( # type: ignore[call-arg, call-arg] + model=MODEL_NAME, + anthropic_api_key="secret-api-key", + ) + bound = chat_model.bind_tools( + [{"type": "mcp_toolset", "mcp_server_name": "notion"}], + tool_choice="notion_create_attachment", + ) + assert cast("RunnableBinding", bound).kwargs["tool_choice"] == { + "type": "tool", + "name": "notion_create_attachment", + } + + +def test_with_structured_output_root_combinator_raises_actionable_error() -> None: + """A root-combinator schema fails at bind time, naming the real remedy.""" + chat_model = ChatAnthropic( # type: ignore[call-arg, call-arg] + model=MODEL_NAME, + anthropic_api_key="secret-api-key", + ) + + class _Left(BaseModel): + a: int + + class _Right(BaseModel): + b: str + + class _Either(RootModel): + root: _Left | _Right + + with ( + pytest.warns(UserWarning, match="Dropping tool"), + pytest.raises(ValueError, match="method='json_schema'"), + ): + chat_model.with_structured_output(_Either, method="function_calling") + + +def test_with_structured_output_root_combinator_raises_when_thinking_enabled() -> None: + """The thinking path must not degrade to a toolless request and a wrong error. + + Without the bind-time raise, this spends an API call and then reports the + failure as a `thinking` limitation rather than a schema problem. + """ + chat_model = ChatAnthropic( # type: ignore[call-arg, call-arg] + model=MODEL_NAME, + anthropic_api_key="secret-api-key", + thinking={"type": "enabled", "budget_tokens": 1024}, + ) + + class _Left(BaseModel): + a: int + + class _Right(BaseModel): + b: str + + class _Either(RootModel): + root: _Left | _Right + + with ( + pytest.warns(UserWarning, match="Dropping tool"), + pytest.raises(ValueError, match="All 1 bound tool"), + ): + chat_model.with_structured_output(_Either, method="function_calling") + + +def test_get_num_tokens_from_messages_filters_unsupported_tools() -> None: + """Token counting and sending agree on which tools the API will accept.""" + chat_model = ChatAnthropic( # type: ignore[call-arg, call-arg] + model=MODEL_NAME, + anthropic_api_key="secret-api-key", + ) + counted: dict[str, Any] = {} + + def _count_tokens(**kwargs: Any) -> Any: + counted.update(kwargs) + return SimpleNamespace(input_tokens=42) + + with ( + patch.object(chat_model._client.messages, "count_tokens", _count_tokens), + pytest.warns(UserWarning, match="Dropping tool"), + ): + chat_model.get_num_tokens_from_messages( + [HumanMessage("hi")], + tools=[_plain_tool("search"), _composition_tool("attach")], + ) + + assert [t["name"] for t in counted["tools"]] == ["search"] + + def test_fine_grained_tool_streaming_beta() -> None: """Test that fine-grained tool streaming beta can be enabled.""" # Test with betas parameter at initialization @@ -3773,6 +4110,49 @@ def test_bind_tools_drops_forced_tool_choice_when_thinking_enabled() -> None: assert len(w) == 1 +@pytest.mark.parametrize( + "thinking", + [ + pytest.param({"type": "enabled", "budget_tokens": 5000}, id="enabled"), + pytest.param({"type": "adaptive"}, id="adaptive"), + ], +) +def test_bind_tools_drops_forced_choice_for_filtered_tool_when_thinking_enabled( + thinking: dict[str, Any], +) -> None: + """Thinking takes precedence over forced-choice validation. + + Thinking discards a forced `tool_choice` before the request is sent, so the + choice need not survive filtering. A usable tool must remain, though -- + otherwise the all-tools-dropped guard applies regardless of thinking. + """ + chat_model = ChatAnthropic( + model=MODEL_NAME, + anthropic_api_key="secret-api-key", + thinking=thinking, + ) + unsupported_tool = { + "name": "unsupported_tool", + "description": "A tool with an unsupported root schema composition.", + "input_schema": {"oneOf": [{"type": "string"}, {"type": "number"}]}, + } + + with warnings.catch_warnings(record=True) as w: + warnings.simplefilter("always") + result = chat_model.bind_tools( + [_plain_tool("search"), unsupported_tool], + tool_choice="unsupported_tool", + ) + + assert "tool_choice" not in cast("RunnableBinding", result).kwargs + assert [t["name"] for t in cast("RunnableBinding", result).kwargs["tools"]] == [ + "search" + ] + assert len(w) == 2 + assert "unsupported_tool" in str(w[0].message) + assert "thinking is enabled" in str(w[1].message) + + def test_bind_tools_drops_forced_tool_choice_when_adaptive_thinking() -> None: """Adaptive thinking has the same forced tool_choice restriction as enabled.""" chat_model = ChatAnthropic( diff --git a/libs/partners/anthropic/uv.lock b/libs/partners/anthropic/uv.lock index df6ff2397c..3bbed224b0 100644 --- a/libs/partners/anthropic/uv.lock +++ b/libs/partners/anthropic/uv.lock @@ -343,7 +343,7 @@ name = "exceptiongroup" version = "1.3.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "typing-extensions", marker = "python_full_version < '3.11'" }, + { name = "typing-extensions" }, ] sdist = { url = "https://files.pythonhosted.org/packages/0b/9f/a65090624ecf468cdca03533906e7c69ed7588582240cfe7cc9e770b50eb/exceptiongroup-1.3.0.tar.gz", hash = "sha256:b241f5885f560bc56a59ee63ca4c6a8bfa46ae4ad651af316d4e81817bb9fd88", size = 29749, upload-time = "2025-05-10T17:42:51.123Z" } wheels = [ @@ -587,7 +587,7 @@ requires-dist = [ provides-extras = ["community", "anthropic", "openai", "azure-ai", "google-vertexai", "google-genai", "fireworks", "ollama", "together", "mistralai", "huggingface", "groq", "aws", "baseten", "deepseek", "xai", "perplexity", "meta"] [package.metadata.requires-dev] -lint = [{ name = "ruff", specifier = ">=0.15.0,<0.16.0" }] +lint = [{ name = "ruff", specifier = ">=0.15.0,<0.17.0" }] test = [ { name = "blockbuster", specifier = ">=1.5.26,<1.6.0" }, { name = "langchain-openai", editable = "../openai" }, @@ -617,7 +617,7 @@ typing = [ [[package]] name = "langchain-anthropic" -version = "1.5.3" +version = "1.5.4" source = { editable = "." } dependencies = [ { name = "anthropic" }, @@ -663,7 +663,7 @@ requires-dist = [ ] [package.metadata.requires-dev] -lint = [{ name = "ruff", specifier = ">=0.13.1,<0.16.0" }] +lint = [{ name = "ruff", specifier = ">=0.13.1,<0.17.0" }] test = [ { name = "blockbuster", specifier = ">=1.5.5,<1.6" }, { name = "defusedxml", specifier = ">=0.7.1,<1.0.0" }, @@ -690,7 +690,7 @@ typing = [ [[package]] name = "langchain-core" -version = "1.5.2" +version = "1.5.3" source = { editable = "../../core" } dependencies = [ { name = "jsonpatch" }, @@ -723,7 +723,7 @@ dev = [ { name = "jupyter", specifier = ">=1.0.0,<2.0.0" }, { name = "setuptools", specifier = ">=67.6.1,<84.0.0" }, ] -lint = [{ name = "ruff", specifier = ">=0.15.0,<0.16.0" }] +lint = [{ name = "ruff", specifier = ">=0.15.0,<0.17.0" }] test = [ { name = "blockbuster", specifier = ">=1.5.18,<1.6.0" }, { name = "freezegun", specifier = ">=1.2.2,<2.0.0" }, @@ -798,11 +798,11 @@ requires-dist = [ ] [package.metadata.requires-dev] -lint = [{ name = "ruff", specifier = ">=0.15.0,<0.16.0" }] +lint = [{ name = "ruff", specifier = ">=0.15.0,<0.17.0" }] test = [] test-integration = [] typing = [ - { name = "mypy", specifier = ">=2.1.0,<2.2.0" }, + { name = "mypy", specifier = ">=2.1.0,<2.4.0" }, { name = "types-pyyaml", specifier = ">=6.0.12.2,<7.0.0.0" }, ]