test(proxy): drive search endpoint tests through a real router and cached team

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
mrinal 2026-09-29 05:24:52 +00:00
parent b0b6678543
commit 9b5d82e651

View file

@ -1,14 +1,17 @@
from unittest.mock import AsyncMock, MagicMock
from collections.abc import Iterator
from typing import Final
import pytest
import respx
from fastapi import FastAPI
from fastapi.testclient import TestClient
import litellm.proxy.proxy_server as proxy_server
from litellm.llms.base_llm.search.transformation import SearchResponse
import litellm
from litellm import Router
from litellm.proxy import proxy_server
from litellm.proxy._types import (
LiteLLM_ObjectPermissionTable,
LiteLLM_TeamTable,
LiteLLM_TeamTableCachedObj,
LiteLLM_UserTable,
LitellmUserRoles,
ProxyException,
@ -18,14 +21,18 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.search_endpoints.endpoints import router
TAVILY_SEARCH_URL: Final = "https://api.tavily.com/search"
TAVILY_RESULT: Final = {"title": "LiteLLM", "url": "https://docs.litellm.ai", "content": "LLM gateway"}
def _search_router() -> MagicMock:
llm_router = MagicMock()
llm_router.search_tools = [
{"search_tool_name": "search-a", "litellm_params": {"search_provider": "tavily", "api_key": "fake"}},
]
llm_router.asearch = AsyncMock(return_value=SearchResponse(object="search", results=[]))
return llm_router
def _search_router() -> Router:
return Router(
model_list=[],
search_tools=[
{"search_tool_name": "search-a", "litellm_params": {"search_provider": "tavily", "api_key": "fake"}},
],
num_retries=0,
)
def _client(caller: UserAPIKeyAuth) -> TestClient:
@ -37,23 +44,34 @@ def _client(caller: UserAPIKeyAuth) -> TestClient:
@pytest.fixture
def user_without_search_grants(monkeypatch: pytest.MonkeyPatch) -> None:
cache = UserApiKeyCache()
cache.set_cache(key="user-1", value=LiteLLM_UserTable(user_id="user-1"))
monkeypatch.setattr(proxy_server, "user_api_key_cache", cache)
monkeypatch.setattr(proxy_server, "prisma_client", MagicMock())
def tavily(monkeypatch: pytest.MonkeyPatch) -> Iterator[respx.Route]:
monkeypatch.setattr( # test-quality-ok: respx needs HTTPX enabled to fake the provider HTTP boundary.
litellm,
"disable_aiohttp_transport",
True,
)
litellm.in_memory_llm_clients_cache.flush_cache()
with respx.mock(assert_all_called=False) as mock:
yield mock.post(TAVILY_SEARCH_URL).respond(200, json={"results": [TAVILY_RESULT]})
litellm.in_memory_llm_clients_cache.flush_cache()
@pytest.fixture
def team_with_search_tools(monkeypatch: pytest.MonkeyPatch, user_without_search_grants: None):
def set_team_search_tools(search_tools: list[str]) -> None:
team = LiteLLM_TeamTable(
team_id="team-1",
object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-team", search_tools=search_tools),
)
monkeypatch.setattr("litellm.proxy.auth.auth_checks.get_team_object", AsyncMock(return_value=team))
def cache(monkeypatch: pytest.MonkeyPatch) -> UserApiKeyCache:
user_api_key_cache = UserApiKeyCache()
user_api_key_cache.set_cache(key="user-1", value=LiteLLM_UserTable(user_id="user-1"))
monkeypatch.setattr(proxy_server, "user_api_key_cache", user_api_key_cache)
monkeypatch.setattr(proxy_server, "prisma_client", object())
monkeypatch.setattr(proxy_server, "llm_router", _search_router())
return user_api_key_cache
return set_team_search_tools
def _cache_team(cache: UserApiKeyCache, search_tools: list[str]) -> None:
team = LiteLLM_TeamTableCachedObj(
team_id="team-1",
object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-team", search_tools=search_tools),
)
cache.set_cache(key="team_id:team-1", value=team)
@pytest.mark.parametrize(
@ -68,32 +86,28 @@ def team_with_search_tools(monkeypatch: pytest.MonkeyPatch, user_without_search_
)
@pytest.mark.parametrize("path", ["/v1/search/search-a", "/search/search-a"])
def test_direct_search_team_key_follows_default_search_list_deny(
monkeypatch, team_with_search_tools, path, general_settings, team_search_tools, expected_status
monkeypatch, cache, tavily, path, general_settings, team_search_tools, expected_status
):
llm_router = _search_router()
monkeypatch.setattr(proxy_server, "llm_router", llm_router)
monkeypatch.setattr(proxy_server, "general_settings", general_settings)
team_with_search_tools(team_search_tools)
_cache_team(cache, team_search_tools)
caller = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="user-1", team_id="team-1")
response = _client(caller).post(path, json={"query": "what is litellm"})
assert response.status_code == expected_status, response.text
if expected_status == 200:
assert response.json()["object"] == "search"
llm_router.asearch.assert_awaited_once()
assert response.json()["results"][0]["url"] == TAVILY_RESULT["url"]
assert tavily.call_count == 1
else:
assert "search-a" in response.text
llm_router.asearch.assert_not_awaited()
assert tavily.call_count == 0
def test_direct_search_body_tool_name_is_denied_under_default_search_list_deny(monkeypatch):
llm_router = _search_router()
monkeypatch.setattr(proxy_server, "llm_router", llm_router)
def test_direct_search_body_tool_name_is_denied_under_default_search_list_deny(monkeypatch, cache, tavily):
monkeypatch.setattr(proxy_server, "general_settings", {"default_search_list_deny": True})
caller = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER)
response = _client(caller).post("/v1/search", json={"search_tool_name": "search-a", "query": "what is litellm"})
assert response.status_code == 403, response.text
llm_router.asearch.assert_not_awaited()
assert tavily.call_count == 0