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)