mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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>
This commit is contained in:
parent
e82721fd72
commit
bda87054fc
6 changed files with 302 additions and 46 deletions
|
|
@ -15,12 +15,13 @@ 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
|
||||
|
||||
from fastapi import HTTPException, Request, status
|
||||
from pydantic import BaseModel, TypeAdapter, ValidationError
|
||||
from typing_extensions import NotRequired, ReadOnly, Required, TypedDict, Unpack
|
||||
from typing_extensions import NotRequired, ReadOnly, Required, TypedDict, TypeIs, Unpack
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -6635,6 +6636,32 @@ 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
|
||||
|
||||
|
||||
def _is_object_list(value: object) -> TypeIs[list[object]]: # guard-ok: trivial isinstance narrowing
|
||||
return isinstance(value, list)
|
||||
|
||||
|
||||
def _is_object_mapping(value: object) -> TypeIs[Mapping[str, object]]: # guard-ok: request JSON keys are str
|
||||
return isinstance(value, dict)
|
||||
|
||||
|
||||
def _object_items(value: object) -> tuple[object, ...]:
|
||||
return tuple(value) if _is_object_list(value) else ()
|
||||
|
||||
|
||||
def _tool_vector_store_ids(tool: object) -> tuple[object, ...]:
|
||||
return _object_items(tool.get("vector_store_ids")) if _is_object_mapping(tool) else ()
|
||||
|
||||
|
||||
def _get_requested_vector_store_ids(request_body: Mapping[str, object]) -> tuple[str, ...]:
|
||||
candidate_ids: Final = (
|
||||
*_object_items(request_body.get("vector_store_ids")),
|
||||
*chain.from_iterable(_tool_vector_store_ids(tool) for tool in _object_items(request_body.get("tools"))),
|
||||
)
|
||||
return tuple(
|
||||
vector_store_id for vector_store_id in candidate_ids if isinstance(vector_store_id, str) and vector_store_id
|
||||
)
|
||||
|
||||
|
||||
def _is_strict_vector_store_identity(valid_token: UserAPIKeyAuth | None) -> bool:
|
||||
return (
|
||||
valid_token is not None
|
||||
|
|
@ -6675,7 +6702,7 @@ async def vector_store_access_check(
|
|||
|
||||
Raises ProxyException if the object (key, team, org) cannot access the specific vector store.
|
||||
"""
|
||||
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
|
||||
|
|
@ -6685,12 +6712,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)
|
||||
_get_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)))
|
||||
|
|
@ -6742,8 +6774,10 @@ async def vector_store_access_check(
|
|||
|
||||
if strict_identity and not strict_key and not has_team:
|
||||
user_object_permission: Final = (
|
||||
await _object_permission_table(ObjectPermissionRepository(prisma_client)).find_unique(
|
||||
where={"object_permission_id": user_object.object_permission_id},
|
||||
await get_object_permission(
|
||||
object_permission_id=user_object.object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
if user_object is not None and user_object.object_permission_id is not None
|
||||
else None
|
||||
|
|
|
|||
|
|
@ -3054,6 +3054,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:
|
||||
|
|
|
|||
|
|
@ -182,34 +182,44 @@ def _strict_denied_request(
|
|||
granted: Final = _permission_for_stores(store_id)
|
||||
empty: Final = _permission_for_stores()
|
||||
if case == "standalone_key_no_permission":
|
||||
key = scenario.key(models=models)
|
||||
return strict.gateway.request("POST", "/v1/rag/query", _rag_query_body(model, marker, store_id), key=key), (
|
||||
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 = scenario.team(models=models, object_permission=empty if key_grants_store else granted)
|
||||
key = scenario.key(team_id=team, models=models, object_permission=granted if key_grants_store else empty)
|
||||
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=key), (
|
||||
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":
|
||||
key = scenario.key(models=models, object_permission=_permission_for_stores(store_id))
|
||||
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=key), "key_vector_store_access_denied"
|
||||
return strict.gateway.request(
|
||||
"POST", "/v1/chat/completions", body, key=partial_key
|
||||
), "key_vector_store_access_denied"
|
||||
if case == "rag_alias_no_permission":
|
||||
key = scenario.key(models=models)
|
||||
return strict.gateway.request("POST", "/rag/query", _rag_query_body(model, marker, store_id), key=key), (
|
||||
alias_key: Final = scenario.key(models=models)
|
||||
return strict.gateway.request("POST", "/rag/query", _rag_query_body(model, marker, store_id), key=alias_key), (
|
||||
"key_vector_store_access_denied"
|
||||
)
|
||||
if case == "chat_retrieval_config_no_permission":
|
||||
key = scenario.key(models=models)
|
||||
chat_key: Final = scenario.key(models=models)
|
||||
return (
|
||||
strict.gateway.request("POST", "/v1/chat/completions", _rag_query_body(model, marker, store_id), key=key),
|
||||
strict.gateway.request(
|
||||
"POST", "/v1/chat/completions", _rag_query_body(model, marker, store_id), key=chat_key
|
||||
),
|
||||
"key_vector_store_access_denied",
|
||||
)
|
||||
user: Final = scenario.user(user_role="internal_user")
|
||||
|
|
@ -267,13 +277,15 @@ def _strict_granted_request(
|
|||
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"):
|
||||
team = scenario.team(models=models, object_permission=granted) if case.startswith("team") else None
|
||||
key = scenario.key(models=models, object_permission=granted, **({} if team is None else {"team_id": team}))
|
||||
return strict.gateway.request("POST", "/v1/rag/query", body, key=key)
|
||||
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":
|
||||
team = scenario.team(models=models, object_permission=granted)
|
||||
member: Final = scenario.member(team)
|
||||
return strict.gateway.request("POST", "/v1/rag/query", body, key=strict.jwt(member, (team,)))
|
||||
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))
|
||||
|
|
@ -457,3 +469,73 @@ def test_explicit_false_flag_keeps_legacy_vector_store_outcomes(tmp_path: Path)
|
|||
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
|
||||
|
|
|
|||
|
|
@ -92,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,
|
||||
)
|
||||
|
|
@ -1827,7 +1828,7 @@ def _virtual_key(
|
|||
user_role: LitellmUserRoles | None = None,
|
||||
team_id: str | None = None,
|
||||
) -> UserAPIKeyAuth:
|
||||
key = UserAPIKeyAuth(
|
||||
key: Final = UserAPIKeyAuth(
|
||||
api_key=api_key,
|
||||
user_id="key-owner",
|
||||
user_role=user_role,
|
||||
|
|
@ -1847,14 +1848,21 @@ async def _common_checks_for_rag_query(
|
|||
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": key_permission,
|
||||
"team-permission": team_permission,
|
||||
"user-permission": user_permission,
|
||||
"user-permission": (
|
||||
LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="user-permission", vector_stores=user_permission.vector_stores
|
||||
)
|
||||
if isinstance(user_permission, SimpleNamespace)
|
||||
else user_permission
|
||||
),
|
||||
}
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client: Final = MagicMock()
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(
|
||||
side_effect=lambda where: permissions.get(where["object_permission_id"])
|
||||
)
|
||||
|
|
@ -1875,6 +1883,10 @@ async def _common_checks_for_rag_query(
|
|||
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,
|
||||
|
|
@ -1958,7 +1970,7 @@ async def test_standalone_key_vector_store_permission_record_under_deny_by_defau
|
|||
@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 = _virtual_key(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
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
|
||||
|
|
@ -1978,10 +1990,10 @@ async def test_master_key_rag_query_is_unchanged_by_deny_by_default(deny_by_defa
|
|||
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 = _virtual_key(
|
||||
admin_key: Final = _virtual_key(
|
||||
object_permission_id=None if vector_stores is None else "key-permission", user_role=LitellmUserRoles.PROXY_ADMIN
|
||||
)
|
||||
key_permission = None if vector_stores is None else SimpleNamespace(vector_stores=vector_stores)
|
||||
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)
|
||||
|
||||
|
|
@ -2033,7 +2045,7 @@ async def test_team_key_keeps_legacy_vector_store_behavior_when_flag_off(
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_owned_team_key_without_key_grant_is_denied_despite_team_membership():
|
||||
team_member = LiteLLM_UserTable(user_id="key-owner", teams=["team-1"], user_role=LitellmUserRoles.INTERNAL_USER)
|
||||
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
|
||||
|
|
@ -2136,7 +2148,7 @@ async def test_keyless_team_member_needs_only_team_grant_under_deny_by_default(
|
|||
async def test_session_token_is_not_a_virtual_key_under_deny_by_default(
|
||||
team_id: str | None, denied_by: ProxyErrorTypes | None
|
||||
):
|
||||
session = UserAPIKeyAuth(
|
||||
session: Final = UserAPIKeyAuth(
|
||||
api_key="hashed-session-token",
|
||||
user_id="user-1",
|
||||
team_id=team_id,
|
||||
|
|
@ -2228,6 +2240,62 @@ async def test_request_without_vector_stores_is_unaffected_by_deny_by_default(va
|
|||
)
|
||||
|
||||
|
||||
@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
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
def test_can_object_call_model_with_alias():
|
||||
"""Test that can_object_call_model works with model aliases"""
|
||||
from litellm import Router
|
||||
|
|
|
|||
|
|
@ -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,6 +22,7 @@ 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,
|
||||
|
|
@ -5318,19 +5320,19 @@ async def test_team_key_vector_store_access_when_team_cannot_be_resolved(
|
|||
from fastapi import Request
|
||||
from starlette.datastructures import URL
|
||||
|
||||
token = UserAPIKeyAuth(
|
||||
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 = Request(scope={"type": "http"})
|
||||
request: Final = Request(scope={"type": "http"})
|
||||
request._url = URL(url="/v1/rag/query")
|
||||
database = MagicMock()
|
||||
database: Final = MagicMock()
|
||||
database.db.litellm_objectpermissiontable.find_unique = AsyncMock(
|
||||
side_effect=lambda where: SimpleNamespace(vector_stores=["KBSTOREA"])
|
||||
if where["object_permission_id"] == "key-permission"
|
||||
else None
|
||||
)
|
||||
attrs = {
|
||||
attrs: Final = {
|
||||
**_proxy_attrs_for_centralized_checks(),
|
||||
"prisma_client": database,
|
||||
"general_settings": {"vector_store_deny_by_default": deny_by_default},
|
||||
|
|
@ -5342,7 +5344,7 @@ async def test_team_key_vector_store_access_when_team_cannot_be_resolved(
|
|||
"litellm.proxy.auth.user_api_key_auth.get_team_object", AsyncMock(side_effect=team_lookup_error)
|
||||
)
|
||||
|
||||
checks = _run_centralized_common_checks(
|
||||
checks: Final = _run_centralized_common_checks(
|
||||
user_api_key_auth_obj=token,
|
||||
request=request,
|
||||
request_data={
|
||||
|
|
@ -5380,25 +5382,25 @@ async def test_keyless_team_member_vector_store_access_uses_only_the_resolved_te
|
|||
from fastapi import Request
|
||||
from starlette.datastructures import URL
|
||||
|
||||
token = UserAPIKeyAuth(
|
||||
token: Final = UserAPIKeyAuth(
|
||||
user_id="user-1",
|
||||
team_id="team-1",
|
||||
team_models=["gpt-4o-mini"],
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
)
|
||||
request = Request(scope={"type": "http"})
|
||||
request: Final = Request(scope={"type": "http"})
|
||||
request._url = URL(url="/v1/rag/query")
|
||||
grants = {
|
||||
grants: Final = {
|
||||
"team-1-grants-a": ["KBSTOREA"],
|
||||
"team-1-grants-none": [],
|
||||
"team-2-grants-a": ["KBSTOREA"],
|
||||
"user-grants-a": ["KBSTOREA"],
|
||||
}
|
||||
database = MagicMock()
|
||||
database: Final = MagicMock()
|
||||
database.db.litellm_objectpermissiontable.find_unique = AsyncMock(
|
||||
side_effect=lambda where: SimpleNamespace(vector_stores=grants[where["object_permission_id"]])
|
||||
)
|
||||
attrs = {
|
||||
attrs: Final = {
|
||||
**_proxy_attrs_for_centralized_checks(),
|
||||
"prisma_client": database,
|
||||
"general_settings": {"vector_store_deny_by_default": True},
|
||||
|
|
@ -5423,7 +5425,7 @@ async def test_keyless_team_member_vector_store_access_uses_only_the_resolved_te
|
|||
),
|
||||
)
|
||||
|
||||
checks = _run_centralized_common_checks(
|
||||
checks: Final = _run_centralized_common_checks(
|
||||
user_api_key_auth_obj=token,
|
||||
request=request,
|
||||
request_data={
|
||||
|
|
@ -5445,6 +5447,73 @@ async def test_keyless_team_member_vector_store_access_uses_only_the_resolved_te
|
|||
)
|
||||
|
||||
|
||||
@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``
|
||||
|
|
|
|||
|
|
@ -2279,7 +2279,7 @@ async def test_load_config_logs_disabled_budget_reservation_once(tmp_path, monke
|
|||
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 = tmp_path / "vector_store.yaml"
|
||||
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"
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue