From 0d5e14fcd1a6a618e7a827314ce332089b7517c6 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Thu, 4 Jun 2026 20:19:46 -0700 Subject: [PATCH] feat(proxy): auth_v2 slice 1 - virtual-key authn + casbin model RBAC behind a flag Introduces auth_v2 as a clean-slate, flag-gated auth path (general_settings auth_version: v2). When on, the existing auth is bypassed entirely and requests flow through a new authenticator chain plus a casbin authorization engine. Slice 1 scope: - Entry point fork in user_api_key_auth; v1 untouched when the flag is off - Authenticator chain with a virtual-key node resolving identity via the existing key store (reuses get_key_object; no parallel identity storage) - casbin engine: RBAC policy rows for the control plane, role bridged from the key's existing user_role; per-resource-id objects supported - Governs the model-deployment management plane only (/model/new, /model/update, /model/delete, /model/info); every other route is loud-open and logs a warning so unprotected surfaces are never silent - Policies/groupings stored in LiteLLM_CasbinRule, loaded on cold routes with a short snapshot cache; a bootstrap policy keeps proxy_admin fully authorized - Decision core (enforcer, route map, principal, authorizer, policy store) holds no framework imports, so it is unit-testable in isolation Tests cover the allow/deny matrix, deny-override, domain scoping, per-id granularity, loud-open behavior, and policy loading. Data plane (inference-time model access) and additional resources/mechanisms are deferred to later slices. casbin governs everything eventually via ABAC matchers over cached attributes; this slice lays the control-plane foundation. --- .../litellm_proxy_extras/schema.prisma | 16 ++++ litellm/proxy/auth/user_api_key_auth.py | 6 ++ litellm/proxy/auth/v2/__init__.py | 3 + litellm/proxy/auth/v2/authenticators.py | 63 ++++++++++++++ litellm/proxy/auth/v2/authorizer.py | 55 ++++++++++++ litellm/proxy/auth/v2/enforcer.py | 29 +++++++ litellm/proxy/auth/v2/entry.py | 84 +++++++++++++++++++ litellm/proxy/auth/v2/model.conf | 14 ++++ litellm/proxy/auth/v2/policy_store.py | 67 +++++++++++++++ litellm/proxy/auth/v2/principal.py | 50 +++++++++++ litellm/proxy/auth/v2/route_map.py | 30 +++++++ litellm/proxy/schema.prisma | 16 ++++ pyproject.toml | 1 + schema.prisma | 16 ++++ .../proxy/auth/v2/test_authorizer.py | 57 +++++++++++++ .../proxy/auth/v2/test_enforcer.py | 61 ++++++++++++++ .../proxy/auth/v2/test_policy_store.py | 79 +++++++++++++++++ .../proxy/auth/v2/test_principal.py | 49 +++++++++++ .../proxy/auth/v2/test_route_map.py | 28 +++++++ 19 files changed, 724 insertions(+) create mode 100644 litellm/proxy/auth/v2/__init__.py create mode 100644 litellm/proxy/auth/v2/authenticators.py create mode 100644 litellm/proxy/auth/v2/authorizer.py create mode 100644 litellm/proxy/auth/v2/enforcer.py create mode 100644 litellm/proxy/auth/v2/entry.py create mode 100644 litellm/proxy/auth/v2/model.conf create mode 100644 litellm/proxy/auth/v2/policy_store.py create mode 100644 litellm/proxy/auth/v2/principal.py create mode 100644 litellm/proxy/auth/v2/route_map.py create mode 100644 tests/test_litellm/proxy/auth/v2/test_authorizer.py create mode 100644 tests/test_litellm/proxy/auth/v2/test_enforcer.py create mode 100644 tests/test_litellm/proxy/auth/v2/test_policy_store.py create mode 100644 tests/test_litellm/proxy/auth/v2/test_principal.py create mode 100644 tests/test_litellm/proxy/auth/v2/test_route_map.py diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index c4754ef6117..d17df3122c3 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -1377,3 +1377,19 @@ model LiteLLM_WorkflowMessage { @@unique([run_id, sequence_number]) @@index([run_id]) } + +// auth_v2 (casbin) policy + grouping rules. `ptype` is "p" for policies and "g" +// for role groupings; v0..v5 are the casbin rule columns. Read on cold +// management routes only; never on the inference path. +model LiteLLM_CasbinRule { + id Int @id @default(autoincrement()) + ptype String + v0 String? + v1 String? + v2 String? + v3 String? + v4 String? + v5 String? + + @@index([ptype]) +} diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 9d4efbaeeee..c2abfe14e84 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -2190,6 +2190,12 @@ async def user_api_key_auth( """ Parent function to authenticate user api key / jwt token. """ + from litellm.proxy.proxy_server import general_settings + + if general_settings.get("auth_version") == "v2": + from litellm.proxy.auth.v2 import user_api_key_auth_v2 + + return await user_api_key_auth_v2(request=request, api_key=api_key) # Create the SERVER span and stash it on request.state BEFORE reading the # body. _read_request_body can raise ProxyException for malformed JSON; diff --git a/litellm/proxy/auth/v2/__init__.py b/litellm/proxy/auth/v2/__init__.py new file mode 100644 index 00000000000..8f7bb6ff767 --- /dev/null +++ b/litellm/proxy/auth/v2/__init__.py @@ -0,0 +1,3 @@ +from .entry import user_api_key_auth_v2 + +__all__ = ["user_api_key_auth_v2"] diff --git a/litellm/proxy/auth/v2/authenticators.py b/litellm/proxy/auth/v2/authenticators.py new file mode 100644 index 00000000000..7627d15ff43 --- /dev/null +++ b/litellm/proxy/auth/v2/authenticators.py @@ -0,0 +1,63 @@ +from typing import Any, List, Optional, Protocol, runtime_checkable + +from fastapi import HTTPException, status + + +@runtime_checkable +class Authenticator(Protocol): + def can_handle(self, api_key: Optional[str]) -> bool: + ... + + async def authenticate(self, api_key: str, ctx: "AuthContext") -> Any: + ... + + +class AuthContext: + """Carries the proxy dependencies an authenticator needs to resolve identity.""" + + def __init__( + self, + prisma_client: Any, + user_api_key_cache: Any, + proxy_logging_obj: Any, + parent_otel_span: Any = None, + ): + self.prisma_client = prisma_client + self.user_api_key_cache = user_api_key_cache + self.proxy_logging_obj = proxy_logging_obj + self.parent_otel_span = parent_otel_span + + +class VirtualKeyAuthenticator: + """Resolves a ``sk-`` virtual key to its identity via the existing key store.""" + + def can_handle(self, api_key: Optional[str]) -> bool: + return isinstance(api_key, str) and api_key.startswith("sk-") + + async def authenticate(self, api_key: str, ctx: AuthContext) -> Any: + from litellm.proxy._types import hash_token + from litellm.proxy.auth.auth_checks import get_key_object + + return await get_key_object( + hashed_token=hash_token(token=api_key), + prisma_client=ctx.prisma_client, + user_api_key_cache=ctx.user_api_key_cache, + parent_otel_span=ctx.parent_otel_span, + proxy_logging_obj=ctx.proxy_logging_obj, + ) + + +# Slice 1: virtual keys only. authlib-backed JWT / OAuth2 nodes slot in here next, +# implementing the same interface. +AUTHENTICATORS: List[Authenticator] = [VirtualKeyAuthenticator()] + + +async def authenticate(api_key: Optional[str], ctx: AuthContext) -> Any: + """Dispatch by credential shape to the first authenticator that handles it.""" + for authenticator in AUTHENTICATORS: + if authenticator.can_handle(api_key): + return await authenticator.authenticate(api_key, ctx) + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail="auth_v2: no authenticator for the supplied credential", + ) diff --git a/litellm/proxy/auth/v2/authorizer.py b/litellm/proxy/auth/v2/authorizer.py new file mode 100644 index 00000000000..b6700a47a8f --- /dev/null +++ b/litellm/proxy/auth/v2/authorizer.py @@ -0,0 +1,55 @@ +import logging +from typing import Any, Dict, Optional + +from .principal import Principal +from .route_map import GovernedRoute, match_route + +logger = logging.getLogger("litellm.proxy.auth.v2") + + +class AuthorizationDenied(Exception): + """Raised when a governed route is denied by policy. Translated to a 403 at the edge.""" + + def __init__(self, subject: str, obj: str, action: str): + self.subject = subject + self.obj = obj + self.action = action + super().__init__( + f"auth_v2: {subject} is not permitted to '{action}' on '{obj}'" + ) + + +def _build_object(rule: GovernedRoute, request_data: Optional[Dict[str, Any]]) -> str: + data = request_data or {} + for key in rule.id_fields: + value = data.get(key) + if value: + return f"{rule.resource}:{value}" + return f"{rule.resource}:*" + + +def authorize( + principal: Principal, + route: str, + request_data: Optional[Dict[str, Any]], + enforcer: Any, +) -> None: + """Enforce policy for ``route``. No-op (loud) for routes v2 doesn't yet govern. + + Raises :class:`AuthorizationDenied` when a governed route is denied. + ``enforcer`` is anything exposing ``enforce(subject, domain, obj, action)``. + """ + rule = match_route(route) + if rule is None: + logger.warning( + "auth_v2: route '%s' is not yet protected by auth_v2; allowing. " + "This must not reach production with auth_v2 enabled.", + route, + ) + return + + obj = _build_object(rule, request_data) + if not enforcer.enforce(principal.subject, principal.domain, obj, rule.action): + raise AuthorizationDenied( + subject=principal.subject, obj=obj, action=rule.action + ) diff --git a/litellm/proxy/auth/v2/enforcer.py b/litellm/proxy/auth/v2/enforcer.py new file mode 100644 index 00000000000..84a6cf2833c --- /dev/null +++ b/litellm/proxy/auth/v2/enforcer.py @@ -0,0 +1,29 @@ +import os +from typing import List, Sequence + +import casbin + +_MODEL_PATH = os.path.join(os.path.dirname(__file__), "model.conf") + +Rule = Sequence[str] + + +class CasbinEnforcer: + """Wraps a casbin Enforcer built from explicit in-memory rules. + + Control-plane access is expressed as ``p`` policy rows; identity-to-role + bridges are ``g`` grouping rows. Built per policy snapshot so swapping the + snapshot (e.g. after a reload) is a fresh, side-effect-free object. Holds no + litellm imports so the decision logic is testable in isolation. + """ + + def __init__(self, policies: List[Rule], groupings: List[Rule]): + self._enforcer = casbin.Enforcer(_MODEL_PATH) + self._enforcer.enable_auto_save(False) + for rule in policies: + self._enforcer.add_policy(*rule) + for rule in groupings: + self._enforcer.add_named_grouping_policy("g", *rule) + + def enforce(self, subject: str, domain: str, obj: str, action: str) -> bool: + return self._enforcer.enforce(subject, domain, obj, action) diff --git a/litellm/proxy/auth/v2/entry.py b/litellm/proxy/auth/v2/entry.py new file mode 100644 index 00000000000..dbc0b321b59 --- /dev/null +++ b/litellm/proxy/auth/v2/entry.py @@ -0,0 +1,84 @@ +from typing import Any, Optional + +from fastapi import HTTPException, Request, status + +from .authenticators import AuthContext, authenticate +from .authorizer import AuthorizationDenied, authorize +from .enforcer import CasbinEnforcer +from .policy_store import load_policy_snapshot +from .principal import build_principal +from .route_map import match_route + + +async def _anonymous_identity(api_key: Optional[str]) -> Any: + from litellm.proxy._types import UserAPIKeyAuth + + return UserAPIKeyAuth(api_key=api_key) + + +async def _best_effort_identity(api_key: Optional[str], ctx: AuthContext) -> Any: + """On loud-open routes, use the real identity if a usable key is present, + otherwise fall back to an anonymous principal. Never fails the request.""" + if isinstance(api_key, str) and api_key.startswith("sk-"): + try: + return await authenticate(api_key, ctx) + except Exception: + pass + return await _anonymous_identity(api_key) + + +async def user_api_key_auth_v2( + request: Request, + api_key: str = "", +) -> Any: + """auth_v2 entry point: authenticator chain establishes identity, casbin + authorizes governed routes. Routes v2 doesn't yet own are loud-open.""" + from litellm.proxy.auth.user_api_key_auth import _get_bearer_token + from litellm.proxy.auth.auth_utils import get_request_route + from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) + + token = _get_bearer_token(api_key=api_key) if api_key else api_key + route = get_request_route(request=request) + ctx = AuthContext( + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + rule = match_route(route) + if rule is None: + # Loud-open handled inside authorize(); no identity required here. + identity = await _best_effort_identity(token, ctx) + authorize(build_principal(identity), route, None, _DENY_ALL) + identity.request_route = route + return identity + + identity = await authenticate(token, ctx) + request_data = await _read_request_body(request=request) + principal = build_principal(identity) + + policies, groupings = await load_policy_snapshot(prisma_client) + enforcer = CasbinEnforcer(policies, groupings + principal.groupings) + + try: + authorize(principal, route, request_data, enforcer) + except AuthorizationDenied as e: + raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=str(e)) + + identity.request_route = route + return identity + + +class _DenyAll: + def enforce(self, *_args: Any) -> bool: + return False + + +# Only ever consulted on ungoverned routes, where authorize() short-circuits to +# loud-open before calling enforce(); present so the call signature is uniform. +_DENY_ALL = _DenyAll() diff --git a/litellm/proxy/auth/v2/model.conf b/litellm/proxy/auth/v2/model.conf new file mode 100644 index 00000000000..7271443642d --- /dev/null +++ b/litellm/proxy/auth/v2/model.conf @@ -0,0 +1,14 @@ +[request_definition] +r = sub, dom, obj, act + +[policy_definition] +p = sub, dom, obj, act, eft + +[role_definition] +g = _, _ + +[policy_effect] +e = some(where (p.eft == allow)) && !some(where (p.eft == deny)) + +[matchers] +m = g(r.sub, p.sub) && (p.dom == "*" || r.dom == p.dom) && keyMatch(r.obj, p.obj) && (p.act == "*" || r.act == p.act) diff --git a/litellm/proxy/auth/v2/policy_store.py b/litellm/proxy/auth/v2/policy_store.py new file mode 100644 index 00000000000..b63a768cbea --- /dev/null +++ b/litellm/proxy/auth/v2/policy_store.py @@ -0,0 +1,67 @@ +import time +from typing import Any, List, Optional, Tuple + +# Always-present bootstrap: principals carrying the proxy_admin role keep full +# access so enabling auth_v2 never locks admins out. Granular custom roles are +# layered on top via rows in LiteLLM_CasbinRule. +DEFAULT_POLICIES: List[List[str]] = [ + ["role:proxy_admin", "*", "*", "*", "allow"], +] + +# Casbin only runs on cold management routes, so a short snapshot cache is enough +# to absorb bursts while keeping policy edits visible within seconds across pods. +_CACHE_TTL_SECONDS = 5.0 + +_cache: Optional[Tuple[float, List[List[str]], List[List[str]]]] = None + + +def _row_values(row: Any) -> List[str]: + values = [ + getattr(row, "v0", None), + getattr(row, "v1", None), + getattr(row, "v2", None), + getattr(row, "v3", None), + getattr(row, "v4", None), + getattr(row, "v5", None), + ] + return [v for v in values if v is not None and v != ""] + + +def _split_rules(rows: List[Any]) -> Tuple[List[List[str]], List[List[str]]]: + policies: List[List[str]] = [list(p) for p in DEFAULT_POLICIES] + groupings: List[List[str]] = [] + for row in rows: + ptype = getattr(row, "ptype", None) + values = _row_values(row) + if ptype == "p": + policies.append(values) + elif ptype is not None and ptype.startswith("g"): + groupings.append(values) + return policies, groupings + + +def reset_cache() -> None: + global _cache + _cache = None + + +async def load_policy_snapshot( + prisma_client: Any, +) -> Tuple[List[List[str]], List[List[str]]]: + """Load (policies, groupings) from LiteLLM_CasbinRule, with a short TTL cache. + + Returns only the bootstrap defaults when no DB is connected, so the engine is + always constructible. + """ + global _cache + now = time.monotonic() + if _cache is not None and (now - _cache[0]) < _CACHE_TTL_SECONDS: + return _cache[1], _cache[2] + + rows: List[Any] = [] + if prisma_client is not None: + rows = await prisma_client.db.litellm_casbinrule.find_many() + + policies, groupings = _split_rules(rows) + _cache = (now, policies, groupings) + return policies, groupings diff --git a/litellm/proxy/auth/v2/principal.py b/litellm/proxy/auth/v2/principal.py new file mode 100644 index 00000000000..ca611c44db4 --- /dev/null +++ b/litellm/proxy/auth/v2/principal.py @@ -0,0 +1,50 @@ +from dataclasses import dataclass +from typing import Any, List, Optional + + +@dataclass(frozen=True) +class Principal: + """The casbin-facing view of an authenticated identity. + + ``subject`` and ``domain`` are the request coordinates; ``groupings`` are the + ``g`` rows bridging this identity to its casbin roles, derived from identity + data that already exists (the key's ``user_role``). New decision logic, + existing identity data. + """ + + subject: str + domain: str + groupings: List[List[str]] + + +def _role_to_str(role: Any) -> Optional[str]: + if role is None: + return None + return getattr(role, "value", role) + + +def build_principal(identity: Any) -> Principal: + """Derive a :class:`Principal` from an authenticated identity object. + + Duck-typed on ``user_id`` / ``team_id`` / ``token`` / ``user_role`` so it + stays decoupled from the full ``UserAPIKeyAuth`` import. + """ + user_id = getattr(identity, "user_id", None) + team_id = getattr(identity, "team_id", None) + token = getattr(identity, "token", None) + + if user_id: + subject = f"user:{user_id}" + elif token: + subject = f"key:{token}" + else: + subject = "anonymous" + + domain = f"team:{team_id}" if team_id else "*" + + groupings: List[List[str]] = [] + role = _role_to_str(getattr(identity, "user_role", None)) + if role: + groupings.append([subject, f"role:{role}"]) + + return Principal(subject=subject, domain=domain, groupings=groupings) diff --git a/litellm/proxy/auth/v2/route_map.py b/litellm/proxy/auth/v2/route_map.py new file mode 100644 index 00000000000..12c5c4ae94f --- /dev/null +++ b/litellm/proxy/auth/v2/route_map.py @@ -0,0 +1,30 @@ +from dataclasses import dataclass, field +from typing import Dict, List, Optional + + +@dataclass(frozen=True) +class GovernedRoute: + resource: str + action: str + # Request-data keys searched, in order, for the concrete resource id. The + # first present, non-empty value becomes the casbin object ``:``. + # When none are found the object falls back to ``:*``. + id_fields: List[str] = field(default_factory=list) + + +# Slice 1 governs only the model-deployment management plane. Every other route +# is intentionally left ungoverned (loud-open) until later slices wire it in. +_MODEL_ID_FIELDS = ["model_id", "id"] + +_GOVERNED: Dict[str, GovernedRoute] = { + "/model/new": GovernedRoute("model", "write"), + "/model/update": GovernedRoute("model", "write", _MODEL_ID_FIELDS), + "/model/delete": GovernedRoute("model", "delete", _MODEL_ID_FIELDS), + "/model/info": GovernedRoute("model", "read", _MODEL_ID_FIELDS), +} + + +def match_route(route: str) -> Optional[GovernedRoute]: + """Return the governance rule for ``route``, or None if v2 doesn't yet own it.""" + normalized = route.rstrip("/") or "/" + return _GOVERNED.get(normalized) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index c4754ef6117..d17df3122c3 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1377,3 +1377,19 @@ model LiteLLM_WorkflowMessage { @@unique([run_id, sequence_number]) @@index([run_id]) } + +// auth_v2 (casbin) policy + grouping rules. `ptype` is "p" for policies and "g" +// for role groupings; v0..v5 are the casbin rule columns. Read on cold +// management routes only; never on the inference path. +model LiteLLM_CasbinRule { + id Int @id @default(autoincrement()) + ptype String + v0 String? + v1 String? + v2 String? + v3 String? + v4 String? + v5 String? + + @@index([ptype]) +} diff --git a/pyproject.toml b/pyproject.toml index bc252d3b135..100ac8e7805 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -53,6 +53,7 @@ proxy = [ "orjson>=3.11.6,<4.0", "apscheduler>=3.11.2,<4.0", "fastapi-sso>=0.19.0,<1.0", + "casbin>=1.43.0,<2.0", "PyJWT>=2.12.0,<3.0", "python-multipart>=0.0.27,<1.0", "cryptography>=46.0.7,<47.0", diff --git a/schema.prisma b/schema.prisma index c4754ef6117..d17df3122c3 100644 --- a/schema.prisma +++ b/schema.prisma @@ -1377,3 +1377,19 @@ model LiteLLM_WorkflowMessage { @@unique([run_id, sequence_number]) @@index([run_id]) } + +// auth_v2 (casbin) policy + grouping rules. `ptype` is "p" for policies and "g" +// for role groupings; v0..v5 are the casbin rule columns. Read on cold +// management routes only; never on the inference path. +model LiteLLM_CasbinRule { + id Int @id @default(autoincrement()) + ptype String + v0 String? + v1 String? + v2 String? + v3 String? + v4 String? + v5 String? + + @@index([ptype]) +} diff --git a/tests/test_litellm/proxy/auth/v2/test_authorizer.py b/tests/test_litellm/proxy/auth/v2/test_authorizer.py new file mode 100644 index 00000000000..cee139d12d6 --- /dev/null +++ b/tests/test_litellm/proxy/auth/v2/test_authorizer.py @@ -0,0 +1,57 @@ +import logging + +import pytest + +from litellm.proxy.auth.v2.authorizer import AuthorizationDenied, authorize +from litellm.proxy.auth.v2.principal import Principal + +PRINCIPAL = Principal(subject="user:u1", domain="*", groupings=[]) + + +class _Enforcer: + def __init__(self, ret): + self.ret = ret + self.calls = [] + + def enforce(self, sub, dom, obj, act): + self.calls.append((sub, dom, obj, act)) + return self.ret + + +class _Exploding: + def enforce(self, *_): + raise AssertionError("enforce must not be called on a loud-open route") + + +def test_denied_governed_route_raises(): + enforcer = _Enforcer(ret=False) + with pytest.raises(AuthorizationDenied) as exc: + authorize(PRINCIPAL, "/model/delete", {"model_id": "m9"}, enforcer) + assert "delete" in str(exc.value) + assert "model:m9" in str(exc.value) + + +def test_allowed_governed_route_passes(): + enforcer = _Enforcer(ret=True) + authorize(PRINCIPAL, "/model/new", {}, enforcer) # must not raise + assert enforcer.calls == [("user:u1", "*", "model:*", "write")] + + +def test_object_is_built_from_request_id_field(): + enforcer = _Enforcer(ret=True) + authorize(PRINCIPAL, "/model/update", {"model_id": "abc123"}, enforcer) + assert enforcer.calls[0][2] == "model:abc123" + + +def test_object_falls_back_to_wildcard_without_id(): + enforcer = _Enforcer(ret=True) + authorize(PRINCIPAL, "/model/info", {}, enforcer) + assert enforcer.calls[0][2] == "model:*" + + +def test_ungoverned_route_is_loud_open_and_never_enforces(caplog): + with caplog.at_level(logging.WARNING, logger="litellm.proxy.auth.v2"): + # _Exploding asserts enforce() is not reached; no raise expected. + authorize(PRINCIPAL, "/chat/completions", {}, _Exploding()) + assert "not yet protected" in caplog.text + assert "/chat/completions" in caplog.text diff --git a/tests/test_litellm/proxy/auth/v2/test_enforcer.py b/tests/test_litellm/proxy/auth/v2/test_enforcer.py new file mode 100644 index 00000000000..ae3244862e0 --- /dev/null +++ b/tests/test_litellm/proxy/auth/v2/test_enforcer.py @@ -0,0 +1,61 @@ +from litellm.proxy.auth.v2.enforcer import CasbinEnforcer + +READER_POLICY = ["role:model_reader", "*", "model:*", "read", "allow"] +ADMIN_POLICY = ["role:proxy_admin", "*", "*", "*", "allow"] +TEAM_POLICY = ["role:team_eng", "team:eng", "model:*", "write", "allow"] + + +def _enforcer(policies, groupings): + return CasbinEnforcer(policies, groupings) + + +def test_role_grants_only_its_action(): + e = _enforcer([READER_POLICY], [["user:alice", "role:model_reader"]]) + assert e.enforce("user:alice", "*", "model:gpt4", "read") is True + # The whole point of granular RBAC: read does NOT imply write. + assert e.enforce("user:alice", "*", "model:gpt4", "write") is False + assert e.enforce("user:alice", "*", "model:gpt4", "delete") is False + + +def test_subject_without_role_is_denied(): + e = _enforcer([READER_POLICY], [["user:alice", "role:model_reader"]]) + assert e.enforce("user:bob", "*", "model:gpt4", "read") is False + + +def test_wildcard_admin_can_do_everything(): + e = _enforcer([ADMIN_POLICY], [["user:root", "role:proxy_admin"]]) + for action in ("read", "write", "delete"): + assert e.enforce("user:root", "*", "model:anything", action) is True + + +def test_specific_resource_id_is_scoped(): + policy = ["role:gpt_owner", "*", "model:gpt-4o", "write", "allow"] + e = _enforcer([policy], [["user:carol", "role:gpt_owner"]]) + # Permitted on the exact id... + assert e.enforce("user:carol", "*", "model:gpt-4o", "write") is True + # ...denied on a different id (per-resource granularity). + assert e.enforce("user:carol", "*", "model:claude", "write") is False + + +def test_domain_scoped_policy_only_applies_in_its_domain(): + e = _enforcer([TEAM_POLICY], [["user:dan", "role:team_eng"]]) + assert e.enforce("user:dan", "team:eng", "model:gpt4", "write") is True + # Same subject+role, wrong domain -> denied. + assert e.enforce("user:dan", "team:sales", "model:gpt4", "write") is False + + +def test_explicit_deny_overrides_allow(): + e = _enforcer( + [ + ["role:model_reader", "*", "model:*", "read", "allow"], + ["role:model_reader", "*", "model:secret", "read", "deny"], + ], + [["user:eve", "role:model_reader"]], + ) + assert e.enforce("user:eve", "*", "model:public", "read") is True + assert e.enforce("user:eve", "*", "model:secret", "read") is False + + +def test_empty_policy_denies_everything(): + e = _enforcer([], []) + assert e.enforce("user:anyone", "*", "model:gpt4", "read") is False diff --git a/tests/test_litellm/proxy/auth/v2/test_policy_store.py b/tests/test_litellm/proxy/auth/v2/test_policy_store.py new file mode 100644 index 00000000000..ebb934f2fb3 --- /dev/null +++ b/tests/test_litellm/proxy/auth/v2/test_policy_store.py @@ -0,0 +1,79 @@ +import pytest + +from litellm.proxy.auth.v2 import policy_store +from litellm.proxy.auth.v2.policy_store import ( + DEFAULT_POLICIES, + load_policy_snapshot, + reset_cache, +) + + +class _Row: + def __init__(self, ptype, *values): + self.ptype = ptype + for i in range(6): + setattr(self, f"v{i}", values[i] if i < len(values) else None) + + +class _DB: + def __init__(self, rows): + self.litellm_casbinrule = self + self._rows = rows + + async def find_many(self): + return self._rows + + +class _Prisma: + def __init__(self, rows): + self.db = _DB(rows) + + +@pytest.fixture(autouse=True) +def _clear_cache(): + reset_cache() + yield + reset_cache() + + +def test_bootstrap_admin_policy_always_present(): + policies, _ = policy_store._split_rules([]) + assert ["role:proxy_admin", "*", "*", "*", "allow"] in policies + assert DEFAULT_POLICIES[0] == ["role:proxy_admin", "*", "*", "*", "allow"] + + +def test_p_rows_become_policies_and_g_rows_become_groupings(): + rows = [ + _Row("p", "role:model_reader", "*", "model:*", "read", "allow"), + _Row("g", "user:alice", "role:model_reader"), + ] + policies, groupings = policy_store._split_rules(rows) + assert ["role:model_reader", "*", "model:*", "read", "allow"] in policies + assert ["user:alice", "role:model_reader"] in groupings + + +def test_empty_trailing_columns_are_trimmed(): + policies, _ = policy_store._split_rules( + [_Row("p", "role:x", "*", "model:*", "read", "allow")] + ) + assert policies[-1] == ["role:x", "*", "model:*", "read", "allow"] + + +@pytest.mark.asyncio +async def test_load_snapshot_without_db_returns_only_bootstrap(): + policies, groupings = await load_policy_snapshot(prisma_client=None) + assert policies == [list(p) for p in DEFAULT_POLICIES] + assert groupings == [] + + +@pytest.mark.asyncio +async def test_load_snapshot_reads_db_rows(): + prisma = _Prisma( + [ + _Row("p", "role:model_reader", "*", "model:*", "read", "allow"), + _Row("g", "user:alice", "role:model_reader"), + ] + ) + policies, groupings = await load_policy_snapshot(prisma_client=prisma) + assert ["role:model_reader", "*", "model:*", "read", "allow"] in policies + assert ["user:alice", "role:model_reader"] in groupings diff --git a/tests/test_litellm/proxy/auth/v2/test_principal.py b/tests/test_litellm/proxy/auth/v2/test_principal.py new file mode 100644 index 00000000000..24e7c60eb2f --- /dev/null +++ b/tests/test_litellm/proxy/auth/v2/test_principal.py @@ -0,0 +1,49 @@ +from dataclasses import dataclass +from typing import Any, Optional + +from litellm.proxy.auth.v2.principal import build_principal + + +@dataclass +class _Identity: + user_id: Optional[str] = None + team_id: Optional[str] = None + token: Optional[str] = None + user_role: Any = None + + +class _Role: + def __init__(self, value): + self.value = value + + +def test_user_id_becomes_subject_and_team_becomes_domain(): + p = build_principal(_Identity(user_id="u1", team_id="t1")) + assert p.subject == "user:u1" + assert p.domain == "team:t1" + + +def test_role_is_bridged_into_a_grouping(): + p = build_principal(_Identity(user_id="u1", user_role=_Role("proxy_admin"))) + assert p.groupings == [["user:u1", "role:proxy_admin"]] + + +def test_plain_string_role_is_supported(): + p = build_principal(_Identity(user_id="u1", user_role="internal_user")) + assert p.groupings == [["user:u1", "role:internal_user"]] + + +def test_no_role_yields_no_grouping(): + p = build_principal(_Identity(user_id="u1")) + assert p.groupings == [] + + +def test_falls_back_to_key_subject_and_global_domain(): + p = build_principal(_Identity(token="hashed", team_id=None)) + assert p.subject == "key:hashed" + assert p.domain == "*" + + +def test_anonymous_when_no_identifiers(): + p = build_principal(_Identity()) + assert p.subject == "anonymous" diff --git a/tests/test_litellm/proxy/auth/v2/test_route_map.py b/tests/test_litellm/proxy/auth/v2/test_route_map.py new file mode 100644 index 00000000000..72820c25517 --- /dev/null +++ b/tests/test_litellm/proxy/auth/v2/test_route_map.py @@ -0,0 +1,28 @@ +from litellm.proxy.auth.v2.route_map import match_route + + +def test_model_routes_map_to_resource_and_action(): + assert match_route("/model/new").resource == "model" + assert match_route("/model/new").action == "write" + assert match_route("/model/update").action == "write" + assert match_route("/model/delete").action == "delete" + assert match_route("/model/info").action == "read" + + +def test_update_and_delete_carry_id_fields(): + assert match_route("/model/update").id_fields == ["model_id", "id"] + assert match_route("/model/delete").id_fields == ["model_id", "id"] + + +def test_create_has_no_id_field(): + assert match_route("/model/new").id_fields == [] + + +def test_trailing_slash_is_normalized(): + assert match_route("/model/info/").resource == "model" + + +def test_ungoverned_routes_return_none(): + # These are loud-open in slice 1 and must not be governed yet. + for route in ("/chat/completions", "/key/generate", "/team/new", "/v1/models", "/"): + assert match_route(route) is None