diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260506231700_access_group_pass_through_vector_store_resources/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260506231700_access_group_pass_through_vector_store_resources/migration.sql new file mode 100644 index 00000000000..1dd7aeaacc3 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260506231700_access_group_pass_through_vector_store_resources/migration.sql @@ -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[]; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 84ce99557e3..5d13d24d4c8 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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([]) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 2c976479798..057d8e7becb 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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): diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index f6f99eb62c8..038d02a43a0 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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], diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py index 8bcfbb67539..39ec3a6d33c 100644 --- a/litellm/proxy/auth/route_checks.py +++ b/litellm/proxy/auth/route_checks.py @@ -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: diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 9d3c06e641f..3dfa4692b56 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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) diff --git a/litellm/proxy/management_endpoints/access_group_endpoints.py b/litellm/proxy/management_endpoints/access_group_endpoints.py index 62a770f46ae..79831777501 100644 --- a/litellm/proxy/management_endpoints/access_group_endpoints.py +++ b/litellm/proxy/management_endpoints/access_group_endpoints.py @@ -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 ): diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 259624f1e18..cf50fa4f556 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -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 diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index cc6c26fdf90..15ae862453c 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -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 diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 84ce99557e3..55d7131651d 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -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([]) diff --git a/litellm/proxy/vector_store_endpoints/utils.py b/litellm/proxy/vector_store_endpoints/utils.py index 657b520b271..e2d2a95d6f0 100644 --- a/litellm/proxy/vector_store_endpoints/utils.py +++ b/litellm/proxy/vector_store_endpoints/utils.py @@ -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 diff --git a/litellm/types/access_group.py b/litellm/types/access_group.py index e26ebe00625..dba720e266f 100644 --- a/litellm/types/access_group.py +++ b/litellm/types/access_group.py @@ -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 diff --git a/schema.prisma b/schema.prisma index 84ce99557e3..55d7131651d 100644 --- a/schema.prisma +++ b/schema.prisma @@ -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([]) diff --git a/tests/proxy_unit_tests/test_auth_checks.py b/tests/proxy_unit_tests/test_auth_checks.py index d9f4a6e56b8..2d2968363b1 100644 --- a/tests/proxy_unit_tests/test_auth_checks.py +++ b/tests/proxy_unit_tests/test_auth_checks.py @@ -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) diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index cf6feabf85f..559e73b8f49 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -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""" diff --git a/tests/test_litellm/proxy/management_endpoints/test_access_group_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_access_group_endpoints.py index 016e10859b6..f26d0d08e66 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_access_group_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_access_group_endpoints.py @@ -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"] # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_access_control.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_access_control.py index 7d72121456a..b01ce6849be 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_access_control.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_access_control.py @@ -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"""