diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 60100e8c2fdd..00c67e9f0fb3 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -91,6 +91,7 @@ async def try_short_circuit_search( messages: List[Dict], tools: Optional[List[Dict]], custom_llm_provider: Optional[str], + kwargs: Optional[dict[str, Any]] = None, ) -> Optional[Dict[str, Any]]: """ Short-circuit web-search-only requests by executing the search directly. @@ -176,7 +177,10 @@ async def try_short_circuit_search( # Execute search — keep the structured SearchResponse so the native # block can carry per-result url/title/page_age. try: - search_result_text, structured = await self._execute_search(query) + if kwargs is None: + search_result_text, structured = await self._execute_search(query) + else: + search_result_text, structured = await self._execute_search(query, kwargs=kwargs) except Exception as e: verbose_logger.error(f"WebSearchInterception: Short-circuit search failed: {e}") search_result_text, structured = f"Search failed: {e}", None @@ -936,7 +940,7 @@ async def _build_anthropic_request_patch( query = tool_call["input"].get("query") if query: verbose_logger.debug(f"WebSearchInterception: Queuing search for query='{query}'") - search_tasks.append(self._execute_search(query)) + search_tasks.append(self._execute_search(query, kwargs=kwargs)) else: verbose_logger.debug(f"WebSearchInterception: Tool call {tool_call['id']} has no query") # Add empty result for tools without query @@ -1009,7 +1013,9 @@ async def _build_anthropic_request_patch( ) return patch, structured_results - async def _execute_search(self, query: str) -> Tuple[str, Optional[SearchResponse]]: + async def _execute_search( + self, query: str, kwargs: Optional[dict[str, Any]] = None + ) -> Tuple[str, Optional[SearchResponse]]: """ Execute a single web search using router's search tools. @@ -1031,36 +1037,13 @@ async def _execute_search(self, query: str) -> Tuple[str, Optional[SearchRespons ) llm_router = None - # Determine search provider from router's search_tools + search_tool = self._select_search_tool_from_router(llm_router=llm_router) search_provider: Optional[str] = None - if llm_router is not None and hasattr(llm_router, "search_tools"): - if self.search_tool_name: - # Find specific search tool by name - matching_tools = [ - tool - for tool in llm_router.search_tools - if tool.get("search_tool_name") == self.search_tool_name - ] - if matching_tools: - search_tool = matching_tools[0] - search_provider = search_tool.get("litellm_params", {}).get("search_provider") - verbose_logger.debug( - f"WebSearchInterception: Found search tool '{self.search_tool_name}' " - f"with provider '{search_provider}'" - ) - else: - verbose_logger.debug( - f"WebSearchInterception: Search tool '{self.search_tool_name}' not found in router, " - "falling back to first available or perplexity" - ) - - # If no specific tool or not found, use first available - if not search_provider and llm_router.search_tools: - first_tool = llm_router.search_tools[0] - search_provider = first_tool.get("litellm_params", {}).get("search_provider") - verbose_logger.debug( - f"WebSearchInterception: Using first available search tool with provider '{search_provider}'" - ) + search_litellm_params: dict[str, Any] = {} + if search_tool is not None: + await self._authorize_search_tool(search_tool=search_tool, kwargs=kwargs) + search_litellm_params = dict(search_tool.get("litellm_params", {}) or {}) + search_provider = search_litellm_params.get("search_provider") # Fallback to perplexity if no router or no search tools configured if not search_provider: @@ -1073,7 +1056,12 @@ async def _execute_search(self, query: str) -> Tuple[str, Optional[SearchRespons verbose_logger.debug( f"WebSearchInterception: Executing search for '{query}' using provider '{search_provider}'" ) - result = await litellm.asearch(query=query, search_provider=search_provider) + search_kwargs = { + key: value + for key, value in search_litellm_params.items() + if key != "search_provider" and value is not None + } + result = await litellm.asearch(query=query, search_provider=search_provider, **search_kwargs) # Format using transformation function search_result_text = WebSearchTransformation.format_search_response(result) @@ -1086,6 +1074,107 @@ async def _execute_search(self, query: str) -> Tuple[str, Optional[SearchRespons verbose_logger.error(f"WebSearchInterception: Search failed for '{query}': {str(e)}") raise + async def _authorize_search_tool( + self, + search_tool: dict[str, Any], + kwargs: Optional[dict[str, Any]], + ) -> None: + search_tool_name = search_tool.get("search_tool_name") + if not isinstance(search_tool_name, str) or not search_tool_name: + return + + user_api_key_auth = self._get_user_api_key_auth_from_kwargs(kwargs) + if user_api_key_auth is None: + return + + from litellm.proxy.auth.auth_checks import ( + can_key_call_search_tool, + can_team_call_search_tool, + get_team_object, + ) + + await can_key_call_search_tool( + search_tool_name=search_tool_name, + valid_token=user_api_key_auth, + ) + + team_id = getattr(user_api_key_auth, "team_id", None) + if team_id: + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) + + team_object = await get_team_object( + team_id=team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=getattr(user_api_key_auth, "parent_otel_span", None), + proxy_logging_obj=proxy_logging_obj, + ) + await can_team_call_search_tool( + search_tool_name=search_tool_name, + team_object=team_object, + ) + + @staticmethod + def _get_user_api_key_auth_from_kwargs(kwargs: Optional[dict[str, Any]]) -> Any: + if not kwargs: + return None + + for metadata_key in ("metadata", "litellm_metadata"): + metadata = kwargs.get(metadata_key) + if isinstance(metadata, dict) and metadata.get("user_api_key_auth") is not None: + return metadata["user_api_key_auth"] + + litellm_params = kwargs.get("litellm_params") + if not isinstance(litellm_params, dict): + return None + + for metadata_key in ("metadata", "litellm_metadata"): + metadata = litellm_params.get(metadata_key) + if isinstance(metadata, dict) and metadata.get("user_api_key_auth") is not None: + return metadata["user_api_key_auth"] + + return None + + def _select_search_tool_from_router(self, llm_router: Any) -> Optional[dict[str, Any]]: + if llm_router is None or not hasattr(llm_router, "search_tools"): + return None + search_tools = list(getattr(llm_router, "search_tools") or []) + return self._select_search_tool_from_list(search_tools=search_tools, source="router") + + def _select_search_tool_from_list( + self, + search_tools: list[dict[str, Any]], + source: str, + ) -> Optional[dict[str, Any]]: + if self.search_tool_name: + matching_tools = [tool for tool in search_tools if tool.get("search_tool_name") == self.search_tool_name] + if matching_tools: + search_provider = (matching_tools[0].get("litellm_params", {}) or {}).get("search_provider") + verbose_logger.debug( + f"WebSearchInterception: Found search tool '{self.search_tool_name}' " + f"from {source} with provider '{search_provider}'" + ) + return matching_tools[0] + verbose_logger.debug( + f"WebSearchInterception: Search tool '{self.search_tool_name}' not found in {source}, " + "falling back to first available or perplexity" + ) + + if search_tools: + first_tool = search_tools[0] + search_provider = (first_tool.get("litellm_params", {}) or {}).get("search_provider") + verbose_logger.debug( + f"WebSearchInterception: Using first available search tool from {source} " + f"with provider '{search_provider}'" + ) + return first_tool + + return None + async def _execute_chat_completion_agentic_loop( self, model: str, @@ -1145,7 +1234,7 @@ async def _build_chat_completion_request_patch( if query: verbose_logger.debug(f"WebSearchInterception: Queuing search for query='{query}'") - search_tasks.append(self._execute_search(query)) + search_tasks.append(self._execute_search(query, kwargs=kwargs)) else: verbose_logger.debug(f"WebSearchInterception: Tool call {tool_call.get('id')} has no query") # Add empty result for tools without query diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py index effd7dda6a06..dd983f0c344d 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py @@ -148,6 +148,7 @@ async def _try_websearch_short_circuit( tools: Optional[List[Dict]], custom_llm_provider: Optional[str], stream: Optional[bool], + kwargs: Optional[dict] = None, ) -> Optional[Union[AnthropicMessagesResponse, AsyncIterator]]: """ Attempt to short-circuit a web-search-only request. @@ -177,6 +178,7 @@ async def _try_websearch_short_circuit( messages=messages, tools=tools, custom_llm_provider=custom_llm_provider, + kwargs=kwargs, ) if response is not None: anthropic_response = cast(AnthropicMessagesResponse, response) @@ -292,6 +294,7 @@ async def anthropic_messages( tools=tools, custom_llm_provider=custom_llm_provider, stream=original_stream, + kwargs={**kwargs, "metadata": metadata}, ) if short_circuit_response is not None: return short_circuit_response diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 9f4c610dc478..f6e5c1180855 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6382,7 +6382,6 @@ async def _init_agents_in_db(self, prisma_client: PrismaClient): async def _init_search_tools_in_db(self, prisma_client: PrismaClient): """ Initialize search tools from database into the router on startup. - Only updates router if there are tools in the database, otherwise preserves config-loaded tools. """ global llm_router @@ -6392,26 +6391,29 @@ async def _init_search_tools_in_db(self, prisma_client: PrismaClient): from litellm.router_utils.search_api_router import SearchAPIRouter try: - search_tools = await SearchToolRegistry.get_all_search_tools_from_db(prisma_client=prisma_client) + db_search_tools = await SearchToolRegistry.get_all_search_tools_from_db(prisma_client=prisma_client) - verbose_proxy_logger.info(f"Loading {len(search_tools)} search tool(s) from database into router") + parsed_tools = self.parse_search_tools(self.get_config_state()) + config_search_tools = parsed_tools or [] - # Only update router if there are tools in the database - # This prevents overwriting config-loaded tools with an empty list - if len(search_tools) > 0: - if llm_router is not None: - # Add search tools to the router - await SearchAPIRouter.update_router_search_tools( - router_instance=llm_router, search_tools=search_tools - ) - verbose_proxy_logger.info(f"Successfully loaded {len(search_tools)} search tool(s) into router") - else: - verbose_proxy_logger.debug( - "Router not initialized yet, search tools will be added when router is created" - ) + search_tools = self._merge_config_and_db_search_tools( + config_search_tools=config_search_tools, + db_search_tools=[dict(tool) for tool in db_search_tools], + ) + + verbose_proxy_logger.info( + f"Loading {len(search_tools)} search tool(s) into router " + f"({len(config_search_tools)} from config, {len(db_search_tools)} from database)" + ) + + if llm_router is not None and search_tools: + await SearchAPIRouter.update_router_search_tools(router_instance=llm_router, search_tools=search_tools) + verbose_proxy_logger.info(f"Successfully loaded {len(search_tools)} search tool(s) into router") + elif llm_router is not None: + verbose_proxy_logger.debug("No search tools found in config or database, skipping router update") else: verbose_proxy_logger.debug( - "No search tools found in database, keeping config-loaded search tools (if any)" + "Router not initialized yet, search tools will be added when router is created" ) except Exception as e: @@ -6419,6 +6421,21 @@ async def _init_search_tools_in_db(self, prisma_client: PrismaClient): "litellm.proxy.proxy_server.py::ProxyConfig:_init_search_tools_in_db - {}".format(str(e)) ) + @staticmethod + def _merge_config_and_db_search_tools( + config_search_tools: list[SearchToolTypedDict], + db_search_tools: list[dict[str, Any]], + ) -> list[dict[str, Any]]: + db_tool_names = {tool.get("search_tool_name") for tool in db_search_tools} + return [ + *[ + dict(config_search_tool) + for config_search_tool in config_search_tools + if config_search_tool.get("search_tool_name") not in db_tool_names + ], + *db_search_tools, + ] + async def _init_pass_through_endpoints_in_db(self): from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( initialize_pass_through_endpoints_in_db, diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py b/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py index ea117380edf5..b6ff3b70a4d3 100644 --- a/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py +++ b/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py @@ -12,6 +12,8 @@ from litellm.integrations.websearch_interception.handler import ( WebSearchInterceptionLogger, ) +from litellm.llms.base_llm.search.transformation import SearchResponse +from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, ProxyException, UserAPIKeyAuth from litellm.types.utils import LlmProviders @@ -56,9 +58,7 @@ def test_initialize_from_proxy_config_honors_dict_callback_specific_params(): """A valid dict under callback_settings.websearch_interception is applied.""" logger = WebSearchInterceptionLogger.initialize_from_proxy_config( litellm_settings={}, - callback_specific_params={ - "websearch_interception": {"search_tool_name": "ws-tool"} - }, + callback_specific_params={"websearch_interception": {"search_tool_name": "ws-tool"}}, ) assert logger.search_tool_name == "ws-tool" @@ -119,9 +119,7 @@ async def test_async_build_agentic_loop_plan_returns_request_patch(): "response_format": "anthropic", } logging_obj = MagicMock() - logging_obj.model_call_details = { - "agentic_loop_params": {"model": "bedrock/invoke/claude-3-5-sonnet"} - } + logging_obj.model_call_details = {"agentic_loop_params": {"model": "bedrock/invoke/claude-3-5-sonnet"}} kwargs = { "temperature": 0.2, "_websearch_interception_converted_stream": True, @@ -162,8 +160,6 @@ async def test_internal_flags_filtered_from_followup_kwargs(): to the follow-up LLM request, causing "Extra inputs are not permitted" errors from providers like Bedrock that use strict parameter validation. """ - logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"]) - # Simulate kwargs that would be passed during agentic loop execution kwargs_with_internal_flags = { "_websearch_interception_converted_stream": True, @@ -174,9 +170,7 @@ async def test_internal_flags_filtered_from_followup_kwargs(): # Apply the same filtering logic used in _execute_agentic_loop kwargs_for_followup = { - k: v - for k, v in kwargs_with_internal_flags.items() - if not k.startswith("_websearch_interception") + k: v for k, v in kwargs_with_internal_flags.items() if not k.startswith("_websearch_interception") } # Verify internal flags are filtered out @@ -188,6 +182,138 @@ async def test_internal_flags_filtered_from_followup_kwargs(): assert kwargs_for_followup["max_tokens"] == 1024 +@pytest.mark.asyncio +async def test_execute_search_passes_selected_search_tool_litellm_params(monkeypatch): + import litellm + from litellm.proxy import proxy_server + + logger = WebSearchInterceptionLogger( + enabled_providers=["bedrock"], + search_tool_name="ui-tavily", + ) + router = MagicMock() + router.search_tools = [ + { + "search_tool_name": "ui-tavily", + "litellm_params": { + "search_provider": "tavily", + "api_key": "fake-ui-key", + "api_base": "https://api.tavily.com", + "timeout": 10.0, + "max_retries": 2, + "country": None, + }, + } + ] + mock_asearch = AsyncMock(return_value=SearchResponse(object="search", results=[])) + user_api_key_auth = UserAPIKeyAuth( + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="op-allowed-search", + search_tools=["ui-tavily"], + ) + ) + + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(litellm, "asearch", mock_asearch) + + await logger._execute_search( + "what is litellm", + kwargs={"litellm_params": {"metadata": {"user_api_key_auth": user_api_key_auth}}}, + ) + + mock_asearch.assert_awaited_once_with( + query="what is litellm", + search_provider="tavily", + api_key="fake-ui-key", + api_base="https://api.tavily.com", + timeout=10.0, + max_retries=2, + ) + + +@pytest.mark.asyncio +async def test_execute_search_enforces_key_search_tool_permission(monkeypatch): + import litellm + from litellm.proxy import proxy_server + + logger = WebSearchInterceptionLogger( + enabled_providers=["bedrock"], + search_tool_name="blocked-search", + ) + router = MagicMock() + router.search_tools = [ + { + "search_tool_name": "blocked-search", + "litellm_params": { + "search_provider": "tavily", + "api_key": "fake-ui-key", + }, + } + ] + mock_asearch = AsyncMock(return_value=SearchResponse(object="search", results=[])) + user_api_key_auth = UserAPIKeyAuth( + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="op-key-search", + search_tools=["allowed-search"], + ) + ) + + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(litellm, "asearch", mock_asearch) + + with pytest.raises(ProxyException): + await logger._execute_search( + "what is litellm", + kwargs={"metadata": {"user_api_key_auth": user_api_key_auth}}, + ) + + mock_asearch.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_execute_search_enforces_team_search_tool_permission(monkeypatch): + import litellm + from litellm.proxy import proxy_server + + logger = WebSearchInterceptionLogger( + enabled_providers=["bedrock"], + search_tool_name="blocked-search", + ) + router = MagicMock() + router.search_tools = [ + { + "search_tool_name": "blocked-search", + "litellm_params": { + "search_provider": "tavily", + "api_key": "fake-ui-key", + }, + } + ] + mock_asearch = AsyncMock(return_value=SearchResponse(object="search", results=[])) + team_key_auth = UserAPIKeyAuth(team_id="team-1") + team_object = LiteLLM_TeamTable( + team_id="team-1", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="op-team-search", + search_tools=["allowed-search"], + ), + ) + mock_get_team_object = AsyncMock(return_value=team_object) + + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(litellm, "asearch", mock_asearch) + monkeypatch.setattr("litellm.proxy.auth.auth_checks.get_team_object", mock_get_team_object) + + with pytest.raises(ProxyException): + await logger._execute_search( + "what is litellm", + kwargs={"metadata": {"user_api_key_auth": team_key_auth}}, + ) + + mock_get_team_object.assert_awaited_once() + mock_asearch.assert_not_awaited() + + @pytest.mark.asyncio async def test_async_pre_call_deployment_hook_provider_from_top_level_kwargs(): """Test that async_pre_call_deployment_hook finds custom_llm_provider at top-level kwargs. @@ -216,15 +342,12 @@ async def test_async_pre_call_deployment_hook_provider_from_top_level_kwargs(): assert result is not None # The web_search tool should be converted to litellm_web_search (OpenAI format) assert any( - t.get("type") == "function" - and t.get("function", {}).get("name") == "litellm_web_search" + t.get("type") == "function" and t.get("function", {}).get("name") == "litellm_web_search" for t in result["tools"] ) # The non-web-search tool should be preserved assert any( - t.get("type") == "function" - and t.get("function", {}).get("name") == "other_tool" - for t in result["tools"] + t.get("type") == "function" and t.get("function", {}).get("name") == "other_tool" for t in result["tools"] ) @@ -261,8 +384,7 @@ async def test_async_pre_call_deployment_hook_returns_full_kwargs(): assert result["custom_llm_provider"] == "openai" # Tools should be converted assert any( - t.get("type") == "function" - and t.get("function", {}).get("name") == "litellm_web_search" + t.get("type") == "function" and t.get("function", {}).get("name") == "litellm_web_search" for t in result["tools"] ) @@ -323,8 +445,7 @@ async def test_async_pre_call_deployment_hook_nested_litellm_params_fallback(): assert result is not None assert any( - t.get("type") == "function" - and t.get("function", {}).get("name") == "litellm_web_search" + t.get("type") == "function" and t.get("function", {}).get("name") == "litellm_web_search" for t in result["tools"] ) # Full kwargs preserved @@ -357,8 +478,7 @@ async def test_async_pre_call_deployment_hook_provider_derived_from_model_name() # Should NOT be None — the hook should derive "openai" from "openai/gpt-4o-mini" assert result is not None assert any( - t.get("type") == "function" - and t.get("function", {}).get("name") == "litellm_web_search" + t.get("type") == "function" and t.get("function", {}).get("name") == "litellm_web_search" for t in result["tools"] ) # Full kwargs preserved @@ -478,9 +598,7 @@ def test_sync_forced_tool_choice_leaves_non_forced_untouched(tool_choice): and None pass through unchanged.""" converted_tools = [{"name": LITELLM_WEB_SEARCH_TOOL_NAME}] - result = WebSearchInterceptionLogger._sync_forced_tool_choice( - tool_choice, converted_tools - ) + result = WebSearchInterceptionLogger._sync_forced_tool_choice(tool_choice, converted_tools) assert result == tool_choice diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index e7c036dd039d..7550432fcc06 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -189,9 +189,7 @@ def test_ProxyConfig__load_yaml_file_raises_on_missing_file(): @pytest.mark.asyncio async def test_ProxyConfig__get_config_from_file_loads_yaml(tmp_path): f = tmp_path / "c.yaml" - f.write_text( - "model_list: []\ngeneral_settings: {}\nlitellm_settings:\n drop_params: true\n" - ) + f.write_text("model_list: []\ngeneral_settings: {}\nlitellm_settings:\n drop_params: true\n") pc = ProxyConfig() result = await pc._get_config_from_file(config_file_path=str(f)) assert result == { @@ -534,6 +532,139 @@ def test_ProxyConfig_parse_search_tools_missing_returns_none(): assert pc.parse_search_tools({}) is None +def test_ProxyConfig_merge_config_and_db_search_tools_returns_superset(): + config_tools = [ + { + "search_tool_name": "config-search", + "litellm_params": {"search_provider": "tavily"}, + } + ] + db_tools = [ + { + "search_tool_name": "db-search", + "litellm_params": { + "search_provider": "exa_ai", + "api_key": "fake-db-key", + }, + } + ] + + merged = ProxyConfig._merge_config_and_db_search_tools( + config_search_tools=config_tools, + db_search_tools=db_tools, + ) + + assert [tool["search_tool_name"] for tool in merged] == ["config-search", "db-search"] + assert merged[1]["litellm_params"]["api_key"] == "fake-db-key" + + +def test_ProxyConfig_merge_config_and_db_search_tools_prefers_db_duplicate(): + config_tools = [ + { + "search_tool_name": "shared-search", + "litellm_params": {"search_provider": "tavily"}, + }, + { + "search_tool_name": "config-only", + "litellm_params": {"search_provider": "perplexity"}, + }, + ] + db_tools = [ + { + "search_tool_name": "shared-search", + "litellm_params": { + "search_provider": "exa_ai", + "api_key": "fake-db-key", + }, + } + ] + + merged = ProxyConfig._merge_config_and_db_search_tools( + config_search_tools=config_tools, + db_search_tools=db_tools, + ) + + assert [tool["search_tool_name"] for tool in merged] == ["config-only", "shared-search"] + assert merged[1]["litellm_params"]["search_provider"] == "exa_ai" + assert merged[1]["litellm_params"]["api_key"] == "fake-db-key" + + +@pytest.mark.asyncio +async def test_ProxyConfig__init_search_tools_in_db_loads_merged_tools(monkeypatch): + from litellm.proxy import proxy_server + from litellm.router_utils.search_api_router import SearchAPIRouter + + pc = ProxyConfig() + pc.update_config_state( + { + "search_tools": [ + { + "search_tool_name": "shared-search", + "litellm_params": {"search_provider": "tavily"}, + }, + { + "search_tool_name": "config-only", + "litellm_params": {"search_provider": "perplexity"}, + }, + ] + } + ) + db_tools = [ + { + "search_tool_name": "shared-search", + "litellm_params": { + "search_provider": "exa_ai", + "api_key": "fake-db-key", + }, + } + ] + fake_router = MagicMock() + mock_get_db_tools = AsyncMock(return_value=db_tools) + mock_update_router = AsyncMock() + + monkeypatch.setattr(proxy_server, "llm_router", fake_router) + monkeypatch.setattr( + "litellm.proxy.search_endpoints.search_tool_registry.SearchToolRegistry.get_all_search_tools_from_db", + mock_get_db_tools, + ) + monkeypatch.setattr(SearchAPIRouter, "update_router_search_tools", mock_update_router) + + await pc._init_search_tools_in_db(prisma_client=MagicMock()) + + mock_get_db_tools.assert_awaited_once() + mock_update_router.assert_awaited_once() + update_kwargs = mock_update_router.await_args.kwargs + assert update_kwargs["router_instance"] is fake_router + assert [tool["search_tool_name"] for tool in update_kwargs["search_tools"]] == [ + "config-only", + "shared-search", + ] + assert update_kwargs["search_tools"][1]["litellm_params"]["api_key"] == "fake-db-key" + + +@pytest.mark.asyncio +async def test_ProxyConfig__init_search_tools_in_db_skips_empty_router_update(monkeypatch): + from litellm.proxy import proxy_server + from litellm.router_utils.search_api_router import SearchAPIRouter + + pc = ProxyConfig() + pc.update_config_state({}) + mock_get_db_tools = AsyncMock(return_value=[]) + mock_update_router = AsyncMock() + + monkeypatch.setattr(proxy_server, "llm_router", MagicMock()) + monkeypatch.setattr( + "litellm.proxy.search_endpoints.search_tool_registry.SearchToolRegistry.get_all_search_tools_from_db", + mock_get_db_tools, + ) + monkeypatch.setattr(SearchAPIRouter, "update_router_search_tools", mock_update_router) + + await pc._init_search_tools_in_db(prisma_client=MagicMock()) + + mock_get_db_tools.assert_awaited_once() + mock_update_router.assert_not_awaited() + + # --------------------------------------------------------------------------- # ProxyConfig._load_environment_variables # --------------------------------------------------------------------------- @@ -542,9 +673,7 @@ def test_ProxyConfig_parse_search_tools_missing_returns_none(): def test_ProxyConfig__load_environment_variables_sets_env(monkeypatch): monkeypatch.delenv("TEST_LOAD_ENV_X", raising=False) pc = ProxyConfig() - pc._load_environment_variables( - {"environment_variables": {"TEST_LOAD_ENV_X": "hello"}} - ) + pc._load_environment_variables({"environment_variables": {"TEST_LOAD_ENV_X": "hello"}}) result = { "TEST_LOAD_ENV_X": os.environ.get("TEST_LOAD_ENV_X"), "set": True, @@ -602,9 +731,7 @@ async def test_ProxyConfig_load_config_missing_file_raises(monkeypatch): @pytest.mark.asyncio -async def test_ProxyConfig_load_config_forwards_callback_specific_params( - tmp_path, monkeypatch -): +async def test_ProxyConfig_load_config_forwards_callback_specific_params(tmp_path, monkeypatch): """Regression: callback_settings from config must be forwarded to initialize_callbacks_on_proxy as callback_specific_params. @@ -645,16 +772,12 @@ def _fake_initialize_callbacks_on_proxy(**kwargs): # The callbacks branch must forward the loaded callback_settings. assert captured.get("callback_specific_params") == { - "datadog_cost_management": { - "cost_tag_keys": ["capability", "platform", "ai_product"] - } + "datadog_cost_management": {"cost_tag_keys": ["capability", "platform", "ai_product"]} } @pytest.mark.asyncio -async def test_ProxyConfig_load_config_blank_callback_settings_does_not_crash( - tmp_path, monkeypatch -): +async def test_ProxyConfig_load_config_blank_callback_settings_does_not_crash(tmp_path, monkeypatch): """Regression: `callback_settings:` with no body loads as None because dict.get() only falls back to the default when the key is absent. The None was forwarded verbatim to initialize_callbacks_on_proxy, where the first @@ -678,17 +801,13 @@ async def test_ProxyConfig_load_config_blank_callback_settings_does_not_crash( CompressionInterceptionLogger, ) - original_callbacks = ( - list(litellm.callbacks) if isinstance(litellm.callbacks, list) else [] - ) + original_callbacks = list(litellm.callbacks) if isinstance(litellm.callbacks, list) else [] litellm.callbacks = [] try: pc = ProxyConfig() await pc.load_config(router=None, config_file_path=str(f)) - assert any( - isinstance(c, CompressionInterceptionLogger) for c in litellm.callbacks - ) + assert any(isinstance(c, CompressionInterceptionLogger) for c in litellm.callbacks) finally: litellm.callbacks = original_callbacks @@ -1071,7 +1190,9 @@ def test_ProxyConfig_decrypt_model_list_from_db_resolves_env_refs_after_db_decry lambda value, key, return_original_value: ( "os.environ/LITELLM_DB_MODEL_API_KEY" if key == "api_key" - else "os.environ/LITELLM_MASTER_KEY" if key == "api_base" else value + else "os.environ/LITELLM_MASTER_KEY" + if key == "api_base" + else value ), ) pc = ProxyConfig() @@ -1102,9 +1223,7 @@ def fail_on_call(secret_name, *args, **kwargs): monkeypatch.setenv("LITELLM_MASTER_KEY", "master-secret") monkeypatch.setattr( "litellm.proxy.proxy_server.decrypt_value_helper", - lambda value, key, return_original_value: ( - "os.environ/LITELLM_MASTER_KEY" if key == "api_key" else value - ), + lambda value, key, return_original_value: "os.environ/LITELLM_MASTER_KEY" if key == "api_key" else value, ) monkeypatch.setattr("litellm.proxy.proxy_server.get_secret", fail_on_call) pc = ProxyConfig() @@ -1128,9 +1247,7 @@ def fail_on_call(secret_name, *args, **kwargs): def test_ProxyConfig_decrypt_model_list_from_db_invalid_params_skips(): pc = ProxyConfig() - bad = SimpleNamespace( - model_id="m-1", model_name="x", model_info={}, litellm_params="not-a-dict" - ) + bad = SimpleNamespace(model_id="m-1", model_name="x", model_info={}, litellm_params="not-a-dict") out = pc.decrypt_model_list_from_db(new_models=[bad]) # Invalid entries skipped — empty list returned. assert out == [] @@ -1180,9 +1297,7 @@ async def fake_get_config(): monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", fake_router) monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-x") monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) - monkeypatch.setattr( - "litellm.proxy.proxy_server.general_settings", {"alerting": ["email"]} - ) + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {"alerting": ["email"]}) monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", pc) # Passing None for proxy_logging_obj triggers AttributeError in _add_general_settings_from_db_config # when it calls proxy_logging_obj.update_values. @@ -1433,9 +1548,7 @@ async def test_ProxyConfig__add_router_settings_from_db_config_updates_router(): fake_router.update_settings = MagicMock() fake_prisma = MagicMock() fake_prisma.db.litellm_config.find_first = AsyncMock( - return_value=SimpleNamespace( - param_value={"timeout": 30, "retries": 2, "fallbacks": []} - ) + return_value=SimpleNamespace(param_value={"timeout": 30, "retries": 2, "fallbacks": []}) ) config_data = {"router_settings": {"timeout": 10}} await pc._add_router_settings_from_db_config( @@ -1446,9 +1559,7 @@ async def test_ProxyConfig__add_router_settings_from_db_config_updates_router(): snapshot = { "called": fake_router.update_settings.called, "call_count": fake_router.update_settings.call_count, - "kwargs_keys": sorted( - list(fake_router.update_settings.call_args.kwargs.keys()) - ), + "kwargs_keys": sorted(list(fake_router.update_settings.call_args.kwargs.keys())), } assert snapshot == { "called": True, @@ -1461,9 +1572,7 @@ async def test_ProxyConfig__add_router_settings_from_db_config_updates_router(): async def test_ProxyConfig__add_router_settings_from_db_config_none_router_noop(): pc = ProxyConfig() # No router and no prisma — should silently return. - await pc._add_router_settings_from_db_config( - config_data={}, llm_router=None, prisma_client=None - ) + await pc._add_router_settings_from_db_config(config_data={}, llm_router=None, prisma_client=None) # Error-style: bad call signature raises. with pytest.raises(TypeError): await pc._add_router_settings_from_db_config() # type: ignore[call-arg] @@ -1568,9 +1677,7 @@ async def test_ProxyConfig__update_general_settings_updates_max_parallel(monkeyp snapshot = { "max_parallel_requests": ps.general_settings.get("max_parallel_requests"), - "global_max_parallel_requests": ps.general_settings.get( - "global_max_parallel_requests" - ), + "global_max_parallel_requests": ps.general_settings.get("global_max_parallel_requests"), "ui_access_mode": ps.general_settings.get("ui_access_mode"), } assert snapshot == { @@ -1712,9 +1819,7 @@ def emit(self, record: logging.LogRecord) -> None: @pytest.mark.asyncio -async def test_ProxyConfig_load_config_redacts_secret_litellm_setting_keeps_plain( - tmp_path, monkeypatch -): +async def test_ProxyConfig_load_config_redacts_secret_litellm_setting_keeps_plain(tmp_path, monkeypatch): """Regression for LIT-4152 on the ``litellm_settings`` apply loop. ``load_config`` logged ``setting litellm.=`` verbatim at DEBUG, @@ -1735,11 +1840,7 @@ async def test_ProxyConfig_load_config_redacts_secret_litellm_setting_keeps_plai api_key_secret = "sk-lit4152-litellm-settings-secret-abcdef1234567890" f = tmp_path / "c.yaml" f.write_text( - "model_list: []\n" - "general_settings: {}\n" - "litellm_settings:\n" - f" api_key: {api_key_secret}\n" - " num_retries: 7\n" + f"model_list: []\ngeneral_settings: {{}}\nlitellm_settings:\n api_key: {api_key_secret}\n num_retries: 7\n" ) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)