mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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>
This commit is contained in:
parent
f17918c250
commit
5a892fc548
4 changed files with 200 additions and 51 deletions
|
|
@ -2968,7 +2968,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
)
|
||||
vector_store_deny_by_default: bool = Field(
|
||||
default=False,
|
||||
description="When True, a virtual key without a team may only use vector stores explicitly listed in its object_permission.vector_stores. A key with no permission record or an empty list is denied. Team keys and non-key callers are not yet covered",
|
||||
description="When True, a virtual key may only use vector stores explicitly listed in its object_permission.vector_stores, and a key on a team also needs the team to list them. A missing permission record, an empty list, or an unresolved team grants nothing. Dashboard session keys and non-key callers are not yet covered",
|
||||
)
|
||||
missing_session_id: Literal["generate", "reject", "omit"] | None = Field(
|
||||
None,
|
||||
|
|
|
|||
|
|
@ -44,6 +44,7 @@ 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,
|
||||
|
|
@ -6633,30 +6634,31 @@ 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_standalone_virtual_key(valid_token: UserAPIKeyAuth | None, team_object: LiteLLM_TeamTable | None) -> bool:
|
||||
def _is_strict_vector_store_virtual_key(valid_token: UserAPIKeyAuth | None) -> bool:
|
||||
return (
|
||||
valid_token is not None
|
||||
and valid_token.via_virtual_key
|
||||
and valid_token.api_key != LITELLM_PROXY_MASTER_KEY_ALIAS
|
||||
and valid_token.team_id is None
|
||||
and team_object is None
|
||||
and valid_token.team_id != UI_TEAM_ID
|
||||
)
|
||||
|
||||
|
||||
def _require_key_vector_store_grant(
|
||||
vector_store_ids_to_run: Sequence[str], key_object_permission: _VectorStorePermissionsRow | None
|
||||
def _require_vector_store_grant(
|
||||
object_type: Literal["key", "team"],
|
||||
vector_store_ids_to_run: Sequence[str],
|
||||
object_permission: _VectorStorePermissionsRow | None,
|
||||
) -> None:
|
||||
if key_object_permission is None or not key_object_permission.vector_stores:
|
||||
if object_permission is None or not object_permission.vector_stores:
|
||||
raise ProxyException(
|
||||
message=f"Key not allowed to access vector store. Tried to access {vector_store_ids_to_run[0]}. vector_store_deny_by_default is enabled and the key has no vector store grants",
|
||||
type=ProxyErrorTypes.key_vector_store_access_denied,
|
||||
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="key",
|
||||
object_type=object_type,
|
||||
vector_store_ids_to_run=vector_store_ids_to_run,
|
||||
object_permissions=key_object_permission,
|
||||
object_permissions=object_permission,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -6706,8 +6708,9 @@ async def vector_store_access_check(
|
|||
if valid_token is not None and valid_token.object_permission_id is not None
|
||||
else None
|
||||
)
|
||||
if deny_by_default and _is_standalone_virtual_key(valid_token, team_object):
|
||||
_require_key_vector_store_grant(vector_store_ids_to_run, key_object_permission)
|
||||
strict_key: Final = deny_by_default and _is_strict_vector_store_virtual_key(valid_token)
|
||||
if strict_key:
|
||||
_require_vector_store_grant("key", vector_store_ids_to_run, key_object_permission)
|
||||
elif key_object_permission is not None:
|
||||
_can_object_call_vector_stores(
|
||||
object_type="key",
|
||||
|
|
@ -6716,18 +6719,21 @@ async def vector_store_access_check(
|
|||
)
|
||||
|
||||
# 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(
|
||||
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,
|
||||
)
|
||||
if team_object is not None and team_object.object_permission_id is not None
|
||||
else None
|
||||
)
|
||||
if strict_key and (team_object is not None or (valid_token is not None and valid_token.team_id is not None)):
|
||||
_require_vector_store_grant("team", vector_store_ids_to_run, team_object_permission)
|
||||
elif 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,
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1816,6 +1816,10 @@ async def test_vector_store_access_check_enforces_team_allowlist_for_rag_query(
|
|||
|
||||
|
||||
|
||||
_KEY_DENIED: Final = ProxyErrorTypes.key_vector_store_access_denied
|
||||
_TEAM_DENIED: Final = ProxyErrorTypes.team_vector_store_access_denied
|
||||
|
||||
|
||||
def _virtual_key(
|
||||
object_permission_id: str | None = None,
|
||||
api_key: str = "sk-standalone",
|
||||
|
|
@ -1838,9 +1842,14 @@ async def _common_checks_for_rag_query(
|
|||
valid_token: UserAPIKeyAuth,
|
||||
key_permission: SimpleNamespace | None,
|
||||
team_object: LiteLLM_TeamTable | None = None,
|
||||
team_permission: SimpleNamespace | None = None,
|
||||
user_object: LiteLLM_UserTable | None = None,
|
||||
) -> bool:
|
||||
permissions: Final = {"key-permission": key_permission, "team-permission": team_permission}
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=key_permission)
|
||||
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)
|
||||
request_body = {
|
||||
"model": "gpt-4o-mini",
|
||||
|
|
@ -1858,7 +1867,7 @@ async def _common_checks_for_rag_query(
|
|||
return await common_checks(
|
||||
request_body=request_body,
|
||||
team_object=team_object,
|
||||
user_object=None,
|
||||
user_object=user_object,
|
||||
end_user_object=None,
|
||||
global_proxy_spend=None,
|
||||
general_settings=dict(general_settings),
|
||||
|
|
@ -1870,45 +1879,59 @@ async def _common_checks_for_rag_query(
|
|||
)
|
||||
|
||||
|
||||
async def _assert_rag_query_outcome(request: Awaitable[bool], allowed: bool) -> None:
|
||||
if allowed:
|
||||
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) == (
|
||||
ProxyErrorTypes.key_vector_store_access_denied,
|
||||
"vector_store",
|
||||
"401",
|
||||
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", "allowed"),
|
||||
[({}, True), ({"vector_store_deny_by_default": False}, True), ({"vector_store_deny_by_default": True}, False)],
|
||||
("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], allowed: bool
|
||||
general_settings: Mapping[str, object], denied_by: ProxyErrorTypes | None
|
||||
):
|
||||
await _assert_rag_query_outcome(_common_checks_for_rag_query(general_settings, _virtual_key(), None), allowed)
|
||||
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", "allowed"),
|
||||
("deny_by_default", "vector_stores", "denied_by"),
|
||||
[
|
||||
(True, [], False),
|
||||
(True, ["KBSTOREA"], True),
|
||||
(True, ["KBSTOREB"], False),
|
||||
(True, None, False),
|
||||
(False, ["KBSTOREB"], False),
|
||||
(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, allowed: bool
|
||||
deny_by_default: bool, vector_stores: list[str] | None, denied_by: ProxyErrorTypes | None
|
||||
):
|
||||
await _assert_rag_query_outcome(
|
||||
_common_checks_for_rag_query(
|
||||
|
|
@ -1916,7 +1939,7 @@ async def test_standalone_key_vector_store_permission_record_under_deny_by_defau
|
|||
_virtual_key(object_permission_id="key-permission"),
|
||||
SimpleNamespace(vector_stores=vector_stores),
|
||||
),
|
||||
allowed,
|
||||
denied_by,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -1926,24 +1949,82 @@ async def test_master_key_rag_query_is_unchanged_by_deny_by_default(deny_by_defa
|
|||
master_key = _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), True
|
||||
_common_checks_for_rag_query({"vector_store_deny_by_default": deny_by_default}, master_key, None), None
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("valid_token", "team_object", "allowed"),
|
||||
("general_settings", "vector_stores", "denied_by"),
|
||||
[
|
||||
(_virtual_key(user_role=LitellmUserRoles.PROXY_ADMIN), None, False),
|
||||
(_virtual_key(team_id="team-1"), LiteLLM_TeamTable(team_id="team-1"), True),
|
||||
({}, None, None),
|
||||
({"vector_store_deny_by_default": False}, None, None),
|
||||
({"vector_store_deny_by_default": False}, ["KBSTOREB"], _KEY_DENIED),
|
||||
],
|
||||
ids=["admin-owned-standalone-key-denied", "team-key-deferred"],
|
||||
ids=["flag-omitted-no-record", "flag-false-no-record", "flag-false-excludes"],
|
||||
)
|
||||
async def test_deny_by_default_scope_is_standalone_virtual_keys(
|
||||
valid_token: UserAPIKeyAuth, team_object: LiteLLM_TeamTable | None, allowed: bool
|
||||
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(
|
||||
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)
|
||||
|
||||
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(
|
||||
_common_checks_for_rag_query({"vector_store_deny_by_default": True}, valid_token, None, team_object), allowed
|
||||
_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 = 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
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -4854,6 +4854,68 @@ 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 = 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._url = URL(url="/v1/rag/query")
|
||||
database = 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 = {
|
||||
**_proxy_attrs_for_centralized_checks(),
|
||||
"prisma_client": database,
|
||||
"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 = _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
|
||||
async def test_centralized_common_checks_propagates_end_user_budget_error():
|
||||
"""Regression: ``get_end_user_object`` raises ``litellm.BudgetExceededError``
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue