From 9b5d82e65109f76b46e5cf17649eb6488527bdfb Mon Sep 17 00:00:00 2001 From: mrinal Date: Tue, 29 Sep 2026 05:24:52 +0000 Subject: [PATCH] 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> --- .../proxy/search_endpoints/test_endpoints.py | 84 +++++++++++-------- 1 file changed, 49 insertions(+), 35 deletions(-) diff --git a/tests/unit/proxy/search_endpoints/test_endpoints.py b/tests/unit/proxy/search_endpoints/test_endpoints.py index 3b8a55310ee..03b5ed92005 100644 --- a/tests/unit/proxy/search_endpoints/test_endpoints.py +++ b/tests/unit/proxy/search_endpoints/test_endpoints.py @@ -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