diff --git a/litellm/search/main.py b/litellm/search/main.py index 25379332b05..148bba31b48 100644 --- a/litellm/search/main.py +++ b/litellm/search/main.py @@ -15,6 +15,7 @@ from litellm._logging import verbose_logger from litellm.constants import request_timeout from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.search.transformation import BaseSearchConfig, SearchResponse +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.types.utils import SearchProviders from litellm.utils import ProviderConfigManager, client, filter_out_litellm_params @@ -233,6 +234,17 @@ def search( client: Final = kwargs.get("client", None) _is_async: Final = kwargs.pop("asearch", False) is True + if _is_async: + if client is not None and not isinstance(client, AsyncHTTPHandler): + raise ValueError( + f"client must be an instance of AsyncHTTPHandler for asynchronous search, got {type(client)}" + ) + else: + if client is not None and not isinstance(client, HTTPHandler): + raise ValueError( + f"client must be an instance of HTTPHandler for synchronous search, got {type(client)}" + ) + # Validate query parameter if not isinstance(query, (str, list)): raise ValueError(f"query must be a string or list of strings, got {type(query)}") diff --git a/tests/test_litellm/llms/base_llm/search/test_base_search_transformation.py b/tests/test_litellm/llms/base_llm/search/test_base_search_transformation.py index 787acf42c14..cd8b78e7910 100644 --- a/tests/test_litellm/llms/base_llm/search/test_base_search_transformation.py +++ b/tests/test_litellm/llms/base_llm/search/test_base_search_transformation.py @@ -450,3 +450,26 @@ class TestSearchHTTPErrorHandling: search_provider="tavily", api_key="tvly-testkey", ) + + def test_sync_search_rejects_async_client(self) -> None: + """Verify that passing an AsyncHTTPHandler to synchronous search raises a clear exception.""" + async_client = AsyncHTTPHandler() + with pytest.raises(Exception, match="client must be an instance of HTTPHandler"): + litellm.search( + query="test query", + search_provider="tavily", + api_key="tvly-testkey", + client=async_client, + ) + + @pytest.mark.asyncio + async def test_async_search_rejects_sync_client(self) -> None: + """Verify that passing a sync HTTPHandler to asynchronous search raises a clear exception.""" + sync_client = HTTPHandler() + with pytest.raises(Exception, match="client must be an instance of AsyncHTTPHandler"): + await litellm.asearch( + query="test query", + search_provider="tavily", + api_key="tvly-testkey", + client=sync_client, + )