From 47214be317f45a2d2ac08258362e0481411eb427 Mon Sep 17 00:00:00 2001 From: user <70670632+stuxf@users.noreply.github.com> Date: Thu, 16 Apr 2026 00:45:57 +0000 Subject: [PATCH] fix(proxy): harden request parameter handling Tighten validation of request body parameters in the proxy routing layer. Use context variables for internal call state management instead of passing flags through request kwargs. Clean up metadata handling at the proxy boundary. --- litellm/_internal_context.py | 13 ++ litellm/proxy/auth/auth_utils.py | 2 +- litellm/proxy/litellm_pre_call_utils.py | 50 ++++-- .../responses/file_search/emulated_handler.py | 46 +++-- litellm/utils.py | 68 +++---- .../llms/test_file_search_responses.py | 167 +++++++++++------- 6 files changed, 220 insertions(+), 126 deletions(-) create mode 100644 litellm/_internal_context.py diff --git a/litellm/_internal_context.py b/litellm/_internal_context.py new file mode 100644 index 00000000000..727ca87d3e8 --- /dev/null +++ b/litellm/_internal_context.py @@ -0,0 +1,13 @@ +""" +Internal request context for LiteLLM. + +Provides a ContextVar-based mechanism for internal signals that must not +be settable from user input. Context variables are scoped to the current +asyncio task and cannot be injected via HTTP request bodies. +""" + +from contextvars import ContextVar + +# When True, suppresses async logging and billing for internal sub-calls +# (e.g., emulated file-search steps that make nested LLM calls). +is_internal_call: ContextVar[bool] = ContextVar("is_internal_call", default=False) diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index d2f8320668e..64766bbaadd 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -150,7 +150,7 @@ def is_request_body_safe( A malicious user can set the api_base to their own domain and invoke POST /chat/completions to intercept and steal the OpenAI API key. Relevant issue: https://huntr.com/bounties/4001e1a2-7b7a-4776-a3ae-e6692ec3d997 """ - banned_params = ["api_base", "base_url"] + banned_params = ["api_base", "base_url", "user_config"] for param in banned_params: if ( diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 4ec31925ea2..b0adf7aa6ee 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -43,6 +43,8 @@ def _sanitize_for_log(value: Any) -> str: text = repr(value) # Strip CR/LF characters commonly used for log injection return text.replace("\r", "").replace("\n", "") + + from litellm.router import Router from litellm.secret_managers.main import get_secret_bool from litellm.types.llms.anthropic import ANTHROPIC_API_HEADERS @@ -220,12 +222,12 @@ def _get_dynamic_logging_metadata( user_api_key_dict: UserAPIKeyAuth, proxy_config: ProxyConfig ) -> Optional[TeamCallbackMetadata]: callback_settings_obj: Optional[TeamCallbackMetadata] = None - key_dynamic_logging_settings: Optional[ - dict - ] = KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings(user_api_key_dict) - team_dynamic_logging_settings: Optional[ - dict - ] = KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings(user_api_key_dict) + key_dynamic_logging_settings: Optional[dict] = ( + KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings(user_api_key_dict) + ) + team_dynamic_logging_settings: Optional[dict] = ( + KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings(user_api_key_dict) + ) ######################################################################################### # Key-based callbacks ######################################################################################### @@ -779,11 +781,11 @@ class LiteLLMProxyRequestSetup: ## KEY-LEVEL SPEND LOGS / TAGS if "tags" in key_metadata and key_metadata["tags"] is not None: - data[_metadata_variable_name][ - "tags" - ] = LiteLLMProxyRequestSetup._merge_tags( - request_tags=data[_metadata_variable_name].get("tags"), - tags_to_add=key_metadata["tags"], + data[_metadata_variable_name]["tags"] = ( + LiteLLMProxyRequestSetup._merge_tags( + request_tags=data[_metadata_variable_name].get("tags"), + tags_to_add=key_metadata["tags"], + ) ) if "disable_global_guardrails" in key_metadata and isinstance( key_metadata["disable_global_guardrails"], bool @@ -918,6 +920,22 @@ async def add_litellm_data_to_request( # noqa: PLR0915 from litellm.proxy.proxy_server import llm_router, premium_user from litellm.types.proxy.litellm_pre_call_utils import RedactedDict, SecretFields + # Strip internal-only keys from user input before the proxy sets its own. + # These keys are injected by the proxy itself below — user-supplied values + # must not be trusted. + for _internal_key in ( + "proxy_server_request", + "standard_logging_object", + "secret_fields", + ): + data.pop(_internal_key, None) + # Strip spoofable auth metadata from user-supplied metadata dict + _user_metadata = data.get("metadata") + if isinstance(_user_metadata, dict): + for _mk in list(_user_metadata.keys()): + if _mk.startswith("user_api_key_"): + del _user_metadata[_mk] + _raw_headers: Dict[str, str] = RedactedDict(_safe_get_request_headers(request)) forward_llm_auth = False @@ -1079,9 +1097,9 @@ async def add_litellm_data_to_request( # noqa: PLR0915 data[_metadata_variable_name]["litellm_api_version"] = version if general_settings is not None: - data[_metadata_variable_name][ - "global_max_parallel_requests" - ] = general_settings.get("global_max_parallel_requests", None) + data[_metadata_variable_name]["global_max_parallel_requests"] = ( + general_settings.get("global_max_parallel_requests", None) + ) ### KEY-LEVEL Controls key_metadata = user_api_key_dict.metadata @@ -1881,7 +1899,9 @@ async def move_guardrails_to_metadata( ) # Only check policy engine if no local config (avoid import + registry lookup) - if not (has_key_config or has_team_config or has_project_config or has_request_config): + if not ( + has_key_config or has_team_config or has_project_config or has_request_config + ): from litellm.proxy.policy_engine.policy_registry import get_policy_registry if not get_policy_registry().is_initialized(): diff --git a/litellm/responses/file_search/emulated_handler.py b/litellm/responses/file_search/emulated_handler.py index 74a7d443c6c..532ac8f10d3 100644 --- a/litellm/responses/file_search/emulated_handler.py +++ b/litellm/responses/file_search/emulated_handler.py @@ -16,6 +16,7 @@ import time import uuid from typing import Any, Dict, Iterable, List, Optional, Tuple, Union, cast +from litellm._internal_context import is_internal_call from litellm._logging import verbose_logger from litellm.types.llms.openai import ResponseOutputItem, ResponsesAPIResponse from litellm.types.vector_stores import VectorStoreSearchResult @@ -449,15 +450,20 @@ async def aresponses_with_emulated_file_search( # 2. First provider call — provider will call the file_search function. # Mark as an internal sub-call so wrapper_async skips billing callbacks; # the parent litellm_logging_obj (propagated via kwargs) fires once at the end. - first_response: ResponsesAPIResponse = cast( - ResponsesAPIResponse, - await _call_aresponses( - input=input, - model=model, - tools=transformed_tools or None, - **{**kwargs, "_is_litellm_internal_call": True}, - ), - ) + _prev_internal = is_internal_call.get() + is_internal_call.set(True) + try: + first_response: ResponsesAPIResponse = cast( + ResponsesAPIResponse, + await _call_aresponses( + input=input, + model=model, + tools=transformed_tools or None, + **kwargs, + ), + ) + finally: + is_internal_call.set(_prev_internal) # 3. Look for a file_search function_call in the output file_search_calls = [ @@ -566,15 +572,19 @@ async def aresponses_with_emulated_file_search( # 6. Follow-up call — provider writes the final answer given search results. # Also an internal sub-call; billing is suppressed so the outer call fires once. - final_response: ResponsesAPIResponse = cast( - ResponsesAPIResponse, - await _call_aresponses( - input=follow_up_input, - model=model, - tools=None, # no tools needed for the answer step - **{**kwargs, "_is_litellm_internal_call": True}, - ), - ) + is_internal_call.set(True) + try: + final_response: ResponsesAPIResponse = cast( + ResponsesAPIResponse, + await _call_aresponses( + input=follow_up_input, + model=model, + tools=None, # no tools needed for the answer step + **kwargs, + ), + ) + finally: + is_internal_call.set(_prev_internal) # 7. Synthesize OpenAI-format output response_text = _extract_text_from_responses_output(final_response) diff --git a/litellm/utils.py b/litellm/utils.py index 09df88f0ceb..8d6ea5b7e9e 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -51,6 +51,7 @@ import litellm.litellm_core_utils # audio_utils.utils is lazy-loaded - only imported when needed for transcription calls import litellm.litellm_core_utils.json_validation_rule +from litellm._internal_context import is_internal_call from litellm._lazy_imports import ( _get_default_encoding, _get_modified_max_tokens, @@ -213,9 +214,11 @@ _CALL_TYPE_ENUM_MAP: dict = {ct.value: ct for ct in CallTypes} try: # Python 3.9+ - with resources.files("litellm.litellm_core_utils.tokenizers").joinpath( - "anthropic_tokenizer.json" - ).open("r", encoding="utf-8") as f: + with ( + resources.files("litellm.litellm_core_utils.tokenizers") + .joinpath("anthropic_tokenizer.json") + .open("r", encoding="utf-8") as f + ): json_data = json.load(f) except (ImportError, AttributeError, TypeError): with resources.open_text( @@ -781,9 +784,9 @@ def function_setup( # noqa: PLR0915 coroutine_checker = get_coroutine_checker_fn() ## DYNAMIC CALLBACKS ## - dynamic_callbacks: Optional[ - List[Union[str, Callable, "CustomLogger"]] - ] = kwargs.pop("callbacks", None) + dynamic_callbacks: Optional[List[Union[str, Callable, "CustomLogger"]]] = ( + kwargs.pop("callbacks", None) + ) all_callbacks = get_dynamic_callbacks(dynamic_callbacks=dynamic_callbacks) if len(all_callbacks) > 0: @@ -1689,9 +1692,9 @@ def client(original_function): # noqa: PLR0915 exception=e, retry_policy=kwargs.get("retry_policy"), ) - kwargs[ - "retry_policy" - ] = reset_retry_policy() # prevent infinite loops + kwargs["retry_policy"] = ( + reset_retry_policy() + ) # prevent infinite loops litellm.num_retries = ( None # set retries to None to prevent infinite loops ) @@ -1738,9 +1741,9 @@ def client(original_function): # noqa: PLR0915 exception=e, retry_policy=kwargs.get("retry_policy"), ) - kwargs[ - "retry_policy" - ] = reset_retry_policy() # prevent infinite loops + kwargs["retry_policy"] = ( + reset_retry_policy() + ) # prevent infinite loops litellm.num_retries = ( None # set retries to None to prevent infinite loops ) @@ -1792,7 +1795,8 @@ def client(original_function): # noqa: PLR0915 model: Optional[str] = args[0] if len(args) > 0 else kwargs.get("model", None) is_completion_with_fallbacks = kwargs.get("fallbacks") is not None - _is_litellm_internal_call = kwargs.pop("_is_litellm_internal_call", False) + kwargs.pop("_is_litellm_internal_call", None) # discard if injected + _is_litellm_internal_call = is_internal_call.get() try: if logging_obj is None: @@ -3774,10 +3778,10 @@ def pre_process_non_default_params( if "response_format" in non_default_params: if provider_config is not None: - non_default_params[ - "response_format" - ] = provider_config.get_json_schema_from_pydantic_object( - response_format=non_default_params["response_format"] + non_default_params["response_format"] = ( + provider_config.get_json_schema_from_pydantic_object( + response_format=non_default_params["response_format"] + ) ) else: non_default_params["response_format"] = type_to_response_format_param( @@ -3906,16 +3910,16 @@ def pre_process_optional_params( True # so that main.py adds the function call to the prompt ) if "tools" in non_default_params: - optional_params[ - "functions_unsupported_model" - ] = non_default_params.pop("tools") + optional_params["functions_unsupported_model"] = ( + non_default_params.pop("tools") + ) non_default_params.pop( "tool_choice", None ) # causes ollama requests to hang elif "functions" in non_default_params: - optional_params[ - "functions_unsupported_model" - ] = non_default_params.pop("functions") + optional_params["functions_unsupported_model"] = ( + non_default_params.pop("functions") + ) elif ( litellm.add_function_to_prompt ): # if user opts to add it to prompt instead @@ -4896,9 +4900,7 @@ def _get_order_filtered_deployments( ) -> List: if target_order is not None: filtered = [ - d - for d in healthy_deployments - if _get_deployment_order(d) == target_order + d for d in healthy_deployments if _get_deployment_order(d) == target_order ] if filtered: return filtered @@ -5878,8 +5880,12 @@ def _get_model_info_helper( # noqa: PLR0915 supports_web_search=_model_info.get("supports_web_search", None), supports_url_context=_model_info.get("supports_url_context", None), supports_reasoning=_model_info.get("supports_reasoning", None), - supports_none_reasoning_effort=_model_info.get("supports_none_reasoning_effort", None), - supports_xhigh_reasoning_effort=_model_info.get("supports_xhigh_reasoning_effort", None), + supports_none_reasoning_effort=_model_info.get( + "supports_none_reasoning_effort", None + ), + supports_xhigh_reasoning_effort=_model_info.get( + "supports_xhigh_reasoning_effort", None + ), supports_computer_use=_model_info.get("supports_computer_use", None), search_context_cost_per_query=_model_info.get( "search_context_cost_per_query", None @@ -7557,9 +7563,9 @@ class ModelResponseIterator: if convert_to_delta is True: _stream_response = ModelResponseStream() _stream_response.choices[0].delta.content = model_response.choices[0].message.content # type: ignore - self.model_response: Union[ - ModelResponse, ModelResponseStream - ] = _stream_response + self.model_response: Union[ModelResponse, ModelResponseStream] = ( + _stream_response + ) else: self.model_response = model_response self.is_done = False diff --git a/tests/test_litellm/llms/test_file_search_responses.py b/tests/test_litellm/llms/test_file_search_responses.py index 1864a296eb6..9943b456083 100644 --- a/tests/test_litellm/llms/test_file_search_responses.py +++ b/tests/test_litellm/llms/test_file_search_responses.py @@ -28,6 +28,7 @@ from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfi # Helpers # --------------------------------------------------------------------------- + def _make_unified_vs_id( unified_uuid: str = "abc-123", provider_resource_id: str = "vs_provider_native", @@ -61,6 +62,7 @@ def _code_interpreter_tool(file_ids: Optional[List[str]] = None) -> Dict[str, An # A-series: _decode_vector_store_ids_in_tools # --------------------------------------------------------------------------- + class TestDecodeVectorStoreIdsInTools: def test_A1_none_input_returns_none(self): assert _decode_vector_store_ids_in_tools(None) is None @@ -109,6 +111,7 @@ class TestDecodeVectorStoreIdsInTools: # B-series: update_responses_tools_with_model_file_ids # --------------------------------------------------------------------------- + class TestUpdateResponsesToolsWithModelFileIds: def test_B1_file_search_decode_runs_without_mapping(self): """Decode pass executes even when model_file_id_mapping is None.""" @@ -164,6 +167,7 @@ class TestUpdateResponsesToolsWithModelFileIds: # C/D-series: supports_native_file_search # --------------------------------------------------------------------------- + class TestSupportsNativeFileSearch: def test_C1_base_class_default_is_false(self): # Access the unbound method directly — no need to instantiate an abstract class @@ -177,6 +181,7 @@ class TestSupportsNativeFileSearch: # E-series: file_search guard in responses/main.py # --------------------------------------------------------------------------- + class TestFileSearchGuardInResponsesMain: """Tests for _has_file_search_tool helper and emulated routing guard.""" @@ -238,7 +243,9 @@ class TestFileSearchGuardInResponsesMain: "litellm.responses.main.ResponsesAPIRequestUtils.get_requested_response_api_optional_param", return_value={}, ), - patch("litellm.responses.main.run_async_function", return_value=expected) as run_async_mock, + patch( + "litellm.responses.main.run_async_function", return_value=expected + ) as run_async_mock, ): result = responses( input="hello", @@ -287,7 +294,9 @@ class TestFileSearchGuardInResponsesMain: "litellm.responses.main.ResponsesAPIRequestUtils.get_requested_response_api_optional_param", return_value={}, ), - patch("litellm.responses.main.run_async_function", return_value=expected) as run_async_mock, + patch( + "litellm.responses.main.run_async_function", return_value=expected + ) as run_async_mock, ): result = responses( input="hello", @@ -314,6 +323,7 @@ class TestFileSearchGuardInResponsesMain: # F-series: ManagedFiles hook — vector_store_ids access control # --------------------------------------------------------------------------- + class TestManagedFilesVectorStoreAccess: def _make_hook(self): """Return a ManagedFiles instance with prisma_client mocked.""" @@ -372,15 +382,20 @@ class TestManagedFilesVectorStoreAccess: mock_row = self._make_vs_row(vector_store_id="uuid-001", team_id="team-other") - async def mock_get_rows(uuids, prisma_client, user_api_key_cache, proxy_logging_obj=None): + async def mock_get_rows( + uuids, prisma_client, user_api_key_cache, proxy_logging_obj=None + ): return [mock_row] - with patch( - "litellm.proxy.proxy_server.prisma_client", - MagicMock(), - ), patch( - "litellm.proxy.auth.auth_checks.get_managed_vector_store_rows_by_uuids", - side_effect=mock_get_rows, + with ( + patch( + "litellm.proxy.proxy_server.prisma_client", + MagicMock(), + ), + patch( + "litellm.proxy.auth.auth_checks.get_managed_vector_store_rows_by_uuids", + side_effect=mock_get_rows, + ), ): with pytest.raises(HTTPException) as exc_info: await hook.check_vector_store_ids_access( @@ -396,15 +411,20 @@ class TestManagedFilesVectorStoreAccess: mock_row = self._make_vs_row(vector_store_id="uuid-002", team_id=None) - async def mock_get_rows(uuids, prisma_client, user_api_key_cache, proxy_logging_obj=None): + async def mock_get_rows( + uuids, prisma_client, user_api_key_cache, proxy_logging_obj=None + ): return [mock_row] - with patch( - "litellm.proxy.proxy_server.prisma_client", - MagicMock(), - ), patch( - "litellm.proxy.auth.auth_checks.get_managed_vector_store_rows_by_uuids", - side_effect=mock_get_rows, + with ( + patch( + "litellm.proxy.proxy_server.prisma_client", + MagicMock(), + ), + patch( + "litellm.proxy.auth.auth_checks.get_managed_vector_store_rows_by_uuids", + side_effect=mock_get_rows, + ), ): await hook.check_vector_store_ids_access( [unified_id], self._make_user(team_id="team-caller") @@ -415,7 +435,9 @@ class TestManagedFilesVectorStoreAccess: """Multiple unified IDs resolved in a single DB call (no N+1).""" hook = self._make_hook() ids = [ - _make_unified_vs_id(unified_uuid=f"uuid-{i}", provider_resource_id=f"vs_{i}") + _make_unified_vs_id( + unified_uuid=f"uuid-{i}", provider_resource_id=f"vs_{i}" + ) for i in range(3) ] @@ -426,18 +448,25 @@ class TestManagedFilesVectorStoreAccess: get_rows_mock = AsyncMock(return_value=rows) - with patch( - "litellm.proxy.proxy_server.prisma_client", - MagicMock(), - ), patch( - "litellm.proxy.auth.auth_checks.get_managed_vector_store_rows_by_uuids", - get_rows_mock, + with ( + patch( + "litellm.proxy.proxy_server.prisma_client", + MagicMock(), + ), + patch( + "litellm.proxy.auth.auth_checks.get_managed_vector_store_rows_by_uuids", + get_rows_mock, + ), ): await hook.check_vector_store_ids_access(ids, self._make_user("team-abc")) get_rows_mock.assert_called_once() call_args = get_rows_mock.call_args - assert set(call_args.kwargs["uuids"] or call_args.args[0]) == {"uuid-0", "uuid-1", "uuid-2"} + assert set(call_args.kwargs["uuids"] or call_args.args[0]) == { + "uuid-0", + "uuid-1", + "uuid-2", + } @pytest.mark.asyncio async def test_F6_non_responses_call_type_skipped(self): @@ -455,7 +484,9 @@ class TestManagedFilesVectorStoreAccess: await hook.async_pre_call_hook( user_api_key_dict=self._make_user(), cache=MagicMock(), - data={"tools": [{"type": "file_search", "vector_store_ids": ["vs_native"]}]}, + data={ + "tools": [{"type": "file_search", "vector_store_ids": ["vs_native"]}] + }, call_type=CallTypes.acompletion.value, ) hook.async_pre_call_hook.assert_called_once() @@ -465,6 +496,7 @@ class TestManagedFilesVectorStoreAccess: # G-series: get_vector_store_ids_from_file_search_tools helper # --------------------------------------------------------------------------- + class TestGetVectorStoreIdsFromFileSearchTools: def _make_hook(self): from enterprise.litellm_enterprise.proxy.hooks.managed_files import ( @@ -494,10 +526,12 @@ class TestGetVectorStoreIdsFromFileSearchTools: # Only the unified ID is included; native IDs are filtered assert result == [unified_id] + # --------------------------------------------------------------------------- # Phase 2: Emulated file_search handler # --------------------------------------------------------------------------- + class TestEmulatedFileSearchHandler: """Tests for litellm/responses/file_search/emulated_handler.py""" @@ -694,7 +728,9 @@ class TestEmulatedFileSearchHandler: r2.content = [{"type": "text", "text": "second hit"}] search_results = _build_search_results_for_include([r1, r2]) - assert len(search_results) == 2, "Both chunks should be returned, not deduplicated" + assert ( + len(search_results) == 2 + ), "Both chunks should be returned, not deduplicated" assert search_results[0]["text"] == "first hit" assert search_results[1]["text"] == "second hit" @@ -708,7 +744,9 @@ class TestEmulatedFileSearchHandler: ) first_resp = self._make_mock_responses_api_response(include_function_call=True) - final_resp = self._make_mock_responses_api_response(text="Deep research enables multi-step queries.") + final_resp = self._make_mock_responses_api_response( + text="Deep research enables multi-step queries." + ) search_result = MagicMock() search_result.file_id = "file-xyz" @@ -719,12 +757,15 @@ class TestEmulatedFileSearchHandler: mock_search_response = MagicMock() mock_search_response.data = [search_result] - with patch( - "litellm.responses.file_search.emulated_handler._call_aresponses", - new=AsyncMock(side_effect=[first_resp, final_resp]), - ), patch( - "litellm.vector_stores.main.asearch", - new=AsyncMock(return_value=mock_search_response), + with ( + patch( + "litellm.responses.file_search.emulated_handler._call_aresponses", + new=AsyncMock(side_effect=[first_resp, final_resp]), + ), + patch( + "litellm.vector_stores.main.asearch", + new=AsyncMock(return_value=mock_search_response), + ), ): result = await aresponses_with_emulated_file_search( input="What is deep research?", @@ -767,7 +808,9 @@ class TestEmulatedFileSearchHandler: first_resp_plural.model = "claude-3-5-sonnet" first_resp_plural.usage = None - final_resp = self._make_mock_responses_api_response(text="Deep research uses multiple queries.") + final_resp = self._make_mock_responses_api_response( + text="Deep research uses multiple queries." + ) search_result = MagicMock() search_result.file_id = "file-multi" @@ -777,12 +820,15 @@ class TestEmulatedFileSearchHandler: mock_search_response = MagicMock() mock_search_response.data = [search_result] - with patch( - "litellm.responses.file_search.emulated_handler._call_aresponses", - new=AsyncMock(side_effect=[first_resp_plural, final_resp]), - ), patch( - "litellm.vector_stores.main.asearch", - new=AsyncMock(return_value=mock_search_response), + with ( + patch( + "litellm.responses.file_search.emulated_handler._call_aresponses", + new=AsyncMock(side_effect=[first_resp_plural, final_resp]), + ), + patch( + "litellm.vector_stores.main.asearch", + new=AsyncMock(return_value=mock_search_response), + ), ): result = await aresponses_with_emulated_file_search( input="What is deep research?", @@ -805,7 +851,9 @@ class TestEmulatedFileSearchHandler: aresponses_with_emulated_file_search, ) - direct_resp = self._make_mock_responses_api_response(text="I already know the answer.") + direct_resp = self._make_mock_responses_api_response( + text="I already know the answer." + ) with patch( "litellm.responses.file_search.emulated_handler._call_aresponses", @@ -835,11 +883,12 @@ class TestEmulatedFileSearchHandler: @pytest.mark.asyncio async def test_H15_sub_calls_carry_internal_call_flag(self): - """Both internal aresponses sub-calls receive _is_litellm_internal_call=True. + """Both internal aresponses sub-calls run with is_internal_call context var True. This ensures wrapper_async skips success/failure callbacks for sub-calls so billing fires exactly once (on the outer call) with the synthesized result. """ + from litellm._internal_context import is_internal_call from litellm.responses.file_search.emulated_handler import ( aresponses_with_emulated_file_search, ) @@ -855,25 +904,21 @@ class TestEmulatedFileSearchHandler: mock_search_response = MagicMock() mock_search_response.data = [search_result] - captured_kwargs: list = [] - - async def _capture(*args, **kwargs): - captured_kwargs.append(dict(kwargs)) - return captured_kwargs.__len__() == 1 and first_resp or final_resp - - with patch( - "litellm.responses.file_search.emulated_handler._call_aresponses", - new=AsyncMock(side_effect=[first_resp, final_resp]), - ) as mock_call, patch( - "litellm.vector_stores.main.asearch", - new=AsyncMock(return_value=mock_search_response), + with ( + patch( + "litellm.responses.file_search.emulated_handler._call_aresponses", + new=AsyncMock(side_effect=[first_resp, final_resp]), + ) as mock_call, + patch( + "litellm.vector_stores.main.asearch", + new=AsyncMock(return_value=mock_search_response), + ), ): - # Intercept kwargs before the mock returns + captured_ctx: list = [] original_side_effect = [first_resp, final_resp] - call_kwargs: list = [] async def _intercept(**kwargs): # type: ignore[misc] - call_kwargs.append(dict(kwargs)) + captured_ctx.append(is_internal_call.get()) return original_side_effect.pop(0) mock_call.side_effect = _intercept @@ -884,9 +929,9 @@ class TestEmulatedFileSearchHandler: tools=[{"type": "file_search", "vector_store_ids": ["vs_h15"]}], ) - assert len(call_kwargs) == 2, "Expected exactly 2 sub-calls" - for i, kw in enumerate(call_kwargs): - assert kw.get("_is_litellm_internal_call") is True, ( - f"Sub-call {i} must carry _is_litellm_internal_call=True to suppress " + assert len(captured_ctx) == 2, "Expected exactly 2 sub-calls" + for i, ctx_val in enumerate(captured_ctx): + assert ctx_val is True, ( + f"Sub-call {i} must run with is_internal_call=True to suppress " "billing callbacks in wrapper_async" )