mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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>
This commit is contained in:
parent
4c648f181a
commit
f17918c250
4 changed files with 212 additions and 12 deletions
|
|
@ -2966,6 +2966,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 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",
|
||||
)
|
||||
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.",
|
||||
|
|
|
|||
|
|
@ -33,6 +33,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,6 +45,7 @@ from litellm.models.project import LiteLLM_ProjectTable
|
|||
from litellm.proxy._types import (
|
||||
RBAC_ROLES,
|
||||
CallInfo,
|
||||
ConfigGeneralSettings,
|
||||
LiteLLM_AccessGroupTable,
|
||||
LiteLLM_BudgetTable,
|
||||
LiteLLM_EndUserTable,
|
||||
|
|
@ -1281,6 +1283,15 @@ async def common_checks(
|
|||
request_body=request_body,
|
||||
team_object=team_object,
|
||||
valid_token=valid_token,
|
||||
deny_by_default=ConfigGeneralSettings.model_validate(
|
||||
MappingProxyType(
|
||||
{
|
||||
"vector_store_deny_by_default": _typed_request_body(general_settings).get(
|
||||
"vector_store_deny_by_default", False
|
||||
)
|
||||
}
|
||||
)
|
||||
).vector_store_deny_by_default,
|
||||
)
|
||||
|
||||
# 12. [OPTIONAL] Tool allowlist - key/team allowed_tools (no DB in hot path)
|
||||
|
|
@ -6622,10 +6633,39 @@ 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:
|
||||
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
|
||||
)
|
||||
|
||||
|
||||
def _require_key_vector_store_grant(
|
||||
vector_store_ids_to_run: Sequence[str], key_object_permission: _VectorStorePermissionsRow | None
|
||||
) -> None:
|
||||
if key_object_permission is None or not key_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,
|
||||
param="vector_store",
|
||||
code=status.HTTP_401_UNAUTHORIZED,
|
||||
)
|
||||
_can_object_call_vector_stores(
|
||||
object_type="key",
|
||||
vector_store_ids_to_run=vector_store_ids_to_run,
|
||||
object_permissions=key_object_permission,
|
||||
)
|
||||
|
||||
|
||||
async def vector_store_access_check(
|
||||
request_body: dict,
|
||||
team_object: LiteLLM_TeamTable | None,
|
||||
valid_token: UserAPIKeyAuth | None,
|
||||
*,
|
||||
deny_by_default: bool = False,
|
||||
):
|
||||
"""
|
||||
Checks if the object (key, team, org) has access to the vector store.
|
||||
|
|
@ -6659,18 +6699,21 @@ 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(
|
||||
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:
|
||||
_can_object_call_vector_stores(
|
||||
object_type="key",
|
||||
vector_store_ids_to_run=vector_store_ids_to_run,
|
||||
object_permissions=key_object_permission,
|
||||
)
|
||||
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)
|
||||
elif key_object_permission is not None:
|
||||
_can_object_call_vector_stores(
|
||||
object_type="key",
|
||||
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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
@ -1813,6 +1815,138 @@ async def test_vector_store_access_check_enforces_team_allowlist_for_rag_query(
|
|||
assert exc_info.value.type == expected_error_type
|
||||
|
||||
|
||||
|
||||
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 = 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
|
||||
|
||||
|
||||
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,
|
||||
) -> bool:
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=key_permission)
|
||||
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=None)
|
||||
request_body = {
|
||||
"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", None
|
||||
),
|
||||
):
|
||||
return await common_checks(
|
||||
request_body=request_body,
|
||||
team_object=team_object,
|
||||
user_object=None,
|
||||
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], allowed: bool) -> None:
|
||||
if allowed:
|
||||
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",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("general_settings", "allowed"),
|
||||
[({}, True), ({"vector_store_deny_by_default": False}, True), ({"vector_store_deny_by_default": True}, False)],
|
||||
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
|
||||
):
|
||||
await _assert_rag_query_outcome(_common_checks_for_rag_query(general_settings, _virtual_key(), None), allowed)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("deny_by_default", "vector_stores", "allowed"),
|
||||
[
|
||||
(True, [], False),
|
||||
(True, ["KBSTOREA"], True),
|
||||
(True, ["KBSTOREB"], False),
|
||||
(True, None, False),
|
||||
(False, ["KBSTOREB"], False),
|
||||
],
|
||||
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
|
||||
):
|
||||
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),
|
||||
),
|
||||
allowed,
|
||||
)
|
||||
|
||||
|
||||
@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)
|
||||
|
||||
await _assert_rag_query_outcome(
|
||||
_common_checks_for_rag_query({"vector_store_deny_by_default": deny_by_default}, master_key, None), True
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("valid_token", "team_object", "allowed"),
|
||||
[
|
||||
(_virtual_key(user_role=LitellmUserRoles.PROXY_ADMIN), None, False),
|
||||
(_virtual_key(team_id="team-1"), LiteLLM_TeamTable(team_id="team-1"), True),
|
||||
],
|
||||
ids=["admin-owned-standalone-key-denied", "team-key-deferred"],
|
||||
)
|
||||
async def test_deny_by_default_scope_is_standalone_virtual_keys(
|
||||
valid_token: UserAPIKeyAuth, team_object: LiteLLM_TeamTable | None, allowed: bool
|
||||
):
|
||||
await _assert_rag_query_outcome(
|
||||
_common_checks_for_rag_query({"vector_store_deny_by_default": True}, valid_token, None, team_object), allowed
|
||||
)
|
||||
|
||||
|
||||
def test_can_object_call_model_with_alias():
|
||||
"""Test that can_object_call_model works with model aliases"""
|
||||
from litellm import Router
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
@ -2274,6 +2274,25 @@ 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 = 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
|
||||
async def test_ProxyConfig_load_config_resolves_router_settings_plugins(tmp_path, monkeypatch):
|
||||
"""Regression: router_settings.plugins dotted-path strings must be resolved to
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue