From 6483b71b3c466cd9f26691299b3b37672cbc2425 Mon Sep 17 00:00:00 2001 From: linhongyu510 Date: Sat, 29 Aug 2026 14:20:38 +0800 Subject: [PATCH] fix(search): keep provider credentials out of logs --- litellm/llms/custom_httpx/llm_http_handler.py | 7 +- litellm/search/main.py | 12 +- .../router_utils/test_search_api_router.py | 105 +++++++++++++++++- 3 files changed, 115 insertions(+), 9 deletions(-) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 834f7d564a2..19f0e212902 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -1777,6 +1777,7 @@ class BaseLLMHTTPHandler: asearch: bool = False, headers: dict[str, object] | None = None, provider_config: BaseSearchConfig | None = None, + auth_params: dict[str, object] | None = None, ) -> SearchResponse | Coroutine[object, object, SearchResponse]: """ Sync Search handler. @@ -1796,6 +1797,7 @@ class BaseLLMHTTPHandler: client=client, headers=headers, provider_config=provider_config, + auth_params=auth_params, ) # Validate environment and get headers @@ -1821,7 +1823,7 @@ class BaseLLMHTTPHandler: signed_headers, signed_json_body = provider_config.sign_request( headers=headers, - optional_params=optional_params, + optional_params={**optional_params, **(auth_params or {})}, request_data=data, api_base=complete_url, api_key=api_key, @@ -1882,6 +1884,7 @@ class BaseLLMHTTPHandler: client: HTTPHandler | AsyncHTTPHandler | None = None, headers: dict[str, object] | None = None, provider_config: BaseSearchConfig | None = None, + auth_params: dict[str, object] | None = None, ) -> SearchResponse: """ Async Search handler. @@ -1915,7 +1918,7 @@ class BaseLLMHTTPHandler: signed_headers, signed_json_body = provider_config.sign_request( headers=headers, - optional_params=optional_params, + optional_params={**optional_params, **(auth_params or {})}, request_data=data, api_base=complete_url, api_key=api_key, diff --git a/litellm/search/main.py b/litellm/search/main.py index b2dd51799a1..29addce1094 100644 --- a/litellm/search/main.py +++ b/litellm/search/main.py @@ -6,7 +6,7 @@ import asyncio import contextvars from collections.abc import Coroutine from functools import partial -from typing import Any, Final +from typing import Any, Final, cast import httpx @@ -249,6 +249,15 @@ def search( verbose_logger.debug("Search call - provider: %s", search_provider) + auth_param_names: Final[frozenset[str]] = frozenset( + cast( # cast-ok: BaseAWSLLM exposes authentication parameter names as list[str] + list[str], getattr(search_provider_config, "aws_authentication_params", ()) + ) + ) + auth_params: Final[dict[str, object]] = { # mutable-ok: isolated provider signing params + key: kwargs.pop(key) for key in tuple(kwargs) if key in auth_param_names + } + # Build optional_params from explicit parameters optional_params: Final = _build_search_optional_params( max_results=max_results, @@ -306,6 +315,7 @@ def search( asearch=_is_async, headers=headers, provider_config=search_provider_config, + auth_params=auth_params, ) return response diff --git a/tests/test_litellm/router_utils/test_search_api_router.py b/tests/test_litellm/router_utils/test_search_api_router.py index 914b3b576bc..a817ab8472d 100644 --- a/tests/test_litellm/router_utils/test_search_api_router.py +++ b/tests/test_litellm/router_utils/test_search_api_router.py @@ -8,12 +8,22 @@ was ignored with no warning. These tests are provider-free: they capture the kwargs handed to the search function via a fake `original_generic_function`. """ +import importlib from typing import Any +from unittest.mock import MagicMock import pytest +from litellm.llms.base_llm.search.transformation import ( + BaseSearchConfig, + SearchResponse, +) +from litellm.llms.bedrock.search.transformation import AgentCoreSearchConfig +from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.router_utils.search_api_router import SearchAPIRouter +search_module = importlib.import_module("litellm.search.main") + class _FakeRouter: """Minimal stand-in exposing only the `search_tools` the helper reads.""" @@ -23,9 +33,7 @@ class _FakeRouter: def _make_router(litellm_params: dict[str, Any], name: str = "agentcore-search") -> _FakeRouter: - return _FakeRouter( - [{"search_tool_name": name, "litellm_params": {**litellm_params}}] - ) + return _FakeRouter([{"search_tool_name": name, "litellm_params": {**litellm_params}}]) @pytest.mark.asyncio @@ -96,9 +104,7 @@ async def test_credentials_are_not_double_passed(): captured: dict[str, Any] = {} async def fake_search(*, search_provider, api_key, api_base, **kwargs: Any): - captured.update( - {"search_provider": search_provider, "api_key": api_key, "api_base": api_base, **kwargs} - ) + captured.update({"search_provider": search_provider, "api_key": api_key, "api_base": api_base, **kwargs}) return {"object": "search", "results": []} router = _make_router( @@ -139,3 +145,90 @@ async def test_missing_search_provider_still_raises(): original_generic_function=fake_search, query="anything", ) + + +def test_agentcore_auth_params_are_excluded_from_logging(monkeypatch: pytest.MonkeyPatch): + config = AgentCoreSearchConfig() + monkeypatch.setattr( + search_module.ProviderConfigManager, + "get_provider_search_config", + lambda provider: config, + ) + monkeypatch.setattr(config, "validate_environment", lambda **kwargs: {}) + monkeypatch.setattr( + config, + "get_complete_url", + lambda **kwargs: "https://gateway.example/mcp", + ) + + captured: dict[str, Any] = {} + + def fake_search(**kwargs: Any): + captured.update(kwargs) + return SearchResponse(results=[]) + + monkeypatch.setattr(search_module.base_llm_http_handler, "search", fake_search) + logging_obj = MagicMock() + + search_module.search.__wrapped__( + query="test", + search_provider="agentcore", + litellm_logging_obj=logging_obj, + tool_name="target___WebSearch", + aws_access_key_id="access", + aws_secret_access_key="secret", + aws_session_token="token", + aws_region_name="us-east-1", + ) + + assert captured["optional_params"] == {"tool_name": "target___WebSearch"} + assert captured["auth_params"] == { + "aws_access_key_id": "access", + "aws_secret_access_key": "secret", + "aws_session_token": "token", + "aws_region_name": "us-east-1", + } + logged = logging_obj.update_from_kwargs.call_args.kwargs + assert logged["kwargs"] == {"tool_name": "target___WebSearch"} + assert logged["optional_params"] == {"tool_name": "target___WebSearch"} + + +def test_search_handler_only_exposes_auth_params_to_signing(): + class CapturingConfig(BaseSearchConfig): + def __init__(self): + self.transformed_params: dict[str, object] | None = None + self.signing_params: dict[str, object] | None = None + + def validate_environment(self, **kwargs): + return {} + + def transform_search_request(self, query, optional_params, **kwargs): + self.transformed_params = optional_params + return {} + + def get_complete_url(self, **kwargs): + return "https://gateway.example/mcp" + + def sign_request(self, *, optional_params, **kwargs): + self.signing_params = optional_params + raise RuntimeError("stop before network") + + config = CapturingConfig() + with pytest.raises(RuntimeError, match="stop before network"): + BaseLLMHTTPHandler().search( + query="test", + optional_params={"tool_name": "target___WebSearch"}, + auth_params={"aws_secret_access_key": "secret"}, + timeout=1, + logging_obj=MagicMock(), + api_key=None, + api_base="https://gateway.example/mcp", + custom_llm_provider="agentcore", + provider_config=config, + ) + + assert config.transformed_params == {"tool_name": "target___WebSearch"} + assert config.signing_params == { + "tool_name": "target___WebSearch", + "aws_secret_access_key": "secret", + }