mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
chore(vector stores): address access review followups
This commit is contained in:
parent
aef71ae2d5
commit
ce0c557012
5 changed files with 118 additions and 57 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue