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 bd295f9e822..3ee320994d4 100644 --- a/tests/integration/authorization/test_rag_query_vector_store_allowlist.py +++ b/tests/integration/authorization/test_rag_query_vector_store_allowlist.py @@ -1,16 +1,20 @@ from __future__ import annotations +import asyncio import json import os import time import uuid from collections.abc import Iterator, Mapping +from concurrent.futures import ThreadPoolExecutor from pathlib import Path from types import MappingProxyType from typing import Final, Literal, TypeAlias +import anthropic import httpx import jwt +import openai import pytest import yaml from cryptography.hazmat.primitives.asymmetric import rsa @@ -21,6 +25,9 @@ from integration.authorization._guardrail_opt_out import upstream_observations from pydantic import JsonValue from redis import Redis +from litellm.proxy._types import LiteLLM_UserTable +from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken + 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",) @@ -66,6 +73,20 @@ def _rag_query( return gateway.request("POST", path, _rag_query_body(model, marker, store_id), key=key) +def _served_model(gateway: Gateway, scenario: Scenario) -> str: + model: Final = scenario.model() + statuses: Final = eventually( + lambda: frozenset( + _rag_query(gateway, model, f"lit6035 warm {uuid.uuid4().hex}", gateway.key).status_code for _ in range(8) + ), + lambda seen: seen == frozenset({200}), + seconds=30, + return_last_on_timeout=True, + ) + assert statuses == frozenset({200}), statuses + return model + + def _searches_for_marker( gateway: Gateway, marker: str, store_id: str = CONFIG_STORE_ID ) -> tuple[Mapping[str, JsonValue], ...]: @@ -174,6 +195,7 @@ def strict_gateway(tmp_path_factory: pytest.TempPathFactory) -> Iterator[StrictG environment, config=config, remove_environment=REMOVE_OPENAI_API_BASE, + workers=2, ) as gateway: yield StrictGateway(gateway, upstream_gateway, signing_key, config, environment) @@ -255,7 +277,7 @@ 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() + model: Final = _served_model(strict_gateway.gateway, scenario) store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" marker: Final = f"lit6035 deny by default {case} {uuid.uuid4().hex}" @@ -314,7 +336,7 @@ def _strict_granted_request( ) 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() + model: Final = _served_model(strict_gateway.gateway, scenario) store_id: Final = ( CONFIG_STORE_ID if case == "standalone_key_granted_registered_store" @@ -584,3 +606,333 @@ def test_revoked_user_grant_stops_working_on_another_proxy(strict_gateway: Stric strict_gateway.gateway.post("/user/update", {"user_id": user, "object_permission": _permission_for_stores()}) assert eventually(peer_status, lambda status: status == 401, seconds=10, return_last_on_timeout=True) == 401 + + +def _store_searches(gateway: Gateway, marker: str) -> tuple[Mapping[str, JsonValue], ...]: + return tuple( + observation + for observation in upstream_observations(gateway) + if str(observation["path"]).startswith("/vector_stores/") and marker in str(observation["body"]) + ) + + +def _marker_observations(gateway: Gateway, marker: str) -> tuple[Mapping[str, JsonValue], ...]: + return tuple(observation for observation in upstream_observations(gateway) if marker in str(observation["body"])) + + +def _messages_body(model: str, marker: str, store_id: str) -> JsonObject: + body: Final[JsonObject] = { + "model": model, + "max_tokens": 16, + "messages": _json_array({"role": "user", "content": marker}), + "vector_store_ids": _json_array(store_id), + } + return body + + +def _file_search_tools(store_id: JsonValue) -> JsonValue: + return _json_array({"type": "file_search", "vector_store_ids": store_id}) + + +RequestShape: TypeAlias = Literal["messages_vector_store_ids", "chat_stream_file_search", "chat_vector_store_ids"] + + +def _shape_request(shape: RequestShape, model: str, marker: str, store_id: str) -> tuple[str, JsonObject]: + messages: Final = _json_array({"role": "user", "content": marker}) + if shape == "messages_vector_store_ids": + return "/v1/messages", _messages_body(model, marker, store_id) + if shape == "chat_stream_file_search": + stream_body: Final[JsonObject] = { + "model": model, + "messages": messages, + "stream": True, + "tools": _file_search_tools(_json_array(store_id)), + } + return "/v1/chat/completions", stream_body + chat_body: Final[JsonObject] = {"model": model, "messages": messages, "vector_store_ids": _json_array(store_id)} + return "/v1/chat/completions", chat_body + + +@pytest.mark.parametrize("shape", ("messages_vector_store_ids", "chat_stream_file_search", "chat_vector_store_ids")) +def test_deny_by_default_covers_messages_streaming_and_top_level_store_ids( + strict_gateway: StrictGateway, shape: RequestShape +) -> None: + with strict_gateway.gateway.scenario() as scenario: + model: Final = _served_model(strict_gateway.gateway, scenario) + store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" + marker: Final = f"lit6035 shape {shape} {uuid.uuid4().hex}" + key: Final = scenario.key(models=_json_array(model)) + path, body = _shape_request(shape, model, marker, store_id) + + response: Final = strict_gateway.gateway.request("POST", path, body, key=key) + + assert response.status_code == 401, response.text + assert response.json()["error"]["type"] == "key_vector_store_access_denied", response.text + assert _marker_observations(strict_gateway.upstream, marker) == () + + +SdkClient: TypeAlias = Literal["openai_sync", "openai_async", "anthropic_sync"] + + +def _sdk_denial(gateway: Gateway, client: SdkClient, key: str, model: str, marker: str, store_id: str) -> int: + base_url: Final = str(gateway.client.base_url) + if client == "anthropic_sync": + with pytest.raises(anthropic.AuthenticationError) as anthropic_denied: + anthropic.Anthropic(base_url=base_url, api_key=key, max_retries=0).messages.create( + model=model, + max_tokens=16, + messages=[{"role": "user", "content": marker}], + extra_body={"vector_store_ids": _json_array(store_id)}, + ) + return anthropic_denied.value.status_code + tools: Final = [{"type": "file_search", "vector_store_ids": [store_id]}] + if client == "openai_sync": + with pytest.raises(openai.AuthenticationError) as sync_denied: + openai.OpenAI(base_url=f"{base_url}/v1", api_key=key, max_retries=0).responses.create( + model=model, input=marker, tools=tools + ) + return sync_denied.value.status_code + + async def create() -> None: + async with openai.AsyncOpenAI(base_url=f"{base_url}/v1", api_key=key, max_retries=0) as async_client: + await async_client.responses.create(model=model, input=marker, tools=tools) + + with pytest.raises(openai.AuthenticationError) as async_denied: + asyncio.run(create()) + return async_denied.value.status_code + + +@pytest.mark.parametrize("client", ("openai_sync", "openai_async", "anthropic_sync")) +def test_deny_by_default_rejects_sdk_clients_without_grant(strict_gateway: StrictGateway, client: SdkClient) -> None: + with strict_gateway.gateway.scenario() as scenario: + model: Final = _served_model(strict_gateway.gateway, scenario) + store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" + marker: Final = f"lit6035 sdk {client} {uuid.uuid4().hex}" + key: Final = scenario.key(models=_json_array(model)) + + status: Final = _sdk_denial(strict_gateway.gateway, client, key, model, marker, store_id) + + assert status == 401 + assert _marker_observations(strict_gateway.upstream, marker) == () + + +MalformedShape: TypeAlias = Literal[ + "ids_string", "ids_int", "ids_empty_string", "tool_ids_string", "tool_is_string", "oversized_id" +] + + +def _malformed_body(shape: MalformedShape, model: str, marker: str) -> JsonObject: + messages: Final = _json_array({"role": "user", "content": marker}) + if shape == "tool_ids_string": + tool_ids_body: Final[JsonObject] = { + "model": model, + "messages": messages, + "tools": _file_search_tools(f"vs_{uuid.uuid4().hex}"), + } + return tool_ids_body + if shape == "tool_is_string": + string_tool_body: Final[JsonObject] = { + "model": model, + "messages": messages, + "tools": _json_array("file_search"), + } + return string_tool_body + ids: Final[Mapping[MalformedShape, JsonValue]] = MappingProxyType( + { + "ids_string": f"vs_{uuid.uuid4().hex}", + "ids_int": _json_array(123), + "ids_empty_string": _json_array(""), + "oversized_id": _json_array("vs_" + "x" * 5000), + } + ) + body: Final[JsonObject] = {"model": model, "messages": messages, "vector_store_ids": ids[shape]} + return body + + +@pytest.mark.parametrize( + "shape", ("ids_string", "ids_int", "ids_empty_string", "tool_ids_string", "tool_is_string", "oversized_id") +) +def test_deny_by_default_malformed_store_ids_never_search_a_store( + strict_gateway: StrictGateway, shape: MalformedShape +) -> None: + with strict_gateway.gateway.scenario() as scenario: + model: Final = _served_model(strict_gateway.gateway, scenario) + marker: Final = f"lit6035 malformed {shape} {uuid.uuid4().hex}" + key: Final = scenario.key(models=_json_array(model)) + + response: Final = strict_gateway.gateway.request( + "POST", "/v1/chat/completions", _malformed_body(shape, model, marker), key=key + ) + searches: Final = _store_searches(strict_gateway.upstream, marker) + healthy: Final = strict_gateway.gateway.request( + "POST", "/v1/rag/query", _rag_query_body(model, f"{marker} master", f"vs_{uuid.uuid4().hex}") + ) + + assert response.status_code < 500, response.text + assert searches == () + assert healthy.status_code == 200, healthy.text + + +def test_deny_by_default_searches_a_duplicated_granted_store_once(strict_gateway: StrictGateway) -> None: + with strict_gateway.gateway.scenario() as scenario: + model: Final = _served_model(strict_gateway.gateway, scenario) + store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" + marker: Final = f"lit6035 duplicate {uuid.uuid4().hex}" + key: Final = scenario.key(models=_json_array(model), object_permission=_permission_for_stores(store_id)) + body: Final[JsonObject] = { + **_rag_query_body(model, marker, store_id), + "vector_store_ids": _json_array(store_id, store_id), + } + + response: Final = strict_gateway.gateway.request("POST", "/v1/rag/query", body, key=key) + + assert response.status_code == 200, response.text + assert len(_searches_for_marker(strict_gateway.upstream, marker, store_id)) == 1 + + +def test_deny_by_default_treats_null_key_grant_list_as_no_grant(strict_gateway: StrictGateway) -> None: + with strict_gateway.gateway.scenario() as scenario: + model: Final = _served_model(strict_gateway.gateway, scenario) + store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" + marker: Final = f"lit6035 null grants {uuid.uuid4().hex}" + null_grants: Final[JsonObject] = {"vector_stores": None} + key: Final = scenario.key(models=_json_array(model), object_permission=null_grants) + + response: Final = _rag_query(strict_gateway.gateway, model, marker, key, store_id=store_id) + + assert response.status_code == 401, response.text + assert response.json()["error"]["type"] == "key_vector_store_access_denied", response.text + assert _marker_observations(strict_gateway.upstream, marker) == () + + +def _statuses_across_workers(gateway: Gateway, model: str, key: str, store_id: str) -> frozenset[int]: + return frozenset( + _rag_query(gateway, model, f"lit6035 grant change {uuid.uuid4().hex}", key, store_id=store_id).status_code + for _ in range(6) + ) + + +@pytest.mark.parametrize("scope", ("key", "team")) +def test_deny_by_default_grant_and_revoke_take_effect_on_every_worker( + strict_gateway: StrictGateway, scope: Literal["key", "team"] +) -> None: + gateway: Final = strict_gateway.gateway + with gateway.scenario() as scenario: + model: Final = _served_model(gateway, scenario) + store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" + team: Final = scenario.team(models=_json_array(model), object_permission=_permission_for_stores(store_id)) + key: Final = ( + scenario.key(models=_json_array(model)) + if scope == "key" + else scenario.key( + team_id=team, models=_json_array(model), object_permission=_permission_for_stores(store_id) + ) + ) + team_key: Final = scope == "team" + update_path: Final = "/team/update" if team_key else "/key/update" + identity: Final[JsonObject] = {"team_id": team} if team_key else {"key": key} + if team_key: + gateway.post(update_path, {**identity, "object_permission": _permission_for_stores()}) + + def statuses() -> frozenset[int]: + return _statuses_across_workers(gateway, model, key, store_id) + + before: Final = eventually(statuses, lambda seen: seen == frozenset({401}), return_last_on_timeout=True) + gateway.post(update_path, {**identity, "object_permission": _permission_for_stores(store_id)}) + granted: Final = eventually(statuses, lambda seen: seen == frozenset({200}), return_last_on_timeout=True) + gateway.post(update_path, {**identity, "object_permission": _permission_for_stores()}) + revoked: Final = eventually(statuses, lambda seen: seen == frozenset({401}), return_last_on_timeout=True) + + assert (before, granted, revoked) == (frozenset({401}), frozenset({200}), frozenset({401})) + + +def _cli_session_token(user_id: str, team_id: str) -> str: + cli_user: Final = LiteLLM_UserTable(user_id=user_id, user_role="internal_user", teams=[team_id], models=[]) + return ExperimentalUIJWTToken.get_cli_jwt_auth_token(user_info=cli_user, team_id=team_id, team_alias="cli-team") + + +@pytest.mark.parametrize("granted_by", ("team", "user")) +def test_deny_by_default_session_token_uses_only_the_resolved_team_grant( + strict_gateway: StrictGateway, granted_by: Literal["team", "user"], monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt")) + with strict_gateway.gateway.scenario() as scenario: + model: Final = _served_model(strict_gateway.gateway, scenario) + store_id: Final = f"vs_unregistered_{uuid.uuid4().hex}" + marker: Final = f"lit6035 session token {granted_by} {uuid.uuid4().hex}" + team_grants: Final = _permission_for_stores(store_id if granted_by == "team" else "vs_some_other_store") + user: Final = scenario.user( + user_role="internal_user", + object_permission=_permission_for_stores(store_id if granted_by == "user" else "vs_some_other_store"), + ) + team: Final = scenario.team( + models=_json_array(model), + object_permission=team_grants, + members_with_roles=_json_array({"role": "user", "user_id": user}), + ) + + response: Final = _rag_query( + strict_gateway.gateway, model, marker, _cli_session_token(user, team), store_id=store_id + ) + + if granted_by == "team": + assert response.status_code == 200, response.text + assert len(_searches_for_marker(strict_gateway.upstream, marker, store_id)) == 1 + return + assert response.status_code == 401, response.text + assert response.json()["error"]["type"] == "team_vector_store_access_denied", response.text + assert _marker_observations(strict_gateway.upstream, marker) == () + + +BurstRoute: TypeAlias = Literal["chat_retrieval_config", "search_route"] + + +def _burst_request( + gateway: Gateway, route: BurstRoute, model: str, key: str, store_id: str, marker: str +) -> httpx.Response: + if route == "search_route": + return gateway.request("POST", f"/v1/vector_stores/{store_id}/search", {"query": marker}, key=key) + return _rag_query(gateway, model, marker, key, store_id=store_id, path="/v1/chat/completions") + + +def test_deny_by_default_concurrent_burst_only_searches_granted_stores(strict_gateway: StrictGateway) -> None: + gateway: Final = strict_gateway.gateway + routes: Final[tuple[BurstRoute, ...]] = ("chat_retrieval_config", "search_route") + with gateway.scenario() as scenario: + model: Final = _served_model(gateway, scenario) + store_id: Final = CONFIG_STORE_ID + granted_key: Final = scenario.key(models=_json_array(model), object_permission=_permission_for_stores(store_id)) + ungranted_key: Final = scenario.key(models=_json_array(model)) + plan: Final = tuple( + (routes[index % 2], index // 2 % 2 == 0, f"lit6035 burst {index} {uuid.uuid4().hex}") for index in range(32) + ) + upstream_observations(strict_gateway.upstream) + + with ThreadPoolExecutor(max_workers=10) as pool: + responses: Final = tuple( + pool.map( + lambda item: _burst_request( + gateway, item[0], model, granted_key if item[1] else ungranted_key, store_id, item[2] + ), + plan, + ) + ) + observed: Final = tuple( + (str(observation["path"]), str(observation["body"])) + for observation in upstream_observations(strict_gateway.upstream) + ) + + def hits(route: BurstRoute, marker: str) -> tuple[int, int]: + expected_path: Final = ( + "/chat/completions" if route == "chat_retrieval_config" else f"/vector_stores/{store_id}/search" + ) + on_route: Final = sum(path.endswith(expected_path) and marker in body for path, body in observed) + return on_route, sum(marker in body for _, body in observed) + + assert tuple(response.status_code for response in responses) == tuple( + 200 if granted else 401 for _, granted, _ in plan + ), tuple(response.text[:300] for response in responses if response.status_code not in (200, 401)) + assert tuple(hits(route, marker)[0] for route, _, marker in plan) == tuple( + 1 if granted else 0 for _, granted, _ in plan + ) + assert tuple(hits(route, marker)[1] for route, granted, marker in plan if not granted) == (0,) * 16