diff --git a/tests/integration/authorization/test_rag_query_vector_store_allowlist.py b/tests/integration/authorization/test_rag_query_vector_store_allowlist.py index 896c88c68bb..4f1270a33aa 100644 --- a/tests/integration/authorization/test_rag_query_vector_store_allowlist.py +++ b/tests/integration/authorization/test_rag_query_vector_store_allowlist.py @@ -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) == () diff --git a/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py b/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py index 436a773b85f..18c73179dda 100644 --- a/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py +++ b/tests/unit/proxy/auth/test_auth_checks_object_access_and_lookup.py @@ -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