feat(auth): deny search tools by default when search_tool_deny_by_default is set (#44490)

Co-authored-by: mrinal <mrinal@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-06 12:50:26 -07:00 • committed by GitHub
parent abc543e701
commit 282d733fb6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
16 changed files with 1497 additions and 289 deletions

View file

@ -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(

View file

@ -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,

View file

@ -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,

View file

@ -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

View file

@ -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)

View file

@ -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

View file

@ -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",

View file

@ -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(

View file

@ -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",)

View file

@ -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()

View file

@ -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

View file

@ -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")

View file

@ -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())

View file

@ -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,

View file

@ -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)

View file

@ -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;
/**