diff --git a/tests/test_litellm/auth_v2/test_rbac.py b/tests/test_litellm/auth_v2/test_rbac.py index 98232fbdc5c..2a83a40ba06 100644 --- a/tests/test_litellm/auth_v2/test_rbac.py +++ b/tests/test_litellm/auth_v2/test_rbac.py @@ -1,9 +1,10 @@ from __future__ import annotations +import pytest from fastapi.security import SecurityScopes from litellm.auth_v2.models import AuthMethod, Principal, PrincipalType -from litellm.auth_v2.rbac import Role, has_any_role, has_required_scopes +from litellm.auth_v2.rbac import RbacEngine, Role, has_required_scopes def _principal(*, scopes=None, roles=None) -> Principal: @@ -16,6 +17,11 @@ def _principal(*, scopes=None, roles=None) -> Principal: ) +# --------------------------------------------------------------------------- # +# Scopes stay a plain SecurityScopes subset check (not Casbin) +# --------------------------------------------------------------------------- # + + def test_required_scopes_is_subset_check(): principal = _principal(scopes=["models:read", "chat:write", "scim:write"]) assert has_required_scopes(SecurityScopes(["models:read"]), principal) @@ -31,15 +37,92 @@ def test_empty_required_scopes_always_passes(): assert has_required_scopes(SecurityScopes([]), _principal()) -def test_has_any_role_matches_one_of_allowed(): - principal = _principal(roles=[Role.TEAM_MEMBER, Role.ORG_VIEWER]) - assert has_any_role(principal, (Role.ORG_VIEWER, Role.PLATFORM_ADMIN)) +# --------------------------------------------------------------------------- # +# RbacEngine.has_role honors the role hierarchy (Casbin g-rules) +# --------------------------------------------------------------------------- # -def test_has_any_role_rejects_when_no_overlap(): - principal = _principal(roles=[Role.TEAM_MEMBER]) - assert not has_any_role(principal, (Role.PLATFORM_ADMIN, Role.ORG_ADMIN)) +@pytest.fixture +def engine() -> RbacEngine: + return RbacEngine() -def test_has_any_role_false_when_principal_has_no_roles(): - assert not has_any_role(_principal(), (Role.PLATFORM_ADMIN,)) +@pytest.mark.parametrize( + "held,gate", + [ + (Role.PLATFORM_ADMIN, Role.ORG_ADMIN), + (Role.PLATFORM_ADMIN, Role.ORG_VIEWER), + (Role.PLATFORM_ADMIN, Role.TEAM_ADMIN), + (Role.PLATFORM_ADMIN, Role.TEAM_MEMBER), + (Role.PLATFORM_ADMIN, Role.PLATFORM_VIEWER), + (Role.ORG_ADMIN, Role.ORG_VIEWER), + (Role.ORG_ADMIN, Role.ORG_ADMIN), # exact match + (Role.TEAM_ADMIN, Role.TEAM_MEMBER), + ], +) +def test_has_role_inherits_down_the_hierarchy(engine, held, gate): + assert engine.has_role(_principal(roles=[held]), (gate,)) + + +@pytest.mark.parametrize( + "held,gate", + [ + (Role.ORG_ADMIN, Role.TEAM_MEMBER), # sideways, no inheritance edge + (Role.TEAM_MEMBER, Role.ORG_ADMIN), # lower cannot reach higher + (Role.ORG_VIEWER, Role.ORG_ADMIN), + ], +) +def test_has_role_does_not_climb_the_hierarchy(engine, held, gate): + assert not engine.has_role(_principal(roles=[held]), (gate,)) + + +def test_has_role_false_without_roles(engine): + assert not engine.has_role(_principal(), (Role.TEAM_MEMBER,)) + + +# --------------------------------------------------------------------------- # +# RbacEngine.enforce against the default policy +# --------------------------------------------------------------------------- # + + +def test_platform_admin_enforces_any_object_and_action(engine): + assert engine.enforce(_principal(roles=[Role.PLATFORM_ADMIN]), "/anything", "POST") + # keyMatch2: /scim/v2/* covers /scim/v2/Users + assert engine.enforce( + _principal(roles=[Role.PLATFORM_ADMIN]), "/scim/v2/Users", "DELETE" + ) + + +def test_platform_viewer_is_read_only(engine): + viewer = _principal(roles=[Role.PLATFORM_VIEWER]) + assert engine.enforce(viewer, "/anything", "GET") + assert not engine.enforce(viewer, "/anything", "POST") + + +def test_org_viewer_has_no_write_grant(engine): + assert not engine.enforce(_principal(roles=[Role.ORG_VIEWER]), "/widgets", "POST") + + +def test_enforce_false_without_roles(engine): + assert not engine.enforce(_principal(), "/anything", "GET") + + +# --------------------------------------------------------------------------- # +# Operator CSV policy fully replaces the in-code defaults +# --------------------------------------------------------------------------- # + + +def test_csv_policy_overrides_defaults(tmp_path): + policy = tmp_path / "policy.csv" + policy.write_text("p, platform_viewer, /reports, POST\n") + engine = RbacEngine(policy_path=str(policy)) + + # the operator rule is honored + assert engine.enforce(_principal(roles=[Role.PLATFORM_VIEWER]), "/reports", "POST") + # the built-in platform_admin "/*" grant is gone, not merged + assert not engine.enforce( + _principal(roles=[Role.PLATFORM_ADMIN]), "/reports", "POST" + ) + assert not engine.enforce( + _principal(roles=[Role.PLATFORM_ADMIN]), "/anything", "GET" + ) diff --git a/tests/test_litellm/auth_v2/test_security.py b/tests/test_litellm/auth_v2/test_security.py index f47466f2d29..e3e0e10acc0 100644 --- a/tests/test_litellm/auth_v2/test_security.py +++ b/tests/test_litellm/auth_v2/test_security.py @@ -18,11 +18,12 @@ from litellm.auth_v2.config import ( OidcProviderConfig, ) from litellm.auth_v2.models import AuthMethod, Principal, PrincipalType -from litellm.auth_v2.rbac import Role +from litellm.auth_v2.rbac import RbacEngine, Role from litellm.auth_v2.resolver import InMemoryIdentityStore, _hash_api_key from litellm.auth_v2.security import ( AuthContext, get_current_principal, + require_permission, require_roles, ) @@ -31,6 +32,8 @@ from auth_v2_helpers import TEST_AUDIENCE, TEST_ISSUER, FakeJwksClient ADMIN_KEY = "sk-admin-key" READER_KEY = "sk-reader-key" NOSCOPE_KEY = "sk-noscope-key" +PLATFORM_ADMIN_KEY = "sk-platform-admin-key" +PLATFORM_VIEWER_KEY = "sk-platform-viewer-key" def _principal(subject: str, *, scopes=None, roles=None) -> Principal: @@ -61,6 +64,12 @@ def _build_app(public_key: Any) -> Tuple[FastAPI, InMemoryIdentityStore]: "reader-principal", scopes=["models:read"] ), _hash_api_key(NOSCOPE_KEY): _principal("noscope-principal"), + _hash_api_key(PLATFORM_ADMIN_KEY): _principal( + "platform-admin-principal", roles=[Role.PLATFORM_ADMIN] + ), + _hash_api_key(PLATFORM_VIEWER_KEY): _principal( + "platform-viewer-principal", roles=[Role.PLATFORM_VIEWER] + ), } ) ctx = AuthContext(AuthConfig(), authenticators, resolver) @@ -92,6 +101,14 @@ def _build_app(public_key: Any) -> Tuple[FastAPI, InMemoryIdentityStore]: ): return {"subject": principal.subject} + @app.post("/perm-widgets") + async def widgets_route( + principal: Annotated[ + Principal, Security(require_permission("/widgets", "POST")) + ], + ): + return {"subject": principal.subject} + return app, resolver @@ -201,3 +218,64 @@ def test_required_role_present_returns_200(client): def test_required_role_missing_returns_403(client): response = client.get("/admin", headers={"x-litellm-api-key": READER_KEY}) assert response.status_code == 403 + + +def test_required_role_honors_hierarchy(client): + # platform_admin inherits org_admin via the Casbin g-rules, so it passes a + # require_roles(ORG_ADMIN) gate without holding org_admin explicitly + response = client.get("/admin", headers={"x-litellm-api-key": PLATFORM_ADMIN_KEY}) + assert response.status_code == 200 + assert response.json()["subject"] == "platform-admin-principal" + + +# --------------------------------------------------------------------------- # +# Permission enforcement (require_permission -> RbacEngine.enforce) +# --------------------------------------------------------------------------- # + + +def test_require_permission_allows_platform_admin(client): + response = client.post( + "/perm-widgets", headers={"x-litellm-api-key": PLATFORM_ADMIN_KEY} + ) + assert response.status_code == 200 + + +def test_require_permission_denies_viewer_on_write(client): + # platform_viewer is GET-only in the default policy -> POST /widgets is denied + response = client.post( + "/perm-widgets", headers={"x-litellm-api-key": PLATFORM_VIEWER_KEY} + ) + assert response.status_code == 403 + assert response.json()["detail"] == "Forbidden" + + +def test_require_permission_unauthenticated_returns_401(client): + response = client.post("/perm-widgets") + assert response.status_code == 401 + assert "WWW-Authenticate" in response.headers + + +def test_injected_rbac_engine_overrides_default_policy(rsa_keypair, tmp_path): + # operator CSV grants only platform_viewer POST /widgets and drops the + # built-in platform_admin "/*" grant; the injected engine governs enforce + policy = tmp_path / "policy.csv" + policy.write_text("p, platform_viewer, /widgets, POST\n") + + _, public_key = rsa_keypair + app, _ = _build_app(public_key) + app.state.auth_v2.rbac = RbacEngine(policy_path=str(policy)) + client = TestClient(app) + + # viewer now passes, platform_admin (default grant removed) now fails + assert ( + client.post( + "/perm-widgets", headers={"x-litellm-api-key": PLATFORM_VIEWER_KEY} + ).status_code + == 200 + ) + assert ( + client.post( + "/perm-widgets", headers={"x-litellm-api-key": PLATFORM_ADMIN_KEY} + ).status_code + == 403 + )