chore(vector stores): address access review followups

This commit is contained in:
user 2026-04-30 16:09:26 -07:00
parent aef71ae2d5
commit ce0c557012
5 changed files with 118 additions and 57 deletions

View file

@ -15,6 +15,7 @@ from fastapi.responses import ORJSONResponse
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.common_utils.http_parsing_utils import (
@ -31,10 +32,17 @@ router = APIRouter()
def _collect_vector_store_ids_from_payload(payload: Any) -> set[str]:
vector_store_ids: set[str] = set()
payload_stack = [payload]
payload_stack = [(payload, 0)]
while payload_stack:
current_payload = payload_stack.pop()
current_payload, depth = payload_stack.pop()
if depth > DEFAULT_MAX_RECURSE_DEPTH:
raise HTTPException(
status_code=400,
detail={
"error": f"Max depth of {DEFAULT_MAX_RECURSE_DEPTH} exceeded while scanning vector_store_id values"
},
)
if isinstance(current_payload, dict):
for key, value in current_payload.items():
@ -49,9 +57,9 @@ def _collect_vector_store_ids_from_payload(payload: Any) -> set[str]:
vector_store_ids.add(value)
continue
if isinstance(value, (dict, list)):
payload_stack.append(value)
payload_stack.append((value, depth + 1))
elif isinstance(current_payload, list):
payload_stack.extend(current_payload)
payload_stack.extend((item, depth + 1) for item in current_payload)
return vector_store_ids

View file

@ -20,29 +20,6 @@ router = APIRouter()
########################################################
async def _check_vector_store_access(
vector_store: LiteLLM_ManagedVectorStore,
user_api_key_dict: UserAPIKeyAuth,
) -> bool:
"""
Check if the user has access to the vector store.
Delegates to :func:`can_user_access_vector_store`, which honors:
- PROXY_ADMIN bypass
- legacy vector stores with no team_id
- key-level and team-level ``object_permission.vector_stores`` allowlists
- team_id match between key and store
"""
try:
await assert_user_can_access_vector_store(
vector_store=vector_store,
user_api_key_dict=user_api_key_dict,
)
return True
except HTTPException:
return False
async def _update_request_data_with_litellm_managed_vector_store_registry(
data: Dict,
vector_store_id: str,

View file

@ -158,11 +158,15 @@ async def get_litellm_managed_vector_store(
if vector_store is not None:
return _normalize_litellm_params(vector_store)
except Exception as e:
verbose_proxy_logger.debug(
verbose_proxy_logger.warning(
"Failed to resolve vector store id=%s from registry: %s",
vector_store_id,
e,
)
raise HTTPException(
status_code=500,
detail="Unable to validate vector store access",
) from e
try:
from litellm.proxy.auth.auth_checks import (
@ -188,12 +192,15 @@ async def get_litellm_managed_vector_store(
LiteLLM_ManagedVectorStore(**rows[0].model_dump())
)
except Exception as e:
verbose_proxy_logger.debug(
verbose_proxy_logger.warning(
"Failed to resolve vector store id=%s from shared cache: %s",
vector_store_id,
e,
)
return None
raise HTTPException(
status_code=500,
detail="Unable to validate vector store access",
) from e
async def assert_user_can_access_vector_store(

View file

@ -21,6 +21,7 @@ from litellm.proxy.vector_store_endpoints.utils import (
is_allowed_to_call_vector_store_files_endpoint,
)
from litellm.types.utils import LlmProviders
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
if TYPE_CHECKING:
from litellm.router import Router
@ -194,6 +195,8 @@ def _update_request_data_with_litellm_managed_vector_store_registry(
data: Dict,
vector_store_id: str,
llm_router: Optional["Router"] = None,
managed_vector_store: Optional[LiteLLM_ManagedVectorStore] = None,
should_lookup_registry: bool = True,
) -> Dict:
"""
Update request data with model routing information from managed vector store.
@ -263,23 +266,27 @@ def _update_request_data_with_litellm_managed_vector_store_registry(
return data
# Legacy path: Check vector store registry for non-managed vector stores
if litellm.vector_store_registry is not None:
# Legacy path: Check vector store registry for non-managed vector stores.
vector_store_to_run = managed_vector_store
if (
vector_store_to_run is None
and should_lookup_registry
and litellm.vector_store_registry is not None
):
vector_store_to_run = litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry(
vector_store_id=vector_store_id
)
if vector_store_to_run is not None:
if "custom_llm_provider" in vector_store_to_run:
data["custom_llm_provider"] = vector_store_to_run.get(
"custom_llm_provider"
)
if "litellm_credential_name" in vector_store_to_run:
data["litellm_credential_name"] = vector_store_to_run.get(
"litellm_credential_name"
)
if "litellm_params" in vector_store_to_run:
litellm_params = vector_store_to_run.get("litellm_params", {}) or {}
data.update(litellm_params)
if vector_store_to_run is not None:
if "custom_llm_provider" in vector_store_to_run:
data["custom_llm_provider"] = vector_store_to_run.get("custom_llm_provider")
if "litellm_credential_name" in vector_store_to_run:
data["litellm_credential_name"] = vector_store_to_run.get(
"litellm_credential_name"
)
if "litellm_params" in vector_store_to_run:
litellm_params = vector_store_to_run.get("litellm_params", {}) or {}
data.update(litellm_params)
return data
@ -365,7 +372,7 @@ async def vector_store_file_create(
data = await _read_request_body(request=request)
data["vector_store_id"] = vector_store_id
await assert_user_can_access_vector_store_id(
managed_vector_store = await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
@ -379,7 +386,11 @@ async def vector_store_file_create(
# Then handle managed vector store IDs
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id, llm_router=llm_router
data=data,
vector_store_id=vector_store_id,
llm_router=llm_router,
managed_vector_store=managed_vector_store,
should_lookup_registry=False,
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -464,13 +475,17 @@ async def vector_store_file_list(
data: Dict[str, Optional[str]] = {"vector_store_id": vector_store_id}
data.update(query_params)
data["vector_store_id"] = vector_store_id
await assert_user_can_access_vector_store_id(
managed_vector_store = await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id, llm_router=llm_router
data=data,
vector_store_id=vector_store_id,
llm_router=llm_router,
managed_vector_store=managed_vector_store,
should_lookup_registry=False,
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -550,7 +565,7 @@ async def vector_store_file_retrieve(
"vector_store_id": vector_store_id,
"file_id": file_id,
}
await assert_user_can_access_vector_store_id(
managed_vector_store = await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
@ -562,7 +577,11 @@ async def vector_store_file_retrieve(
# Then handle managed vector store IDs
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id, llm_router=llm_router
data=data,
vector_store_id=vector_store_id,
llm_router=llm_router,
managed_vector_store=managed_vector_store,
should_lookup_registry=False,
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -648,7 +667,7 @@ async def vector_store_file_content(
"vector_store_id": vector_store_id,
"file_id": file_id,
}
await assert_user_can_access_vector_store_id(
managed_vector_store = await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
@ -660,7 +679,11 @@ async def vector_store_file_content(
# Then handle managed vector store IDs
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id, llm_router=llm_router
data=data,
vector_store_id=vector_store_id,
llm_router=llm_router,
managed_vector_store=managed_vector_store,
should_lookup_registry=False,
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -746,7 +769,7 @@ async def vector_store_file_update(
data = await _read_request_body(request=request)
data["vector_store_id"] = vector_store_id
data["file_id"] = file_id
await assert_user_can_access_vector_store_id(
managed_vector_store = await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
@ -758,7 +781,11 @@ async def vector_store_file_update(
# Then handle managed vector store IDs
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id, llm_router=llm_router
data=data,
vector_store_id=vector_store_id,
llm_router=llm_router,
managed_vector_store=managed_vector_store,
should_lookup_registry=False,
)
provider_enum = await _resolve_provider(data=data, request=request)
@ -844,7 +871,7 @@ async def vector_store_file_delete(
"vector_store_id": vector_store_id,
"file_id": file_id,
}
await assert_user_can_access_vector_store_id(
managed_vector_store = await assert_user_can_access_vector_store_id(
vector_store_id=vector_store_id,
user_api_key_dict=user_api_key_dict,
)
@ -856,7 +883,11 @@ async def vector_store_file_delete(
# Then handle managed vector store IDs
data = _update_request_data_with_litellm_managed_vector_store_registry(
data=data, vector_store_id=vector_store_id, llm_router=llm_router
data=data,
vector_store_id=vector_store_id,
llm_router=llm_router,
managed_vector_store=managed_vector_store,
should_lookup_registry=False,
)
provider_enum = await _resolve_provider(data=data, request=request)

View file

@ -99,7 +99,8 @@ async def test_vector_store_file_create_forces_path_id_over_body_id():
assert response == {"ok": True}
assert captured_data["vector_store_id"] == "vs_path_allowed"
mock_registry.get_litellm_managed_vector_store_from_registry.assert_any_call(
assert captured_data["custom_llm_provider"] == "openai"
mock_registry.get_litellm_managed_vector_store_from_registry.assert_called_once_with(
vector_store_id="vs_path_allowed"
)
@ -227,6 +228,25 @@ async def test_rag_ingest_denies_nested_other_team_vector_store():
mock_aingest.assert_not_called()
def test_rag_payload_scan_rejects_excessive_nesting():
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
from litellm.proxy.rag_endpoints.endpoints import (
_collect_vector_store_ids_from_payload,
)
payload = {}
current = payload
for _ in range(DEFAULT_MAX_RECURSE_DEPTH + 1):
current["nested"] = {}
current = current["nested"]
current["vector_store_id"] = "vs_too_deep"
with pytest.raises(HTTPException) as exc_info:
_collect_vector_store_ids_from_payload(payload)
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_responses_file_search_denies_other_team_vector_store():
from litellm.proxy.common_request_processing import (
@ -333,6 +353,24 @@ async def test_get_managed_vector_store_uses_shared_cache_helper_for_db_fallback
cache_helper.assert_awaited_once()
@pytest.mark.asyncio
async def test_get_managed_vector_store_fails_closed_on_lookup_error():
from litellm.proxy.vector_store_endpoints.utils import (
get_litellm_managed_vector_store,
)
mock_registry = MagicMock()
mock_registry.get_litellm_managed_vector_store_from_registry.side_effect = (
RuntimeError("registry unavailable")
)
with patch.object(litellm, "vector_store_registry", mock_registry):
with pytest.raises(HTTPException) as exc_info:
await get_litellm_managed_vector_store(vector_store_id="vs_registry_only")
assert exc_info.value.status_code == 500
@pytest.mark.asyncio
async def test_vertex_discovery_allows_unregistered_provider_native_datastore_id():
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (