From 282d733fb6515f4dd2acb5b86ad57f6d2f37f075 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 6 Oct 2026 12:50:26 -0700 Subject: [PATCH] feat(auth): deny search tools by default when search_tool_deny_by_default is set (#44490) Co-authored-by: mrinal Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../websearch_interception/handler.py | 45 +-- .../litellm_core_utils/error_normalization.py | 3 + litellm/proxy/_types.py | 29 ++ litellm/proxy/auth/auth_checks.py | 246 ++++++++++---- litellm/proxy/auth/fallback_model_access.py | 35 +- litellm/proxy/proxy_server.py | 13 +- litellm/proxy/search_endpoints/endpoints.py | 52 +-- .../search_tool_management.py | 41 +-- .../test_search_tool_deny_by_default.py | 221 +++++++++++++ .../test_websearch_interception_handler.py | 189 ++++++++++- .../proxy/auth/test_fallback_model_access.py | 69 +++- .../auth/test_search_tool_deny_by_default.py | 270 ++++++++++++++++ .../test_search_tool_management.py | 300 ++++++++++++------ .../proxy/proxy_server/test_proxy_config.py | 32 +- .../proxy/search_endpoints/test_endpoints.py | 235 +++++++++++++- ui/litellm-dashboard/src/lib/http/schema.d.ts | 6 + 16 files changed, 1497 insertions(+), 289 deletions(-) create mode 100644 tests/integration/authorization/test_search_tool_deny_by_default.py create mode 100644 tests/unit/proxy/auth/test_search_tool_deny_by_default.py diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 1db94e82066..15413140d9e 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -1557,8 +1557,8 @@ class WebSearchInterceptionLogger(CustomLogger): search_litellm_params = dict[str, object](tool_params) search_provider = tool_params.get("search_provider") - # Fallback to perplexity if no router or no search tools configured if not search_provider: + self._authorize_unregistered_search_fallback(kwargs=kwargs) search_provider = "perplexity" verbose_logger.debug( "WebSearchInterception: No search tools configured in router, using default provider '%s'", @@ -1623,6 +1623,18 @@ class WebSearchInterceptionLogger(CustomLogger): verbose_logger.error("WebSearchInterception: Search failed for '%s': %s", query, e) raise + def _authorize_unregistered_search_fallback(self, kwargs: Mapping[str, object] | None) -> None: + user_api_key_auth: Final = self._get_user_api_key_auth_from_kwargs(kwargs) + if user_api_key_auth is None: + return + + from litellm.proxy.auth.auth_checks import check_unregistered_search_fallback, typed_general_settings + from litellm.proxy.proxy_server import general_settings + + check_unregistered_search_fallback( + valid_token=user_api_key_auth, general_settings=typed_general_settings(general_settings) + ) + async def _authorize_search_tool( self, search_tool: Mapping[str, object], @@ -1636,36 +1648,9 @@ class WebSearchInterceptionLogger(CustomLogger): if user_api_key_auth is None: return - from litellm.proxy.auth.auth_checks import ( - can_key_call_search_tool, - can_team_call_search_tool, - get_team_object, - ) + from litellm.proxy.auth.auth_checks import can_token_call_search_tool - await can_key_call_search_tool( - search_tool_name=search_tool_name, - valid_token=user_api_key_auth, - ) - - team_id: Final[str | None] = getattr(user_api_key_auth, "team_id", None) - if team_id: - from litellm.proxy.proxy_server import ( - prisma_client, - proxy_logging_obj, - user_api_key_cache, - ) - - team_object: Final = await get_team_object( - team_id=team_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=getattr(user_api_key_auth, "parent_otel_span", None), - proxy_logging_obj=proxy_logging_obj, - ) - await can_team_call_search_tool( - search_tool_name=search_tool_name, - team_object=team_object, - ) + await can_token_call_search_tool(search_tool_name=search_tool_name, valid_token=user_api_key_auth) @staticmethod def _build_search_request_metadata( diff --git a/litellm/litellm_core_utils/error_normalization.py b/litellm/litellm_core_utils/error_normalization.py index 4426d42f767..00317e502ec 100644 --- a/litellm/litellm_core_utils/error_normalization.py +++ b/litellm/litellm_core_utils/error_normalization.py @@ -108,6 +108,9 @@ _PROXY_ERROR_TYPE_MAP: Final[Mapping[str, str]] = MappingProxyType( "team_vector_store_access_denied": PERMISSION_DENIED, "org_vector_store_access_denied": PERMISSION_DENIED, "user_vector_store_access_denied": PERMISSION_DENIED, + "key_search_tool_access_denied": PERMISSION_DENIED, + "team_search_tool_access_denied": PERMISSION_DENIED, + "user_search_tool_access_denied": PERMISSION_DENIED, "tool_access_denied": PERMISSION_DENIED, "team_member_permission_error": PERMISSION_DENIED, "not_found_error": RESOURCE_NOT_FOUND, diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 23cd64b8a49..d02d1f29ded 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -3026,6 +3026,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): default=False, description="When True, a vector store must be explicitly listed in object_permission.vector_stores: a virtual key needs its own grant plus its team's, a keyless team member needs the team's, and a user with neither needs their own. A missing permission record, an empty list, or an unresolved team grants nothing. Dashboard session keys are not yet covered", ) + search_tool_deny_by_default: bool = Field( + default=False, + description="When True, a search tool must be explicitly listed in object_permission.search_tools: a virtual key needs its own grant plus its team's, a keyless team member needs the team's, and a user with neither needs their own. A missing permission record, an empty list, or an unresolved team grants nothing, and the unregistered search fallback is denied. The master key and dashboard sessions are exempt", + ) missing_session_id: Literal["generate", "reject", "omit"] | None = Field( None, description="What to do with LLM API requests that carry no session id (x-litellm-session-id header, metadata.session_id, etc.). 'generate' stamps one id into litellm_session_id, litellm_trace_id and metadata.session_id so SpendLogs and logging callbacks agree; 'reject' returns 400; 'omit' leaves SpendLogs.session_id null, matching callbacks such as Langfuse that only record a client-established metadata.session_id. Unset keeps the legacy behavior where SpendLogs falls back to the trace id while callbacks get no session id.", @@ -4521,6 +4525,21 @@ class ProxyErrorTypes(str, enum.Enum): User does not have access to the vector store """ + key_search_tool_access_denied = "key_search_tool_access_denied" + """ + Key does not have access to the search tool + """ + + team_search_tool_access_denied = "team_search_tool_access_denied" + """ + Team does not have access to the search tool + """ + + user_search_tool_access_denied = "user_search_tool_access_denied" + """ + User does not have access to the search tool + """ + team_member_already_in_team = "team_member_already_in_team" """ Team member is already in team @@ -4569,6 +4588,16 @@ class ProxyErrorTypes(str, enum.Enum): elif object_type == "user": return cls.user_vector_store_access_denied + @classmethod + def get_search_tool_access_error_type_for_object( + cls, object_type: Literal["key", "team", "user"] + ) -> "ProxyErrorTypes": + return { + "key": cls.key_search_tool_access_denied, + "team": cls.team_search_tool_access_denied, + "user": cls.user_search_tool_access_denied, + }[object_type] + DB_CONNECTION_ERROR_TYPES: Final = ( httpx.ConnectError, diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 3ebd1a2cfe0..33fcac6230a 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -14,6 +14,7 @@ import math import re import time from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence +from dataclasses import dataclass from functools import partial from itertools import chain from types import MappingProxyType @@ -312,6 +313,9 @@ def _user_table(repo: _PrismaTableHolder[_PrismaUserRow]) -> _PrismaAuthTable[_P return _DeadlineBoundedTable(repo.table, "user") +GrantLayer = Literal["key", "team", "user"] + + class _VectorStorePermissionsRow(Protocol): @property def vector_stores(self) -> Sequence[str] | None: ... @@ -386,6 +390,9 @@ def _typed_request_body(request_body: dict) -> Mapping[str, object]: return request_body +typed_general_settings: Final = _typed_request_body + + class _JsonLoadsObj(Protocol): def __call__(self, data: str) -> object: ... @@ -5481,10 +5488,10 @@ def _search_tool_names_from_object_permission( def _can_object_call_search_tool( search_tool_name: str, allowed_search_tools: list[str], - object_type: Literal["key", "team", "project"], + object_type: Literal["key", "team", "project", "user"], ) -> Literal[True]: """ - Check if an object (key/team/project) can access a specific search tool. + Check if an object (key/team/project/user) can access a specific search tool. Similar to _can_object_call_model but for search tools. @@ -5517,84 +5524,188 @@ def _can_object_call_search_tool( ) -async def can_key_call_search_tool( - search_tool_name: str, - valid_token: UserAPIKeyAuth, -) -> Literal[True]: +TeamObjectLoader: TypeAlias = Callable[[], Awaitable[LiteLLM_TeamTable | None]] + + +@dataclass(frozen=True) +class SearchToolGrants: """ - Check if a key can access a specific search tool. - - Similar to can_key_call_model but for search tools. - - Args: - search_tool_name: The search tool being requested - valid_token: The authenticated key - - Returns: - True if access is allowed - - Raises: - ProxyException if access is denied + The key, team and user grants that scope one caller's search tools. Every layer in `strict_layers` + must list the tool, and any other loaded grant only narrows when its list is nonempty """ - return _can_object_call_search_tool( - search_tool_name=search_tool_name, - allowed_search_tools=_search_tool_names_from_object_permission(valid_token.object_permission), - object_type="key", - ) + + grants: Mapping[GrantLayer, LiteLLM_ObjectPermissionTable | None] + strict_layers: frozenset[GrantLayer] -async def can_team_call_search_tool( - search_tool_name: str, - team_object: LiteLLM_TeamTable | None, -) -> Literal[True]: +def _search_tool_deny_by_default(general_settings: Mapping[str, object]) -> bool: """ - Check if a team can access a specific search tool. - - Similar to can_team_access_model but for search tools. - - Args: - search_tool_name: The search tool being requested - team_object: The team object - - Returns: - True if access is allowed - - Raises: - ProxyException if access is denied - """ - if team_object is None: - return True - - return _can_object_call_search_tool( - search_tool_name=search_tool_name, - allowed_search_tools=_search_tool_names_from_object_permission(team_object.object_permission), - object_type="team", - ) - - -async def can_user_view_search_tool( - search_tool_name: str, - valid_token: UserAPIKeyAuth, - team_object: LiteLLM_TeamTable | None, -) -> bool: - """ - Boolean variant of the key + team authorization enforced on /search, used to - scope /search_tools/list so a non-admin caller only sees tools it may invoke. + Startup rejects a non-boolean value from the config file. A non-boolean value that reaches + general_settings another way enables the policy, so only search tool requests are denied. """ try: - await can_key_call_search_tool( - search_tool_name=search_tool_name, - valid_token=valid_token, + return ConfigGeneralSettings.model_validate( + MappingProxyType( + {"search_tool_deny_by_default": general_settings.get("search_tool_deny_by_default", False)} + ) + ).search_tool_deny_by_default + except ValidationError: + return True + + +def is_search_tool_deny_by_default_applied(valid_token: UserAPIKeyAuth, general_settings: Mapping[str, object]) -> bool: + return _search_tool_deny_by_default(general_settings) and _is_strict_grant_identity(valid_token) + + +def _search_tool_denied(object_type: GrantLayer, message: str) -> ProxyException: + return ProxyException( + message=message, + type=ProxyErrorTypes.get_search_tool_access_error_type_for_object(object_type), + param="search_tool_name", + code=status.HTTP_403_FORBIDDEN, + ) + + +async def _strict_team_object(load_team_object: TeamObjectLoader) -> LiteLLM_TeamTable | None: + try: + return await load_team_object() + except Exception as e: # noqa: BLE001 # an unresolved team grants nothing under deny-by-default + verbose_proxy_logger.debug("Team lookup failed under search_tool_deny_by_default: %s", e) + return None + + +async def _strict_user_object( + valid_token: UserAPIKeyAuth, prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache +) -> LiteLLM_UserTable | None: + if valid_token.user_id is None: + return None + try: + return await get_user_object( + user_id=valid_token.user_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + user_id_upsert=False, + parent_otel_span=valid_token.parent_otel_span, ) - await can_team_call_search_tool( - search_tool_name=search_tool_name, - team_object=team_object, + except Exception as e: # noqa: BLE001 # an unresolved user grants nothing under deny-by-default + verbose_proxy_logger.debug("User lookup failed under search_tool_deny_by_default: %s", e) + return None + + +async def resolve_search_tool_grants( + valid_token: UserAPIKeyAuth, + general_settings: Mapping[str, object], + load_team_object: TeamObjectLoader, +) -> SearchToolGrants: + """ + Without `general_settings.search_tool_deny_by_default`, or for the master key and dashboard sessions, + the key and team lists only restrict when nonempty. With it, the identities `_strict_grant_layers` names + must each list the tool, and a missing record, `null`, `[]`, or a team or user that fails to load grants nothing + """ + if not is_search_tool_deny_by_default_applied(valid_token, general_settings): + team_object: Final = await load_team_object() + return SearchToolGrants( + grants=MappingProxyType( + {"key": valid_token.object_permission, "team": team_object.object_permission if team_object else None} + ), + strict_layers=frozenset(), ) + + from litellm.proxy.proxy_server import prisma_client, user_api_key_cache + + strict_team_object: Final = await _strict_team_object(load_team_object) + strict_layers: Final = _strict_grant_layers(True, valid_token, strict_team_object) + if prisma_client is None: + return SearchToolGrants(grants=MappingProxyType({}), strict_layers=strict_layers) + user_object: Final = ( + await _strict_user_object(valid_token, prisma_client, user_api_key_cache) if "user" in strict_layers else None + ) + return SearchToolGrants( + grants=MappingProxyType( + dict( + await _identity_grants(valid_token, strict_team_object, user_object, prisma_client, user_api_key_cache) + ) + ), + strict_layers=strict_layers, + ) + + +def check_search_tool_grants(search_tool_name: str, grants: SearchToolGrants) -> Literal[True]: + """Raises a 403 ProxyException naming the first key, team or user layer that does not grant the tool""" + for layer in ("key", "team", "user"): + grant = grants.grants.get(layer) + if layer in grants.strict_layers: + if grant is None or search_tool_name not in (grant.search_tools or ()): + raise _search_tool_denied( + layer, + f"{layer.capitalize()} not allowed to access search tool: {search_tool_name}. " + f"search_tool_deny_by_default is enabled and the {layer} does not grant it", + ) + elif grant is not None: + _can_object_call_search_tool( + search_tool_name=search_tool_name, + allowed_search_tools=_search_tool_names_from_object_permission(grant), + object_type=layer, + ) + return True + + +async def can_caller_call_search_tool( + search_tool_name: str, + valid_token: UserAPIKeyAuth, + general_settings: Mapping[str, object], + load_team_object: TeamObjectLoader, +) -> Literal[True]: + """Key, team and user search tool authorization shared by /search, web search interception and discovery""" + return check_search_tool_grants( + search_tool_name, await resolve_search_tool_grants(valid_token, general_settings, load_team_object) + ) + + +async def can_token_call_search_tool(search_tool_name: str, valid_token: UserAPIKeyAuth) -> Literal[True]: + """`can_caller_call_search_tool` against the proxy's own settings, team cache and database""" + from litellm.proxy.proxy_server import general_settings, prisma_client, proxy_logging_obj, user_api_key_cache + + async def _load_team_object() -> LiteLLM_TeamTable | None: + if not valid_token.team_id: + return None + return await get_team_object( + team_id=valid_token.team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=valid_token.parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + + return await can_caller_call_search_tool( + search_tool_name=search_tool_name, + valid_token=valid_token, + general_settings=typed_general_settings(general_settings), + load_team_object=_load_team_object, + ) + + +def can_grants_view_search_tool(search_tool_name: str, grants: SearchToolGrants) -> bool: + try: + check_search_tool_grants(search_tool_name, grants) except ProxyException: return False return True +def check_unregistered_search_fallback( + valid_token: UserAPIKeyAuth, general_settings: Mapping[str, object] +) -> Literal[True]: + """The unregistered provider fallback names no search tool that a grant could list, so deny-by-default denies it""" + if not is_search_tool_deny_by_default_applied(valid_token, general_settings): + return True + layer: Final = min(_strict_grant_layers(True, valid_token, None), key=("key", "team", "user").index) + raise _search_tool_denied( + layer, + "No registered search tool is available and search_tool_deny_by_default is enabled", + ) + + async def is_valid_fallback_model( model: str, llm_router: Router | None, @@ -6792,9 +6903,6 @@ def _get_rag_query_vector_store_id(request_body: Mapping[str, object]) -> str | return vector_store_id if isinstance(vector_store_id, str) and vector_store_id else None -GrantLayer = Literal["key", "team", "user"] - - def _is_strict_grant_identity(valid_token: UserAPIKeyAuth | None) -> bool: return ( valid_token is not None diff --git a/litellm/proxy/auth/fallback_model_access.py b/litellm/proxy/auth/fallback_model_access.py index 53807f076ea..39346c80b22 100644 --- a/litellm/proxy/auth/fallback_model_access.py +++ b/litellm/proxy/auth/fallback_model_access.py @@ -5,7 +5,8 @@ Authorize router fallback targets against the caller's key, team and project mod Fallbacks configured on the router (`router_settings.fallbacks` and friends) are chosen after auth, inside the router, so this predicate is injected into the router to re-run the same model access checks for each fallback target before it is attempted. Opt-in via -`general_settings.enforce_fallback_model_access: true`. +`general_settings.enforce_fallback_model_access: true`. Fallbacks of a search request always run the +key, team and user search tool grants instead, the same check /search runs on the requested tool. """ from collections.abc import Callable, Mapping @@ -16,8 +17,9 @@ from pydantic import ValidationError from litellm._logging import verbose_proxy_logger from litellm.proxy._types import ProxyException, UserAPIKeyAuth -from litellm.proxy.auth.auth_checks import can_key_call_resolved_model +from litellm.proxy.auth.auth_checks import can_key_call_resolved_model, can_token_call_search_tool from litellm.router import Router +from litellm.search import asearch from litellm.types.llms.base import LiteLLMBaseModel @@ -45,6 +47,23 @@ async def is_model_authorized_for_token(*, model: str, valid_token: UserAPIKeyAu return True +async def is_search_tool_authorized_for_token(*, search_tool_name: str, valid_token: UserAPIKeyAuth) -> bool: + try: + await can_token_call_search_tool(search_tool_name=search_tool_name, valid_token=valid_token) + except ProxyException: + return False + except Exception as e: # noqa: BLE001 # fail closed: a lookup failure must neither run the fallback nor replace the provider error + verbose_proxy_logger.warning( + "Skipping fallback to search tool=%s: authorization lookup failed: %s", search_tool_name, e + ) + return False + return True + + +def _is_search_request(request_kwargs: Mapping[str, object]) -> bool: + return request_kwargs.get("original_generic_function") is asearch + + def _token_in_metadata(metadata: object) -> UserAPIKeyAuth | None: try: return _RequestMetadata.model_validate(metadata).user_api_key_auth @@ -72,19 +91,21 @@ def _enforced_by_general_settings() -> bool: @dataclass(frozen=True, slots=True) class RouterFallbackAccessCheck: """ - `FallbackAccessCheck` for the proxy's router: while `is_enforced()` is true, a fallback target - is attempted only when the key behind the request could have requested it directly. Requests - that carry no key (for example internal health checks) are not restricted. + `FallbackAccessCheck` for the proxy's router: a fallback of a search request, and while `is_enforced()` + is true any other fallback, is attempted only when the key behind the request could have requested it + directly. Requests that carry no key (for example internal health checks) are not restricted. """ is_enforced: Callable[[], bool] async def __call__(self, *, model: str, request_kwargs: Mapping[str, object], llm_router: Router) -> bool: - if not self.is_enforced(): - return True valid_token: Final = _user_api_key_auth_from_request(request_kwargs) if valid_token is None: return True + if _is_search_request(request_kwargs): + return await is_search_tool_authorized_for_token(search_tool_name=model, valid_token=valid_token) + if not self.is_enforced(): + return True return await is_model_authorized_for_token(model=model, valid_token=valid_token, llm_router=llm_router) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 6d522759eb7..5d6ba0f8e13 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6714,12 +6714,15 @@ class ProxyConfig: general_settings = {} typed_general_settings: Final = _GENERAL_SETTINGS_VIEW.validate_python(general_settings) - if "vector_store_deny_by_default" in typed_general_settings: - ConfigGeneralSettings.model_validate( - MappingProxyType( - {"vector_store_deny_by_default": typed_general_settings["vector_store_deny_by_default"]} - ) + ConfigGeneralSettings.model_validate( + MappingProxyType( + { + name: typed_general_settings[name] + for name in ("vector_store_deny_by_default", "search_tool_deny_by_default") + if name in typed_general_settings + } ) + ) if general_settings.get("mcp_advertised_versions") is not None: from litellm.types.mcp import MCPAdvertisedVersions diff --git a/litellm/proxy/search_endpoints/endpoints.py b/litellm/proxy/search_endpoints/endpoints.py index 349eb5b7861..1e37e3243ea 100644 --- a/litellm/proxy/search_endpoints/endpoints.py +++ b/litellm/proxy/search_endpoints/endpoints.py @@ -135,52 +135,24 @@ async def search( if search_tool_name is not None: data["search_tool_name"] = search_tool_name - if not ( - data.get("search_tool_name") or data.get("model") or general_settings.get("completion_model") or user_model - ): + from litellm.proxy.auth.auth_checks import can_token_call_search_tool + from litellm.proxy.common_utils.http_parsing_utils import resolve_inference_model + + routed_search_tool_name: Final = data.get("search_tool_name") or resolve_inference_model( + data.get("model"), general_settings, user_model + ) + if not isinstance(routed_search_tool_name, str) or not routed_search_tool_name: raise ProxyMissingRequiredParamError(route="/search", param="search_tool_name") + try: + await can_token_call_search_tool(search_tool_name=routed_search_tool_name, valid_token=user_api_key_dict) + except ProxyException as e: + verbose_proxy_logger.debug("Search tool authorization denied: %s", e.type) + raise if "search_tool_name" in data and data["search_tool_name"]: data["model"] = data["search_tool_name"] search_tool_name_value: Final = data["search_tool_name"] - # Authorization check: verify key can access this search tool - from litellm.proxy.auth.auth_checks import ( - can_key_call_search_tool, - can_team_call_search_tool, - get_team_object, - ) - - try: - # Check key-level access - await can_key_call_search_tool( - search_tool_name=search_tool_name_value, - valid_token=user_api_key_dict, - ) - - # Check team-level access if key is associated with a team - if user_api_key_dict.team_id: - from litellm.proxy.proxy_server import ( - prisma_client, - proxy_logging_obj, - user_api_key_cache, - ) - - team_object: Final = await get_team_object( - team_id=user_api_key_dict.team_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=user_api_key_dict.parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - ) - await can_team_call_search_tool( - search_tool_name=search_tool_name_value, - team_object=team_object, - ) - except Exception as e: - verbose_proxy_logger.error("Search tool authorization failed for %s: %s", search_tool_name_value, e) - raise - if llm_router is not None and hasattr(llm_router, "search_tools"): verbose_proxy_logger.debug( "Search endpoint - Looking for search_tool_name: %s. Available search tools in router: %s. Total search tools: %s", diff --git a/litellm/proxy/search_endpoints/search_tool_management.py b/litellm/proxy/search_endpoints/search_tool_management.py index 49b868ac71a..1e6ed40bba3 100644 --- a/litellm/proxy/search_endpoints/search_tool_management.py +++ b/litellm/proxy/search_endpoints/search_tool_management.py @@ -2,7 +2,7 @@ CRUD ENDPOINTS FOR SEARCH TOOLS """ -from collections.abc import Awaitable, Callable, Sequence +from collections.abc import Awaitable, Callable, Mapping, Sequence from datetime import datetime from typing import Any, Final, TypeAlias @@ -117,34 +117,39 @@ async def _filter_visible_search_tools( search_tools: list[SearchToolInfoResponse], user_api_key_dict: UserAPIKeyAuth, lookup_team_object: TeamObjectLookup = _team_object_from_db, + general_settings: Mapping[str, object] | None = None, ) -> list[SearchToolInfoResponse]: """ Drop search tools the caller is not authorized to invoke, applying the same - key/team object_permission allowlists enforced on /search. Admins see all tools. + key/team/user grants enforced on /search. Admins see all tools unless + search_tool_deny_by_default applies to their credential. """ + from litellm.proxy.auth.auth_checks import ( + can_grants_view_search_tool, + is_search_tool_deny_by_default_applied, + resolve_search_tool_grants, + typed_general_settings, + ) + from litellm.proxy.proxy_server import general_settings as proxy_general_settings + + settings: Final = typed_general_settings(proxy_general_settings) if general_settings is None else general_settings if user_api_key_dict.user_role in ( LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, - ): + ) and not is_search_tool_deny_by_default_applied(user_api_key_dict, settings): return search_tools - from litellm.proxy.auth.auth_checks import can_user_view_search_tool - allowlist_team_id: Final = _allowlist_team_id(user_api_key_dict) - team_object: Final[LiteLLM_TeamTable | None] = ( - await lookup_team_object(allowlist_team_id, user_api_key_dict) if allowlist_team_id else None - ) - visible: Final[list[SearchToolInfoResponse]] = [] - for tool in search_tools: - tool_name = tool.get("search_tool_name") - if tool_name and await can_user_view_search_tool( - search_tool_name=tool_name, - valid_token=user_api_key_dict, - team_object=team_object, - ): - visible.append(tool) - return visible + async def _load_team_object() -> LiteLLM_TeamTable | None: + return await lookup_team_object(allowlist_team_id, user_api_key_dict) if allowlist_team_id else None + + grants: Final = await resolve_search_tool_grants(user_api_key_dict, settings, _load_team_object) + return [ + tool + for tool in search_tools + if (tool_name := tool.get("search_tool_name")) and can_grants_view_search_tool(tool_name, grants) + ] @router.get( diff --git a/tests/integration/authorization/test_search_tool_deny_by_default.py b/tests/integration/authorization/test_search_tool_deny_by_default.py new file mode 100644 index 00000000000..f29940f68b2 --- /dev/null +++ b/tests/integration/authorization/test_search_tool_deny_by_default.py @@ -0,0 +1,221 @@ +from __future__ import annotations + +import json +import uuid +from collections.abc import Iterator +from pathlib import Path +from typing import Final, Literal, TypeAlias + +import httpx +import pytest +import yaml +from integration._support.client import Gateway, Scenario, gateway_from_environment, object_value +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue + +PROXY_CONFIG: Final = Path(__file__).resolve().parents[1] / "proxy_config.yaml" +SEARCH_TOOL: Final = "integration-search" +OTHER_SEARCH_TOOL: Final = "integration-other-search" +FAILING_SEARCH_TOOL: Final = "integration-failing-search" +SEARCH_RESULT: Final = {"title": "Synthetic result", "url": "https://example.test/result", "content": "snippet"} +JsonObject: TypeAlias = dict[str, JsonValue] + + +def _json_array(*values: JsonValue) -> JsonValue: + return [*values] # mutable-ok: request payloads and YAML sequences require list values + + +def _grant(*search_tools: str) -> JsonObject: + permission: Final[JsonObject] = {"search_tools": _json_array(*search_tools)} + return permission + + +def _respond(request: Request) -> Reply: + if (request.method, request.target) == ("POST", "/failing/search"): + return Reply(status=500, body=b'{"error": "synthetic outage"}') + assert (request.method, request.target) == ("POST", "/tavily/search"), request + query: Final = json.loads(request.body)["query"] + return Reply(body=json.dumps({"query": query, "results": [SEARCH_RESULT]}).encode()) + + +def _config(directory: Path, wire: Wire) -> Path: + config: Final = object_value(yaml.safe_load(PROXY_CONFIG.read_text())) + tool_params: Final[JsonObject] = { + "search_provider": "tavily", + "api_key": "synthetic-tavily-key", + "api_base": f"{wire.url}/tavily", + } + failing_tool_params: Final[JsonObject] = {**tool_params, "api_base": f"{wire.url}/failing"} + strict: Final[JsonObject] = { + **config, + "general_settings": {**object_value(config["general_settings"]), "search_tool_deny_by_default": True}, + "router_settings": { + **object_value(config["router_settings"]), + "fallbacks": _json_array({FAILING_SEARCH_TOOL: _json_array(OTHER_SEARCH_TOOL)}), + }, + "search_tools": _json_array( + {"search_tool_name": SEARCH_TOOL, "litellm_params": tool_params}, + {"search_tool_name": OTHER_SEARCH_TOOL, "litellm_params": tool_params}, + {"search_tool_name": FAILING_SEARCH_TOOL, "litellm_params": failing_tool_params}, + ), + } + path: Final = directory / "proxy_search_tool_deny_by_default.yaml" + path.write_text(yaml.safe_dump(strict)) + return path + + +@pytest.fixture(scope="module") +def strict_search(tmp_path_factory: pytest.TempPathFactory) -> Iterator[tuple[Gateway, Wire]]: + with gateway_from_environment() as upstream, wire_server(_respond) as wire: + directory: Final = tmp_path_factory.mktemp("search_tool_deny_by_default") + with owned_proxy(upstream, directory, {}, config=_config(directory, wire)) as gateway: + yield gateway, wire + + +def _search(gateway: Gateway, key: str, marker: str) -> httpx.Response: + return gateway.request("POST", f"/v1/search/{SEARCH_TOOL}", {"query": marker}, key=key) + + +def _searched(wire: Wire, marker: str) -> tuple[Request, ...]: + return tuple(request for request in wire.drain() if marker in request.body.decode()) + + +Caller: TypeAlias = Literal[ + "standalone_key_no_permission", + "standalone_key_empty_grant", + "standalone_key_other_tool", + "team_key_empty_key_grant", + "team_key_empty_team_grant", + "proxy_admin_key_no_grant", +] + + +def _denied_key(scenario: Scenario, caller: Caller) -> tuple[str, str]: + if caller == "standalone_key_no_permission": + return scenario.key(), "key_search_tool_access_denied" + if caller == "standalone_key_empty_grant": + return scenario.key(object_permission=_grant()), "key_search_tool_access_denied" + if caller == "standalone_key_other_tool": + return scenario.key(object_permission=_grant(OTHER_SEARCH_TOOL)), "key_search_tool_access_denied" + if caller == "team_key_empty_key_grant": + team: Final = scenario.team(object_permission=_grant(SEARCH_TOOL)) + return scenario.key(team_id=team, object_permission=_grant()), "key_search_tool_access_denied" + if caller == "team_key_empty_team_grant": + empty_team: Final = scenario.team(object_permission=_grant()) + return scenario.key(team_id=empty_team, object_permission=_grant(SEARCH_TOOL)), "team_search_tool_access_denied" + admin: Final = scenario.user(user_role="proxy_admin") + return scenario.key(user_id=admin), "key_search_tool_access_denied" + + +@pytest.mark.parametrize( + "caller", + [ + "standalone_key_no_permission", + "standalone_key_empty_grant", + "standalone_key_other_tool", + "team_key_empty_key_grant", + "team_key_empty_team_grant", + "proxy_admin_key_no_grant", + ], +) +def test_deny_by_default_rejects_ungranted_search_before_the_provider_is_called( + strict_search: tuple[Gateway, Wire], caller: Caller +) -> None: + gateway, wire = strict_search + with gateway.scenario() as scenario: + key, error_type = _denied_key(scenario, caller) + marker: Final = f"search deny {caller} {uuid.uuid4().hex}" + + response: Final = _search(gateway, key, marker) + + assert response.status_code == 403, response.text + assert response.json()["error"]["type"] == error_type, response.text + assert _searched(wire, marker) == () + + +@pytest.mark.parametrize("scope", ["standalone_key", "team_key", "master_key"]) +def test_deny_by_default_serves_explicitly_granted_search( + strict_search: tuple[Gateway, Wire], scope: Literal["standalone_key", "team_key", "master_key"] +) -> None: + gateway, wire = strict_search + with gateway.scenario() as scenario: + if scope == "standalone_key": + key = scenario.key(object_permission=_grant(SEARCH_TOOL)) + elif scope == "team_key": + team: Final = scenario.team(object_permission=_grant(SEARCH_TOOL)) + key = scenario.key(team_id=team, object_permission=_grant(SEARCH_TOOL)) + else: + key = gateway.key + marker: Final = f"search allow {scope} {uuid.uuid4().hex}" + + response: Final = _search(gateway, key, marker) + + assert response.status_code == 200, response.text + assert response.json()["results"][0]["url"] == SEARCH_RESULT["url"] + assert len(_searched(wire, marker)) == 1 + + +def test_search_tools_list_shows_only_the_tools_the_key_and_its_team_both_grant( + strict_search: tuple[Gateway, Wire], +) -> None: + gateway, _ = strict_search + with gateway.scenario() as scenario: + team: Final = scenario.team(object_permission=_grant(SEARCH_TOOL, OTHER_SEARCH_TOOL)) + member: Final = scenario.member(team) + key: Final = scenario.key(team_id=team, user_id=member, object_permission=_grant(SEARCH_TOOL)) + + response: Final = gateway.request("GET", "/search_tools/list", key=key) + + assert response.status_code == 200, response.text + assert [tool["search_tool_name"] for tool in response.json()["search_tools"]] == [SEARCH_TOOL] + + +def test_revoking_a_team_search_grant_takes_effect_on_the_next_request(strict_search: tuple[Gateway, Wire]) -> None: + gateway, wire = strict_search + with gateway.scenario() as scenario: + team: Final = scenario.team(object_permission=_grant(SEARCH_TOOL)) + key: Final = scenario.key(team_id=team, object_permission=_grant(SEARCH_TOOL)) + assert _search(gateway, key, f"search warm {uuid.uuid4().hex}").status_code == 200 + + gateway.post("/team/update", {"team_id": team, "object_permission": _grant()}) + marker: Final = f"search revoked {uuid.uuid4().hex}" + response: Final = _search(gateway, key, marker) + + assert response.status_code == 403, response.text + assert response.json()["error"]["type"] == "team_search_tool_access_denied", response.text + assert _searched(wire, marker) == () + + +def test_deny_by_default_authorizes_a_search_tool_named_by_the_model_field(strict_search: tuple[Gateway, Wire]) -> None: + gateway, wire = strict_search + with gateway.scenario() as scenario: + key: Final = scenario.key() + marker: Final = f"search model field {uuid.uuid4().hex}" + + response: Final = gateway.request("POST", "/v1/search", {"model": SEARCH_TOOL, "query": marker}, key=key) + + assert response.status_code == 403, response.text + assert response.json()["error"]["type"] == "key_search_tool_access_denied", response.text + assert _searched(wire, marker) == () + + +@pytest.mark.parametrize("fallback_granted", [False, True]) +def test_router_search_fallback_only_reaches_a_granted_tool( + strict_search: tuple[Gateway, Wire], fallback_granted: bool +) -> None: + gateway, wire = strict_search + with gateway.scenario() as scenario: + granted: Final = (FAILING_SEARCH_TOOL, OTHER_SEARCH_TOOL) if fallback_granted else (FAILING_SEARCH_TOOL,) + key: Final = scenario.key(object_permission=_grant(*granted)) + marker: Final = f"search fallback {fallback_granted} {uuid.uuid4().hex}" + + response: Final = gateway.request("POST", f"/v1/search/{FAILING_SEARCH_TOOL}", {"query": marker}, key=key) + + searched: Final = tuple(request.target for request in _searched(wire, marker)) + if fallback_granted: + assert response.status_code == 200, response.text + assert searched == ("/failing/search", "/tavily/search") + else: + assert response.status_code == 500, response.text + assert searched == ("/failing/search",) diff --git a/tests/unit/integrations/websearch_interception/test_websearch_interception_handler.py b/tests/unit/integrations/websearch_interception/test_websearch_interception_handler.py index b6a7c2ccb93..b7437b5fd8f 100644 --- a/tests/unit/integrations/websearch_interception/test_websearch_interception_handler.py +++ b/tests/unit/integrations/websearch_interception/test_websearch_interception_handler.py @@ -13,7 +13,13 @@ from litellm.integrations.websearch_interception.handler import ( WebSearchInterceptionLogger, ) from litellm.llms.base_llm.search.transformation import SearchResponse -from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, ProxyException, UserAPIKeyAuth +from litellm.proxy._types import ( + LiteLLM_ObjectPermissionTable, + LiteLLM_TeamTable, + LitellmUserRoles, + ProxyException, + UserAPIKeyAuth, +) from litellm.types.utils import LlmProviders @@ -899,3 +905,184 @@ async def test_pre_request_hook_syncs_forced_tool_choice(): "type": "tool", "name": LITELLM_WEB_SEARCH_TOOL_NAME, } + + +def _single_search_tool_router(search_tool_name): + router = MagicMock() + router.search_tools = [ + { + "search_tool_name": search_tool_name, + "litellm_params": {"search_provider": "tavily", "api_key": "fake-ui-key"}, + } + ] + return router + + +def _virtual_key( + team_id: str | None = None, key_search_tools: list[str] | None = None, api_key: str = "sk-caller" +) -> UserAPIKeyAuth: + token = UserAPIKeyAuth( + api_key=api_key, + user_id="user-1", + team_id=team_id, + object_permission_id=None if key_search_tools is None else "op-key", + object_permission=( + None + if key_search_tools is None + else LiteLLM_ObjectPermissionTable(object_permission_id="op-key", search_tools=key_search_tools) + ), + ) + token.via_virtual_key = True + return token + + +@pytest.mark.parametrize( + "general_settings, key_search_tools, team_search_tools, expect_search", + [ + ({}, None, [], True), + ({"search_tool_deny_by_default": True}, ["team-search"], [], False), + ({"search_tool_deny_by_default": True}, ["team-search"], None, False), + ({"search_tool_deny_by_default": True}, [], ["team-search"], False), + ({"search_tool_deny_by_default": True}, ["team-search"], ["team-search"], True), + ], +) +@pytest.mark.asyncio +async def test_execute_search_team_key_follows_search_tool_deny_by_default( + monkeypatch: pytest.MonkeyPatch, + general_settings: dict[str, bool], + key_search_tools: list[str] | None, + team_search_tools: list[str] | None, + expect_search: bool, +): + import litellm + from litellm.proxy import proxy_server + + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"], search_tool_name="team-search") + mock_asearch = AsyncMock(return_value=SearchResponse(object="search", results=[])) + team_object = LiteLLM_TeamTable( + team_id="team-1", + object_permission_id=None if team_search_tools is None else "op-team", + object_permission=( + None + if team_search_tools is None + else LiteLLM_ObjectPermissionTable(object_permission_id="op-team", search_tools=team_search_tools) + ), + ) + + monkeypatch.setattr(proxy_server, "llm_router", _single_search_tool_router("team-search")) + monkeypatch.setattr(proxy_server, "general_settings", general_settings) + monkeypatch.setattr(proxy_server, "prisma_client", MagicMock()) + monkeypatch.setattr(litellm, "asearch", mock_asearch) + monkeypatch.setattr("litellm.proxy.auth.auth_checks.get_team_object", AsyncMock(return_value=team_object)) + kwargs = {"metadata": {"user_api_key_auth": _virtual_key("team-1", key_search_tools)}} + + if expect_search: + await logger._execute_search("what is litellm", kwargs=kwargs) + mock_asearch.assert_awaited_once() + else: + with pytest.raises(ProxyException) as exc_info: + await logger._execute_search("what is litellm", kwargs=kwargs) + assert exc_info.value.code == "403" + mock_asearch.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_execute_search_key_without_grant_is_denied_under_search_tool_deny_by_default( + monkeypatch: pytest.MonkeyPatch, +): + import litellm + from litellm.proxy import proxy_server + + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"], search_tool_name="any-search") + mock_asearch = AsyncMock(return_value=SearchResponse(object="search", results=[])) + + monkeypatch.setattr(proxy_server, "llm_router", _single_search_tool_router("any-search")) + monkeypatch.setattr(proxy_server, "general_settings", {"search_tool_deny_by_default": True}) + monkeypatch.setattr(proxy_server, "prisma_client", MagicMock()) + monkeypatch.setattr(litellm, "asearch", mock_asearch) + + with pytest.raises(ProxyException) as exc_info: + await logger._execute_search("what is litellm", kwargs={"metadata": {"user_api_key_auth": _virtual_key()}}) + + assert exc_info.value.code == "403" + mock_asearch.assert_not_awaited() + + +@pytest.mark.parametrize( + "general_settings, caller, expect_search", + [ + ({}, _virtual_key(), True), + ({"search_tool_deny_by_default": False}, _virtual_key(), True), + ({"search_tool_deny_by_default": True}, _virtual_key(key_search_tools=["any"]), False), + ( + {"search_tool_deny_by_default": True}, + UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN), + False, + ), + ({"search_tool_deny_by_default": True}, _virtual_key(api_key="litellm_proxy_master_key"), True), + ({"search_tool_deny_by_default": True}, None, True), + ], + ids=["omitted", "false", "virtual-key", "proxy-admin-not-exempt", "master-key-exempt", "sdk-without-proxy-auth"], +) +@pytest.mark.asyncio +async def test_execute_search_unregistered_fallback_follows_search_tool_deny_by_default( + monkeypatch: pytest.MonkeyPatch, + general_settings: dict[str, bool], + caller: UserAPIKeyAuth | None, + expect_search: bool, +): + import litellm + from litellm.proxy import proxy_server + + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"]) + router = MagicMock() + router.search_tools = [] + mock_asearch = AsyncMock(return_value=SearchResponse(object="search", results=[])) + + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(proxy_server, "general_settings", general_settings) + monkeypatch.setattr(litellm, "asearch", mock_asearch) + kwargs = {"metadata": {} if caller is None else {"user_api_key_auth": caller}} + + if expect_search: + await logger._execute_search("what is litellm", kwargs=kwargs) + assert mock_asearch.await_args.kwargs["search_provider"] == "perplexity" + else: + with pytest.raises(ProxyException) as exc_info: + await logger._execute_search("what is litellm", kwargs=kwargs) + assert exc_info.value.code == "403" + mock_asearch.assert_not_awaited() + + +@pytest.mark.parametrize( + "general_settings, expect_search", + [({}, True), ({"search_tool_deny_by_default": True}, False)], +) +@pytest.mark.asyncio +async def test_execute_search_registered_tool_without_provider_follows_unregistered_fallback_policy( + monkeypatch: pytest.MonkeyPatch, + general_settings: dict[str, bool], + expect_search: bool, +): + import litellm + from litellm.proxy import proxy_server + + logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"], search_tool_name="no-provider") + router = MagicMock() + router.search_tools = [{"search_tool_name": "no-provider", "litellm_params": {}}] + mock_asearch = AsyncMock(return_value=SearchResponse(object="search", results=[])) + + monkeypatch.setattr(proxy_server, "llm_router", router) + monkeypatch.setattr(proxy_server, "general_settings", general_settings) + monkeypatch.setattr(proxy_server, "prisma_client", MagicMock()) + monkeypatch.setattr(litellm, "asearch", mock_asearch) + kwargs = {"metadata": {"user_api_key_auth": _virtual_key(key_search_tools=["no-provider"])}} + + if expect_search: + await logger._execute_search("what is litellm", kwargs=kwargs) + assert mock_asearch.await_args.kwargs["search_provider"] == "perplexity" + else: + with pytest.raises(ProxyException) as exc_info: + await logger._execute_search("what is litellm", kwargs=kwargs) + assert exc_info.value.code == "403" + mock_asearch.assert_not_awaited() diff --git a/tests/unit/proxy/auth/test_fallback_model_access.py b/tests/unit/proxy/auth/test_fallback_model_access.py index 4cbf4474596..24fcd4671c8 100644 --- a/tests/unit/proxy/auth/test_fallback_model_access.py +++ b/tests/unit/proxy/auth/test_fallback_model_access.py @@ -1,12 +1,13 @@ import pytest from litellm import Router -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth from litellm.proxy.auth.fallback_model_access import ( RouterFallbackAccessCheck, is_model_authorized_for_token, router_fallback_access_check, ) +from litellm.search import asearch def _router() -> Router: @@ -105,3 +106,69 @@ async def test_proxy_check_reads_enforce_fallback_model_access_from_general_sett ) is expected ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("check", [ENFORCED, NOT_ENFORCED], ids=["enforced", "not-enforced"]) +async def test_search_tool_fallback_target_follows_the_key_search_tool_grant( + monkeypatch: pytest.MonkeyPatch, check: RouterFallbackAccessCheck +): + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) + router = Router( + model_list=[], + search_tools=[ + {"search_tool_name": name, "litellm_params": {"search_provider": "tavily", "api_key": "k"}} + for name in ("search-a", "search-b") + ], + ) + request_kwargs = { + "original_generic_function": asearch, + "litellm_metadata": { + "user_api_key_auth": UserAPIKeyAuth( + api_key="hashed", + object_permission_id="op-key", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="op-key", search_tools=["search-a"] + ), + ) + }, + } + + assert await check(model="search-a", request_kwargs=request_kwargs, llm_router=router) + assert not await check(model="search-b", request_kwargs=request_kwargs, llm_router=router) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "request_kwargs, expected", + [ + ({"original_generic_function": asearch}, True), + ({}, False), + ], + ids=["search-request", "completion-request"], +) +async def test_fallback_named_like_both_a_model_and_a_search_tool_follows_the_request_kind( + request_kwargs: dict, expected: bool +): + router = Router( + model_list=[ + {"model_name": "shared-name", "litellm_params": {"model": "openai/secret", "api_key": "k"}}, + ], + search_tools=[ + {"search_tool_name": "shared-name", "litellm_params": {"search_provider": "tavily", "api_key": "k"}}, + ], + ) + key = UserAPIKeyAuth( + api_key="hashed", + models=["open-model"], + object_permission_id="op-key", + object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-key", search_tools=["shared-name"]), + ) + + allowed = await ENFORCED( + model="shared-name", + request_kwargs={**request_kwargs, "metadata": {"user_api_key_auth": key}}, + llm_router=router, + ) + + assert allowed is expected diff --git a/tests/unit/proxy/auth/test_search_tool_deny_by_default.py b/tests/unit/proxy/auth/test_search_tool_deny_by_default.py new file mode 100644 index 00000000000..cd285e0386f --- /dev/null +++ b/tests/unit/proxy/auth/test_search_tool_deny_by_default.py @@ -0,0 +1,270 @@ +from collections.abc import Mapping +from typing import Final, TypedDict +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS +from litellm.proxy import proxy_server +from litellm.proxy._types import ( + UI_TEAM_ID, + LiteLLM_ObjectPermissionTable, + LiteLLM_TeamTable, + LiteLLM_UserTable, + LitellmUserRoles, + ProxyErrorTypes, + ProxyException, + UserAPIKeyAuth, +) +from litellm.proxy.auth.auth_checks import ( + TeamObjectLoader, + can_caller_call_search_tool, + check_unregistered_search_fallback, +) +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key + +_DENY_ON: Final = {"search_tool_deny_by_default": True} +_KEY: Final = ProxyErrorTypes.key_search_tool_access_denied +_TEAM: Final = ProxyErrorTypes.team_search_tool_access_denied +_USER: Final = ProxyErrorTypes.user_search_tool_access_denied +_LEGACY: Final = ProxyErrorTypes.key_model_access_denied +_TEAM_LOAD_FAILS: Final = "team-load-fails" +_NO_ROW: Final = "no-row" + + +class CallerFields(TypedDict, total=False): + virtual_key: bool + key_tools: list[str] | None | str + team_id: str | None + user_id: str | None + user_role: LitellmUserRoles | None + api_key: str + + +def _grant(search_tools: list[str] | None, object_permission_id: str = "op") -> LiteLLM_ObjectPermissionTable: + return LiteLLM_ObjectPermissionTable(object_permission_id=object_permission_id, search_tools=search_tools) + + +def _caller( + virtual_key: bool = True, + key_tools: list[str] | None | str = _NO_ROW, + team_id: str | None = None, + user_id: str | None = "user-1", + user_role: LitellmUserRoles | None = LitellmUserRoles.INTERNAL_USER, + api_key: str = "sk-caller", +) -> UserAPIKeyAuth: + token: Final = UserAPIKeyAuth( + api_key=api_key, + user_id=user_id, + user_role=user_role, + team_id=team_id, + object_permission_id=None if key_tools == _NO_ROW else "op-key", + object_permission=None if key_tools == _NO_ROW else _grant(key_tools, "op-key"), + ) + token.via_virtual_key = virtual_key + return token + + +def _team_loader(team_tools: list[str] | None | str) -> TeamObjectLoader: + async def load() -> LiteLLM_TeamTable | None: + if team_tools == _TEAM_LOAD_FAILS: + raise ProxyException(message="team lookup failed", type="auth_error", param="team_id", code=404) + if team_tools == _NO_ROW: + return LiteLLM_TeamTable(team_id="team-1") + return LiteLLM_TeamTable( + team_id="team-1", object_permission_id="op-team", object_permission=_grant(team_tools, "op-team") + ) + + return load + + +async def _no_team() -> None: + return None + + +@pytest.fixture +def cache(monkeypatch: pytest.MonkeyPatch) -> UserApiKeyCache: + user_api_key_cache: Final = UserApiKeyCache() + monkeypatch.setattr(proxy_server, "user_api_key_cache", user_api_key_cache) + monkeypatch.setattr(proxy_server, "prisma_client", MagicMock()) + return user_api_key_cache + + +def _cache_user(cache: UserApiKeyCache, user_tools: list[str] | None) -> None: + cache.set_cache(key="user-1", value=LiteLLM_UserTable(user_id="user-1", object_permission_id="op-user")) + cache.set_cache(key=object_permission_cache_key("op-user"), value=_grant(user_tools, "op-user")) + + +async def _denied_by( + general_settings: Mapping[str, object], caller: UserAPIKeyAuth, load_team: TeamObjectLoader = _no_team +) -> ProxyErrorTypes | None: + try: + await can_caller_call_search_tool("search-a", caller, general_settings, load_team) + except ProxyException as e: + denial = e + else: + return None + assert (denial.code, denial.param) == ("403", "search_tool_name") + return ProxyErrorTypes(denial.type) + + +@pytest.mark.parametrize( + "general_settings", + [{}, {"search_tool_deny_by_default": False}], + ids=["omitted", "false"], +) +@pytest.mark.parametrize( + "caller, team_tools, expected", + [ + pytest.param(CallerFields(), None, None, id="no grants anywhere"), + pytest.param(CallerFields(key_tools=[]), None, None, id="empty key list"), + pytest.param(CallerFields(team_id="team-1"), [], None, id="empty team list"), + pytest.param(CallerFields(key_tools=["search-b"]), None, _LEGACY, id="key allowlist excludes"), + pytest.param(CallerFields(team_id="team-1"), ["search-b"], _LEGACY, id="team allowlist excludes"), + pytest.param(CallerFields(key_tools=["search-a"], team_id="team-1"), ["search-a"], None, id="both allow"), + ], +) +@pytest.mark.asyncio +async def test_search_tool_access_is_unchanged_without_search_tool_deny_by_default( + cache: UserApiKeyCache, + general_settings: dict[str, bool], + caller: CallerFields, + team_tools: list[str] | None, + expected: ProxyErrorTypes | None, +): + _cache_user(cache, ["search-b"]) + load_team: Final = _no_team if team_tools is None else _team_loader(team_tools) + assert await _denied_by(general_settings, _caller(**caller), load_team) == expected + + +@pytest.mark.parametrize( + "caller, team_tools, user_tools, expected", + [ + pytest.param(CallerFields(), None, ["search-a"], _KEY, id="standalone key, no permission row"), + pytest.param(CallerFields(key_tools=None), None, None, _KEY, id="standalone key, search_tools null"), + pytest.param(CallerFields(key_tools=[]), None, None, _KEY, id="standalone key, empty grant"), + pytest.param(CallerFields(key_tools=["search-b"]), None, None, _KEY, id="standalone key, grants another tool"), + pytest.param(CallerFields(key_tools=["search-a"]), None, [], None, id="standalone key, key grant is enough"), + pytest.param(CallerFields(key_tools=[]), None, ["search-a"], _KEY, id="standalone key, user cannot stand in"), + pytest.param( + CallerFields(key_tools=["search-a"], team_id="team-1"), ["search-a"], None, None, id="team key, both" + ), + pytest.param(CallerFields(key_tools=[], team_id="team-1"), ["search-a"], None, _KEY, id="team key, empty key"), + pytest.param( + CallerFields(key_tools=["search-a"], team_id="team-1"), [], None, _TEAM, id="team key, empty team" + ), + pytest.param( + CallerFields(key_tools=["search-a"], team_id="team-1"), None, None, _TEAM, id="team key, team null" + ), + pytest.param( + CallerFields(key_tools=["search-a"], team_id="team-1"), _NO_ROW, None, _TEAM, id="team key, team has no row" + ), + pytest.param( + CallerFields(key_tools=["search-a"], team_id="team-1"), + _TEAM_LOAD_FAILS, + None, + _TEAM, + id="team key, team fails to load", + ), + pytest.param( + CallerFields(key_tools=["search-a"], team_id="team-1"), + ["search-b"], + ["search-a"], + _TEAM, + id="team key, user cannot stand in for team", + ), + pytest.param( + CallerFields(virtual_key=False, team_id="team-1"), ["search-a"], [], None, id="keyless team member" + ), + pytest.param( + CallerFields(virtual_key=False, team_id="team-1"), [], ["search-a"], _TEAM, id="keyless, empty team" + ), + pytest.param(CallerFields(virtual_key=False), None, ["search-a"], None, id="keyless user, user grants"), + pytest.param(CallerFields(virtual_key=False), None, None, _USER, id="keyless user, search_tools null"), + pytest.param(CallerFields(virtual_key=False), None, [], _USER, id="keyless user, empty grant"), + pytest.param( + CallerFields(virtual_key=False, user_id="user-2"), None, None, _USER, id="keyless user fails to load" + ), + pytest.param(CallerFields(virtual_key=False, user_id=None), None, None, _USER, id="keyless caller, no user"), + pytest.param( + CallerFields(user_role=LitellmUserRoles.PROXY_ADMIN), + None, + None, + _KEY, + id="proxy admin virtual key not exempt", + ), + pytest.param(CallerFields(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS), None, None, None, id="master key exempt"), + pytest.param( + CallerFields(team_id=UI_TEAM_ID, virtual_key=False), None, None, None, id="dashboard session exempt" + ), + ], +) +@pytest.mark.asyncio +async def test_search_tool_deny_by_default_requires_every_owning_identity_to_grant( + cache: UserApiKeyCache, + caller: CallerFields, + team_tools: list[str] | str | None, + user_tools: list[str] | None, + expected: ProxyErrorTypes | None, +): + if user_tools is not None: + _cache_user(cache, user_tools) + elif caller.get("user_id", "user-1") == "user-1": + _cache_user(cache, None) + load_team: Final = _no_team if team_tools is None and "team_id" not in caller else _team_loader(team_tools) + assert await _denied_by(_DENY_ON, _caller(**caller), load_team) == expected + + +@pytest.mark.parametrize("value", ["true", "enabled", 1], ids=["string-true", "string", "int"]) +@pytest.mark.asyncio +async def test_non_boolean_search_tool_deny_by_default_enables_the_policy(cache: UserApiKeyCache, value: object): + assert await _denied_by({"search_tool_deny_by_default": value}, _caller(key_tools=[])) == _KEY + + +@pytest.mark.asyncio +async def test_search_tool_deny_by_default_denies_when_no_database_is_connected( + cache: UserApiKeyCache, monkeypatch: pytest.MonkeyPatch +): + monkeypatch.setattr(proxy_server, "prisma_client", None) + assert await _denied_by(_DENY_ON, _caller(key_tools=["search-a"])) == _KEY + + +@pytest.mark.asyncio +async def test_search_tool_deny_by_default_reads_the_user_grant_on_every_call(cache: UserApiKeyCache): + caller: Final = _caller(virtual_key=False) + _cache_user(cache, ["search-a"]) + assert await _denied_by(_DENY_ON, caller) is None + cache.set_cache(key=object_permission_cache_key("op-user"), value=_grant([], "op-user")) + assert await _denied_by(_DENY_ON, caller) == _USER + + +@pytest.mark.asyncio +async def test_search_tool_deny_by_default_does_not_load_the_user_when_off( + cache: UserApiKeyCache, monkeypatch: pytest.MonkeyPatch +): + get_user_object: Final = AsyncMock() + monkeypatch.setattr("litellm.proxy.auth.auth_checks.get_user_object", get_user_object) + assert await _denied_by({}, _caller(virtual_key=False)) is None + get_user_object.assert_not_awaited() + + +@pytest.mark.parametrize( + "general_settings, caller, expected", + [ + pytest.param({}, _caller(), None, id="flag off"), + pytest.param(_DENY_ON, _caller(), _KEY, id="virtual key"), + pytest.param(_DENY_ON, _caller(virtual_key=False, team_id="team-1"), _TEAM, id="keyless team member"), + pytest.param(_DENY_ON, _caller(virtual_key=False), _USER, id="keyless user"), + pytest.param(_DENY_ON, _caller(user_role=LitellmUserRoles.PROXY_ADMIN), _KEY, id="proxy admin key"), + pytest.param(_DENY_ON, _caller(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS), None, id="master key"), + ], +) +def test_unregistered_search_fallback_follows_search_tool_deny_by_default( + general_settings: Mapping[str, object], caller: UserAPIKeyAuth, expected: ProxyErrorTypes | None +): + if expected is None: + assert check_unregistered_search_fallback(caller, general_settings) is True + return + with pytest.raises(ProxyException) as exc_info: + check_unregistered_search_fallback(caller, general_settings) + assert (exc_info.value.type, exc_info.value.code) == (expected, "403") diff --git a/tests/unit/proxy/management_endpoints/search_endpoints/test_search_tool_management.py b/tests/unit/proxy/management_endpoints/search_endpoints/test_search_tool_management.py index 582f6e82f11..49bc2d0f64e 100644 --- a/tests/unit/proxy/management_endpoints/search_endpoints/test_search_tool_management.py +++ b/tests/unit/proxy/management_endpoints/search_endpoints/test_search_tool_management.py @@ -1,14 +1,16 @@ import contextlib import json -from types import SimpleNamespace +from collections.abc import Callable from datetime import datetime +from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import HTTPException from fastapi.testclient import TestClient - +# Import proxy_server module first to ensure it's initialized +import litellm.proxy.proxy_server as ps from litellm.proxy._types import ( LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, @@ -16,9 +18,6 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) -# Import proxy_server module first to ensure it's initialized -import litellm.proxy.proxy_server as ps - # Now we can safely import app from litellm.proxy.proxy_server import app from litellm.types.search import SearchToolInfoResponse @@ -116,9 +115,7 @@ async def test_list_search_tools_config_only(monkeypatch): with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): # Mock proxy_config mock_proxy_config = MagicMock() - mock_proxy_config.get_config = AsyncMock( - return_value={"search_tools": config_tools} - ) + mock_proxy_config.get_config = AsyncMock(return_value={"search_tools": config_tools}) mock_proxy_config.parse_search_tools = MagicMock(return_value=config_tools) with patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config): # Mock auth @@ -192,9 +189,7 @@ async def test_list_search_tools_filters_duplicate_config_tools(monkeypatch): with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): # Mock proxy_config mock_proxy_config = MagicMock() - mock_proxy_config.get_config = AsyncMock( - return_value={"search_tools": config_tools} - ) + mock_proxy_config.get_config = AsyncMock(return_value={"search_tools": config_tools}) mock_proxy_config.parse_search_tools = MagicMock(return_value=config_tools) with patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config): # Mock auth @@ -215,11 +210,7 @@ async def test_list_search_tools_filters_duplicate_config_tools(monkeypatch): # Verify DB tool is present db_tool = next( - ( - t - for t in data["search_tools"] - if t["search_tool_name"] == "existing-tool" - ), + (t for t in data["search_tools"] if t["search_tool_name"] == "existing-tool"), None, ) assert db_tool is not None @@ -232,11 +223,7 @@ async def test_list_search_tools_filters_duplicate_config_tools(monkeypatch): # Verify unique config tool is present config_tool = next( - ( - t - for t in data["search_tools"] - if t["search_tool_name"] == "unique-config-tool" - ), + (t for t in data["search_tools"] if t["search_tool_name"] == "unique-config-tool"), None, ) assert config_tool is not None @@ -247,8 +234,7 @@ async def test_list_search_tools_filters_duplicate_config_tools(monkeypatch): ( t for t in data["search_tools"] - if t["search_tool_name"] == "existing-tool" - and t["is_from_config"] is True + if t["search_tool_name"] == "existing-tool" and t["is_from_config"] is True ), None, ) @@ -323,11 +309,7 @@ async def test_list_search_tools_datetime_conversion(monkeypatch): # Test datetime conversion for tool 1 tool1 = next( - ( - t - for t in data["search_tools"] - if t["search_tool_name"] == "datetime-test-tool" - ), + (t for t in data["search_tools"] if t["search_tool_name"] == "datetime-test-tool"), None, ) assert tool1 is not None @@ -341,11 +323,7 @@ async def test_list_search_tools_datetime_conversion(monkeypatch): # Test None handling for tool 2 tool2 = next( - ( - t - for t in data["search_tools"] - if t["search_tool_name"] == "null-datetime-tool" - ), + (t for t in data["search_tools"] if t["search_tool_name"] == "null-datetime-tool"), None, ) assert tool2 is not None @@ -357,11 +335,7 @@ async def test_list_search_tools_datetime_conversion(monkeypatch): # Test string passthrough for tool 3 tool3 = next( - ( - t - for t in data["search_tools"] - if t["search_tool_name"] == "string-datetime-tool" - ), + (t for t in data["search_tools"] if t["search_tool_name"] == "string-datetime-tool"), None, ) assert tool3 is not None @@ -401,9 +375,7 @@ async def test_list_search_tools_config_error_handling(monkeypatch): with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): # Mock proxy_config to raise an error mock_proxy_config = MagicMock() - mock_proxy_config.get_config = AsyncMock( - side_effect=Exception("Config error") - ) + mock_proxy_config.get_config = AsyncMock(side_effect=Exception("Config error")) with patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config): # Mock auth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth @@ -422,13 +394,8 @@ async def test_list_search_tools_config_error_handling(monkeypatch): assert len(data["search_tools"]) == 1 assert data["search_tools"][0]["search_tool_name"] == "db-tool-1" # Verify masking of sensitive values - assert ( - data["search_tools"][0]["litellm_params"]["api_key"] - != "sk-test" - ) - assert ( - "****" in data["search_tools"][0]["litellm_params"]["api_key"] - ) + assert data["search_tools"][0]["litellm_params"]["api_key"] != "sk-test" + assert "****" in data["search_tools"][0]["litellm_params"]["api_key"] finally: app.dependency_overrides.pop(user_api_key_auth, None) @@ -543,31 +510,18 @@ async def test_list_search_tools_db_masking_sensitive_values(monkeypatch): # Test tool 1: api_key should be masked tool1 = next( - ( - t - for t in data["search_tools"] - if t["search_tool_name"] == "perplexity-tool" - ), + (t for t in data["search_tools"] if t["search_tool_name"] == "perplexity-tool"), None, ) assert tool1 is not None - assert ( - tool1["litellm_params"]["api_key"] != "pplx-sk-9876567890abcdef" - ) + assert tool1["litellm_params"]["api_key"] != "pplx-sk-9876567890abcdef" assert "****" in tool1["litellm_params"]["api_key"] assert tool1["litellm_params"]["search_provider"] == "perplexity" - assert ( - tool1["litellm_params"]["api_base"] - == "https://api.perplexity.ai" - ) + assert tool1["litellm_params"]["api_base"] == "https://api.perplexity.ai" # Test tool 2: api_key should be masked tool2 = next( - ( - t - for t in data["search_tools"] - if t["search_tool_name"] == "tavily-tool" - ), + (t for t in data["search_tools"] if t["search_tool_name"] == "tavily-tool"), None, ) assert tool2 is not None @@ -577,29 +531,18 @@ async def test_list_search_tools_db_masking_sensitive_values(monkeypatch): # Test tool 3: access_token and secret_key should be masked tool3 = next( - ( - t - for t in data["search_tools"] - if t["search_tool_name"] == "tool-with-token" - ), + (t for t in data["search_tools"] if t["search_tool_name"] == "tool-with-token"), None, ) assert tool3 is not None - assert ( - tool3["litellm_params"]["access_token"] - != "token-abcdefghijklmnop" - ) + assert tool3["litellm_params"]["access_token"] != "token-abcdefghijklmnop" assert "****" in tool3["litellm_params"]["access_token"] assert tool3["litellm_params"]["secret_key"] != "secret-xyz123" assert "****" in tool3["litellm_params"]["secret_key"] # Test tool 4: non-sensitive fields should remain unmasked tool4 = next( - ( - t - for t in data["search_tools"] - if t["search_tool_name"] == "tool-with-non-sensitive" - ), + (t for t in data["search_tools"] if t["search_tool_name"] == "tool-with-non-sensitive"), None, ) assert tool4 is not None @@ -615,6 +558,7 @@ async def test_get_all_search_tools_from_db_retries_on_transport_error(): """`SearchToolRegistry.get_all_search_tools_from_db` self-heals across one ClientNotConnectedError via call_with_db_reconnect_retry.""" import prisma + from litellm.proxy.search_endpoints.search_tool_registry import ( SearchToolRegistry, ) @@ -628,25 +572,18 @@ async def test_get_all_search_tools_from_db_retries_on_transport_error(): return [] mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_searchtoolstable.find_many = AsyncMock( - side_effect=_flaky_find_many - ) + mock_prisma_client.db.litellm_searchtoolstable.find_many = AsyncMock(side_effect=_flaky_find_many) mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=True) mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0 mock_prisma_client._db_auth_reconnect_lock_timeout_seconds = 0.1 - result = await SearchToolRegistry.get_all_search_tools_from_db( - prisma_client=mock_prisma_client - ) + result = await SearchToolRegistry.get_all_search_tools_from_db(prisma_client=mock_prisma_client) assert result == [] assert len(invocations) == 2 mock_prisma_client.attempt_db_reconnect.assert_awaited_once() reconnect_kwargs = mock_prisma_client.attempt_db_reconnect.await_args.kwargs - assert ( - reconnect_kwargs["reason"] - == "get_all_search_tools_from_db_lookup_failure" - ) + assert reconnect_kwargs["reason"] == "get_all_search_tools_from_db_lookup_failure" @contextlib.contextmanager @@ -749,9 +686,7 @@ async def test_list_search_tools_scoped_to_key_object_permission(): @pytest.mark.asyncio async def test_list_search_tools_unrestricted_internal_user_sees_all(): """An internal user with no search_tools allowlist is unrestricted and sees every tool.""" - unrestricted_user = UserAPIKeyAuth( - user_role=LitellmUserRoles.INTERNAL_USER, user_id="internal_user" - ) + unrestricted_user = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="internal_user") with ( _mock_search_tool_backend(_scoping_db_tools()), @@ -791,9 +726,7 @@ async def test_list_search_tools_scoped_to_team_object_permission(): response = TestClient(app).get("/search_tools/list") assert response.status_code == 200 - assert [t["search_tool_name"] for t in response.json()["search_tools"]] == [ - "db-tool-2" - ] + assert [t["search_tool_name"] for t in response.json()["search_tools"]] == ["db-tool-2"] @pytest.mark.asyncio @@ -1061,9 +994,15 @@ def _live_router_and_db(db_rows: list): fake_router.search_tools = list(db_rows) with contextlib.ExitStack() as stack: - stack.enter_context(patch("litellm.proxy.proxy_server.prisma_client", MagicMock())) # test-quality-ok: proxy globals are the only seam; see the module note above - stack.enter_context(patch("litellm.proxy.proxy_server.proxy_config", proxy_config)) # test-quality-ok: proxy globals are the only seam; see the module note above - stack.enter_context(patch("litellm.proxy.proxy_server.llm_router", fake_router)) # test-quality-ok: proxy globals are the only seam; see the module note above + stack.enter_context( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()) + ) # test-quality-ok: proxy globals are the only seam; see the module note above + stack.enter_context( + patch("litellm.proxy.proxy_server.proxy_config", proxy_config) + ) # test-quality-ok: proxy globals are the only seam; see the module note above + stack.enter_context( + patch("litellm.proxy.proxy_server.llm_router", fake_router) + ) # test-quality-ok: proxy globals are the only seam; see the module note above stack.enter_context( patch( # test-quality-ok: proxy globals are the only seam; see the module note above "litellm.proxy.search_endpoints.search_tool_management.SEARCH_TOOL_REGISTRY", @@ -1152,6 +1091,171 @@ async def test_create_search_tool_survives_a_failing_router_refresh(): assert response.json()["search_tool_name"] == "tavily-search" +def _search_caller( + virtual_key: bool, + key_search_tools: list[str] | None = None, + team_id: str | None = None, + user_role: LitellmUserRoles = LitellmUserRoles.INTERNAL_USER, + api_key: str = "sk-caller", +) -> UserAPIKeyAuth: + caller = UserAPIKeyAuth( + api_key=api_key, + user_role=user_role, + user_id="internal_user", + team_id=team_id, + object_permission_id=None if key_search_tools is None else "op-key", + object_permission=( + None + if key_search_tools is None + else LiteLLM_ObjectPermissionTable(object_permission_id="op-key", search_tools=key_search_tools) + ), + ) + caller.via_virtual_key = virtual_key + return caller + + +@pytest.fixture +def search_permission_cache(monkeypatch: pytest.MonkeyPatch) -> Callable[[list[str] | None], None]: + from litellm.proxy._types import LiteLLM_UserTable + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key + + cache = UserApiKeyCache() + monkeypatch.setattr(ps, "user_api_key_cache", cache) + monkeypatch.setattr(ps, "prisma_client", MagicMock()) + + def grant_user(search_tools: list[str] | None) -> None: + cache.set_cache( + key="internal_user", value=LiteLLM_UserTable(user_id="internal_user", object_permission_id="op-user") + ) + cache.set_cache( + key=object_permission_cache_key("op-user"), + value=LiteLLM_ObjectPermissionTable(object_permission_id="op-user", search_tools=search_tools), + ) + + return grant_user + + +@pytest.mark.parametrize( + "general_settings, user_search_tools, expected", + [ + ({}, None, ["db-tool-1", "db-tool-2"]), + ({"search_tool_deny_by_default": False}, [], ["db-tool-1", "db-tool-2"]), + ({"search_tool_deny_by_default": True}, None, []), + ({"search_tool_deny_by_default": True}, [], []), + ({"search_tool_deny_by_default": True}, ["db-tool-2"], ["db-tool-2"]), + ], +) +@pytest.mark.asyncio +async def test_filter_visible_search_tools_keyless_user_follows_search_tool_deny_by_default( + search_permission_cache: Callable[[list[str] | None], None], + general_settings: dict[str, bool], + user_search_tools: list[str] | None, + expected: list[str], +): + from litellm.proxy.search_endpoints.search_tool_management import _filter_visible_search_tools + + search_permission_cache(user_search_tools) + team_lookup = AsyncMock() + + visible = await _filter_visible_search_tools( + _search_tool_responses("db-tool-1", "db-tool-2"), + _search_caller(virtual_key=False), + team_lookup, + general_settings, + ) + + assert [t["search_tool_name"] for t in visible] == expected + team_lookup.assert_not_awaited() + + +@pytest.mark.parametrize( + "key_search_tools, expected", + [(None, []), ([], []), (["db-tool-1"], ["db-tool-1"])], +) +@pytest.mark.asyncio +async def test_filter_visible_search_tools_standalone_key_needs_its_own_grant( + search_permission_cache: Callable[[list[str] | None], None], key_search_tools: list[str] | None, expected: list[str] +): + from litellm.proxy.search_endpoints.search_tool_management import _filter_visible_search_tools + + search_permission_cache(["db-tool-1", "db-tool-2"]) + + visible = await _filter_visible_search_tools( + _search_tool_responses("db-tool-1", "db-tool-2"), + _search_caller(virtual_key=True, key_search_tools=key_search_tools), + AsyncMock(), + {"search_tool_deny_by_default": True}, + ) + + assert [t["search_tool_name"] for t in visible] == expected + + +@pytest.mark.parametrize( + "key_search_tools, team_search_tools, expected", + [ + (["db-tool-1"], [], []), + ([], ["db-tool-1"], []), + (["db-tool-1", "db-tool-2"], ["db-tool-1"], ["db-tool-1"]), + ], +) +@pytest.mark.asyncio +async def test_filter_visible_search_tools_team_key_needs_key_and_team_grants( + search_permission_cache: Callable[[list[str] | None], None], + key_search_tools: list[str], + team_search_tools: list[str], + expected: list[str], +): + from litellm.proxy.search_endpoints.search_tool_management import _filter_visible_search_tools + + team_lookup = AsyncMock( + return_value=LiteLLM_TeamTable( + team_id="team-1", + object_permission_id="op-team", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="op-team", search_tools=team_search_tools + ), + ) + ) + + visible = await _filter_visible_search_tools( + _search_tool_responses("db-tool-1", "db-tool-2"), + _search_caller(virtual_key=True, key_search_tools=key_search_tools, team_id="team-1"), + team_lookup, + {"search_tool_deny_by_default": True}, + ) + + assert [t["search_tool_name"] for t in visible] == expected + + +@pytest.mark.parametrize( + "caller, expected", + [ + ( + _search_caller( + virtual_key=False, user_role=LitellmUserRoles.PROXY_ADMIN, api_key="litellm_proxy_master_key" + ), + ["db-tool-1", "db-tool-2"], + ), + (_search_caller(virtual_key=True, user_role=LitellmUserRoles.PROXY_ADMIN), []), + ], + ids=["master-key-sees-every-tool", "admin-virtual-key-is-filtered"], +) +@pytest.mark.asyncio +async def test_filter_visible_search_tools_admin_under_search_tool_deny_by_default( + search_permission_cache: Callable[[list[str] | None], None], caller: UserAPIKeyAuth, expected: list[str] +): + from litellm.proxy.search_endpoints.search_tool_management import _filter_visible_search_tools + + visible = await _filter_visible_search_tools( + _search_tool_responses("db-tool-1", "db-tool-2"), + caller, + AsyncMock(), + {"search_tool_deny_by_default": True}, + ) + + assert [t["search_tool_name"] for t in visible] == expected + + class _StoredSearchToolRow(SimpleNamespace): def __iter__(self): return iter(self.__dict__.items()) diff --git a/tests/unit/proxy/proxy_server/test_proxy_config.py b/tests/unit/proxy/proxy_server/test_proxy_config.py index b37576dab80..0b58e5d141f 100644 --- a/tests/unit/proxy/proxy_server/test_proxy_config.py +++ b/tests/unit/proxy/proxy_server/test_proxy_config.py @@ -2330,38 +2330,36 @@ async def test_load_config_logs_disabled_budget_reservation_once(tmp_path, monke @pytest.mark.asyncio +@pytest.mark.parametrize("setting", ["vector_store_deny_by_default", "search_tool_deny_by_default"]) @pytest.mark.parametrize(("yaml_value", "expected"), [("true", True), ("false", False)]) -async def test_load_config_yaml_vector_store_deny_by_default_is_boolean( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch, yaml_value: str, expected: bool +async def test_load_config_yaml_deny_by_default_is_boolean( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, setting: str, yaml_value: str, expected: bool ): - config_file: Final = tmp_path / "vector_store.yaml" - config_file.write_text( - f"model_list: []\nlitellm_settings: {{}}\ngeneral_settings:\n vector_store_deny_by_default: {yaml_value}\n" - ) + config_file: Final = tmp_path / "deny_by_default.yaml" + config_file.write_text(f"model_list: []\nlitellm_settings: {{}}\ngeneral_settings:\n {setting}: {yaml_value}\n") monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False) _, _, general_settings = await ProxyConfig().load_config(router=None, config_file_path=str(config_file)) - assert general_settings["vector_store_deny_by_default"] is expected - assert ConfigGeneralSettings.model_validate(dict(general_settings)).vector_store_deny_by_default is expected + assert general_settings[setting] is expected + assert getattr(ConfigGeneralSettings.model_validate(dict(general_settings)), setting) is expected @pytest.mark.asyncio +@pytest.mark.parametrize("setting", ["vector_store_deny_by_default", "search_tool_deny_by_default"]) @pytest.mark.parametrize("yaml_value", ["", "enabled"], ids=["null", "string"]) -async def test_load_config_rejects_non_boolean_vector_store_deny_by_default( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch, yaml_value: str +async def test_load_config_rejects_non_boolean_deny_by_default( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, setting: str, yaml_value: str ): - config_file: Final = tmp_path / "vector_store.yaml" - config_file.write_text( - f"model_list: []\nlitellm_settings: {{}}\ngeneral_settings:\n vector_store_deny_by_default: {yaml_value}\n" - ) + config_file: Final = tmp_path / "deny_by_default.yaml" + config_file.write_text(f"model_list: []\nlitellm_settings: {{}}\ngeneral_settings:\n {setting}: {yaml_value}\n") monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False) - with pytest.raises(ValidationError, match="vector_store_deny_by_default"): + with pytest.raises(ValidationError, match=setting): await ProxyConfig().load_config(router=None, config_file_path=str(config_file)) @@ -3652,9 +3650,7 @@ def test_ProxyConfig__decrypt_and_set_db_env_variables_cannot_enable_mcp_stdio(m assert os.environ.get("LITELLM_ENABLE_MCP_STDIO") is None -def test_ProxyConfig__decrypt_and_set_db_env_variables_warns_once_about_the_ignored_mcp_stdio_flag( - monkeypatch, caplog -): +def test_ProxyConfig__decrypt_and_set_db_env_variables_warns_once_about_the_ignored_mcp_stdio_flag(monkeypatch, caplog): monkeypatch.setattr( "litellm.proxy.proxy_server.decrypt_value_helper", lambda value, key, return_original_value=False: value, diff --git a/tests/unit/proxy/search_endpoints/test_endpoints.py b/tests/unit/proxy/search_endpoints/test_endpoints.py index bd6460e3dfb..e8212b2a8b1 100644 --- a/tests/unit/proxy/search_endpoints/test_endpoints.py +++ b/tests/unit/proxy/search_endpoints/test_endpoints.py @@ -1,12 +1,142 @@ +from collections.abc import Iterator +from typing import Final from unittest.mock import AsyncMock, MagicMock import orjson import pytest +import respx +from fastapi import FastAPI +from fastapi.testclient import TestClient +import litellm +from litellm import Router from litellm.proxy import proxy_server -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import ( + LiteLLM_ObjectPermissionTable, + LiteLLM_TeamTableCachedObj, + LitellmUserRoles, + ProxyException, + UserAPIKeyAuth, +) +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.route_llm_request import ProxyMissingRequiredParamError -from litellm.proxy.search_endpoints.endpoints import search +from litellm.proxy.search_endpoints.endpoints import router, search + +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() -> 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: + app: Final = FastAPI() + app.include_router(router) + app.add_exception_handler(ProxyException, proxy_server.openai_exception_handler) + app.dependency_overrides[user_api_key_auth] = lambda: caller + return TestClient(app, raise_server_exceptions=False) + + +def _team_key(key_search_tools: list[str]) -> UserAPIKeyAuth: + caller: Final = UserAPIKeyAuth( + api_key="sk-team-key", + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="user-1", + team_id="team-1", + object_permission_id="op-key", + object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-key", search_tools=key_search_tools), + ) + caller.via_virtual_key = True + return caller + + +@pytest.fixture +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 cache(monkeypatch: pytest.MonkeyPatch) -> UserApiKeyCache: + user_api_key_cache: Final = UserApiKeyCache() + monkeypatch.setattr(proxy_server, "user_api_key_cache", user_api_key_cache) + monkeypatch.setattr(proxy_server, "prisma_client", MagicMock()) + monkeypatch.setattr(proxy_server, "llm_router", _search_router()) + return user_api_key_cache + + +def _cache_team(cache: UserApiKeyCache, search_tools: list[str]) -> None: + team: Final = LiteLLM_TeamTableCachedObj( + team_id="team-1", + object_permission_id="op-team", + 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( + "general_settings, key_search_tools, team_search_tools, expected_status", + [ + ({}, [], [], 200), + ({"search_tool_deny_by_default": False}, [], [], 200), + ({"search_tool_deny_by_default": True}, ["search-a"], [], 403), + ({"search_tool_deny_by_default": True}, [], ["search-a"], 403), + ({"search_tool_deny_by_default": True}, ["search-a"], ["search-b"], 403), + ({"search_tool_deny_by_default": True}, ["search-a"], ["search-a"], 200), + ], +) +@pytest.mark.parametrize("path", ["/v1/search/search-a", "/search/search-a"]) +def test_direct_search_team_key_follows_search_tool_deny_by_default( + monkeypatch: pytest.MonkeyPatch, + cache: UserApiKeyCache, + tavily: respx.Route, + path: str, + general_settings: dict[str, bool], + key_search_tools: list[str], + team_search_tools: list[str], + expected_status: int, +): + monkeypatch.setattr(proxy_server, "general_settings", general_settings) + _cache_team(cache, team_search_tools) + + response: Final = _client(_team_key(key_search_tools)).post(path, json={"query": "what is litellm"}) + + assert response.status_code == expected_status, response.text + if expected_status == 200: + assert response.json()["results"][0]["url"] == TAVILY_RESULT["url"] + assert tavily.call_count == 1 + else: + assert "search-a" in response.text + assert tavily.call_count == 0 + + +def test_direct_search_body_tool_name_is_denied_under_search_tool_deny_by_default(monkeypatch, cache, tavily): + monkeypatch.setattr(proxy_server, "general_settings", {"search_tool_deny_by_default": True}) + caller: Final = UserAPIKeyAuth(api_key="sk-standalone", user_role=LitellmUserRoles.INTERNAL_USER) + caller.via_virtual_key = True + + response: Final = _client(caller).post( + "/v1/search", json={"search_tool_name": "search-a", "query": "what is litellm"} + ) + + assert response.status_code == 403, response.text + assert response.json()["error"]["type"] == "key_search_tool_access_denied" + assert tavily.call_count == 0 def _json_request(body: dict[str, object]) -> MagicMock: @@ -52,3 +182,104 @@ async def test_search_with_only_a_query_falls_back_to_the_proxy_default_model(mo router.asearch.assert_awaited_once() assert router.asearch.await_args.kwargs["query"] == "litellm" assert router.asearch.await_args.kwargs["model"] == "perplexity-search" + + +def _standalone_key(key_search_tools: list[str] | None) -> UserAPIKeyAuth: + caller: Final = UserAPIKeyAuth( + api_key="sk-standalone", + user_role=LitellmUserRoles.INTERNAL_USER, + object_permission_id=None if key_search_tools is None else "op-key", + object_permission=( + None + if key_search_tools is None + else LiteLLM_ObjectPermissionTable(object_permission_id="op-key", search_tools=key_search_tools) + ), + ) + caller.via_virtual_key = True + return caller + + +@pytest.mark.parametrize( + "general_settings, body", + [ + ({"search_tool_deny_by_default": True}, {"model": "search-a", "query": "what is litellm"}), + ( + {"search_tool_deny_by_default": True, "completion_model": "search-a"}, + {"query": "what is litellm"}, + ), + ], + ids=["body-model", "completion-model"], +) +@pytest.mark.parametrize("key_search_tools, expected_status", [(None, 403), (["search-a"], 200)]) +def test_direct_search_authorizes_the_tool_resolved_from_model_settings( + monkeypatch, cache, tavily, general_settings, body, key_search_tools, expected_status +): + monkeypatch.setattr(proxy_server, "general_settings", general_settings) + + response: Final = _client(_standalone_key(key_search_tools)).post("/v1/search", json=body) + + assert response.status_code == expected_status, response.text + assert tavily.call_count == (1 if expected_status == 200 else 0) + + +@pytest.mark.parametrize( + "key_search_tools, expected_status, fallback_calls", + [(["search-a"], 500, 0), (["search-a", "search-b"], 200, 1)], +) +def test_router_search_fallback_target_must_be_granted( + monkeypatch, cache, key_search_tools, expected_status, fallback_calls +): + from litellm.proxy.auth.fallback_model_access import router_fallback_access_check + + monkeypatch.setattr(proxy_server, "general_settings", {"search_tool_deny_by_default": True}) + monkeypatch.setattr( # test-quality-ok: respx needs HTTPX enabled to fake the provider HTTP boundary. + litellm, + "disable_aiohttp_transport", + True, + ) + monkeypatch.setattr( + proxy_server, + "llm_router", + Router( + model_list=[], + search_tools=[ + {"search_tool_name": "search-a", "litellm_params": {"search_provider": "exa_ai", "api_key": "fake"}}, + {"search_tool_name": "search-b", "litellm_params": {"search_provider": "tavily", "api_key": "fake"}}, + ], + fallbacks=[{"search-a": ["search-b"]}], + fallback_access_check=router_fallback_access_check, + num_retries=0, + ), + ) + litellm.in_memory_llm_clients_cache.flush_cache() + with respx.mock(assert_all_called=False) as mock: + failing_tool: Final = mock.post(url__regex=r"https://api\.exa\.ai/.*").respond(500, json={"error": "down"}) + fallback_tool: Final = mock.post(TAVILY_SEARCH_URL).respond(200, json={"results": [TAVILY_RESULT]}) + response: Final = _client(_standalone_key(key_search_tools)).post( + "/v1/search/search-a", json={"query": "what is litellm"} + ) + litellm.in_memory_llm_clients_cache.flush_cache() + + assert response.status_code == expected_status, response.text + assert failing_tool.call_count == 1 + assert fallback_tool.call_count == fallback_calls + + +@pytest.mark.parametrize( + "requested_tool, completion_model, expected_status", + [("search-b", "search-a", 403), ("search-a", "search-b", 200)], + ids=["requested-tool-ungranted", "completion-model-ungranted"], +) +def test_direct_search_authorizes_the_requested_tool_over_completion_model( + monkeypatch, cache, tavily, requested_tool, completion_model, expected_status +): + monkeypatch.setattr( + proxy_server, "general_settings", {"search_tool_deny_by_default": True, "completion_model": completion_model} + ) + + response: Final = _client(_standalone_key(["search-a"])).post( + f"/v1/search/{requested_tool}", json={"query": "what is litellm"} + ) + + assert response.status_code == expected_status, response.text + assert tavily.call_count == (1 if expected_status == 200 else 0) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index c632ad01203..201746f1873 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -30008,6 +30008,12 @@ export interface components { reject_clientside_metadata_tags?: boolean | null; /** @description Spreads the proxy's scheduled background jobs (spend flushes, budget resets, config reloads, exports) across a window instead of firing them together on every replica. On by default; set to tune the window, pin a job, or turn it off. */ scheduled_job_stagger?: components["schemas"]["ScheduledJobStaggerSettings"] | null; + /** + * Search Tool Deny By Default + * @description When True, a search tool must be explicitly listed in object_permission.search_tools: a virtual key needs its own grant plus its team's, a keyless team member needs the team's, and a user with neither needs their own. A missing permission record, an empty list, or an unresolved team grants nothing, and the unregistered search fallback is denied. The master key and dashboard sessions are exempt + * @default false + */ + search_tool_deny_by_default: boolean; /** @description Daily check of the spend LiteLLM captured against the provider's own bill (OpenAI via OPENAI_ADMIN_KEY). Publishes litellm_spend_capture_rate per provider and alerts when the ratio over the lookback window falls under the threshold (default 0.9). Off unless set. */ spend_capture_rate_check?: components["schemas"]["SpendCaptureRateCheckSettings"] | null; /**