Merge branch 'master' into nm/langchain/pii-state-hook-parity

This commit is contained in:
Nishitha M authored and GitHub committed 2026-08-05 15:18:54 -04:00
commit fb077ca2e3
10 files changed
+1102 -143

No files matched your search

@@ -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)
@@ -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)
@@ -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)
@@ -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:
@@ -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}
@@ -1,3 +1,3 @@
"""Version information for `langchain-anthropic`."""
__version__ = "1.5.3"
__version__ = "1.5.4"
@@ -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,
*,
+1 -1
View File
@@ -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",
@@ -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(
+8 -8
View File
@@ -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" },
]