From e9c34d6d36904cc29ca45e131efe3513f707ac3e Mon Sep 17 00:00:00 2001 From: Id545 Date: Fri, 2 Oct 2026 22:07:15 +0200 Subject: [PATCH] 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 --- litellm/responses/main.py | 10 +++++- litellm/responses/tool_search/lowering.py | 34 ++++++++++++++---- .../responses/tool_search/test_handler.py | 16 +++++++++ .../responses/tool_search/test_lowering.py | 35 +++++++++++++++++++ 4 files changed, 88 insertions(+), 7 deletions(-) diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 9459aba6661..bb84066a6a7 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -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 ( diff --git a/litellm/responses/tool_search/lowering.py b/litellm/responses/tool_search/lowering.py index 59ef2bda084..4f6bfc01c86 100644 --- a/litellm/responses/tool_search/lowering.py +++ b/litellm/responses/tool_search/lowering.py @@ -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))) diff --git a/tests/unit/responses/tool_search/test_handler.py b/tests/unit/responses/tool_search/test_handler.py index 93b6398ebe0..c898a247ff8 100644 --- a/tests/unit/responses/tool_search/test_handler.py +++ b/tests/unit/responses/tool_search/test_handler.py @@ -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") diff --git a/tests/unit/responses/tool_search/test_lowering.py b/tests/unit/responses/tool_search/test_lowering.py index cf0d22cbc89..095bddf6daf 100644 --- a/tests/unit/responses/tool_search/test_lowering.py +++ b/tests/unit/responses/tool_search/test_lowering.py @@ -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)