mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(vector_stores): only list vector stores the caller was granted (#39612)
* fix(vector_stores): only list vector stores the caller was granted /vector_store/list returned every managed vector store with no team_id to any key, and let a dashboard session see stores created from the dashboard because every session shares the litellm-dashboard team id. Non-admin listings now show a store only when the key or one of the caller's real teams is allowlisted for it via object_permission.vector_stores, or the team owns it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(vector_stores): keep a dashboard session key's own grants when the user has no teams Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin <yassin@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
e4b8caeb36
commit
1f20b38115
3 changed files with 192 additions and 16 deletions
|
|
@ -30,7 +30,10 @@ from litellm.proxy._types import (
|
|||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user
|
||||
from litellm.proxy.vector_store_endpoints.utils import can_user_access_vector_store
|
||||
from litellm.proxy.vector_store_endpoints.utils import (
|
||||
can_user_access_vector_store,
|
||||
filter_listable_vector_stores,
|
||||
)
|
||||
from litellm.repositories.prisma_protocols import TableActions
|
||||
from litellm.repositories.table_repositories import ManagedVectorStoresRepository
|
||||
from litellm.types.vector_stores import (
|
||||
|
|
@ -390,11 +393,10 @@ async def list_vector_stores(
|
|||
|
||||
# Filter vector stores based on access control
|
||||
accessible_vector_stores: Final = []
|
||||
for vs in vector_store_map.values():
|
||||
if await _check_vector_store_access(vs, user_api_key_dict):
|
||||
redacted = LiteLLM_ManagedVectorStore(**vs)
|
||||
redacted["litellm_params"] = _redact_sensitive_litellm_params(vs.get("litellm_params"))
|
||||
accessible_vector_stores.append(redacted)
|
||||
for vs in await filter_listable_vector_stores(vector_store_map.values(), user_api_key_dict):
|
||||
redacted = LiteLLM_ManagedVectorStore(**vs)
|
||||
redacted["litellm_params"] = _redact_sensitive_litellm_params(vs.get("litellm_params"))
|
||||
accessible_vector_stores.append(redacted)
|
||||
|
||||
total_count: Final = len(accessible_vector_stores)
|
||||
total_pages: Final = (total_count + page_size - 1) // page_size
|
||||
|
|
|
|||
|
|
@ -1,11 +1,17 @@
|
|||
import json
|
||||
import re
|
||||
from collections.abc import Iterable
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final, Literal
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._experimental.mcp_server.ui_session_utils import (
|
||||
is_ui_session_credential,
|
||||
resolve_ui_session_team_ids,
|
||||
)
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
LitellmUserRoles,
|
||||
|
|
@ -160,10 +166,16 @@ async def can_user_access_vector_store(
|
|||
if _is_proxy_admin(user_api_key_dict):
|
||||
return True
|
||||
|
||||
vector_store_team_id: Final = vector_store.get("team_id")
|
||||
if vector_store_team_id is None:
|
||||
if vector_store.get("team_id") is None:
|
||||
return True
|
||||
|
||||
return await _is_vector_store_granted(vector_store, user_api_key_dict)
|
||||
|
||||
|
||||
async def _is_vector_store_granted(
|
||||
vector_store: LiteLLM_ManagedVectorStore,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> bool:
|
||||
vector_store_id: Final = vector_store.get("vector_store_id") or ""
|
||||
|
||||
key_object_permission = user_api_key_dict.object_permission
|
||||
|
|
@ -178,12 +190,70 @@ async def can_user_access_vector_store(
|
|||
if _object_permission_allows_vector_store(team_object_permission, vector_store_id):
|
||||
return True
|
||||
|
||||
if user_api_key_dict.team_id is not None and user_api_key_dict.team_id == vector_store_team_id:
|
||||
return True
|
||||
return user_api_key_dict.team_id is not None and user_api_key_dict.team_id == vector_store.get("team_id")
|
||||
|
||||
|
||||
async def _team_auth_context(team_id: str, user_api_key_dict: UserAPIKeyAuth) -> UserAPIKeyAuth:
|
||||
from litellm.proxy.auth.auth_checks import get_team_object
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
team: Final = await get_team_object(
|
||||
team_id=team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_dict.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
return user_api_key_dict.model_copy(
|
||||
update=MappingProxyType(
|
||||
{
|
||||
"team_id": team_id,
|
||||
"team_object_permission": team.object_permission,
|
||||
"team_object_permission_id": team.object_permission_id,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
async def _vector_store_listing_auth_contexts(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> tuple[UserAPIKeyAuth, ...]:
|
||||
if not is_ui_session_credential(user_api_key_dict):
|
||||
return (user_api_key_dict,)
|
||||
session_key_context: Final = user_api_key_dict.model_copy(
|
||||
update=MappingProxyType({"team_id": None, "team_object_permission": None, "team_object_permission_id": None})
|
||||
)
|
||||
team_ids: Final = await resolve_ui_session_team_ids(user_api_key_dict)
|
||||
team_contexts: Final = tuple([await _team_auth_context(team_id, user_api_key_dict) for team_id in team_ids])
|
||||
return (session_key_context, *team_contexts)
|
||||
|
||||
|
||||
async def _is_vector_store_granted_to_any(
|
||||
vector_store: LiteLLM_ManagedVectorStore,
|
||||
auth_contexts: tuple[UserAPIKeyAuth, ...],
|
||||
) -> bool:
|
||||
for auth_context in auth_contexts:
|
||||
if await _is_vector_store_granted(vector_store, auth_context):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
async def filter_listable_vector_stores(
|
||||
vector_stores: Iterable[LiteLLM_ManagedVectorStore],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> tuple[LiteLLM_ManagedVectorStore, ...]:
|
||||
"""Non-admins only see stores their key, one of their teams' object_permission, or team ownership grants."""
|
||||
if _is_proxy_admin(user_api_key_dict):
|
||||
return tuple(vector_stores)
|
||||
|
||||
auth_contexts: Final = await _vector_store_listing_auth_contexts(user_api_key_dict)
|
||||
return tuple([vs for vs in vector_stores if await _is_vector_store_granted_to_any(vs, auth_contexts)])
|
||||
|
||||
|
||||
async def get_litellm_managed_vector_store(
|
||||
vector_store_id: str,
|
||||
) -> LiteLLM_ManagedVectorStore | None:
|
||||
|
|
|
|||
|
|
@ -120,9 +120,7 @@ async def test_delete_vector_store_checks_access():
|
|||
"team_id": "team_456",
|
||||
}
|
||||
)
|
||||
mock_prisma.db.litellm_managedvectorstorestable.find_unique = AsyncMock(
|
||||
return_value=mock_vector_store
|
||||
)
|
||||
mock_prisma.db.litellm_managedvectorstorestable.find_unique = AsyncMock(return_value=mock_vector_store)
|
||||
|
||||
# User from different team should get 403
|
||||
user_api_key_dict = UserAPIKeyAuth(team_id="team_789")
|
||||
|
|
@ -134,9 +132,115 @@ async def test_delete_vector_store_checks_access():
|
|||
):
|
||||
with patch("litellm.vector_store_registry", None):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await delete_vector_store(
|
||||
data=request, user_api_key_dict=user_api_key_dict
|
||||
)
|
||||
await delete_vector_store(data=request, user_api_key_dict=user_api_key_dict)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "Access denied" in exc_info.value.detail
|
||||
|
||||
|
||||
_UNSCOPED: LiteLLM_ManagedVectorStore = {
|
||||
"vector_store_id": "vs_unscoped",
|
||||
"custom_llm_provider": "openai",
|
||||
"team_id": None,
|
||||
}
|
||||
_TEAM_A_OWNED: LiteLLM_ManagedVectorStore = {
|
||||
"vector_store_id": "vs_team_a",
|
||||
"custom_llm_provider": "openai",
|
||||
"team_id": "team_a",
|
||||
}
|
||||
_UI_CREATED: LiteLLM_ManagedVectorStore = {
|
||||
"vector_store_id": "vs_ui_created",
|
||||
"custom_llm_provider": "openai",
|
||||
"team_id": "litellm-dashboard",
|
||||
}
|
||||
|
||||
|
||||
async def _listed_ids(user_api_key_dict: UserAPIKeyAuth) -> list[str]:
|
||||
from litellm.proxy.vector_store_endpoints.management_endpoints import (
|
||||
list_vector_stores,
|
||||
)
|
||||
|
||||
with patch( # test-quality-ok: the list route reads rows through this module-level DB helper, no injection seam
|
||||
"litellm.proxy.vector_store_endpoints.management_endpoints.VectorStoreRegistry._get_vector_stores_from_db",
|
||||
new=AsyncMock(return_value=[_UNSCOPED, _TEAM_A_OWNED, _UI_CREATED]),
|
||||
):
|
||||
response = await list_vector_stores(user_api_key_dict=user_api_key_dict)
|
||||
return sorted(vs["vector_store_id"] for vs in response["data"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_vector_stores_hides_ungranted_stores_from_non_admin_keys():
|
||||
"""A store with no team_id and no allowlist entry is not listed for a key it was never granted to;
|
||||
only team ownership or an explicit object_permission grant makes a store visible."""
|
||||
assert await _listed_ids(UserAPIKeyAuth()) == []
|
||||
assert await _listed_ids(UserAPIKeyAuth(team_id="team_a")) == ["vs_team_a"]
|
||||
assert await _listed_ids(
|
||||
UserAPIKeyAuth(
|
||||
team_id="team_b",
|
||||
object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-1", vector_stores=["vs_unscoped"]),
|
||||
)
|
||||
) == ["vs_unscoped"]
|
||||
assert await _listed_ids(
|
||||
UserAPIKeyAuth(
|
||||
team_id="team_b",
|
||||
team_object_permission=LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="op-2", vector_stores=["vs_unscoped"]
|
||||
),
|
||||
)
|
||||
) == ["vs_unscoped"]
|
||||
assert await _listed_ids(UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)) == [
|
||||
"vs_team_a",
|
||||
"vs_ui_created",
|
||||
"vs_unscoped",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("user_team_ids", "session_key_grants", "expected"),
|
||||
[
|
||||
([], None, []),
|
||||
([], ["vs_unscoped"], ["vs_unscoped"]),
|
||||
(["team_a"], None, ["vs_team_a"]),
|
||||
(["team_a", "team_granted"], None, ["vs_team_a", "vs_unscoped"]),
|
||||
],
|
||||
)
|
||||
async def test_list_vector_stores_dashboard_session_resolves_real_teams(
|
||||
user_team_ids: list[str], session_key_grants: list[str] | None, expected: list[str]
|
||||
):
|
||||
"""A dashboard session lists through the user's real teams plus the session key's own grants: stores created
|
||||
from the dashboard (team_id litellm-dashboard) are not visible just because every session shares that team id,
|
||||
while stores owned by or granted to one of the user's teams, or granted to the session key itself, are."""
|
||||
from litellm.models.team import LiteLLM_TeamTableCachedObj
|
||||
|
||||
alice = UserAPIKeyAuth(
|
||||
team_id="litellm-dashboard",
|
||||
user_id="alice",
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
object_permission=(
|
||||
LiteLLM_ObjectPermissionTable(object_permission_id="op-4", vector_stores=session_key_grants)
|
||||
if session_key_grants is not None
|
||||
else None
|
||||
),
|
||||
)
|
||||
teams = {
|
||||
"team_a": LiteLLM_TeamTableCachedObj(team_id="team_a"),
|
||||
"team_granted": LiteLLM_TeamTableCachedObj(
|
||||
team_id="team_granted",
|
||||
object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-3", vector_stores=["vs_unscoped"]),
|
||||
),
|
||||
}
|
||||
|
||||
async def fake_get_team_object(team_id: str, **_kwargs: object) -> LiteLLM_TeamTableCachedObj:
|
||||
return teams[team_id]
|
||||
|
||||
with (
|
||||
patch( # test-quality-ok: team rows come from the module-level prisma client, no injection seam
|
||||
"litellm.proxy.auth.auth_checks.get_team_object", new=fake_get_team_object
|
||||
),
|
||||
patch( # test-quality-ok: the user row comes from the module-level prisma client, no injection seam
|
||||
"litellm.proxy.vector_store_endpoints.utils.resolve_ui_session_team_ids",
|
||||
new=AsyncMock(return_value=user_team_ids),
|
||||
),
|
||||
):
|
||||
assert await _listed_ids(alice) == expected
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue