mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
b318231fe9
commit
5b959fb7a7
17 changed files with 370 additions and 8 deletions
|
|
@ -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[];
|
||||
|
|
@ -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([])
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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([])
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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([])
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue