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>
This commit is contained in:
mrinal 2026-10-02 20:19:10 +00:00
parent b94047f2ac
commit d3c21f716e
2 changed files with 301 additions and 7 deletions

View file

@ -1,5 +1,7 @@
from __future__ import annotations
import json
import time
import uuid
from collections.abc import Iterator, Mapping
from pathlib import Path
@ -7,16 +9,20 @@ from types import MappingProxyType
from typing import Final, Literal, TypeAlias
import httpx
import jwt
import pytest
import yaml
from cryptography.hazmat.primitives.asymmetric import rsa
from integration._support.client import Gateway, Scenario, gateway_from_environment, object_value
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, wire_server
from integration.authorization._guardrail_opt_out import upstream_observations
from pydantic import JsonValue
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"
JsonObject: TypeAlias = dict[str, JsonValue]
@ -99,6 +105,209 @@ 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) -> None:
self.gateway: Final = gateway
self.upstream: Final = upstream
self._signing_key: Final = signing_key
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")
with owned_proxy(
upstream_gateway,
directory,
{**_openai_environment(upstream_gateway), "JWT_PUBLIC_KEY_URL": jwks.url},
config=_strict_config(directory),
remove_environment=REMOVE_OPENAI_API_BASE,
) as gateway:
yield StrictGateway(gateway, upstream_gateway, signing_key)
StrictCase: TypeAlias = Literal[
"standalone_key_no_permission",
"team_key_empty_key_grants",
"team_key_empty_team_grants",
"multi_store_one_ungranted",
"rag_alias_no_permission",
"chat_retrieval_config_no_permission",
"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":
key = scenario.key(models=models)
return strict.gateway.request("POST", "/v1/rag/query", _rag_query_body(model, marker, store_id), key=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)
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), (
error_type
)
if case == "multi_store_one_ungranted":
key = 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"
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), (
"key_vector_store_access_denied"
)
if case == "chat_retrieval_config_no_permission":
key = scenario.key(models=models)
return (
strict.gateway.request("POST", "/v1/chat/completions", _rag_query_body(model, marker, store_id), key=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",
"rag_alias_no_permission",
"chat_retrieval_config_no_permission",
"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"):
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)
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,)))
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 +429,31 @@ 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) == ()

View file

@ -1846,6 +1846,8 @@ async def _common_checks_for_rag_query(
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,
vector_store_registry: VectorStoreRegistry | None = None,
) -> bool:
permissions: Final = {
"key-permission": key_permission,
@ -1857,21 +1859,25 @@ async def _common_checks_for_rag_query(
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",
"messages": [{"role": "user", "content": "what is in this KB?"}],
"retrieval_config": {"vector_store_id": "KBSTOREA", "custom_llm_provider": "bedrock"},
}
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", None
"litellm.vector_store_registry", vector_store_registry
),
):
return await common_checks(
request_body=request_body,
request_body=body,
team_object=team_object,
user_object=user_object,
end_user_object=None,
@ -2171,6 +2177,57 @@ async def test_user_owned_standalone_key_cannot_use_owner_grants_under_deny_by_d
)
@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,
)
def test_can_object_call_model_with_alias():
"""Test that can_object_call_model works with model aliases"""
from litellm import Router