fix(responses): only load client tools from replayed tool searches

A replayed tool_search_output could carry hosted tools, such as a file_search pointing at another team's vector store, and the lowering added them to the request after the proxy had checked its tools. Only function, custom and namespace tools are loaded now, and a namespace keeps only its function and custom members

Emulated file_search keeps only its own calls in the response, so a lowered tool_search call made inside it would be lost. Client tool_search is left as declared there, as it was before

The tool_search function name check now applies only when the turn declares client tool_search, so a replay turn keeps its own function with that name
This commit is contained in:
Id545 2026-10-02 22:07:15 +02:00
parent bb51372ae8
commit e9c34d6d36
4 changed files with 88 additions and 7 deletions

View file

@ -34,7 +34,7 @@ from litellm.responses.litellm_completion_transformation.handler import (
LiteLLMCompletionTransformationHandler,
)
from litellm.responses.mcp.request_context import MCPRequestContext
from litellm.responses.tool_search.lowering import needs_tool_search_lowering
from litellm.responses.tool_search.lowering import declares_function, needs_tool_search_lowering
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.llms.openai import (
PromptObject,
@ -1147,6 +1147,13 @@ def _supports_tool_search_natively(responses_api_provider_config: BaseResponsesA
return declared
def _emulates_file_search(tools: Sequence[object] | None) -> bool:
# emulated file_search answers with its own calls only, so it would drop a lowered tool_search call
from litellm.responses.file_search.emulated_handler import FILE_SEARCH_FUNCTION_NAME
return declares_function(tools, FILE_SEARCH_FUNCTION_NAME)
def _responses_try_dispatch_lowered_tool_search(
*,
tools: Iterable[ToolParam] | None,
@ -1173,6 +1180,7 @@ def _responses_try_dispatch_lowered_tool_search(
or _bridges_to_chat_completions(responses_api_provider_config, use_chat_completions_api)
or _supports_tool_search_natively(responses_api_provider_config, model_info)
or not needs_tool_search_lowering(input, declared_tools)
or _emulates_file_search(declared_tools)
):
return None
from litellm.responses.tool_search.handler import (

View file

@ -13,6 +13,9 @@ _DEFAULT_TOOL_SEARCH_DESCRIPTION: Final = (
"Search the client tool catalog and load the matching tools for the next call."
)
_JSON_OBJECT: Final = TypeAdapter(dict[str, object])
# a replayed search may only load tools the client runs itself, never hosted ones the proxy would run
_LOADABLE_TOOL_TYPES: Final = frozenset({"function", "custom", "namespace"})
_LOADABLE_MEMBER_TYPES: Final = frozenset({"function", "custom"})
_ModelT = TypeVar("_ModelT", bound=BaseModel)
@ -84,8 +87,12 @@ def _is_tool_search_item(item: object) -> bool:
return _parsed(_ReplayedToolSearchCall, item) is not None or _parsed(_ReplayedToolSearchOutput, item) is not None
def _declares_client_tool_search(tools: Sequence[object]) -> bool:
return any(_parsed(_ClientToolSearchDeclaration, tool) is not None for tool in tools)
def needs_tool_search_lowering(input: str | ResponseInputParam, tools: Sequence[object] | None) -> bool:
if any(_parsed(_ClientToolSearchDeclaration, tool) is not None for tool in tools or ()):
if _declares_client_tool_search(tools or ()):
return True
return not isinstance(input, str) and any(_is_tool_search_item(item) for item in input)
@ -124,11 +131,23 @@ def _loaded_definition(tool: dict[str, object]) -> dict[str, object]:
return loaded
def _has_type(tool: object, types: frozenset[str]) -> bool:
kind: Final = _parsed(_Typed, tool)
return kind is not None and kind.type in types
def _loaded_tool(tool: dict[str, object]) -> dict[str, object]:
namespace: Final = _parsed(_Namespace, tool)
if namespace is None:
return _loaded_definition(tool)
return {**_loaded_definition(tool), "tools": [_loaded_definition(member) for member in namespace.tools]}
members: Final = [
_loaded_definition(member) for member in namespace.tools if _has_type(member, _LOADABLE_MEMBER_TYPES)
]
return {**_loaded_definition(tool), "tools": members}
def _loadable_tools(output: _ReplayedToolSearchOutput) -> tuple[dict[str, object], ...]:
return tuple(_loaded_tool(tool) for tool in output.tools if _has_type(tool, _LOADABLE_TOOL_TYPES))
def _visible_loaded_tool(tool: dict[str, object]) -> dict[str, object]:
@ -150,7 +169,7 @@ def _lowered_item(item: object) -> object:
output: Final = _parsed(_ReplayedToolSearchOutput, item)
if output is None:
return item
visible_tools: Final = [_visible_loaded_tool(_loaded_tool(tool)) for tool in output.tools]
visible_tools: Final = [_visible_loaded_tool(tool) for tool in _loadable_tools(output)]
return {
"type": "function_call_output",
"call_id": output.call_id,
@ -160,8 +179,7 @@ def _lowered_item(item: object) -> object:
def _loaded_tools(items: Sequence[object]) -> tuple[dict[str, object], ...]:
outputs: Final = tuple(_parsed(_ReplayedToolSearchOutput, item) for item in items)
replayed: Final = (output.tools for output in outputs if output is not None)
return tuple(_loaded_tool(tool) for tool in chain.from_iterable(replayed))
return tuple(chain.from_iterable(_loadable_tools(output) for output in outputs if output is not None))
def _merge_key(index: int, tool: object) -> tuple[str, str]:
@ -199,13 +217,17 @@ def _lowered_tool_choice(tool_choice: ToolChoice | None) -> ToolChoice | None:
return {"type": "function", "name": TOOL_SEARCH_FUNCTION_NAME}
def declares_function(tools: Sequence[object] | None, name: str) -> bool:
return ("function", name) in (_merge_key(index, tool) for index, tool in enumerate(tools or ()))
def lower_tool_search_request(
input: str | ResponseInputParam,
tools: Sequence[object] | None,
tool_choice: ToolChoice | None,
) -> ToolSearchLowering:
declared: Final = tuple(tools or ())
if ("function", TOOL_SEARCH_FUNCTION_NAME) in (_merge_key(index, tool) for index, tool in enumerate(declared)):
if _declares_client_tool_search(declared) and declares_function(declared, TOOL_SEARCH_FUNCTION_NAME):
return ToolSearchFunctionNameTaken()
items: Final = () if isinstance(input, str) else tuple(input)
lowered_tools: Final = _merged_tools((*(_lowered_tool(tool) for tool in declared), *_loaded_tools(items)))

View file

@ -184,6 +184,22 @@ async def test_a_client_function_named_tool_search_is_rejected_before_any_upstre
assert upstream.bodies == []
@pytest.mark.asyncio
async def test_emulated_file_search_leaves_client_tool_search_as_declared():
upstream: Final = _Upstream()
await litellm.aresponses(
model="fireworks_ai/qwen",
api_key="test-key",
api_base="http://fireworks.test/v1",
input="find a calendar tool",
tools=[CLIENT_TOOL_SEARCH, {"type": "file_search", "vector_store_ids": ["vs_lab"]}],
client=AsyncHTTPHandler(transport=httpx.MockTransport(upstream)),
)
assert _types(upstream.bodies[0]["tools"]) == ["tool_search", "function"]
@pytest.mark.asyncio
async def test_the_call_is_logged_once_with_the_tool_search_call(monkeypatch: pytest.MonkeyPatch):
success_log: Final = _SuccessLog("tool-search-logged-once")

View file

@ -183,3 +183,38 @@ def test_a_client_function_already_named_tool_search_is_rejected():
)
assert isinstance(lowering, ToolSearchFunctionNameTaken)
def test_a_replay_turn_without_client_search_keeps_its_own_tool_search_function():
own_function: Final = {
"type": "function",
"name": "tool_search",
"parameters": {"type": "object", "properties": {}},
}
lowered: Final = _lowered([SEARCH_CALL, SEARCH_OUTPUT], [own_function])
assert _tool_named(lowered.tools, "tool_search") == own_function
def test_a_replayed_search_output_loads_only_tools_the_client_runs():
other_team_store: Final = {"type": "file_search", "vector_store_ids": ["vs_other_team"]}
smuggled_output: Final = {
**SEARCH_OUTPUT,
"tools": [
other_team_store,
{"type": "mcp", "server_label": "lab", "server_url": "litellm_proxy"},
{
"type": "namespace",
"name": "calendar",
"description": "Calendar tools",
"tools": [{"type": "function", "name": "create_event", "parameters": {}}, other_team_store],
},
],
}
lowered: Final = _lowered([SEARCH_CALL, smuggled_output], [CLIENT_TOOL_SEARCH])
assert [tool["type"] for tool in lowered.tools if isinstance(tool, dict)] == ["function", "namespace"]
assert [member["type"] for member in _tool_named(lowered.tools, "calendar")["tools"]] == ["function"]
assert "vs_other_team" not in json.dumps(lowered.input)