mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(proxy): add opt-in vector_store_deny_by_default for least-privilege vector store access (#44244)
* feat(proxy): add opt-in vector_store_deny_by_default for standalone virtual keys
Adds general_settings.vector_store_deny_by_default (typed bool, default false). When enabled, a virtual key
with no team must list the requested vector store in its object_permission.vector_stores; no permission
record, null, or an empty list is denied with key_vector_store_access_denied. Omitted or false keeps the
existing behavior, including nonempty allowlist enforcement. The master key is unchanged in both modes.
Team keys and keyless callers are deferred to later increments of LIT-6035
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* feat(proxy): require key and team vector store grants for team keys under vector_store_deny_by_default
With the flag enabled, a virtual key on a team needs both its own grant and its team's grant for every requested vector store. A missing permission record, an empty list or an unresolved team grants nothing. Dashboard session keys and the master key keep their existing behavior, and flag-off behavior is unchanged
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* feat(proxy): require user or team vector store grants for keyless requests under vector_store_deny_by_default
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(proxy): cover vector_store_deny_by_default through a real proxy with key, team and JWT identities
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* chore(ui): regenerate schema.d.ts for vector_store_deny_by_default
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(auth): cover path and file_search vector store ids under vector_store_deny_by_default
Strict mode now reads vector_store_ids and tools[].vector_store_ids from the request body without needing a vector store registry, so /v1/vector_stores/{id}/search and Responses file_search are checked. User grants load through the object permission cache, and the proxy admin user rebuild keeps object_permission_id
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(proxy): broadcast user entitlement cache eviction to every worker
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* refactor(auth): reuse VectorStoreRegistry id extraction for vector_store_deny_by_default
Strict mode now uses get_vector_store_ids_to_run on an empty registry when none is loaded, instead of a parallel set of request-shape helpers, and vector_store_access_check documents the policy
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(integration): add strict vector store audit cells for routes, SDKs, workers and concurrency
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(integration): consolidate vector_store_deny_by_default coverage to core cases
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(proxy): reject invalid vector_store_deny_by_default at config load and return 400 for malformed vector_store_ids
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(auth): read vector store key and team grants through the object permission cache
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(auth): return typed permission rows in request flow vector store tests
Co-authored-by: mrinal <mrinal@berri.ai>
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* refactor(auth): validate vector store ids and tools as immutable sequences
Co-authored-by: mrinal <mrinal@berri.ai>
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---------
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
26a9b02f7b
commit
abc543e701
15 changed files with 1515 additions and 68 deletions
|
|
@ -107,6 +107,7 @@ _PROXY_ERROR_TYPE_MAP: Final[Mapping[str, str]] = MappingProxyType(
|
|||
"key_vector_store_access_denied": PERMISSION_DENIED,
|
||||
"team_vector_store_access_denied": PERMISSION_DENIED,
|
||||
"org_vector_store_access_denied": PERMISSION_DENIED,
|
||||
"user_vector_store_access_denied": PERMISSION_DENIED,
|
||||
"tool_access_denied": PERMISSION_DENIED,
|
||||
"team_member_permission_error": PERMISSION_DENIED,
|
||||
"not_found_error": RESOURCE_NOT_FOUND,
|
||||
|
|
|
|||
|
|
@ -3022,6 +3022,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
None,
|
||||
description="When set to True, rejects requests that contain client-side 'metadata.tags' to prevent users from influencing budgets by sending different tags. Tags can only be inherited from the API key metadata.",
|
||||
)
|
||||
vector_store_deny_by_default: bool = Field(
|
||||
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",
|
||||
)
|
||||
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.",
|
||||
|
|
@ -4512,6 +4516,11 @@ class ProxyErrorTypes(str, enum.Enum):
|
|||
Organization does not have access to the vector store
|
||||
"""
|
||||
|
||||
user_vector_store_access_denied = "user_vector_store_access_denied"
|
||||
"""
|
||||
User does not have access to the vector store
|
||||
"""
|
||||
|
||||
team_member_already_in_team = "team_member_already_in_team"
|
||||
"""
|
||||
Team member is already in team
|
||||
|
|
@ -4546,7 +4555,7 @@ class ProxyErrorTypes(str, enum.Enum):
|
|||
|
||||
@classmethod
|
||||
def get_vector_store_access_error_type_for_object(
|
||||
cls, object_type: Literal["key", "team", "org"]
|
||||
cls, object_type: Literal["key", "team", "org", "user"]
|
||||
) -> "ProxyErrorTypes":
|
||||
"""
|
||||
Get the vector store access error type for object_type
|
||||
|
|
@ -4557,6 +4566,8 @@ class ProxyErrorTypes(str, enum.Enum):
|
|||
return cls.team_vector_store_access_denied
|
||||
elif object_type == "org":
|
||||
return cls.org_vector_store_access_denied
|
||||
elif object_type == "user":
|
||||
return cls.user_vector_store_access_denied
|
||||
|
||||
|
||||
DB_CONNECTION_ERROR_TYPES: Final = (
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ import re
|
|||
import time
|
||||
from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence
|
||||
from functools import partial
|
||||
from itertools import chain
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Generic, Literal, Optional, Protocol, TypeAlias
|
||||
|
||||
|
|
@ -34,6 +35,7 @@ from litellm.constants import (
|
|||
DEFAULT_MAX_RECURSE_DEPTH,
|
||||
EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE,
|
||||
END_USER_RESTRICTED_REGISTRY_MAX_SIZE,
|
||||
LITELLM_PROXY_MASTER_KEY_ALIAS,
|
||||
MODEL_ACCESS_GROUP_REGISTRY_MAX_SIZE,
|
||||
REGISTRY_ERROR_NEGATIVE_CACHE_TTL,
|
||||
TAG_REGISTRY_MAX_SIZE,
|
||||
|
|
@ -44,7 +46,9 @@ from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
|||
from litellm.models.project import LiteLLM_ProjectTable
|
||||
from litellm.proxy._types import (
|
||||
RBAC_ROLES,
|
||||
UI_TEAM_ID,
|
||||
CallInfo,
|
||||
ConfigGeneralSettings,
|
||||
LiteLLM_AccessGroupTable,
|
||||
LiteLLM_BudgetTable,
|
||||
LiteLLM_EndUserTable,
|
||||
|
|
@ -313,12 +317,6 @@ class _VectorStorePermissionsRow(Protocol):
|
|||
def vector_stores(self) -> Sequence[str] | None: ...
|
||||
|
||||
|
||||
def _object_permission_table(
|
||||
repo: _PrismaTableHolder[_VectorStorePermissionsRow],
|
||||
) -> _PrismaAuthTable[_VectorStorePermissionsRow]:
|
||||
return _DeadlineBoundedTable(repo.table, "object_permission")
|
||||
|
||||
|
||||
class _PrismaTagRow(Protocol):
|
||||
tag_name: str
|
||||
|
||||
|
|
@ -1353,6 +1351,8 @@ async def common_checks(
|
|||
request_body=request_body,
|
||||
team_object=team_object,
|
||||
valid_token=valid_token,
|
||||
user_object=user_object,
|
||||
deny_by_default=_vector_store_deny_by_default(_typed_request_body(general_settings)),
|
||||
)
|
||||
|
||||
# 12. [OPTIONAL] Tool allowlist - key/team allowed_tools (no DB in hot path)
|
||||
|
|
@ -6792,17 +6792,210 @@ 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
|
||||
and valid_token.api_key != LITELLM_PROXY_MASTER_KEY_ALIAS
|
||||
and valid_token.team_id != UI_TEAM_ID
|
||||
)
|
||||
|
||||
|
||||
def _strict_grant_layers(
|
||||
deny_by_default: bool, valid_token: UserAPIKeyAuth | None, team_object: LiteLLM_TeamTable | None
|
||||
) -> frozenset[GrantLayer]:
|
||||
"""
|
||||
The identities that must each grant an object under a deny-by-default policy: a virtual key and its team,
|
||||
a keyless team member's team, or a keyless user's own grant. The master key and dashboard sessions need none
|
||||
"""
|
||||
if not deny_by_default or valid_token is None or not _is_strict_grant_identity(valid_token):
|
||||
return frozenset()
|
||||
virtual_key: Final = valid_token.via_virtual_key and not valid_token.is_session_token
|
||||
has_team: Final = team_object is not None or valid_token.team_id is not None
|
||||
if not virtual_key and not has_team:
|
||||
return frozenset(("user",))
|
||||
return frozenset(layer for layer, required in (("key", virtual_key), ("team", has_team)) if required)
|
||||
|
||||
|
||||
async def _identity_grants(
|
||||
valid_token: UserAPIKeyAuth | None,
|
||||
team_object: LiteLLM_TeamTable | None,
|
||||
user_object: LiteLLM_UserTable | None,
|
||||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
) -> tuple[tuple[GrantLayer, LiteLLM_ObjectPermissionTable | None], ...]:
|
||||
return (
|
||||
(
|
||||
"key",
|
||||
None
|
||||
if valid_token is None
|
||||
else await _cached_object_permission(
|
||||
valid_token.object_permission_id, valid_token.object_permission, prisma_client, user_api_key_cache
|
||||
),
|
||||
),
|
||||
(
|
||||
"team",
|
||||
None
|
||||
if team_object is None
|
||||
else await _cached_object_permission(
|
||||
team_object.object_permission_id, team_object.object_permission, prisma_client, user_api_key_cache
|
||||
),
|
||||
),
|
||||
(
|
||||
"user",
|
||||
None
|
||||
if user_object is None
|
||||
else await _cached_object_permission(
|
||||
user_object.object_permission_id, user_object.object_permission, prisma_client, user_api_key_cache
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _vector_store_deny_by_default(general_settings: Mapping[str, object]) -> bool:
|
||||
"""
|
||||
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 vector store requests are denied.
|
||||
"""
|
||||
try:
|
||||
return ConfigGeneralSettings.model_validate(
|
||||
MappingProxyType(
|
||||
{"vector_store_deny_by_default": general_settings.get("vector_store_deny_by_default", False)}
|
||||
)
|
||||
).vector_store_deny_by_default
|
||||
except ValidationError:
|
||||
return True
|
||||
|
||||
|
||||
_VECTOR_STORE_IDS_ADAPTER: Final[TypeAdapter[Sequence[str]]] = TypeAdapter(Sequence[str])
|
||||
_TOOLS_ADAPTER: Final[TypeAdapter[Sequence[object]]] = TypeAdapter(Sequence[object])
|
||||
_TOOL_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(Mapping[str, object])
|
||||
|
||||
|
||||
def _validated_vector_store_ids(value: object) -> tuple[str, ...]:
|
||||
if value is None:
|
||||
return ()
|
||||
try:
|
||||
return tuple(_VECTOR_STORE_IDS_ADAPTER.validate_python(value, strict=True))
|
||||
except ValidationError:
|
||||
raise _malformed_vector_store_ids() from None
|
||||
|
||||
|
||||
def _malformed_vector_store_ids() -> ProxyException:
|
||||
return ProxyException(
|
||||
message="vector_store_ids must be a list of strings",
|
||||
type="invalid_request_error",
|
||||
param="vector_store_ids",
|
||||
code=status.HTTP_400_BAD_REQUEST,
|
||||
)
|
||||
|
||||
|
||||
def _tools(tools: object) -> tuple[object, ...]:
|
||||
try:
|
||||
return tuple(_TOOLS_ADAPTER.validate_python(tools, strict=True))
|
||||
except ValidationError:
|
||||
return ()
|
||||
|
||||
|
||||
def _tool_vector_store_ids(tool: object) -> tuple[str, ...]:
|
||||
try:
|
||||
tool_fields: Final = _TOOL_ADAPTER.validate_python(tool, strict=True)
|
||||
except ValidationError:
|
||||
return ()
|
||||
return _validated_vector_store_ids(tool_fields.get("vector_store_ids"))
|
||||
|
||||
|
||||
def _strict_requested_vector_store_ids(request_body: Mapping[str, object]) -> tuple[str, ...]:
|
||||
"""
|
||||
Same fields VectorStoreRegistry.get_vector_store_ids_to_run reads, but a vector_store_ids that is
|
||||
not a list of strings is a 400 instead of being skipped or iterated, and tools that are not
|
||||
objects name no store.
|
||||
"""
|
||||
return tuple(
|
||||
chain(
|
||||
_validated_vector_store_ids(request_body.get("vector_store_ids")),
|
||||
chain.from_iterable(_tool_vector_store_ids(tool) for tool in _tools(request_body.get("tools"))),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _require_vector_store_grant(
|
||||
object_type: GrantLayer,
|
||||
vector_store_ids_to_run: Sequence[str],
|
||||
object_permission: _VectorStorePermissionsRow | None,
|
||||
) -> None:
|
||||
if object_permission is None or not object_permission.vector_stores:
|
||||
raise ProxyException(
|
||||
message=f"{object_type.capitalize()} not allowed to access vector store. Tried to access {vector_store_ids_to_run[0]}. vector_store_deny_by_default is enabled and the {object_type} has no vector store grants",
|
||||
type=ProxyErrorTypes.get_vector_store_access_error_type_for_object(object_type),
|
||||
param="vector_store",
|
||||
code=status.HTTP_401_UNAUTHORIZED,
|
||||
)
|
||||
_can_object_call_vector_stores(
|
||||
object_type=object_type,
|
||||
vector_store_ids_to_run=vector_store_ids_to_run,
|
||||
object_permissions=object_permission,
|
||||
)
|
||||
|
||||
|
||||
async def _cached_object_permission(
|
||||
object_permission_id: str | None,
|
||||
loaded: LiteLLM_ObjectPermissionTable | None,
|
||||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
) -> LiteLLM_ObjectPermissionTable | None:
|
||||
"""
|
||||
The grant row auth already attached to the key, team or user, else the cached row by id. Both are
|
||||
evicted on every worker when the grant changes, so the request path never reads the table directly.
|
||||
"""
|
||||
if object_permission_id is None:
|
||||
return None
|
||||
if loaded is not None:
|
||||
return loaded
|
||||
return await get_object_permission(
|
||||
object_permission_id=object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
|
||||
|
||||
async def vector_store_access_check(
|
||||
request_body: dict,
|
||||
team_object: LiteLLM_TeamTable | None,
|
||||
valid_token: UserAPIKeyAuth | None,
|
||||
*,
|
||||
user_object: LiteLLM_UserTable | None = None,
|
||||
deny_by_default: bool = False,
|
||||
):
|
||||
"""
|
||||
Checks if the object (key, team, org) has access to the vector store.
|
||||
Checks whether the caller may use every vector store the request names.
|
||||
|
||||
Raises ProxyException if the object (key, team, org) cannot access the specific vector store.
|
||||
Requested stores come from `vector_store_ids`, `tools[].vector_store_ids` and the RAG
|
||||
`retrieval_config.vector_store_id`. Grants come from each identity's
|
||||
`object_permission.vector_stores`.
|
||||
|
||||
With `deny_by_default=False` (legacy), a key or team only restricts access when its list is
|
||||
nonempty, and stores are only read from the request when a vector store registry is loaded.
|
||||
|
||||
With `deny_by_default=True` (`general_settings.vector_store_deny_by_default`), stores are read
|
||||
even without a registry, and a missing record, `null` or `[]` grants nothing:
|
||||
|
||||
- virtual key: the key must grant every store, and so must its team when it has one, even if
|
||||
the team failed to load
|
||||
- keyless team member (JWT, `lite login` session token): only the resolved team is checked
|
||||
- keyless user with no team: the user's own grant is checked
|
||||
- master key and dashboard sessions: legacy behavior
|
||||
|
||||
The user's personal grant is only consulted in the keyless no-team case, so it can neither
|
||||
rescue nor restrict a key or team request.
|
||||
|
||||
Raises ProxyException (401, `{key,team,user}_vector_store_access_denied`) on the first identity
|
||||
that does not grant a requested store, and with the flag on, ProxyException (400,
|
||||
`invalid_request_error`) when a `vector_store_ids` field is not a list of strings.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
|
||||
|
||||
#########################################################
|
||||
# Get the vector store the user is trying to access
|
||||
|
|
@ -6812,12 +7005,17 @@ async def vector_store_access_check(
|
|||
return True
|
||||
|
||||
registry_ids: Final = (
|
||||
litellm.vector_store_registry.get_vector_store_ids_to_run(
|
||||
non_default_params=request_body, tools=request_body.get("tools", None)
|
||||
_strict_requested_vector_store_ids(_typed_request_body(request_body))
|
||||
if deny_by_default
|
||||
else (
|
||||
litellm.vector_store_registry.get_vector_store_ids_to_run(
|
||||
non_default_params=request_body, tools=request_body.get("tools", None)
|
||||
)
|
||||
if litellm.vector_store_registry is not None
|
||||
else None
|
||||
)
|
||||
if litellm.vector_store_registry is not None
|
||||
else None
|
||||
) or ()
|
||||
or ()
|
||||
)
|
||||
rag_vector_store_id: Final = _get_rag_query_vector_store_id(_typed_request_body(request_body))
|
||||
rag_ids: Final = (rag_vector_store_id,) if rag_vector_store_id is not None else ()
|
||||
vector_store_ids_to_run: Final = tuple(dict.fromkeys((*registry_ids, *rag_ids)))
|
||||
|
|
@ -6828,38 +7026,28 @@ async def vector_store_access_check(
|
|||
#########################################################
|
||||
# Check if the object (key, team, org) has access to the vector store
|
||||
#########################################################
|
||||
# Check if the key can access the vector store
|
||||
if valid_token is not None and valid_token.object_permission_id is not None:
|
||||
key_object_permission: Final = await _object_permission_table(
|
||||
ObjectPermissionRepository(prisma_client)
|
||||
).find_unique(
|
||||
where={"object_permission_id": valid_token.object_permission_id},
|
||||
)
|
||||
if key_object_permission is not None:
|
||||
strict_layers: Final = _strict_grant_layers(deny_by_default, valid_token, team_object)
|
||||
grants: Final = await _identity_grants(
|
||||
valid_token,
|
||||
team_object,
|
||||
user_object if "user" in strict_layers else None,
|
||||
prisma_client,
|
||||
user_api_key_cache,
|
||||
)
|
||||
for layer, grant in grants:
|
||||
if layer in strict_layers:
|
||||
_require_vector_store_grant(layer, vector_store_ids_to_run, grant)
|
||||
elif grant is not None:
|
||||
_can_object_call_vector_stores(
|
||||
object_type="key",
|
||||
object_type=layer,
|
||||
vector_store_ids_to_run=vector_store_ids_to_run,
|
||||
object_permissions=key_object_permission,
|
||||
)
|
||||
|
||||
# Check if the team can access the vector store
|
||||
if team_object is not None and team_object.object_permission_id is not None:
|
||||
team_object_permission: Final = await _object_permission_table(
|
||||
ObjectPermissionRepository(prisma_client)
|
||||
).find_unique(
|
||||
where={"object_permission_id": team_object.object_permission_id},
|
||||
)
|
||||
if team_object_permission is not None:
|
||||
_can_object_call_vector_stores(
|
||||
object_type="team",
|
||||
vector_store_ids_to_run=vector_store_ids_to_run,
|
||||
object_permissions=team_object_permission,
|
||||
object_permissions=grant,
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
def _can_object_call_vector_stores(
|
||||
object_type: Literal["key", "team", "org"],
|
||||
object_type: Literal["key", "team", "org", "user"],
|
||||
vector_store_ids_to_run: Sequence[str],
|
||||
object_permissions: _VectorStorePermissionsRow | None,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -3063,6 +3063,9 @@ async def _run_centralized_common_checks(
|
|||
user_id=user_api_key_auth_obj.user_id or litellm_proxy_admin_name,
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
spend=user_object.spend if user_object is not None else 0.0,
|
||||
object_permission_id=(
|
||||
user_object.object_permission_id if isinstance(user_object, LiteLLM_UserTable) else None
|
||||
),
|
||||
)
|
||||
|
||||
if project_object is not None:
|
||||
|
|
|
|||
|
|
@ -1458,11 +1458,7 @@ async def _invalidate_cached_user_entitlement(user_id: str | None, object_permis
|
|||
*(object_permission_cache_key(permission_id) for permission_id in dict.fromkeys(object_permission_ids)),
|
||||
*((user_object_permission_id_cache_key(user_id), user_id) if user_id is not None else ()),
|
||||
)
|
||||
for key in keys:
|
||||
try:
|
||||
await user_api_key_cache.async_delete_cache(key=key)
|
||||
except Exception as e: # noqa: BLE001 # a cache we cannot clear still expires; never fail the write
|
||||
verbose_proxy_logger.warning("Failed to invalidate cached entitlement key %r: %s", key, e)
|
||||
await evict_and_broadcast(cache_keys=keys, user_api_key_cache=user_api_key_cache)
|
||||
|
||||
|
||||
async def _update_single_user_helper(
|
||||
|
|
|
|||
|
|
@ -165,6 +165,7 @@ from litellm.proxy.management_helpers.object_permission_utils import (
|
|||
_set_object_permission,
|
||||
enforce_all_proxy_mcp_servers_grant_is_admin_only,
|
||||
handle_update_object_permission_common,
|
||||
invalidate_cached_object_permissions,
|
||||
)
|
||||
from litellm.proxy.management_helpers.team_member_permission_checks import (
|
||||
TeamMemberPermissionChecks,
|
||||
|
|
@ -2546,6 +2547,10 @@ async def update_team(
|
|||
|
||||
verbose_proxy_logger.info("Successfully updated team - %s, info", team_row.team_id)
|
||||
await sync_team_access_group_membership(prisma_client=prisma_client, team_id=team_row.team_id)
|
||||
await invalidate_cached_object_permissions(
|
||||
object_permission_ids=(existing_team.object_permission_id, team_row.object_permission_id),
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
await _refresh_cached_team(
|
||||
team_row=team_row,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
|
|
|
|||
|
|
@ -6713,6 +6713,14 @@ class ProxyConfig:
|
|||
if general_settings is None:
|
||||
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"]}
|
||||
)
|
||||
)
|
||||
|
||||
if general_settings.get("mcp_advertised_versions") is not None:
|
||||
from litellm.types.mcp import MCPAdvertisedVersions
|
||||
|
||||
|
|
|
|||
|
|
@ -57,18 +57,18 @@ def _assert_forbidden_vector_store_denied(gateway: Gateway, model: str, key: str
|
|||
|
||||
|
||||
@pytest.mark.covers("authorization.vector_store.plain_request_skips_object_permission_lookup")
|
||||
def test_chat_request_without_vector_stores_does_not_read_object_permission_table(gateway: Gateway) -> None:
|
||||
def test_requests_do_not_read_object_permission_table_once_the_grant_is_cached(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model()
|
||||
key: Final = scenario.key(models=[model], object_permission={"vector_stores": ["vs_allowed"]})
|
||||
_assert_plain_chat_served(gateway, model, key)
|
||||
before_control: Final = _object_permission_reads()
|
||||
before_first_use: Final = _object_permission_reads()
|
||||
_assert_forbidden_vector_store_denied(gateway, model, key)
|
||||
eventually(_object_permission_reads, lambda reads: reads > before_control, seconds=15)
|
||||
eventually(_object_permission_reads, lambda reads: reads > before_first_use, seconds=15)
|
||||
baseline: Final = _settled_object_permission_reads(_object_permission_reads(), time.monotonic())
|
||||
for _ in range(PLAIN_REQUESTS):
|
||||
_assert_plain_chat_served(gateway, model, key)
|
||||
_assert_forbidden_vector_store_denied(gateway, model, key)
|
||||
after: Final = _settled_object_permission_reads(_object_permission_reads(), time.monotonic())
|
||||
assert after - baseline < PLAIN_REQUESTS, (
|
||||
f"{PLAIN_REQUESTS} plain chat requests added {after - baseline} object permission reads"
|
||||
f"{2 * PLAIN_REQUESTS} requests after the first added {after - baseline} object permission reads"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,5 +1,8 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Iterator, Mapping
|
||||
from pathlib import Path
|
||||
|
|
@ -7,16 +10,22 @@ from types import MappingProxyType
|
|||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
import httpx
|
||||
import jwt
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import Gateway, Scenario, gateway_from_environment, object_value
|
||||
from cryptography.hazmat.primitives.asymmetric import rsa
|
||||
from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment, object_value
|
||||
from integration._support.process import owned_proxy
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from integration.authorization._guardrail_opt_out import upstream_observations
|
||||
from pydantic import JsonValue
|
||||
from redis import Redis
|
||||
|
||||
CONFIG_STORE_ID: Final = "vs_integration_config_store"
|
||||
PROXY_CONFIG: Final = Path(__file__).resolve().parents[1] / "proxy_config.yaml"
|
||||
REMOVE_OPENAI_API_BASE: Final = ("OPENAI_API_BASE",)
|
||||
JWT_KEY_ID: Final = "integration-vector-store-jwt-key"
|
||||
AUTH_CACHE_INVALIDATION_CHANNEL: Final = "litellm_proxy.auth_cache_invalidation"
|
||||
JsonObject: TypeAlias = dict[str, JsonValue]
|
||||
|
||||
|
||||
|
|
@ -99,6 +108,211 @@ def no_registry_gateways(tmp_path_factory: pytest.TempPathFactory) -> Iterator[t
|
|||
yield no_registry_gateway, upstream_gateway
|
||||
|
||||
|
||||
class StrictGateway:
|
||||
def __init__(
|
||||
self,
|
||||
gateway: Gateway,
|
||||
upstream: Gateway,
|
||||
signing_key: rsa.RSAPrivateKey,
|
||||
config: Path,
|
||||
environment: Mapping[str, str],
|
||||
) -> None:
|
||||
self.gateway: Final = gateway
|
||||
self.upstream: Final = upstream
|
||||
self._signing_key: Final = signing_key
|
||||
self.config: Final = config
|
||||
self.environment: Final = environment
|
||||
|
||||
def jwt(self, subject: str, groups: tuple[str, ...] = ()) -> str:
|
||||
claims: Final[JsonObject] = {
|
||||
"sub": subject,
|
||||
"groups": _json_array(*groups),
|
||||
"iat": int(time.time()),
|
||||
"exp": int(time.time()) + 300,
|
||||
}
|
||||
return jwt.encode(claims, self._signing_key, algorithm="RS256", headers={"kid": JWT_KEY_ID})
|
||||
|
||||
|
||||
def _strict_config(directory: Path, *, deny_by_default: bool = True) -> Path:
|
||||
config: Final = object_value(yaml.safe_load(PROXY_CONFIG.read_text()))
|
||||
general_settings: Final = object_value(config["general_settings"])
|
||||
strict: Final[JsonObject] = {
|
||||
**config,
|
||||
"general_settings": {
|
||||
**general_settings,
|
||||
"vector_store_deny_by_default": deny_by_default,
|
||||
"enable_jwt_auth": True,
|
||||
"litellm_jwtauth": {
|
||||
"user_id_jwt_field": "sub",
|
||||
"team_ids_jwt_field": "groups",
|
||||
"team_allowed_routes": _json_array("openai_routes", "info_routes", "/v1/rag/query"),
|
||||
},
|
||||
},
|
||||
}
|
||||
path: Final = directory / f"proxy_vector_store_deny_by_default_{deny_by_default}.yaml"
|
||||
path.write_text(yaml.safe_dump(strict))
|
||||
return path
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def strict_gateway(tmp_path_factory: pytest.TempPathFactory) -> Iterator[StrictGateway]:
|
||||
signing_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
||||
public_jwk: Final = jwt.algorithms.RSAAlgorithm.to_jwk(signing_key.public_key())
|
||||
jwks_body: Final = json.dumps({"keys": [{**json.loads(public_jwk), "kid": JWT_KEY_ID}]}).encode()
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "GET", request
|
||||
return Reply(body=jwks_body)
|
||||
|
||||
with gateway_from_environment() as upstream_gateway, wire_server(respond) as jwks:
|
||||
directory: Final = tmp_path_factory.mktemp("rag_query_deny_by_default")
|
||||
config: Final = _strict_config(directory)
|
||||
environment: Final = MappingProxyType({**_openai_environment(upstream_gateway), "JWT_PUBLIC_KEY_URL": jwks.url})
|
||||
with owned_proxy(
|
||||
upstream_gateway,
|
||||
directory,
|
||||
environment,
|
||||
config=config,
|
||||
remove_environment=REMOVE_OPENAI_API_BASE,
|
||||
) as gateway:
|
||||
yield StrictGateway(gateway, upstream_gateway, signing_key, config, environment)
|
||||
|
||||
|
||||
StrictCase: TypeAlias = Literal[
|
||||
"standalone_key_no_permission",
|
||||
"team_key_empty_key_grants",
|
||||
"team_key_empty_team_grants",
|
||||
"multi_store_one_ungranted",
|
||||
"jwt_user_without_grant",
|
||||
]
|
||||
|
||||
|
||||
def _strict_denied_request(
|
||||
strict: StrictGateway, scenario: Scenario, case: StrictCase, model: str, marker: str, store_id: str
|
||||
) -> tuple[httpx.Response, str]:
|
||||
models: Final = _json_array(model)
|
||||
granted: Final = _permission_for_stores(store_id)
|
||||
empty: Final = _permission_for_stores()
|
||||
if case == "standalone_key_no_permission":
|
||||
standalone_key: Final = scenario.key(models=models)
|
||||
return strict.gateway.request(
|
||||
"POST", "/v1/rag/query", _rag_query_body(model, marker, store_id), key=standalone_key
|
||||
), ("key_vector_store_access_denied")
|
||||
if case in ("team_key_empty_key_grants", "team_key_empty_team_grants"):
|
||||
key_grants_store: Final = case == "team_key_empty_team_grants"
|
||||
team: Final = scenario.team(models=models, object_permission=empty if key_grants_store else granted)
|
||||
team_key: Final = scenario.key(
|
||||
team_id=team, models=models, object_permission=granted if key_grants_store else empty
|
||||
)
|
||||
error_type: Final = "team_vector_store_access_denied" if key_grants_store else "key_vector_store_access_denied"
|
||||
return strict.gateway.request(
|
||||
"POST", "/v1/rag/query", _rag_query_body(model, marker, store_id), key=team_key
|
||||
), (error_type)
|
||||
if case == "multi_store_one_ungranted":
|
||||
partial_key: Final = scenario.key(models=models, object_permission=_permission_for_stores(store_id))
|
||||
body: Final[JsonObject] = {
|
||||
**_rag_query_body(model, marker, store_id),
|
||||
"tools": _json_array({"type": "file_search", "vector_store_ids": _json_array(CONFIG_STORE_ID)}),
|
||||
}
|
||||
return strict.gateway.request(
|
||||
"POST", "/v1/chat/completions", body, key=partial_key
|
||||
), "key_vector_store_access_denied"
|
||||
user: Final = scenario.user(user_role="internal_user")
|
||||
return (
|
||||
strict.gateway.request("POST", "/v1/rag/query", _rag_query_body(model, marker, store_id), key=strict.jwt(user)),
|
||||
"user_vector_store_access_denied",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"case",
|
||||
(
|
||||
"standalone_key_no_permission",
|
||||
"team_key_empty_key_grants",
|
||||
"team_key_empty_team_grants",
|
||||
"multi_store_one_ungranted",
|
||||
"jwt_user_without_grant",
|
||||
),
|
||||
)
|
||||
def test_deny_by_default_rejects_ungranted_store_before_upstream_search(
|
||||
strict_gateway: StrictGateway, case: StrictCase
|
||||
) -> None:
|
||||
with strict_gateway.gateway.scenario() as scenario:
|
||||
model: Final = scenario.model()
|
||||
store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}"
|
||||
marker: Final = f"lit6035 deny by default {case} {uuid.uuid4().hex}"
|
||||
|
||||
response, error_type = _strict_denied_request(strict_gateway, scenario, case, model, marker, store_id)
|
||||
observations: Final = tuple(
|
||||
observation
|
||||
for observation in upstream_observations(strict_gateway.upstream)
|
||||
if marker in str(observation["body"])
|
||||
)
|
||||
|
||||
assert response.status_code == 401, f"{response.text}; scripted_upstream_observations={observations!r}"
|
||||
assert response.json()["error"]["type"] == error_type, response.text
|
||||
assert observations == ()
|
||||
|
||||
|
||||
GrantedCase: TypeAlias = Literal[
|
||||
"standalone_key_granted_registered_store",
|
||||
"team_key_both_grant_unregistered_store",
|
||||
"jwt_team_member_team_grant_only",
|
||||
"jwt_user_personal_grant",
|
||||
"master_key_without_grants",
|
||||
]
|
||||
|
||||
|
||||
def _strict_granted_request(
|
||||
strict: StrictGateway, scenario: Scenario, case: GrantedCase, model: str, marker: str, store_id: str
|
||||
) -> httpx.Response:
|
||||
models: Final = _json_array(model)
|
||||
granted: Final = _permission_for_stores(store_id)
|
||||
body: Final = _rag_query_body(model, marker, store_id)
|
||||
if case in ("standalone_key_granted_registered_store", "team_key_both_grant_unregistered_store"):
|
||||
key_team: Final = scenario.team(models=models, object_permission=granted) if case.startswith("team") else None
|
||||
granted_key: Final = scenario.key(
|
||||
models=models, object_permission=granted, **({} if key_team is None else {"team_id": key_team})
|
||||
)
|
||||
return strict.gateway.request("POST", "/v1/rag/query", body, key=granted_key)
|
||||
if case == "jwt_team_member_team_grant_only":
|
||||
member_team: Final = scenario.team(models=models, object_permission=granted)
|
||||
member: Final = scenario.member(member_team)
|
||||
return strict.gateway.request("POST", "/v1/rag/query", body, key=strict.jwt(member, (member_team,)))
|
||||
if case == "jwt_user_personal_grant":
|
||||
user: Final = scenario.user(user_role="internal_user", object_permission=granted)
|
||||
return strict.gateway.request("POST", "/v1/rag/query", body, key=strict.jwt(user))
|
||||
return strict.gateway.request("POST", "/v1/rag/query", body)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"case",
|
||||
(
|
||||
"standalone_key_granted_registered_store",
|
||||
"team_key_both_grant_unregistered_store",
|
||||
"jwt_team_member_team_grant_only",
|
||||
"jwt_user_personal_grant",
|
||||
"master_key_without_grants",
|
||||
),
|
||||
)
|
||||
def test_deny_by_default_searches_explicitly_granted_store(strict_gateway: StrictGateway, case: GrantedCase) -> None:
|
||||
with strict_gateway.gateway.scenario() as scenario:
|
||||
model: Final = scenario.model()
|
||||
store_id: Final = (
|
||||
CONFIG_STORE_ID
|
||||
if case == "standalone_key_granted_registered_store"
|
||||
else f"vs_unregistered_{uuid.uuid4().hex}"
|
||||
)
|
||||
marker: Final = f"lit6035 granted {case} {uuid.uuid4().hex}"
|
||||
|
||||
response: Final = _strict_granted_request(strict_gateway, scenario, case, model, marker, store_id)
|
||||
assert response.status_code == 200, response.text
|
||||
|
||||
searches: Final = _searches_for_marker(strict_gateway.upstream, marker, store_id)
|
||||
assert len(searches) == 1, searches
|
||||
assert marker in str(object_value(searches[0]["body"])["query"]), searches
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("scope", "error_type"),
|
||||
(("key", "key_vector_store_access_denied"), ("team", "team_vector_store_access_denied")),
|
||||
|
|
@ -220,3 +434,136 @@ def test_rag_query_alias_denies_store_when_team_allowlist_excludes(gateway: Gate
|
|||
assert response.status_code == 401, response.text
|
||||
assert response.json()["error"]["type"] == "team_vector_store_access_denied", response.text
|
||||
assert _searches_for_marker(gateway, marker) == ()
|
||||
|
||||
|
||||
def test_explicit_false_flag_keeps_legacy_vector_store_outcomes(tmp_path: Path) -> None:
|
||||
with gateway_from_environment() as upstream_gateway:
|
||||
with owned_proxy(
|
||||
upstream_gateway,
|
||||
tmp_path,
|
||||
_openai_environment(upstream_gateway),
|
||||
config=_strict_config(tmp_path, deny_by_default=False),
|
||||
remove_environment=REMOVE_OPENAI_API_BASE,
|
||||
) as gateway:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model()
|
||||
allowed_marker: Final = f"flag-false-allowed-{uuid.uuid4().hex}"
|
||||
denied_marker: Final = f"flag-false-denied-{uuid.uuid4().hex}"
|
||||
no_permission_key: Final = scenario.key(models=_json_array(model))
|
||||
excluding_key: Final = scenario.key(
|
||||
models=_json_array(model), object_permission=_permission_for_stores("vs_some_other_store")
|
||||
)
|
||||
|
||||
allowed: Final = _rag_query(gateway, model, allowed_marker, no_permission_key)
|
||||
denied: Final = _rag_query(gateway, model, denied_marker, excluding_key)
|
||||
|
||||
assert allowed.status_code == 200, allowed.text
|
||||
assert len(_searches_for_marker(upstream_gateway, allowed_marker)) == 1
|
||||
assert denied.status_code == 401, denied.text
|
||||
assert denied.json()["error"]["type"] == "key_vector_store_access_denied"
|
||||
assert _searches_for_marker(upstream_gateway, denied_marker) == ()
|
||||
|
||||
|
||||
def _strict_no_registry_config(directory: Path) -> Path:
|
||||
config: Final = object_value(yaml.safe_load(_no_registry_config(directory).read_text()))
|
||||
general_settings: Final = object_value(config["general_settings"])
|
||||
strict: Final[JsonObject] = {
|
||||
**config,
|
||||
"general_settings": {**general_settings, "vector_store_deny_by_default": True},
|
||||
}
|
||||
path: Final = directory / "proxy_vector_store_deny_by_default_no_registry.yaml"
|
||||
path.write_text(yaml.safe_dump(strict))
|
||||
return path
|
||||
|
||||
|
||||
def test_deny_by_default_without_registry_checks_search_route_and_file_search_tools(tmp_path: Path) -> None:
|
||||
with gateway_from_environment() as upstream_gateway:
|
||||
with owned_proxy(
|
||||
upstream_gateway,
|
||||
tmp_path,
|
||||
_openai_environment(upstream_gateway),
|
||||
config=_strict_no_registry_config(tmp_path),
|
||||
remove_environment=REMOVE_OPENAI_API_BASE,
|
||||
) as gateway:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model()
|
||||
store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}"
|
||||
search_marker: Final = f"lit6035 no registry search {uuid.uuid4().hex}"
|
||||
responses_marker: Final = f"lit6035 no registry responses {uuid.uuid4().hex}"
|
||||
granted_marker: Final = f"lit6035 no registry granted search {uuid.uuid4().hex}"
|
||||
ungranted_key: Final = scenario.key(models=_json_array(model))
|
||||
granted_key: Final = scenario.key(
|
||||
models=_json_array(model), object_permission=_permission_for_stores(store_id)
|
||||
)
|
||||
|
||||
search_denied: Final = gateway.request(
|
||||
"POST", f"/v1/vector_stores/{store_id}/search", {"query": search_marker}, key=ungranted_key
|
||||
)
|
||||
responses_denied: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/responses",
|
||||
{
|
||||
"model": model,
|
||||
"input": responses_marker,
|
||||
"tools": _json_array({"type": "file_search", "vector_store_ids": _json_array(store_id)}),
|
||||
},
|
||||
key=ungranted_key,
|
||||
)
|
||||
search_granted: Final = gateway.request(
|
||||
"POST", f"/v1/vector_stores/{store_id}/search", {"query": granted_marker}, key=granted_key
|
||||
)
|
||||
|
||||
observations: Final = upstream_observations(upstream_gateway)
|
||||
denied_observations: Final = tuple(
|
||||
observation
|
||||
for observation in observations
|
||||
if search_marker in str(observation["body"]) or responses_marker in str(observation["body"])
|
||||
)
|
||||
granted_searches: Final = tuple(
|
||||
observation
|
||||
for observation in observations
|
||||
if observation["path"] == f"/vector_stores/{store_id}/search"
|
||||
and granted_marker in str(observation["body"])
|
||||
)
|
||||
assert search_denied.status_code == 401, search_denied.text
|
||||
assert search_denied.json()["error"]["type"] == "key_vector_store_access_denied", search_denied.text
|
||||
assert responses_denied.status_code == 401, responses_denied.text
|
||||
assert responses_denied.json()["error"]["type"] == "key_vector_store_access_denied", responses_denied.text
|
||||
assert denied_observations == ()
|
||||
assert search_granted.status_code == 200, search_granted.text
|
||||
assert len(granted_searches) == 1, granted_searches
|
||||
|
||||
|
||||
def _auth_cache_subscribers() -> int:
|
||||
with Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"])) as cache:
|
||||
return int(cache.pubsub_numsub(AUTH_CACHE_INVALIDATION_CHANNEL)[0][1])
|
||||
|
||||
|
||||
def test_revoked_user_grant_stops_working_on_another_proxy(strict_gateway: StrictGateway, tmp_path: Path) -> None:
|
||||
subscribers_before_peer: Final = _auth_cache_subscribers()
|
||||
with (
|
||||
owned_proxy(
|
||||
strict_gateway.upstream,
|
||||
tmp_path,
|
||||
strict_gateway.environment,
|
||||
config=strict_gateway.config,
|
||||
remove_environment=REMOVE_OPENAI_API_BASE,
|
||||
) as peer,
|
||||
strict_gateway.gateway.scenario() as scenario,
|
||||
):
|
||||
eventually(_auth_cache_subscribers, lambda count: count > subscribers_before_peer)
|
||||
model: Final = scenario.model()
|
||||
store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}"
|
||||
user: Final = scenario.user(user_role="internal_user", object_permission=_permission_for_stores(store_id))
|
||||
token: Final = strict_gateway.jwt(user)
|
||||
|
||||
def peer_status() -> int:
|
||||
marker: Final = f"lit6035 revoked user grant {uuid.uuid4().hex}"
|
||||
return peer.request(
|
||||
"POST", "/v1/rag/query", _rag_query_body(model, marker, store_id), key=token
|
||||
).status_code
|
||||
|
||||
assert eventually(peer_status, lambda status: status == 200, seconds=30, return_last_on_timeout=True) == 200
|
||||
strict_gateway.gateway.post("/user/update", {"user_id": user, "object_permission": _permission_for_stores()})
|
||||
|
||||
assert eventually(peer_status, lambda status: status == 401, seconds=10, return_last_on_timeout=True) == 401
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ import json
|
|||
import re
|
||||
import sys
|
||||
import time
|
||||
from collections.abc import Iterator, Mapping
|
||||
from collections.abc import Awaitable, Iterator, Mapping
|
||||
from types import SimpleNamespace
|
||||
from typing import TYPE_CHECKING, Final, Literal, Optional
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
|
@ -61,6 +61,7 @@ from litellm.proxy.auth.auth_checks import (
|
|||
_check_agent_caller_model_access,
|
||||
_virtual_key_max_budget_check,
|
||||
_virtual_key_soft_budget_check,
|
||||
common_checks,
|
||||
get_key_object,
|
||||
get_user_object,
|
||||
invalidate_team_member_spend_state,
|
||||
|
|
@ -74,6 +75,7 @@ from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
|
|||
from litellm.constants import (
|
||||
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL,
|
||||
END_USER_RESTRICTED_REGISTRY_MAX_SIZE,
|
||||
LITELLM_PROXY_MASTER_KEY_ALIAS,
|
||||
PROXY_DB_LOOKUP_MAX_CONCURRENCY,
|
||||
REGISTRY_ERROR_NEGATIVE_CACHE_TTL,
|
||||
TAG_REGISTRY_MAX_SIZE,
|
||||
|
|
@ -90,6 +92,7 @@ from litellm.proxy.common_utils.user_api_key_cache import (
|
|||
UserApiKeyCache,
|
||||
end_user_cache_key,
|
||||
end_user_restricted_registry_cache_key,
|
||||
object_permission_cache_key,
|
||||
tag_cache_key,
|
||||
tag_registry_cache_key,
|
||||
)
|
||||
|
|
@ -1646,8 +1649,9 @@ async def test_vector_store_access_check_with_permissions():
|
|||
)
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_permissions = MagicMock()
|
||||
mock_permissions.vector_stores = ["store-1", "store-2"]
|
||||
mock_permissions = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="perm-123", vector_stores=["store-1", "store-2"]
|
||||
)
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=mock_permissions)
|
||||
|
||||
mock_vector_store_registry = MagicMock()
|
||||
|
|
@ -1688,12 +1692,12 @@ async def test_vector_store_access_check_with_team_permissions():
|
|||
request_body = {}
|
||||
valid_token = UserAPIKeyAuth(token="team-test-token", object_permission_id=None)
|
||||
|
||||
team_object = MagicMock()
|
||||
team_object.object_permission_id = "team-permission"
|
||||
team_object = LiteLLM_TeamTable(team_id="team-1", object_permission_id="team-permission")
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
team_permissions = MagicMock()
|
||||
team_permissions.vector_stores = ["team-store-allowed"]
|
||||
team_permissions = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="team-permission", vector_stores=["team-store-allowed"]
|
||||
)
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=team_permissions)
|
||||
|
||||
mock_vector_store_registry = MagicMock()
|
||||
|
|
@ -1755,12 +1759,12 @@ async def test_vector_store_access_check_enforces_team_allowlist_for_rag_query(
|
|||
}
|
||||
valid_token = UserAPIKeyAuth(token="team-test-token", object_permission_id=None)
|
||||
|
||||
team_object = MagicMock()
|
||||
team_object.object_permission_id = "team-permission"
|
||||
team_object = LiteLLM_TeamTable(team_id="team-1", object_permission_id="team-permission")
|
||||
|
||||
mock_prisma_client = MagicMock()
|
||||
team_permissions = MagicMock()
|
||||
team_permissions.vector_stores = ["KBALLOWED123"]
|
||||
team_permissions = LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="team-permission", vector_stores=["KBALLOWED123"]
|
||||
)
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=team_permissions)
|
||||
|
||||
with (
|
||||
|
|
@ -1786,6 +1790,592 @@ async def test_vector_store_access_check_enforces_team_allowlist_for_rag_query(
|
|||
assert exc_info.value.type == expected_error_type
|
||||
|
||||
|
||||
_KEY_DENIED: Final = ProxyErrorTypes.key_vector_store_access_denied
|
||||
_TEAM_DENIED: Final = ProxyErrorTypes.team_vector_store_access_denied
|
||||
_USER_DENIED: Final = ProxyErrorTypes.user_vector_store_access_denied
|
||||
|
||||
|
||||
def _virtual_key(
|
||||
object_permission_id: str | None = None,
|
||||
api_key: str = "sk-standalone",
|
||||
user_role: LitellmUserRoles | None = None,
|
||||
team_id: str | None = None,
|
||||
) -> UserAPIKeyAuth:
|
||||
key: Final = UserAPIKeyAuth(
|
||||
api_key=api_key,
|
||||
user_id="key-owner",
|
||||
user_role=user_role,
|
||||
team_id=team_id,
|
||||
object_permission_id=object_permission_id,
|
||||
)
|
||||
key.via_virtual_key = True
|
||||
return key
|
||||
|
||||
|
||||
def _permission_row(
|
||||
object_permission_id: str, permission: SimpleNamespace | LiteLLM_ObjectPermissionTable | None
|
||||
) -> LiteLLM_ObjectPermissionTable | None:
|
||||
if isinstance(permission, SimpleNamespace):
|
||||
return LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id=object_permission_id, vector_stores=permission.vector_stores
|
||||
)
|
||||
return permission
|
||||
|
||||
|
||||
async def _common_checks_for_rag_query(
|
||||
general_settings: Mapping[str, object],
|
||||
valid_token: UserAPIKeyAuth,
|
||||
key_permission: SimpleNamespace | None,
|
||||
team_object: LiteLLM_TeamTable | None = None,
|
||||
team_permission: SimpleNamespace | None = None,
|
||||
user_object: LiteLLM_UserTable | None = None,
|
||||
user_permission: SimpleNamespace | LiteLLM_ObjectPermissionTable | None = None,
|
||||
request_body: Mapping[str, object] | None = None,
|
||||
user_api_key_cache: UserApiKeyCache | None = None,
|
||||
vector_store_registry: VectorStoreRegistry | None = None,
|
||||
) -> bool:
|
||||
permissions: Final = {
|
||||
"key-permission": _permission_row("key-permission", key_permission),
|
||||
"team-permission": _permission_row("team-permission", team_permission),
|
||||
"user-permission": _permission_row("user-permission", user_permission),
|
||||
}
|
||||
mock_prisma_client: Final = MagicMock()
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(
|
||||
side_effect=lambda where: permissions.get(where["object_permission_id"])
|
||||
)
|
||||
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=None)
|
||||
body: Final = (
|
||||
dict(request_body)
|
||||
if request_body is not None
|
||||
else {
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [{"role": "user", "content": "what is in this KB?"}],
|
||||
"retrieval_config": {"vector_store_id": "KBSTOREA", "custom_llm_provider": "bedrock"},
|
||||
}
|
||||
)
|
||||
with (
|
||||
patch( # test-quality-ok: production auth reads these module globals; no dependency injection seam exists
|
||||
"litellm.proxy.proxy_server.prisma_client", mock_prisma_client
|
||||
),
|
||||
patch( # test-quality-ok: production auth reads this module global; no dependency injection seam exists
|
||||
"litellm.vector_store_registry", vector_store_registry
|
||||
),
|
||||
patch( # test-quality-ok: production auth reads this module global; no dependency injection seam exists
|
||||
"litellm.proxy.proxy_server.user_api_key_cache",
|
||||
UserApiKeyCache() if user_api_key_cache is None else user_api_key_cache,
|
||||
),
|
||||
):
|
||||
return await common_checks(
|
||||
request_body=body,
|
||||
team_object=team_object,
|
||||
user_object=user_object,
|
||||
end_user_object=None,
|
||||
global_proxy_spend=None,
|
||||
general_settings=dict(general_settings),
|
||||
route="/v1/rag/query",
|
||||
llm_router=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
valid_token=valid_token,
|
||||
request=MagicMock(spec=Request),
|
||||
)
|
||||
|
||||
|
||||
async def _assert_rag_query_outcome(request: Awaitable[bool], denied_by: ProxyErrorTypes | None) -> None:
|
||||
if denied_by is None:
|
||||
assert await request is True
|
||||
return
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await request
|
||||
assert (exc_info.value.type, exc_info.value.param, exc_info.value.code) == (denied_by, "vector_store", "401")
|
||||
|
||||
|
||||
async def _team_key_rag_query(
|
||||
general_settings: Mapping[str, object],
|
||||
key_vector_stores: list[str] | None,
|
||||
team_vector_stores: list[str] | None,
|
||||
user_object: LiteLLM_UserTable | None = None,
|
||||
) -> bool:
|
||||
return await _common_checks_for_rag_query(
|
||||
general_settings,
|
||||
_virtual_key(object_permission_id=None if key_vector_stores is None else "key-permission", team_id="team-1"),
|
||||
None if key_vector_stores is None else SimpleNamespace(vector_stores=key_vector_stores),
|
||||
LiteLLM_TeamTable(
|
||||
team_id="team-1", object_permission_id=None if team_vector_stores is None else "team-permission"
|
||||
),
|
||||
None if team_vector_stores is None else SimpleNamespace(vector_stores=team_vector_stores),
|
||||
user_object,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("general_settings", "denied_by"),
|
||||
[
|
||||
({}, None),
|
||||
({"vector_store_deny_by_default": False}, None),
|
||||
({"vector_store_deny_by_default": True}, _KEY_DENIED),
|
||||
],
|
||||
ids=["flag-omitted", "flag-false", "flag-true"],
|
||||
)
|
||||
async def test_standalone_key_without_vector_store_permission_follows_deny_by_default(
|
||||
general_settings: Mapping[str, object], denied_by: ProxyErrorTypes | None
|
||||
):
|
||||
await _assert_rag_query_outcome(_common_checks_for_rag_query(general_settings, _virtual_key(), None), denied_by)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("deny_by_default", "vector_stores", "denied_by"),
|
||||
[
|
||||
(True, [], _KEY_DENIED),
|
||||
(True, ["KBSTOREA"], None),
|
||||
(True, ["KBSTOREB"], _KEY_DENIED),
|
||||
(True, None, _KEY_DENIED),
|
||||
(False, ["KBSTOREB"], _KEY_DENIED),
|
||||
],
|
||||
ids=["enabled-empty", "enabled-contains", "enabled-excludes", "enabled-null", "disabled-excludes"],
|
||||
)
|
||||
async def test_standalone_key_vector_store_permission_record_under_deny_by_default(
|
||||
deny_by_default: bool, vector_stores: list[str] | None, denied_by: ProxyErrorTypes | None
|
||||
):
|
||||
await _assert_rag_query_outcome(
|
||||
_common_checks_for_rag_query(
|
||||
{"vector_store_deny_by_default": deny_by_default},
|
||||
_virtual_key(object_permission_id="key-permission"),
|
||||
SimpleNamespace(vector_stores=vector_stores),
|
||||
),
|
||||
denied_by,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("deny_by_default", [False, True], ids=["flag-false", "flag-true"])
|
||||
async def test_master_key_rag_query_is_unchanged_by_deny_by_default(deny_by_default: bool):
|
||||
master_key: Final = _virtual_key(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
|
||||
await _assert_rag_query_outcome(
|
||||
_common_checks_for_rag_query({"vector_store_deny_by_default": deny_by_default}, master_key, None), None
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("general_settings", "vector_stores", "denied_by"),
|
||||
[
|
||||
({}, None, None),
|
||||
({"vector_store_deny_by_default": False}, None, None),
|
||||
({"vector_store_deny_by_default": False}, ["KBSTOREB"], _KEY_DENIED),
|
||||
],
|
||||
ids=["flag-omitted-no-record", "flag-false-no-record", "flag-false-excludes"],
|
||||
)
|
||||
async def test_proxy_admin_virtual_key_keeps_flag_off_vector_store_behavior(
|
||||
general_settings: Mapping[str, object], vector_stores: list[str] | None, denied_by: ProxyErrorTypes | None
|
||||
):
|
||||
admin_key: Final = _virtual_key(
|
||||
object_permission_id=None if vector_stores is None else "key-permission", user_role=LitellmUserRoles.PROXY_ADMIN
|
||||
)
|
||||
key_permission: Final = None if vector_stores is None else SimpleNamespace(vector_stores=vector_stores)
|
||||
|
||||
await _assert_rag_query_outcome(
|
||||
_common_checks_for_rag_query(general_settings, admin_key, key_permission), denied_by
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("key_vector_stores", "team_vector_stores", "denied_by"),
|
||||
[
|
||||
(["KBSTOREA"], ["KBSTOREA"], None),
|
||||
(None, ["KBSTOREA"], _KEY_DENIED),
|
||||
([], ["KBSTOREA"], _KEY_DENIED),
|
||||
(["KBSTOREA"], None, _TEAM_DENIED),
|
||||
(["KBSTOREA"], [], _TEAM_DENIED),
|
||||
(["KBSTOREB"], ["KBSTOREA"], _KEY_DENIED),
|
||||
(["KBSTOREA"], ["KBSTOREB"], _TEAM_DENIED),
|
||||
],
|
||||
ids=["both-grant", "key-no-record", "key-empty", "team-no-record", "team-empty", "key-excludes", "team-excludes"],
|
||||
)
|
||||
async def test_team_key_requires_key_and_team_grant_under_deny_by_default(
|
||||
key_vector_stores: list[str] | None, team_vector_stores: list[str] | None, denied_by: ProxyErrorTypes | None
|
||||
):
|
||||
await _assert_rag_query_outcome(
|
||||
_team_key_rag_query({"vector_store_deny_by_default": True}, key_vector_stores, team_vector_stores), denied_by
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("general_settings", "key_vector_stores", "team_vector_stores", "denied_by"),
|
||||
[
|
||||
({}, [], ["KBSTOREA"], None),
|
||||
({"vector_store_deny_by_default": False}, [], ["KBSTOREA"], None),
|
||||
({"vector_store_deny_by_default": False}, None, None, None),
|
||||
({"vector_store_deny_by_default": False}, ["KBSTOREB"], ["KBSTOREA"], _KEY_DENIED),
|
||||
({"vector_store_deny_by_default": False}, ["KBSTOREA"], ["KBSTOREB"], _TEAM_DENIED),
|
||||
],
|
||||
ids=["omitted-key-empty", "false-key-empty", "false-no-records", "false-key-excludes", "false-team-excludes"],
|
||||
)
|
||||
async def test_team_key_keeps_legacy_vector_store_behavior_when_flag_off(
|
||||
general_settings: Mapping[str, object],
|
||||
key_vector_stores: list[str] | None,
|
||||
team_vector_stores: list[str] | None,
|
||||
denied_by: ProxyErrorTypes | None,
|
||||
):
|
||||
await _assert_rag_query_outcome(
|
||||
_team_key_rag_query(general_settings, key_vector_stores, team_vector_stores), denied_by
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_owned_team_key_without_key_grant_is_denied_despite_team_membership():
|
||||
team_member: Final = LiteLLM_UserTable(
|
||||
user_id="key-owner", teams=["team-1"], user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
|
||||
await _assert_rag_query_outcome(
|
||||
_team_key_rag_query({"vector_store_deny_by_default": True}, [], ["KBSTOREA"], team_member), _KEY_DENIED
|
||||
)
|
||||
|
||||
|
||||
def _user_row(user_vector_stores: list[str] | None, teams: tuple[str, ...] = ()) -> LiteLLM_UserTable:
|
||||
return LiteLLM_UserTable(
|
||||
user_id="user-1",
|
||||
teams=list(teams),
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
object_permission_id=None if user_vector_stores is None else "user-permission",
|
||||
)
|
||||
|
||||
|
||||
async def _keyless_rag_query(
|
||||
general_settings: Mapping[str, object],
|
||||
user_vector_stores: list[str] | None,
|
||||
team_vector_stores: list[str] | None = None,
|
||||
team_id: str | None = None,
|
||||
user_permission: LiteLLM_ObjectPermissionTable | None = None,
|
||||
) -> bool:
|
||||
return await _common_checks_for_rag_query(
|
||||
general_settings,
|
||||
UserAPIKeyAuth(user_id="user-1", team_id=team_id, user_role=LitellmUserRoles.INTERNAL_USER),
|
||||
None,
|
||||
None
|
||||
if team_id is None
|
||||
else LiteLLM_TeamTable(
|
||||
team_id=team_id, object_permission_id=None if team_vector_stores is None else "team-permission"
|
||||
),
|
||||
None if team_vector_stores is None else SimpleNamespace(vector_stores=team_vector_stores),
|
||||
_user_row(user_vector_stores, teams=() if team_id is None else (team_id,)),
|
||||
user_permission or (None if user_vector_stores is None else SimpleNamespace(vector_stores=user_vector_stores)),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("user_vector_stores", "denied_by"),
|
||||
[(["KBSTOREA"], None), (["KBSTOREB"], _USER_DENIED), ([], _USER_DENIED), (None, _USER_DENIED)],
|
||||
ids=["user-grants", "user-excludes", "user-empty", "user-no-record"],
|
||||
)
|
||||
async def test_keyless_user_without_team_needs_user_grant_under_deny_by_default(
|
||||
user_vector_stores: list[str] | None, denied_by: ProxyErrorTypes | None
|
||||
):
|
||||
await _assert_rag_query_outcome(
|
||||
_keyless_rag_query({"vector_store_deny_by_default": True}, user_vector_stores), denied_by
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_keyless_user_null_vector_store_list_grants_nothing_under_deny_by_default():
|
||||
null_list: Final = LiteLLM_ObjectPermissionTable(object_permission_id="user-permission", vector_stores=None)
|
||||
|
||||
await _assert_rag_query_outcome(
|
||||
_keyless_rag_query({"vector_store_deny_by_default": True}, [], user_permission=null_list), _USER_DENIED
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("general_settings", "user_vector_stores"),
|
||||
[({}, ["KBSTOREB"]), ({"vector_store_deny_by_default": False}, ["KBSTOREB"]), ({}, None)],
|
||||
ids=["omitted-user-excludes", "false-user-excludes", "omitted-user-no-record"],
|
||||
)
|
||||
async def test_keyless_user_keeps_legacy_vector_store_behavior_when_flag_off(
|
||||
general_settings: Mapping[str, object], user_vector_stores: list[str] | None
|
||||
):
|
||||
await _assert_rag_query_outcome(_keyless_rag_query(general_settings, user_vector_stores), None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("team_vector_stores", "user_vector_stores", "denied_by"),
|
||||
[
|
||||
(["KBSTOREA"], [], None),
|
||||
(["KBSTOREB"], ["KBSTOREA"], _TEAM_DENIED),
|
||||
([], ["KBSTOREA"], _TEAM_DENIED),
|
||||
(None, ["KBSTOREA"], _TEAM_DENIED),
|
||||
],
|
||||
ids=["team-grants-user-empty", "team-excludes-user-grants", "team-empty-user-grants", "team-no-record-user-grants"],
|
||||
)
|
||||
async def test_keyless_team_member_needs_only_team_grant_under_deny_by_default(
|
||||
team_vector_stores: list[str] | None, user_vector_stores: list[str] | None, denied_by: ProxyErrorTypes | None
|
||||
):
|
||||
await _assert_rag_query_outcome(
|
||||
_keyless_rag_query({"vector_store_deny_by_default": True}, user_vector_stores, team_vector_stores, "team-1"),
|
||||
denied_by,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("team_id", "denied_by"),
|
||||
[("team-1", None), (None, None)],
|
||||
ids=["session-on-team-uses-team-grant", "session-without-team-uses-user-grant"],
|
||||
)
|
||||
async def test_session_token_is_not_a_virtual_key_under_deny_by_default(
|
||||
team_id: str | None, denied_by: ProxyErrorTypes | None
|
||||
):
|
||||
session: Final = UserAPIKeyAuth(
|
||||
api_key="hashed-session-token",
|
||||
user_id="user-1",
|
||||
team_id=team_id,
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
is_session_token=True,
|
||||
)
|
||||
session.via_virtual_key = True
|
||||
|
||||
await _assert_rag_query_outcome(
|
||||
_common_checks_for_rag_query(
|
||||
{"vector_store_deny_by_default": True},
|
||||
session,
|
||||
None,
|
||||
None if team_id is None else LiteLLM_TeamTable(team_id=team_id, object_permission_id="team-permission"),
|
||||
SimpleNamespace(vector_stores=["KBSTOREA"]),
|
||||
_user_row(["KBSTOREA"] if team_id is None else []),
|
||||
SimpleNamespace(vector_stores=["KBSTOREA"] if team_id is None else []),
|
||||
),
|
||||
denied_by,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_owned_standalone_key_cannot_use_owner_grants_under_deny_by_default():
|
||||
owner: Final = LiteLLM_UserTable(
|
||||
user_id="key-owner", user_role=LitellmUserRoles.INTERNAL_USER, object_permission_id="user-permission"
|
||||
)
|
||||
|
||||
await _assert_rag_query_outcome(
|
||||
_common_checks_for_rag_query(
|
||||
{"vector_store_deny_by_default": True},
|
||||
_virtual_key(),
|
||||
None,
|
||||
user_object=owner,
|
||||
user_permission=SimpleNamespace(vector_stores=["KBSTOREA"]),
|
||||
),
|
||||
_KEY_DENIED,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("deny_by_default", "key_vector_stores", "denied_by"),
|
||||
[
|
||||
(True, ["KBSTOREA", "KBSTOREB"], None),
|
||||
(True, ["KBSTOREA"], _KEY_DENIED),
|
||||
(True, ["KBSTOREB"], _KEY_DENIED),
|
||||
(False, ["KBSTOREA"], _KEY_DENIED),
|
||||
],
|
||||
ids=["enabled-grants-both", "enabled-missing-second", "enabled-missing-first", "disabled-missing-second"],
|
||||
)
|
||||
async def test_every_requested_vector_store_needs_a_grant(
|
||||
deny_by_default: bool, key_vector_stores: list[str], denied_by: ProxyErrorTypes | None
|
||||
):
|
||||
await _assert_rag_query_outcome(
|
||||
_common_checks_for_rag_query(
|
||||
{"vector_store_deny_by_default": deny_by_default},
|
||||
_virtual_key(object_permission_id="key-permission"),
|
||||
SimpleNamespace(vector_stores=key_vector_stores),
|
||||
request_body={
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [{"role": "user", "content": "what is in these KBs?"}],
|
||||
"tools": [{"type": "file_search", "vector_store_ids": ["KBSTOREB"]}],
|
||||
"retrieval_config": {"vector_store_id": "KBSTOREA", "custom_llm_provider": "bedrock"},
|
||||
},
|
||||
vector_store_registry=VectorStoreRegistry(),
|
||||
),
|
||||
denied_by,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"valid_token",
|
||||
[_virtual_key(), UserAPIKeyAuth(user_id="user-1", user_role=LitellmUserRoles.INTERNAL_USER)],
|
||||
ids=["standalone-key", "keyless-user"],
|
||||
)
|
||||
async def test_request_without_vector_stores_is_unaffected_by_deny_by_default(valid_token: UserAPIKeyAuth):
|
||||
await _assert_rag_query_outcome(
|
||||
_common_checks_for_rag_query(
|
||||
{"vector_store_deny_by_default": True},
|
||||
valid_token,
|
||||
None,
|
||||
user_object=_user_row(None),
|
||||
request_body={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hello"}]},
|
||||
vector_store_registry=VectorStoreRegistry(),
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"request_body",
|
||||
[
|
||||
{"query": "what is in this KB?", "vector_store_id": "KBSTOREA", "vector_store_ids": ["KBSTOREA"]},
|
||||
{
|
||||
"model": "gpt-4o-mini",
|
||||
"input": "what is in this KB?",
|
||||
"tools": [{"type": "file_search", "vector_store_ids": ["KBSTOREA"]}],
|
||||
},
|
||||
],
|
||||
ids=["vector-store-search-route", "responses-file-search"],
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
("deny_by_default", "key_vector_stores", "denied_by"),
|
||||
[(True, None, _KEY_DENIED), (True, ["KBSTOREA"], None), (False, None, None)],
|
||||
ids=["enabled-no-grant", "enabled-grant", "disabled-no-grant"],
|
||||
)
|
||||
async def test_deny_by_default_reads_requested_vector_stores_without_a_registry(
|
||||
request_body: Mapping[str, object],
|
||||
deny_by_default: bool,
|
||||
key_vector_stores: list[str] | None,
|
||||
denied_by: ProxyErrorTypes | None,
|
||||
):
|
||||
await _assert_rag_query_outcome(
|
||||
_common_checks_for_rag_query(
|
||||
{"vector_store_deny_by_default": deny_by_default},
|
||||
_virtual_key(object_permission_id=None if key_vector_stores is None else "key-permission"),
|
||||
None if key_vector_stores is None else SimpleNamespace(vector_stores=key_vector_stores),
|
||||
request_body=request_body,
|
||||
),
|
||||
denied_by,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("request_body", "key_vector_stores", "denied_by"),
|
||||
[
|
||||
({"tools": [1]}, None, None),
|
||||
({"tools": 1}, None, None),
|
||||
({"tools": [{"type": "file_search", "vector_store_ids": None}]}, None, None),
|
||||
({"tools": [1, {"type": "file_search", "vector_store_ids": ["KBSTOREA"]}]}, ["KBSTOREA"], None),
|
||||
({"tools": [1, {"type": "file_search", "vector_store_ids": ["KBSTOREA"]}]}, ["KBSTOREB"], _KEY_DENIED),
|
||||
],
|
||||
ids=["int-tool", "non-list-tools", "null-tool-ids", "int-tool-beside-granted", "int-tool-beside-ungranted"],
|
||||
)
|
||||
async def test_deny_by_default_ignores_tools_that_name_no_vector_store(
|
||||
request_body: Mapping[str, object], key_vector_stores: list[str] | None, denied_by: ProxyErrorTypes | None
|
||||
):
|
||||
await _assert_rag_query_outcome(
|
||||
_common_checks_for_rag_query(
|
||||
{"vector_store_deny_by_default": True},
|
||||
_virtual_key(object_permission_id=None if key_vector_stores is None else "key-permission"),
|
||||
None if key_vector_stores is None else SimpleNamespace(vector_stores=key_vector_stores),
|
||||
request_body={"model": "gpt-4o-mini", "input": "what is in this KB?", **request_body},
|
||||
),
|
||||
denied_by,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"request_body",
|
||||
[
|
||||
{"tools": [{"type": "file_search", "vector_store_ids": "KBSTOREA"}]},
|
||||
{"tools": [{"type": "file_search", "vector_store_ids": [1]}]},
|
||||
{"vector_store_ids": "KBSTOREA"},
|
||||
],
|
||||
ids=["string-tool-ids", "int-tool-id", "string-top-level-ids"],
|
||||
)
|
||||
async def test_deny_by_default_rejects_malformed_vector_store_ids_as_bad_request(request_body: Mapping[str, object]):
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await _common_checks_for_rag_query(
|
||||
{"vector_store_deny_by_default": True},
|
||||
_virtual_key(object_permission_id="key-permission"),
|
||||
SimpleNamespace(vector_stores=["KBSTOREA"]),
|
||||
request_body={"model": "gpt-4o-mini", "input": "what is in this KB?", **request_body},
|
||||
)
|
||||
assert (exc_info.value.type, exc_info.value.param, exc_info.value.code) == (
|
||||
"invalid_request_error",
|
||||
"vector_store_ids",
|
||||
"400",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("request_body", "denied_by"),
|
||||
[
|
||||
({"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hello"}]}, None),
|
||||
(None, _KEY_DENIED),
|
||||
],
|
||||
ids=["no-vector-store", "ungranted-vector-store"],
|
||||
)
|
||||
@pytest.mark.parametrize("flag_value", [None, "enabled"], ids=["null", "string"])
|
||||
async def test_invalid_deny_by_default_value_only_denies_vector_store_requests(
|
||||
flag_value: object, request_body: Mapping[str, object] | None, denied_by: ProxyErrorTypes | None
|
||||
):
|
||||
await _assert_rag_query_outcome(
|
||||
_common_checks_for_rag_query(
|
||||
{"vector_store_deny_by_default": flag_value}, _virtual_key(), None, request_body=request_body
|
||||
),
|
||||
denied_by,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_keyless_user_grant_is_read_through_the_object_permission_cache():
|
||||
cache: Final = UserApiKeyCache(default_in_memory_ttl=60)
|
||||
await cache.async_set_cache(
|
||||
key=object_permission_cache_key("user-permission"),
|
||||
value=LiteLLM_ObjectPermissionTable(object_permission_id="user-permission", vector_stores=["KBSTOREA"]),
|
||||
model_type=LiteLLM_ObjectPermissionTable,
|
||||
)
|
||||
|
||||
await _assert_rag_query_outcome(
|
||||
_common_checks_for_rag_query(
|
||||
{"vector_store_deny_by_default": True},
|
||||
UserAPIKeyAuth(user_id="user-1", user_role=LitellmUserRoles.INTERNAL_USER),
|
||||
None,
|
||||
user_object=_user_row(["KBSTOREA"]),
|
||||
user_api_key_cache=cache,
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("deny_by_default", [True, False], ids=["strict", "default"])
|
||||
async def test_key_and_team_grants_are_read_through_the_object_permission_cache(deny_by_default: bool):
|
||||
cache: Final = UserApiKeyCache(default_in_memory_ttl=60)
|
||||
for object_permission_id, vector_stores in (("key-permission", ["KBSTOREA"]), ("team-permission", ["KBSTOREB"])):
|
||||
await cache.async_set_cache(
|
||||
key=object_permission_cache_key(object_permission_id),
|
||||
value=LiteLLM_ObjectPermissionTable(object_permission_id=object_permission_id, vector_stores=vector_stores),
|
||||
model_type=LiteLLM_ObjectPermissionTable,
|
||||
)
|
||||
|
||||
await _assert_rag_query_outcome(
|
||||
_common_checks_for_rag_query(
|
||||
{"vector_store_deny_by_default": deny_by_default},
|
||||
_virtual_key(object_permission_id="key-permission", team_id="team-1"),
|
||||
None,
|
||||
LiteLLM_TeamTable(team_id="team-1", object_permission_id="team-permission"),
|
||||
None,
|
||||
user_api_key_cache=cache,
|
||||
),
|
||||
_TEAM_DENIED,
|
||||
)
|
||||
|
||||
|
||||
def test_can_object_call_model_with_alias():
|
||||
"""Test that can_object_call_model works with model aliases"""
|
||||
from litellm import Router
|
||||
|
|
@ -10346,7 +10936,9 @@ async def test_authoritative_access_group_reads_writer_despite_stale_allow_cache
|
|||
from litellm.proxy._types import LiteLLM_AccessGroupTable
|
||||
from litellm.proxy.auth.auth_checks import get_access_object
|
||||
|
||||
stale: Final = LiteLLM_AccessGroupTable(access_group_id="group", access_group_name="Policy", access_model_names=["old"])
|
||||
stale: Final = LiteLLM_AccessGroupTable(
|
||||
access_group_id="group", access_group_name="Policy", access_model_names=["old"]
|
||||
)
|
||||
current: Final = stale.model_copy(update={"access_model_names": ["new"] if allowed else []})
|
||||
client: Final = MagicMock()
|
||||
client.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=current)
|
||||
|
|
@ -10498,9 +11090,12 @@ async def test_authoritative_key_cannot_keep_grants_when_permission_is_unavailab
|
|||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
|
||||
database: Final = MagicMock()
|
||||
database.get_data = AsyncMock(return_value=UserAPIKeyAuth(
|
||||
object_permission_id="grant", object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["allowed"])
|
||||
))
|
||||
database.get_data = AsyncMock(
|
||||
return_value=UserAPIKeyAuth(
|
||||
object_permission_id="grant",
|
||||
object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="grant", agents=["allowed"]),
|
||||
)
|
||||
)
|
||||
database.writer_db.litellm_objectpermissiontable.find_unique = AsyncMock(
|
||||
return_value=None, side_effect=None if missing else RuntimeError("writer unavailable")
|
||||
)
|
||||
|
|
@ -10523,7 +11118,9 @@ async def test_authoritative_group_grants_propagate_policy_outages(
|
|||
|
||||
database: Final = MagicMock()
|
||||
database.db.litellm_accessgrouptable.find_unique = AsyncMock(side_effect=RuntimeError("database unavailable"))
|
||||
database.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(side_effect=RuntimeError("database unavailable"))
|
||||
database.writer_db.litellm_accessgrouptable.find_unique = AsyncMock(
|
||||
side_effect=RuntimeError("database unavailable")
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", database)
|
||||
monkeypatch.setattr(proxy_server, "user_api_key_cache", UserApiKeyCache())
|
||||
if strict:
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ from functools import partial
|
|||
from pathlib import Path
|
||||
from textwrap import dedent
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, patch
|
||||
|
||||
|
||||
|
|
@ -21,11 +22,13 @@ from fastapi import HTTPException, status
|
|||
import litellm
|
||||
import litellm.proxy.proxy_server
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy._types import (
|
||||
LiteLLMRoutes,
|
||||
LiteLLM_JWTAuth,
|
||||
LiteLLM_BudgetTable,
|
||||
LiteLLM_EndUserTable,
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
LiteLLM_OrganizationTable,
|
||||
LiteLLM_TeamTableCachedObj,
|
||||
LiteLLM_UserTable,
|
||||
|
|
@ -5304,6 +5307,222 @@ async def test_centralized_common_checks_tolerates_db_errors_when_fetching_conte
|
|||
setattr(_proxy_server_mod, k, v)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("deny_by_default", "team_lookup_error", "denied"),
|
||||
[
|
||||
(False, RuntimeError("team cache unavailable"), False),
|
||||
(True, RuntimeError("team cache unavailable"), True),
|
||||
(True, HTTPException(status_code=404, detail="team read failed"), True),
|
||||
],
|
||||
ids=["flag-off-lookup-swallowed", "flag-on-lookup-swallowed", "flag-on-team-rebuilt-from-token"],
|
||||
)
|
||||
async def test_team_key_vector_store_access_when_team_cannot_be_resolved(
|
||||
monkeypatch: pytest.MonkeyPatch, deny_by_default: bool, team_lookup_error: Exception, denied: bool
|
||||
):
|
||||
from fastapi import Request
|
||||
from starlette.datastructures import URL
|
||||
|
||||
token: Final = UserAPIKeyAuth(
|
||||
api_key="sk-team-key", team_id="team-1", team_models=["gpt-4o-mini"], object_permission_id="key-permission"
|
||||
)
|
||||
token.via_virtual_key = True
|
||||
request: Final = Request(scope={"type": "http"})
|
||||
request._url = URL(url="/v1/rag/query")
|
||||
database: Final = MagicMock()
|
||||
database.db.litellm_objectpermissiontable.find_unique = AsyncMock(
|
||||
side_effect=lambda where: LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id=where["object_permission_id"], vector_stores=["KBSTOREA"]
|
||||
)
|
||||
if where["object_permission_id"] == "key-permission"
|
||||
else None
|
||||
)
|
||||
attrs: Final = {
|
||||
**_proxy_attrs_for_centralized_checks(),
|
||||
"prisma_client": database,
|
||||
"proxy_logging_obj": MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock())),
|
||||
"general_settings": {"vector_store_deny_by_default": deny_by_default},
|
||||
}
|
||||
for name, value in attrs.items():
|
||||
monkeypatch.setattr(litellm.proxy.proxy_server, name, value)
|
||||
monkeypatch.setattr(litellm, "vector_store_registry", None)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.auth.user_api_key_auth.get_team_object", AsyncMock(side_effect=team_lookup_error)
|
||||
)
|
||||
|
||||
checks: Final = _run_centralized_common_checks(
|
||||
user_api_key_auth_obj=token,
|
||||
request=request,
|
||||
request_data={
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [{"role": "user", "content": "what is in this KB?"}],
|
||||
"retrieval_config": {"vector_store_id": "KBSTOREA", "custom_llm_provider": "bedrock"},
|
||||
},
|
||||
route="/v1/rag/query",
|
||||
)
|
||||
if not denied:
|
||||
await checks
|
||||
return
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await checks
|
||||
assert (exc_info.value.type, exc_info.value.param, exc_info.value.code) == (
|
||||
ProxyErrorTypes.team_vector_store_access_denied,
|
||||
"vector_store",
|
||||
"401",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("team_lookup", "denied"),
|
||||
[
|
||||
({"team-1": "team-1-grants-a"}, False),
|
||||
(RuntimeError("team cache unavailable"), True),
|
||||
({"team-1": "team-1-grants-none", "team-2": "team-2-grants-a"}, True),
|
||||
],
|
||||
ids=["resolved-team-grants", "team-lookup-swallowed-no-personal-fallback", "other-member-team-grants-ignored"],
|
||||
)
|
||||
async def test_keyless_team_member_vector_store_access_uses_only_the_resolved_team(
|
||||
monkeypatch: pytest.MonkeyPatch, team_lookup: dict[str, str] | Exception, denied: bool
|
||||
):
|
||||
from fastapi import Request
|
||||
from starlette.datastructures import URL
|
||||
|
||||
token: Final = UserAPIKeyAuth(
|
||||
user_id="user-1",
|
||||
team_id="team-1",
|
||||
team_models=["gpt-4o-mini"],
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
)
|
||||
request: Final = Request(scope={"type": "http"})
|
||||
request._url = URL(url="/v1/rag/query")
|
||||
grants: Final = {
|
||||
"team-1-grants-a": ["KBSTOREA"],
|
||||
"team-1-grants-none": [],
|
||||
"team-2-grants-a": ["KBSTOREA"],
|
||||
"user-grants-a": ["KBSTOREA"],
|
||||
}
|
||||
database: Final = MagicMock()
|
||||
database.db.litellm_objectpermissiontable.find_unique = AsyncMock(
|
||||
side_effect=lambda where: LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id=where["object_permission_id"], vector_stores=grants[where["object_permission_id"]]
|
||||
)
|
||||
)
|
||||
attrs: Final = {
|
||||
**_proxy_attrs_for_centralized_checks(),
|
||||
"prisma_client": database,
|
||||
"proxy_logging_obj": MagicMock(service_logging_obj=MagicMock(async_service_success_hook=AsyncMock())),
|
||||
"general_settings": {"vector_store_deny_by_default": True},
|
||||
}
|
||||
for name, value in attrs.items():
|
||||
monkeypatch.setattr(litellm.proxy.proxy_server, name, value)
|
||||
monkeypatch.setattr(litellm, "vector_store_registry", None)
|
||||
|
||||
async def get_team(team_id: str, **_: object) -> LiteLLM_TeamTableCachedObj:
|
||||
if isinstance(team_lookup, Exception):
|
||||
raise team_lookup
|
||||
return LiteLLM_TeamTableCachedObj(team_id=team_id, models=["gpt-4o-mini"], object_permission_id=team_lookup[team_id])
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.auth.user_api_key_auth.get_team_object", get_team)
|
||||
monkeypatch.setattr("litellm.proxy.auth.auth_checks.get_team_membership", AsyncMock(return_value=None))
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.auth.user_api_key_auth.get_user_object",
|
||||
AsyncMock(
|
||||
return_value=LiteLLM_UserTable(
|
||||
user_id="user-1", teams=["team-1", "team-2"], object_permission_id="user-grants-a"
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
checks: Final = _run_centralized_common_checks(
|
||||
user_api_key_auth_obj=token,
|
||||
request=request,
|
||||
request_data={
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [{"role": "user", "content": "what is in this KB?"}],
|
||||
"retrieval_config": {"vector_store_id": "KBSTOREA", "custom_llm_provider": "bedrock"},
|
||||
},
|
||||
route="/v1/rag/query",
|
||||
)
|
||||
if not denied:
|
||||
await checks
|
||||
return
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await checks
|
||||
assert (exc_info.value.type, exc_info.value.param, exc_info.value.code) == (
|
||||
ProxyErrorTypes.team_vector_store_access_denied,
|
||||
"vector_store",
|
||||
"401",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("user_permission_id", "denied"),
|
||||
[("admin-grants-a", False), (None, True)],
|
||||
ids=["admin-personal-grant", "admin-without-grant"],
|
||||
)
|
||||
async def test_keyless_proxy_admin_keeps_personal_vector_store_grants_under_deny_by_default(
|
||||
monkeypatch: pytest.MonkeyPatch, user_permission_id: str | None, denied: bool
|
||||
):
|
||||
from fastapi import Request
|
||||
from starlette.datastructures import URL
|
||||
|
||||
token: Final = UserAPIKeyAuth(user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
request: Final = Request(scope={"type": "http"})
|
||||
request._url = URL(url="/v1/rag/query")
|
||||
database: Final = MagicMock()
|
||||
database.db.litellm_objectpermissiontable.find_unique = AsyncMock(
|
||||
side_effect=lambda where: SimpleNamespace(
|
||||
dict=lambda: {"object_permission_id": where["object_permission_id"], "vector_stores": ["KBSTOREA"]},
|
||||
vector_stores=["KBSTOREA"],
|
||||
)
|
||||
)
|
||||
attrs: Final = {
|
||||
**_proxy_attrs_for_centralized_checks(),
|
||||
"prisma_client": database,
|
||||
"general_settings": {"vector_store_deny_by_default": True},
|
||||
"user_api_key_cache": UserApiKeyCache(),
|
||||
"proxy_logging_obj": MagicMock(
|
||||
service_logging_obj=MagicMock(
|
||||
async_service_success_hook=AsyncMock(), async_service_failure_hook=AsyncMock()
|
||||
)
|
||||
),
|
||||
}
|
||||
for name, value in attrs.items():
|
||||
monkeypatch.setattr(litellm.proxy.proxy_server, name, value)
|
||||
monkeypatch.setattr(litellm, "vector_store_registry", None)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.auth.user_api_key_auth.get_user_object",
|
||||
AsyncMock(
|
||||
return_value=LiteLLM_UserTable(
|
||||
user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN, object_permission_id=user_permission_id
|
||||
)
|
||||
),
|
||||
)
|
||||
|
||||
checks: Final = _run_centralized_common_checks(
|
||||
user_api_key_auth_obj=token,
|
||||
request=request,
|
||||
request_data={
|
||||
"model": "gpt-4o-mini",
|
||||
"messages": [{"role": "user", "content": "what is in this KB?"}],
|
||||
"retrieval_config": {"vector_store_id": "KBSTOREA", "custom_llm_provider": "bedrock"},
|
||||
},
|
||||
route="/v1/rag/query",
|
||||
)
|
||||
if not denied:
|
||||
await checks
|
||||
return
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await checks
|
||||
assert (exc_info.value.type, exc_info.value.param, exc_info.value.code) == (
|
||||
ProxyErrorTypes.user_vector_store_access_denied,
|
||||
"vector_store",
|
||||
"401",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_centralized_common_checks_propagates_end_user_budget_error():
|
||||
"""Regression: ``get_end_user_object`` raises ``litellm.BudgetExceededError``
|
||||
|
|
|
|||
|
|
@ -4066,6 +4066,34 @@ async def test_user_update_invalidates_the_cached_entitlement(mocker):
|
|||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_update_broadcasts_the_entitlement_invalidation_to_other_workers(mocker: MockerFixture):
|
||||
from litellm.proxy.management_endpoints.internal_user_endpoints import (
|
||||
_update_single_user_helper,
|
||||
)
|
||||
|
||||
_object_permission_mocks(mocker)
|
||||
cache: Final = mocker.MagicMock()
|
||||
cache.async_delete_cache = mocker.AsyncMock()
|
||||
mocker.patch("litellm.proxy.proxy_server.user_api_key_cache", cache) # test-quality-ok: substitute the cache dependency
|
||||
broadcast: Final = mocker.patch( # test-quality-ok: observe the Redis publication boundary
|
||||
"litellm.proxy.common_utils.auth_cache_invalidation_pubsub.publish_auth_cache_invalidation",
|
||||
new_callable=mocker.AsyncMock,
|
||||
)
|
||||
|
||||
await _update_single_user_helper(
|
||||
user_request=UpdateUserRequest(user_id="target-user", object_permission={"vector_stores": []}),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
broadcast_keys: Final = {call.kwargs["cache_key"] for call in broadcast.await_args_list}
|
||||
assert broadcast_keys == {
|
||||
"object_permission_id:perm-new",
|
||||
"user_object_permission_id:target-user",
|
||||
"target-user",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_admin_can_clear_a_users_mcp_entitlement(mocker):
|
||||
"""An explicit empty object_permission means "no object permission", so it must unlink.
|
||||
|
|
|
|||
|
|
@ -8417,6 +8417,7 @@ async def test_update_team_org_scoped_tpm_rpm_bypasses_user_limit(
|
|||
# Mock team update
|
||||
mock_updated_team = MagicMock(spec=LiteLLM_TeamTable)
|
||||
mock_updated_team.team_id = "org-team-update-bypass-123"
|
||||
mock_updated_team.object_permission_id = None
|
||||
mock_updated_team.tpm_limit = 10000
|
||||
mock_updated_team.rpm_limit = 1000
|
||||
mock_updated_team.access_group_ids = None
|
||||
|
|
@ -8570,6 +8571,7 @@ async def test_update_team_guardrails_with_org_id(
|
|||
# Mock team update
|
||||
mock_updated_team = MagicMock(spec=LiteLLM_TeamTable)
|
||||
mock_updated_team.team_id = "team-guardrails-123"
|
||||
mock_updated_team.object_permission_id = None
|
||||
mock_updated_team.organization_id = "test-org-guardrails"
|
||||
mock_updated_team.metadata = {
|
||||
"guardrails": ["aporia-pre-call", "aporia-post-call"]
|
||||
|
|
|
|||
|
|
@ -26,7 +26,7 @@ import pytest
|
|||
from pydantic import JsonValue, TypeAdapter, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm.proxy._types import CommonProxyErrors
|
||||
from litellm.proxy._types import CommonProxyErrors, ConfigGeneralSettings
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper
|
||||
from litellm.proxy.proxy_server import (
|
||||
ProxyConfig,
|
||||
|
|
@ -2329,6 +2329,42 @@ async def test_load_config_logs_disabled_budget_reservation_once(tmp_path, monke
|
|||
assert [record.levelno for record in records] == ([logging.INFO] if setting == "true" else [])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@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
|
||||
):
|
||||
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"
|
||||
)
|
||||
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
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@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
|
||||
):
|
||||
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"
|
||||
)
|
||||
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"):
|
||||
await ProxyConfig().load_config(router=None, config_file_path=str(config_file))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ProxyConfig_load_config_resolves_router_settings_plugins(tmp_path, monkeypatch):
|
||||
"""Regression: router_settings.plugins dotted-path strings must be resolved to
|
||||
|
|
|
|||
6
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
6
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -30083,6 +30083,12 @@ export interface components {
|
|||
* @description Master switch for the SSRF guard applied to user-supplied URLs (image_url, file_url, MCP/OpenAPI spec URLs, etc). Defaults to True. Set to False to disable DNS/IP validation entirely (not recommended).
|
||||
*/
|
||||
user_url_validation?: boolean | null;
|
||||
/**
|
||||
* Vector Store Deny By Default
|
||||
* @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
|
||||
* @default false
|
||||
*/
|
||||
vector_store_deny_by_default: boolean;
|
||||
};
|
||||
/** ConfigList */
|
||||
ConfigList: {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue