Add access groups for pass-through and vector stores

Co-authored-by: ishaan-berri <ishaan-berri@users.noreply.github.com>
This commit is contained in:
oss-agent-shin 2026-05-06 23:47:14 +00:00
parent b318231fe9
commit 5b959fb7a7
No known key found for this signature in database
17 changed files with 370 additions and 8 deletions

View file

@ -0,0 +1,4 @@
-- AlterTable
ALTER TABLE "LiteLLM_AccessGroupTable"
ADD COLUMN IF NOT EXISTS "access_pass_through_routes" TEXT[] DEFAULT ARRAY[]::TEXT[],
ADD COLUMN IF NOT EXISTS "access_vector_store_ids" TEXT[] DEFAULT ARRAY[]::TEXT[];

View file

@ -1207,6 +1207,8 @@ model LiteLLM_AccessGroupTable {
access_model_names String[] @default([])
access_mcp_server_ids String[] @default([])
access_agent_ids String[] @default([])
access_pass_through_routes String[] @default([])
access_vector_store_ids String[] @default([])
assigned_team_ids String[] @default([])
assigned_key_ids String[] @default([])

View file

@ -2646,6 +2646,11 @@ class UserAPIKeyAuth(
# Team object_permission preloaded in auth (e.g. get_team_object) to avoid
# per-request object_permission fetches in downstream checks (vector stores, etc.)
team_object_permission: Optional[LiteLLM_ObjectPermissionTable] = None
# Access-group resources preloaded during auth for synchronous route checks
# and downstream resource-specific authorization.
team_access_group_ids: Optional[List[str]] = None
access_group_passthrough_routes: Optional[List[str]] = None
access_group_vector_store_ids: Optional[List[str]] = None
# Decoded upstream IdP claims (groups, roles, etc.) propagated by JWT auth machinery
# and forwarded into outbound tokens by guardrails such as MCPJWTSigner.
jwt_claims: Optional[Dict] = None
@ -3103,6 +3108,8 @@ class LiteLLM_AccessGroupTable(LiteLLMPydanticObjectBase):
access_model_names: List[str] = []
access_mcp_server_ids: List[str] = []
access_agent_ids: List[str] = []
access_pass_through_routes: List[str] = []
access_vector_store_ids: List[str] = []
assigned_team_ids: List[str] = []
assigned_key_ids: List[str] = []
created_at: Optional[datetime] = None
@ -4029,6 +4036,8 @@ class TeamInfoResponseObjectTeamTable(LiteLLM_TeamTable):
access_group_models: Optional[List[str]] = None
access_group_mcp_server_ids: Optional[List[str]] = None
access_group_agent_ids: Optional[List[str]] = None
access_group_pass_through_routes: Optional[List[str]] = None
access_group_vector_store_ids: Optional[List[str]] = None
class TeamInfoResponseObject(TypedDict):

View file

@ -2615,7 +2615,11 @@ async def get_org_object(
async def _get_resources_from_access_groups(
access_group_ids: List[str],
resource_field: Literal[
"access_model_names", "access_mcp_server_ids", "access_agent_ids"
"access_model_names",
"access_mcp_server_ids",
"access_agent_ids",
"access_pass_through_routes",
"access_vector_store_ids",
],
prisma_client: Optional[PrismaClient] = None,
user_api_key_cache: Optional[UserApiKeyCache] = None,
@ -2631,6 +2635,8 @@ async def _get_resources_from_access_groups(
- "access_model_names": model names (for model access checks)
- "access_mcp_server_ids": MCP server IDs (for MCP access checks)
- "access_agent_ids": agent IDs (for agent access checks)
- "access_pass_through_routes": pass-through route prefixes
- "access_vector_store_ids": vector store IDs
prisma_client: Optional PrismaClient (lazy-imported from proxy_server if None)
user_api_key_cache: Optional DualCache (lazy-imported from proxy_server if None)
proxy_logging_obj: Optional ProxyLogging (lazy-imported from proxy_server if None)
@ -2730,6 +2736,44 @@ async def _get_agent_ids_from_access_groups(
)
async def _get_pass_through_routes_from_access_groups(
access_group_ids: List[str],
prisma_client: Optional[PrismaClient] = None,
user_api_key_cache: Optional[UserApiKeyCache] = None,
proxy_logging_obj: Optional[ProxyLogging] = None,
) -> List[str]:
"""
Collect pass-through route prefixes from unified access groups.
Pass-through endpoints are matched by route prefix.
"""
return await _get_resources_from_access_groups(
access_group_ids=access_group_ids,
resource_field="access_pass_through_routes",
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
async def _get_vector_store_ids_from_access_groups(
access_group_ids: List[str],
prisma_client: Optional[PrismaClient] = None,
user_api_key_cache: Optional[UserApiKeyCache] = None,
proxy_logging_obj: Optional[ProxyLogging] = None,
) -> List[str]:
"""
Collect vector store IDs from unified access groups.
Vector stores are matched by vector_store_id.
"""
return await _get_resources_from_access_groups(
access_group_ids=access_group_ids,
resource_field="access_vector_store_ids",
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
def _check_model_access_helper(
model: str,
llm_router: Optional[Router],

View file

@ -583,18 +583,21 @@ class RouteChecks:
Check if route is a passthrough route.
Supports both exact match and prefix match.
"""
metadata = user_api_key_dict.metadata
metadata = user_api_key_dict.metadata or {}
team_metadata = user_api_key_dict.team_metadata or {}
if metadata is None and team_metadata is None:
return False
access_group_passthrough_routes = (
user_api_key_dict.access_group_passthrough_routes or []
)
if (
"allowed_passthrough_routes" not in metadata
and "allowed_passthrough_routes" not in team_metadata
and not access_group_passthrough_routes
):
return False
if (
metadata.get("allowed_passthrough_routes") is None
and team_metadata.get("allowed_passthrough_routes") is None
and not access_group_passthrough_routes
):
return False
@ -603,6 +606,9 @@ class RouteChecks:
or team_metadata.get("allowed_passthrough_routes")
or []
)
allowed_passthrough_routes = (
allowed_passthrough_routes + access_group_passthrough_routes
)
# Check if route matches any allowed passthrough route (exact or prefix match)
for allowed_route in allowed_passthrough_routes:

View file

@ -28,6 +28,7 @@ from litellm.proxy.auth.auth_checks import (
ExperimentalUIJWTToken,
_cache_key_object,
_delete_cache_key_object,
_get_pass_through_routes_from_access_groups,
_get_user_role,
_is_model_cost_zero,
_is_user_proxy_admin,
@ -1533,6 +1534,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
if _team_obj is not None:
valid_token.team_object_permission = _team_obj.object_permission
valid_token.team_access_group_ids = _team_obj.access_group_ids
# Keep team_metadata in sync with the freshly fetched team so that
# guardrails (or any other metadata) added after the key was cached
# are picked up on subsequent requests without a cache eviction.
@ -1913,6 +1915,9 @@ async def _run_centralized_common_checks(
None if isinstance(global_spend_result, BaseException) else global_spend_result
)
if team_object is not None:
user_api_key_auth_obj.team_access_group_ids = team_object.access_group_ids or []
# common_checks identifies admin via user_object, not the token
# (non_proxy_admin_allowed_routes_check). JWT admin shortcut and
# master_key tokens get admin from the token; the DB row for the
@ -2014,6 +2019,67 @@ async def _reserve_budget_after_common_checks(
)
async def _hydrate_access_group_route_permissions(
user_api_key_auth_obj: UserAPIKeyAuth,
route: str,
prisma_client: Optional[PrismaClient],
user_api_key_cache: UserApiKeyCache,
parent_otel_span: Span,
proxy_logging_obj: ProxyLogging,
) -> None:
"""
Preload access-group resources needed before synchronous route checks run.
Pass-through route checks happen before centralized common_checks, so team
access group resources must be resolved here for registered pass-throughs.
"""
try:
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
InitPassThroughEndpointHelpers,
)
if not InitPassThroughEndpointHelpers.is_registered_pass_through_route(
route=route
):
return
except Exception:
return
access_group_ids = list(user_api_key_auth_obj.access_group_ids or [])
if user_api_key_auth_obj.team_id is not None:
team_access_group_ids = user_api_key_auth_obj.team_access_group_ids
if team_access_group_ids is None:
try:
team_object = await get_team_object(
team_id=user_api_key_auth_obj.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
team_access_group_ids = team_object.access_group_ids or []
user_api_key_auth_obj.team_access_group_ids = team_access_group_ids
except Exception:
team_access_group_ids = []
access_group_ids.extend(team_access_group_ids or [])
if not access_group_ids:
return
access_group_routes = await _get_pass_through_routes_from_access_groups(
access_group_ids=list(set(access_group_ids)),
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
user_api_key_auth_obj.access_group_passthrough_routes = list(
set(
(user_api_key_auth_obj.access_group_passthrough_routes or [])
+ access_group_routes
)
)
def _should_skip_budget_checks(
request_data: dict,
route: str,
@ -2069,6 +2135,21 @@ async def user_api_key_auth(
)
user_api_key_auth_obj.budget_reservation = None
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
await _hydrate_access_group_route_permissions(
user_api_key_auth_obj=user_api_key_auth_obj,
route=route,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth_obj.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
## ENSURE DISABLE ROUTE WORKS ACROSS ALL USER AUTH FLOWS ##
RouteChecks.should_call_route(route=route, valid_token=user_api_key_auth_obj)

View file

@ -57,6 +57,8 @@ def _record_to_response(record) -> AccessGroupResponse:
access_model_names=record.access_model_names,
access_mcp_server_ids=record.access_mcp_server_ids,
access_agent_ids=record.access_agent_ids,
access_pass_through_routes=record.access_pass_through_routes,
access_vector_store_ids=record.access_vector_store_ids,
assigned_team_ids=record.assigned_team_ids,
assigned_key_ids=record.assigned_key_ids,
created_at=record.created_at,
@ -330,6 +332,8 @@ async def create_access_group(
"access_model_names": data.access_model_names or [],
"access_mcp_server_ids": data.access_mcp_server_ids or [],
"access_agent_ids": data.access_agent_ids or [],
"access_pass_through_routes": data.access_pass_through_routes or [],
"access_vector_store_ids": data.access_vector_store_ids or [],
"assigned_team_ids": data.assigned_team_ids or [],
"assigned_key_ids": data.assigned_key_ids or [],
"created_by": user_api_key_dict.user_id,
@ -441,6 +445,8 @@ async def update_access_group(
"access_model_names",
"access_mcp_server_ids",
"access_agent_ids",
"access_pass_through_routes",
"access_vector_store_ids",
)
and value is None
):

View file

@ -3327,15 +3327,25 @@ async def _resolve_team_access_group_resources(_team_info: Any) -> None:
if not _team_info.access_group_ids:
return
ag_lookup = await _batch_resolve_access_group_resources(_team_info.access_group_ids)
models, mcp_ids, agent_ids = set(), set(), set()
models, mcp_ids, agent_ids, pass_through_routes, vector_store_ids = (
set(),
set(),
set(),
set(),
set(),
)
for ag_id in _team_info.access_group_ids:
if ag_id in ag_lookup:
models.update(ag_lookup[ag_id]["models"])
mcp_ids.update(ag_lookup[ag_id]["mcp_server_ids"])
agent_ids.update(ag_lookup[ag_id]["agent_ids"])
pass_through_routes.update(ag_lookup[ag_id]["pass_through_routes"])
vector_store_ids.update(ag_lookup[ag_id]["vector_store_ids"])
_team_info.access_group_models = list(models)
_team_info.access_group_mcp_server_ids = list(mcp_ids)
_team_info.access_group_agent_ids = list(agent_ids)
_team_info.access_group_pass_through_routes = list(pass_through_routes)
_team_info.access_group_vector_store_ids = list(vector_store_ids)
@router.get(
@ -3882,7 +3892,8 @@ async def _batch_resolve_access_group_resources(
Batch-fetch access groups in a single DB query and return a per-group
resource map.
Returns {ag_id: {"models": [...], "mcp_server_ids": [...], "agent_ids": [...]}}.
Returns {ag_id: {"models": [...], "mcp_server_ids": [...], "agent_ids": [...],
"pass_through_routes": [...], "vector_store_ids": [...]}}.
Missing/invalid groups are silently omitted.
"""
from litellm.proxy.proxy_server import prisma_client as _prisma_client
@ -3901,6 +3912,8 @@ async def _batch_resolve_access_group_resources(
"models": list(row.access_model_names or []),
"mcp_server_ids": list(row.access_mcp_server_ids or []),
"agent_ids": list(row.access_agent_ids or []),
"pass_through_routes": list(row.access_pass_through_routes or []),
"vector_store_ids": list(row.access_vector_store_ids or []),
}
return result

View file

@ -49,6 +49,7 @@ from litellm.proxy._types import (
ProxyException,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import _get_pass_through_routes_from_access_groups
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.common_utils.http_parsing_utils import (
@ -2557,17 +2558,31 @@ async def _filter_endpoints_by_team_allowed_routes(
detail={"error": "Team not found"},
)
allowed_passthrough_routes: List[str] = []
# retrieve team metadata
team_metadata = team.metadata
if (
team_metadata is not None
and team_metadata.get("allowed_passthrough_routes") is not None
):
allowed_passthrough_routes.extend(
team_metadata.get("allowed_passthrough_routes") or []
)
if team.access_group_ids:
allowed_passthrough_routes.extend(
await _get_pass_through_routes_from_access_groups(
access_group_ids=team.access_group_ids,
)
)
if allowed_passthrough_routes:
## FILTER pass_through_endpoints by allowed_passthrough_routes
pass_through_endpoints = [
endpoint
for endpoint in pass_through_endpoints
if endpoint.path in team_metadata.get("allowed_passthrough_routes")
if endpoint.path in allowed_passthrough_routes
]
return pass_through_endpoints

View file

@ -1207,6 +1207,8 @@ model LiteLLM_AccessGroupTable {
access_model_names String[] @default([])
access_mcp_server_ids String[] @default([])
access_agent_ids String[] @default([])
access_pass_through_routes String[] @default([])
access_vector_store_ids String[] @default([])
assigned_team_ids String[] @default([])
assigned_key_ids String[] @default([])

View file

@ -10,6 +10,7 @@ from litellm.proxy._types import (
LitellmUserRoles,
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import _get_vector_store_ids_from_access_groups
from litellm.types.utils import LlmProviders
from litellm.types.vector_stores import LiteLLM_ManagedVectorStore
from litellm.utils import ProviderConfigManager
@ -50,6 +51,29 @@ def _object_permission_allows_vector_store(
return vector_store_id in allowed
async def _access_groups_allow_vector_store(
user_api_key_dict: UserAPIKeyAuth,
vector_store_id: str,
) -> bool:
"""Returns True if key/team access groups explicitly list the vector store."""
access_group_ids = list(user_api_key_dict.access_group_ids or [])
access_group_ids.extend(user_api_key_dict.team_access_group_ids or [])
if not access_group_ids:
return False
allowed_vector_store_ids = set(
user_api_key_dict.access_group_vector_store_ids or []
)
if vector_store_id not in allowed_vector_store_ids:
allowed_vector_store_ids.update(
await _get_vector_store_ids_from_access_groups(
access_group_ids=list(set(access_group_ids)),
)
)
user_api_key_dict.access_group_vector_store_ids = list(allowed_vector_store_ids)
return vector_store_id in allowed_vector_store_ids
async def _get_object_permission_for_id(
object_permission_id: Optional[str],
) -> Optional[LiteLLM_ObjectPermissionTable]:
@ -128,6 +152,9 @@ async def can_user_access_vector_store(
if _object_permission_allows_vector_store(team_object_permission, vector_store_id):
return True
if await _access_groups_allow_vector_store(user_api_key_dict, 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

View file

@ -10,6 +10,8 @@ class AccessGroupCreateRequest(BaseModel):
access_model_names: Optional[List[str]] = None
access_mcp_server_ids: Optional[List[str]] = None
access_agent_ids: Optional[List[str]] = None
access_pass_through_routes: Optional[List[str]] = None
access_vector_store_ids: Optional[List[str]] = None
assigned_team_ids: Optional[List[str]] = None
assigned_key_ids: Optional[List[str]] = None
@ -20,6 +22,8 @@ class AccessGroupUpdateRequest(BaseModel):
access_model_names: Optional[List[str]] = None
access_mcp_server_ids: Optional[List[str]] = None
access_agent_ids: Optional[List[str]] = None
access_pass_through_routes: Optional[List[str]] = None
access_vector_store_ids: Optional[List[str]] = None
assigned_team_ids: Optional[List[str]] = None
assigned_key_ids: Optional[List[str]] = None
@ -31,6 +35,8 @@ class AccessGroupResponse(BaseModel):
access_model_names: List[str]
access_mcp_server_ids: List[str]
access_agent_ids: List[str]
access_pass_through_routes: List[str]
access_vector_store_ids: List[str]
assigned_team_ids: List[str]
assigned_key_ids: List[str]
created_at: datetime

View file

@ -1207,6 +1207,8 @@ model LiteLLM_AccessGroupTable {
access_model_names String[] @default([])
access_mcp_server_ids String[] @default([])
access_agent_ids String[] @default([])
access_pass_through_routes String[] @default([])
access_vector_store_ids String[] @default([])
assigned_team_ids String[] @default([])
assigned_key_ids String[] @default([])

View file

@ -1000,6 +1000,22 @@ async def test_delete_cache_access_object():
{"access_group_id": "ag-2", "access_agent_ids": ["agent-a", "agent-b"]},
["agent-a", "agent-b"],
),
(
"access_pass_through_routes",
{
"access_group_id": "ag-4",
"access_pass_through_routes": ["/shared-passthrough"],
},
["/shared-passthrough"],
),
(
"access_vector_store_ids",
{
"access_group_id": "ag-5",
"access_vector_store_ids": ["vs-shared"],
},
["vs-shared"],
),
(
"access_model_names",
{"access_group_id": "ag-3", "access_model_names": []},
@ -1018,6 +1034,8 @@ async def test_get_resources_from_access_groups(
from litellm.proxy.auth.auth_checks import (
_get_agent_ids_from_access_groups,
_get_models_from_access_groups,
_get_pass_through_routes_from_access_groups,
_get_vector_store_ids_from_access_groups,
)
ag_table = LiteLLM_AccessGroupTable(
@ -1025,6 +1043,10 @@ async def test_get_resources_from_access_groups(
access_group_name="test",
access_model_names=access_group_data.get("access_model_names", []),
access_agent_ids=access_group_data.get("access_agent_ids", []),
access_pass_through_routes=access_group_data.get(
"access_pass_through_routes", []
),
access_vector_store_ids=access_group_data.get("access_vector_store_ids", []),
)
with patch(
@ -1038,12 +1060,24 @@ async def test_get_resources_from_access_groups(
prisma_client=MagicMock(),
user_api_key_cache=DualCache(),
)
else:
elif resource_field == "access_agent_ids":
result = await _get_agent_ids_from_access_groups(
access_group_ids=[access_group_data["access_group_id"]],
prisma_client=MagicMock(),
user_api_key_cache=DualCache(),
)
elif resource_field == "access_pass_through_routes":
result = await _get_pass_through_routes_from_access_groups(
access_group_ids=[access_group_data["access_group_id"]],
prisma_client=MagicMock(),
user_api_key_cache=DualCache(),
)
else:
result = await _get_vector_store_ids_from_access_groups(
access_group_ids=[access_group_data["access_group_id"]],
prisma_client=MagicMock(),
user_api_key_cache=DualCache(),
)
assert sorted(result) == sorted(expected)

View file

@ -834,6 +834,40 @@ def test_check_passthrough_route_access_multiple_routes():
assert result4 is False
def test_check_passthrough_route_access_allows_access_group_routes():
"""Access groups can grant pass-through routes without metadata allowlists."""
valid_token = UserAPIKeyAuth(
user_id="test_user",
access_group_passthrough_routes=["/shared-endpoint"],
)
assert (
RouteChecks.check_passthrough_route_access(
route="/shared-endpoint/v1/messages",
user_api_key_dict=valid_token,
)
is True
)
def test_check_passthrough_route_access_group_routes_prevent_false_prefix_match():
"""Access-group pass-through grants use the same safe prefix matching."""
valid_token = UserAPIKeyAuth(
user_id="test_user",
access_group_passthrough_routes=["/shared-endpoint"],
)
assert (
RouteChecks.check_passthrough_route_access(
route="/shared-endpoint-2",
user_api_key_dict=valid_token,
)
is False
)
def test_check_passthrough_route_access_prevents_false_prefix_match():
"""Test that prefix matching doesn't allow false matches like /endpoint vs /endpoint-2"""

View file

@ -31,6 +31,8 @@ def _make_access_group_record(
access_model_names: list | None = None,
access_mcp_server_ids: list | None = None,
access_agent_ids: list | None = None,
access_pass_through_routes: list | None = None,
access_vector_store_ids: list | None = None,
assigned_team_ids: list | None = None,
assigned_key_ids: list | None = None,
created_by: str | None = "admin-user",
@ -46,6 +48,8 @@ def _make_access_group_record(
"access_model_names": access_model_names or [],
"access_mcp_server_ids": access_mcp_server_ids or [],
"access_agent_ids": access_agent_ids or [],
"access_pass_through_routes": access_pass_through_routes or [],
"access_vector_store_ids": access_vector_store_ids or [],
"assigned_team_ids": assigned_team_ids or [],
"assigned_key_ids": assigned_key_ids or [],
"created_at": created_at_val,
@ -75,6 +79,8 @@ def client_and_mocks(monkeypatch):
access_model_names=data.get("access_model_names", []),
access_mcp_server_ids=data.get("access_mcp_server_ids", []),
access_agent_ids=data.get("access_agent_ids", []),
access_pass_through_routes=data.get("access_pass_through_routes", []),
access_vector_store_ids=data.get("access_vector_store_ids", []),
assigned_team_ids=data.get("assigned_team_ids", []),
assigned_key_ids=data.get("assigned_key_ids", []),
created_by=data.get("created_by"),
@ -92,6 +98,8 @@ def client_and_mocks(monkeypatch):
access_model_names=data.get("access_model_names", []),
access_mcp_server_ids=data.get("access_mcp_server_ids", []),
access_agent_ids=data.get("access_agent_ids", []),
access_pass_through_routes=data.get("access_pass_through_routes", []),
access_vector_store_ids=data.get("access_vector_store_ids", []),
assigned_team_ids=data.get("assigned_team_ids", []),
assigned_key_ids=data.get("assigned_key_ids", []),
updated_by=data.get("updated_by"),
@ -182,6 +190,8 @@ ACCESS_GROUP_PATHS = ["/v1/access_group", "/v1/unified_access_group"]
"description": "Group B description",
"access_model_names": ["model-1"],
"access_mcp_server_ids": ["mcp-1"],
"access_pass_through_routes": ["/shared-passthrough"],
"access_vector_store_ids": ["vs-shared"],
"assigned_team_ids": ["team-1"],
},
],
@ -980,12 +990,16 @@ def test_record_to_access_group_table():
access_group_name="unit-test-group",
access_model_names=["gpt-4", "claude-3"],
access_agent_ids=["agent-1"],
access_pass_through_routes=["/shared-passthrough"],
access_vector_store_ids=["vs-shared"],
)
result = _record_to_access_group_table(record)
assert result.access_group_id == "ag-unit-test"
assert result.access_group_name == "unit-test-group"
assert result.access_model_names == ["gpt-4", "claude-3"]
assert result.access_agent_ids == ["agent-1"]
assert result.access_pass_through_routes == ["/shared-passthrough"]
assert result.access_vector_store_ids == ["vs-shared"]
# ---------------------------------------------------------------------------

View file

@ -104,6 +104,69 @@ async def test_check_vector_store_access_key_object_permission_wrong_store_denie
assert await _check_vector_store_access(vector_store, user) is False
@pytest.mark.asyncio
async def test_check_vector_store_access_key_access_group_grants_access():
"""A key access group can allow a shared vector store outside the key's team."""
vector_store: LiteLLM_ManagedVectorStore = {
"vector_store_id": "vs_shared",
"custom_llm_provider": "openai",
"team_id": "team_456",
}
user = UserAPIKeyAuth(
team_id="team_789",
access_group_ids=["ag-shared-vector-store"],
)
with patch(
"litellm.proxy.vector_store_endpoints.utils._get_vector_store_ids_from_access_groups",
new_callable=AsyncMock,
return_value=["vs_shared"],
):
assert await _check_vector_store_access(vector_store, user) is True
@pytest.mark.asyncio
async def test_check_vector_store_access_team_access_group_grants_access():
"""A team access group can allow a shared vector store outside team ownership."""
vector_store: LiteLLM_ManagedVectorStore = {
"vector_store_id": "vs_team_shared",
"custom_llm_provider": "openai",
"team_id": "team_456",
}
user = UserAPIKeyAuth(
team_id="team_789",
team_access_group_ids=["ag-team-vector-store"],
)
with patch(
"litellm.proxy.vector_store_endpoints.utils._get_vector_store_ids_from_access_groups",
new_callable=AsyncMock,
return_value=["vs_team_shared"],
):
assert await _check_vector_store_access(vector_store, user) is True
@pytest.mark.asyncio
async def test_check_vector_store_access_access_group_wrong_store_denied():
"""Access groups do not widen access unless the target vector store is listed."""
vector_store: LiteLLM_ManagedVectorStore = {
"vector_store_id": "vs_target",
"custom_llm_provider": "openai",
"team_id": "team_456",
}
user = UserAPIKeyAuth(
team_id="team_789",
access_group_ids=["ag-other-vector-store"],
)
with patch(
"litellm.proxy.vector_store_endpoints.utils._get_vector_store_ids_from_access_groups",
new_callable=AsyncMock,
return_value=["vs_other"],
):
assert await _check_vector_store_access(vector_store, user) is False
@pytest.mark.asyncio
async def test_delete_vector_store_checks_access():
"""Test that delete endpoint enforces team access control"""