mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
abc543e701
commit
282d733fb6
16 changed files with 1497 additions and 289 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",)
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
270
tests/unit/proxy/auth/test_search_tool_deny_by_default.py
Normal file
270
tests/unit/proxy/auth/test_search_tool_deny_by_default.py
Normal 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")
|
||||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
6
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
6
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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;
|
||||
/**
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue