diff --git a/litellm/_internal_context.py b/litellm/_internal_context.py index 389add8ed0f..105e0672d57 100644 --- a/litellm/_internal_context.py +++ b/litellm/_internal_context.py @@ -24,6 +24,23 @@ _billing_time: Final[ContextVar[datetime | None]] = ContextVar("billing_time", d _post_response: Final[ContextVar[bool]] = ContextVar("post_response", default=False) +_emulated_file_search: Final[ContextVar[bool]] = ContextVar("emulated_file_search", default=False) + + +@contextmanager +def emulated_file_search_phase() -> Generator[None]: + """Nested calls of emulated file_search, whose answer keeps only its own tool calls.""" + token: Final = _emulated_file_search.set(True) + try: + yield + finally: + _emulated_file_search.reset(token) + + +def in_emulated_file_search() -> bool: + return _emulated_file_search.get() + + @contextmanager def post_response_phase() -> Generator[None]: """Work the caller no longer waits for (success callbacks, response-cache writes), including tasks it spawns.""" diff --git a/litellm/responses/file_search/emulated_handler.py b/litellm/responses/file_search/emulated_handler.py index 887d1a9ff93..d437c2f2186 100644 --- a/litellm/responses/file_search/emulated_handler.py +++ b/litellm/responses/file_search/emulated_handler.py @@ -19,7 +19,7 @@ from typing import TYPE_CHECKING, Any, Final, TypeAlias, cast # noqa: TID251 # from typing_extensions import NotRequired, ReadOnly, TypedDict -from litellm._internal_context import is_internal_call +from litellm._internal_context import emulated_file_search_phase, is_internal_call from litellm._logging import verbose_logger from litellm.types.llms.openai import ResponseOutputItem, ResponsesAPIResponse from litellm.types.vector_stores import VectorStoreSearchResult @@ -537,15 +537,16 @@ async def aresponses_with_emulated_file_search( _prev_internal: Final = is_internal_call.get() is_internal_call.set(True) try: - first_response: Final[ResponsesAPIResponse] = cast( - ResponsesAPIResponse, - await _call_aresponses( - input=input, - model=model, - tools=transformed_tools or None, - **call_kwargs, - ), - ) + with emulated_file_search_phase(): + first_response: Final[ResponsesAPIResponse] = cast( + ResponsesAPIResponse, + await _call_aresponses( + input=input, + model=model, + tools=transformed_tools or None, + **call_kwargs, + ), + ) finally: is_internal_call.set(_prev_internal) @@ -601,15 +602,16 @@ async def aresponses_with_emulated_file_search( # Also an internal sub-call; billing is suppressed so the outer call fires once. is_internal_call.set(True) try: - final_response: Final[ResponsesAPIResponse] = cast( - ResponsesAPIResponse, - await _call_aresponses( - input=follow_up_input, - model=model, - tools=None, # no tools needed for the answer step - **call_kwargs, - ), - ) + with emulated_file_search_phase(): + final_response: Final[ResponsesAPIResponse] = cast( + ResponsesAPIResponse, + await _call_aresponses( + input=follow_up_input, + model=model, + tools=None, # no tools needed for the answer step + **call_kwargs, + ), + ) finally: is_internal_call.set(_prev_internal) diff --git a/litellm/responses/main.py b/litellm/responses/main.py index bb84066a6a7..1618c3e1269 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -13,6 +13,7 @@ from pydantic import BaseModel, TypeAdapter, ValidationError from typing_extensions import assert_never import litellm +from litellm._internal_context import in_emulated_file_search from litellm._logging import verbose_logger from litellm.completion_extras.litellm_responses_transformation.transformation import ( LiteLLMResponsesTransformationHandler, @@ -34,7 +35,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 declares_function, needs_tool_search_lowering +from litellm.responses.tool_search.lowering import needs_tool_search_lowering from litellm.responses.utils import ResponsesAPIRequestUtils from litellm.types.llms.openai import ( PromptObject, @@ -1147,13 +1148,6 @@ 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, @@ -1180,7 +1174,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) + or in_emulated_file_search() ): 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 4f6bfc01c86..c95bf1b8d9f 100644 --- a/litellm/responses/tool_search/lowering.py +++ b/litellm/responses/tool_search/lowering.py @@ -217,7 +217,7 @@ 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: +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 ())) @@ -227,7 +227,7 @@ def lower_tool_search_request( tool_choice: ToolChoice | None, ) -> ToolSearchLowering: declared: Final = tuple(tools or ()) - if _declares_client_tool_search(declared) and declares_function(declared, TOOL_SEARCH_FUNCTION_NAME): + 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 c898a247ff8..1b7a70b237b 100644 --- a/tests/unit/responses/tool_search/test_handler.py +++ b/tests/unit/responses/tool_search/test_handler.py @@ -200,6 +200,20 @@ async def test_emulated_file_search_leaves_client_tool_search_as_declared(): assert _types(upstream.bodies[0]["tools"]) == ["tool_search", "function"] +@pytest.mark.asyncio +async def test_a_client_function_named_like_the_file_search_emulation_keeps_tool_search_lowered(): + upstream: Final = _Upstream() + own_function: Final = { + "type": "function", + "name": "litellm_file_search", + "parameters": {"type": "object", "properties": {}}, + } + + await _hosted_vllm_call(upstream, tools=[CLIENT_TOOL_SEARCH, own_function]) + + assert [tool["name"] for tool in upstream.bodies[0]["tools"]] == ["tool_search", "litellm_file_search"] + + @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")