mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
Merge f344a20e6c into e26a6450c8
This commit is contained in:
commit
c8c832c106
3 changed files with 192 additions and 33 deletions
|
|
@ -2005,7 +2005,6 @@ class BaseLLMHTTPHandler:
|
|||
url=complete_url,
|
||||
headers=signed_headers,
|
||||
)
|
||||
response.raise_for_status()
|
||||
else:
|
||||
# A signed body must be sent verbatim, re-serializing it would break the signature
|
||||
response = client.post(
|
||||
|
|
@ -2015,6 +2014,7 @@ class BaseLLMHTTPHandler:
|
|||
json=data if signed_json_body is None else None,
|
||||
timeout=timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as e:
|
||||
raise provider_config.get_http_error_class(e)
|
||||
except Exception as e:
|
||||
|
|
@ -2109,7 +2109,6 @@ class BaseLLMHTTPHandler:
|
|||
url=complete_url,
|
||||
headers=signed_headers,
|
||||
)
|
||||
response.raise_for_status()
|
||||
else:
|
||||
# A signed body must be sent verbatim, re-serializing it would break the signature
|
||||
response = await async_httpx_client.post(
|
||||
|
|
@ -2119,6 +2118,7 @@ class BaseLLMHTTPHandler:
|
|||
json=data if signed_json_body is None else None,
|
||||
timeout=timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as e:
|
||||
raise provider_config.get_http_error_class(e)
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -230,8 +231,20 @@ def search(
|
|||
try:
|
||||
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.pop("litellm_logging_obj")
|
||||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
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)}")
|
||||
|
|
@ -303,6 +316,7 @@ def search(
|
|||
api_key=api_key,
|
||||
api_base=complete_url,
|
||||
custom_llm_provider=search_provider,
|
||||
client=client,
|
||||
asearch=_is_async,
|
||||
headers=headers,
|
||||
provider_config=search_provider_config,
|
||||
|
|
|
|||
|
|
@ -10,18 +10,23 @@ leaving keyless providers and legitimate operator overrides untouched.
|
|||
"""
|
||||
|
||||
from typing import Dict, Tuple, Type
|
||||
from unittest.mock import AsyncMock, patch
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.apiserpent.search.transformation import APISerpentSearchConfig
|
||||
from litellm.llms.azure.search.transformation import BingGroundingSearchConfig
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.base_llm.search.transformation import (
|
||||
BaseSearchConfig,
|
||||
SearchResponse,
|
||||
_is_trusted_search_api_base,
|
||||
)
|
||||
from litellm.llms.brave.search.transformation import BraveSearchConfig
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.llms.dataforseo.search.transformation import DataForSEOSearchConfig
|
||||
from litellm.llms.exa_ai.search.transformation import ExaAISearchConfig
|
||||
from litellm.llms.fastcrw.search.transformation import FastCRWSearchConfig
|
||||
|
|
@ -107,9 +112,7 @@ PROVIDERS: Tuple[ProviderSpec, ...] = (
|
|||
_IDS = tuple(spec[0].__name__ for spec in PROVIDERS)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"config_cls, server_env, caller_key, extra_env", PROVIDERS, ids=_IDS
|
||||
)
|
||||
@pytest.mark.parametrize("config_cls, server_env, caller_key, extra_env", PROVIDERS, ids=_IDS)
|
||||
def test_server_secret_refused_for_caller_api_base(
|
||||
config_cls: Type[BaseSearchConfig],
|
||||
server_env: Dict[str, str],
|
||||
|
|
@ -124,9 +127,7 @@ def test_server_secret_refused_for_caller_api_base(
|
|||
config_cls().validate_environment(headers={}, api_base=ATTACKER_BASE)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"config_cls, server_env, caller_key, extra_env", PROVIDERS, ids=_IDS
|
||||
)
|
||||
@pytest.mark.parametrize("config_cls, server_env, caller_key, extra_env", PROVIDERS, ids=_IDS)
|
||||
def test_caller_supplied_key_is_honored_for_custom_api_base(
|
||||
config_cls: Type[BaseSearchConfig],
|
||||
server_env: Dict[str, str],
|
||||
|
|
@ -139,14 +140,10 @@ def test_caller_supplied_key_is_honored_for_custom_api_base(
|
|||
|
||||
# An explicit caller key is the caller's own credential, so pointing it at
|
||||
# the caller's own host must be allowed.
|
||||
config_cls().validate_environment(
|
||||
headers={}, api_key=caller_key, api_base=ATTACKER_BASE
|
||||
)
|
||||
config_cls().validate_environment(headers={}, api_key=caller_key, api_base=ATTACKER_BASE)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"config_cls, server_env, caller_key, extra_env", PROVIDERS, ids=_IDS
|
||||
)
|
||||
@pytest.mark.parametrize("config_cls, server_env, caller_key, extra_env", PROVIDERS, ids=_IDS)
|
||||
def test_server_secret_used_without_caller_api_base(
|
||||
config_cls: Type[BaseSearchConfig],
|
||||
server_env: Dict[str, str],
|
||||
|
|
@ -167,9 +164,7 @@ def test_keyless_provider_allows_caller_api_base(
|
|||
) -> None:
|
||||
monkeypatch.delenv("SEARXNG_API_KEY", raising=False)
|
||||
|
||||
headers = SearXNGSearchConfig().validate_environment(
|
||||
headers={}, api_base="https://my-searxng.internal"
|
||||
)
|
||||
headers = SearXNGSearchConfig().validate_environment(headers={}, api_base="https://my-searxng.internal")
|
||||
|
||||
assert "Authorization" not in headers
|
||||
|
||||
|
|
@ -182,9 +177,7 @@ def test_operator_env_base_override_is_trusted(
|
|||
|
||||
# Mirrors the second validate_environment call in the search handler, which
|
||||
# receives the already-resolved operator base as api_base.
|
||||
headers = SerperSearchConfig().validate_environment(
|
||||
headers={}, api_base="https://serper.internal.corp/search"
|
||||
)
|
||||
headers = SerperSearchConfig().validate_environment(headers={}, api_base="https://serper.internal.corp/search")
|
||||
|
||||
assert headers["X-API-KEY"] == "srv"
|
||||
|
||||
|
|
@ -200,9 +193,7 @@ class TestResolveServerApiKey:
|
|||
)
|
||||
assert result == "mine"
|
||||
|
||||
def test_returns_none_when_no_server_secret(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
def test_returns_none_when_no_server_secret(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.delenv("SEARXNG_API_KEY", raising=False)
|
||||
result = BaseSearchConfig().resolve_server_api_key(
|
||||
caller_api_key=None,
|
||||
|
|
@ -228,14 +219,10 @@ class TestResolveServerApiKey:
|
|||
|
||||
class TestIsTrustedSearchApiBase:
|
||||
def test_matches_default_host(self) -> None:
|
||||
assert _is_trusted_search_api_base(
|
||||
"https://google.serper.dev/search", "https://google.serper.dev", None
|
||||
)
|
||||
assert _is_trusted_search_api_base("https://google.serper.dev/search", "https://google.serper.dev", None)
|
||||
|
||||
def test_foreign_host_untrusted(self) -> None:
|
||||
assert not _is_trusted_search_api_base(
|
||||
ATTACKER_BASE, "https://google.serper.dev", None
|
||||
)
|
||||
assert not _is_trusted_search_api_base(ATTACKER_BASE, "https://google.serper.dev", None)
|
||||
|
||||
def test_env_override_host_trusted(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("SERPER_API_BASE", "https://serper.internal.corp")
|
||||
|
|
@ -248,9 +235,7 @@ class TestIsTrustedSearchApiBase:
|
|||
def test_schemeless_candidate_untrusted(self) -> None:
|
||||
# Without a scheme urlsplit puts the value in the path, leaving an empty
|
||||
# netloc; an unparseable host must never be treated as trusted.
|
||||
assert not _is_trusted_search_api_base(
|
||||
"attacker.example.com", "https://google.serper.dev", None
|
||||
)
|
||||
assert not _is_trusted_search_api_base("attacker.example.com", "https://google.serper.dev", None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -333,3 +318,163 @@ async def test_query_param_key_not_leaked_with_dummy_caller_key(
|
|||
assert captured["url"], "expected an outbound request to be attempted"
|
||||
assert server_key not in captured["url"]
|
||||
assert "sk-CALLER-DUMMY" in captured["url"]
|
||||
|
||||
|
||||
class TestSearchHTTPErrorHandling:
|
||||
"""Tests verifying that non-2xx HTTP responses from search providers raise proper exceptions and custom clients are forwarded."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("status_code", "expected_exception"),
|
||||
[
|
||||
(401, litellm.AuthenticationError),
|
||||
(429, litellm.RateLimitError),
|
||||
(500, litellm.InternalServerError),
|
||||
(503, litellm.ServiceUnavailableError),
|
||||
],
|
||||
)
|
||||
def test_sync_search_http_error_raising(
|
||||
self,
|
||||
status_code: int,
|
||||
expected_exception: type[BaseLLMException],
|
||||
) -> None:
|
||||
"""Verify that synchronous search raises mapped exceptions on error status codes."""
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(
|
||||
status_code=status_code,
|
||||
request=request,
|
||||
json={"error": "Test error message"},
|
||||
)
|
||||
|
||||
client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(handler)))
|
||||
with pytest.raises(expected_exception):
|
||||
litellm.search(
|
||||
query="test query",
|
||||
search_provider="tavily",
|
||||
api_key="tvly-testkey",
|
||||
client=client,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("status_code", "expected_exception"),
|
||||
[
|
||||
(401, litellm.AuthenticationError),
|
||||
(429, litellm.RateLimitError),
|
||||
(500, litellm.InternalServerError),
|
||||
(503, litellm.ServiceUnavailableError),
|
||||
],
|
||||
)
|
||||
async def test_async_search_http_error_raising(
|
||||
self,
|
||||
status_code: int,
|
||||
expected_exception: type[BaseLLMException],
|
||||
) -> None:
|
||||
"""Verify that asynchronous search raises mapped exceptions on error status codes."""
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(
|
||||
status_code=status_code,
|
||||
request=request,
|
||||
json={"error": "Async test error message"},
|
||||
)
|
||||
|
||||
client = AsyncHTTPHandler()
|
||||
client._client = httpx.AsyncClient(transport=httpx.MockTransport(handler))
|
||||
with pytest.raises(expected_exception):
|
||||
await litellm.asearch(
|
||||
query="test query",
|
||||
search_provider="tavily",
|
||||
api_key="tvly-testkey",
|
||||
client=client,
|
||||
)
|
||||
|
||||
def test_sync_search_plumbs_custom_client(self) -> None:
|
||||
"""Verify that a custom HTTPHandler passed to litellm.search is forwarded to the underlying handler."""
|
||||
called = False
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
nonlocal called
|
||||
called = True
|
||||
return httpx.Response(
|
||||
status_code=200,
|
||||
request=request,
|
||||
json={"results": [{"title": "Test", "url": "https://example.com", "content": "Sample"}]},
|
||||
)
|
||||
|
||||
custom_client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(handler)))
|
||||
response = litellm.search(
|
||||
query="test query",
|
||||
search_provider="tavily",
|
||||
api_key="tvly-testkey",
|
||||
client=custom_client,
|
||||
)
|
||||
assert isinstance(response, SearchResponse)
|
||||
assert called
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_search_plumbs_custom_client(self) -> None:
|
||||
"""Verify that a custom AsyncHTTPHandler passed to litellm.asearch is forwarded to the underlying handler."""
|
||||
called = False
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
nonlocal called
|
||||
called = True
|
||||
return httpx.Response(
|
||||
status_code=200,
|
||||
request=request,
|
||||
json={"results": [{"title": "Async Test", "url": "https://example.com", "content": "Async Sample"}]},
|
||||
)
|
||||
|
||||
custom_client = AsyncHTTPHandler()
|
||||
custom_client._client = httpx.AsyncClient(transport=httpx.MockTransport(handler))
|
||||
response = await litellm.asearch(
|
||||
query="test query",
|
||||
search_provider="tavily",
|
||||
api_key="tvly-testkey",
|
||||
client=custom_client,
|
||||
)
|
||||
assert isinstance(response, SearchResponse)
|
||||
assert called
|
||||
|
||||
def test_sync_search_handles_empty_error_response_body(self) -> None:
|
||||
"""Verify that 500 responses with empty bodies still raise InternalServerError properly."""
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(
|
||||
status_code=500,
|
||||
request=request,
|
||||
text="",
|
||||
)
|
||||
|
||||
client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(handler)))
|
||||
with pytest.raises(litellm.InternalServerError):
|
||||
litellm.search(
|
||||
query="test query",
|
||||
search_provider="tavily",
|
||||
api_key="tvly-testkey",
|
||||
client=client,
|
||||
)
|
||||
|
||||
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,
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue