mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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.
This commit is contained in:
parent
9344f205a8
commit
0d5e14fcd1
19 changed files with 724 additions and 0 deletions
|
|
@ -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])
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
3
litellm/proxy/auth/v2/__init__.py
Normal file
3
litellm/proxy/auth/v2/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from .entry import user_api_key_auth_v2
|
||||
|
||||
__all__ = ["user_api_key_auth_v2"]
|
||||
63
litellm/proxy/auth/v2/authenticators.py
Normal file
63
litellm/proxy/auth/v2/authenticators.py
Normal file
|
|
@ -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",
|
||||
)
|
||||
55
litellm/proxy/auth/v2/authorizer.py
Normal file
55
litellm/proxy/auth/v2/authorizer.py
Normal file
|
|
@ -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
|
||||
)
|
||||
29
litellm/proxy/auth/v2/enforcer.py
Normal file
29
litellm/proxy/auth/v2/enforcer.py
Normal file
|
|
@ -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)
|
||||
84
litellm/proxy/auth/v2/entry.py
Normal file
84
litellm/proxy/auth/v2/entry.py
Normal file
|
|
@ -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()
|
||||
14
litellm/proxy/auth/v2/model.conf
Normal file
14
litellm/proxy/auth/v2/model.conf
Normal file
|
|
@ -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)
|
||||
67
litellm/proxy/auth/v2/policy_store.py
Normal file
67
litellm/proxy/auth/v2/policy_store.py
Normal file
|
|
@ -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
|
||||
50
litellm/proxy/auth/v2/principal.py
Normal file
50
litellm/proxy/auth/v2/principal.py
Normal file
|
|
@ -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)
|
||||
30
litellm/proxy/auth/v2/route_map.py
Normal file
30
litellm/proxy/auth/v2/route_map.py
Normal file
|
|
@ -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 ``<resource>:<id>``.
|
||||
# When none are found the object falls back to ``<resource>:*``.
|
||||
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)
|
||||
|
|
@ -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])
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
}
|
||||
|
|
|
|||
57
tests/test_litellm/proxy/auth/v2/test_authorizer.py
Normal file
57
tests/test_litellm/proxy/auth/v2/test_authorizer.py
Normal file
|
|
@ -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
|
||||
61
tests/test_litellm/proxy/auth/v2/test_enforcer.py
Normal file
61
tests/test_litellm/proxy/auth/v2/test_enforcer.py
Normal file
|
|
@ -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
|
||||
79
tests/test_litellm/proxy/auth/v2/test_policy_store.py
Normal file
79
tests/test_litellm/proxy/auth/v2/test_policy_store.py
Normal file
|
|
@ -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
|
||||
49
tests/test_litellm/proxy/auth/v2/test_principal.py
Normal file
49
tests/test_litellm/proxy/auth/v2/test_principal.py
Normal file
|
|
@ -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"
|
||||
28
tests/test_litellm/proxy/auth/v2/test_route_map.py
Normal file
28
tests/test_litellm/proxy/auth/v2/test_route_map.py
Normal file
|
|
@ -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
|
||||
Loading…
Add table
Reference in a new issue