mirror of
https://github.com/langchain-ai/langchain.git
synced 2026-10-05 09:25:14 +03:00
Merge branch 'master' into nm/langchain/pii-state-hook-parity
This commit is contained in:
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)
|
||||
+37
-1
@@ -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)
|
||||
+99
-65
@@ -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:
|
||||
|
||||
+221
@@ -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,
|
||||
*,
|
||||
|
||||
@@ -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(
|
||||
|
||||
Generated
+8
-8
@@ -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" },
|
||||
]
|
||||
|
||||
|
||||
Reference in new issue
Block a user