diff --git a/tests/e2e/coverage_registry/other.yaml b/tests/e2e/coverage_registry/other.yaml index 6b183cbf9f3..ffd272be884 100644 --- a/tests/e2e/coverage_registry/other.yaml +++ b/tests/e2e/coverage_registry/other.yaml @@ -6,6 +6,8 @@ - {id: other.auth.jwt.valid_token_allows, module: other, tier: P0, area: auth, assertions: [valid_token_allows], source: "handle_jwt.py:77-150", rationale: "Valid JWT with correct issuer + claims grants access"} - {id: other.auth.jwt.expired_denied, module: other, tier: P0, area: auth, assertions: [expired_denied], source: "handle_jwt.py:125-135", rationale: "Expired JWT rejected even with valid signature"} - {id: other.auth.jwt.invalid_signature_denied, module: other, tier: P0, area: auth, assertions: [invalid_signature_denied], source: "handle_jwt.py:145-150", rationale: "Bad/missing signature fails verification"} +- {id: other.auth.jwt.team_model_allowed, module: other, tier: P1, area: auth, assertions: [team_model_allowed], source: "handle_jwt.py:1324-1408", rationale: "A JWT whose team_id claim maps to a team can call that team's models"} +- {id: other.auth.jwt.team_model_denied, module: other, tier: P1, area: auth, assertions: [team_model_denied], source: "handle_jwt.py:1400-1405", rationale: "The same team JWT is denied a model outside the team's allow-list"} - {id: other.auth.virtual_key.route_permission_enforced, module: other, tier: P0, area: auth, assertions: [route_permission_enforced], source: "route_checks.py:89-151", rationale: "allowed_routes whitelist denies disallowed routes"} - {id: other.auth.virtual_key.route_group_allowed, module: other, tier: P1, area: auth, assertions: [route_group_allowed], source: "route_checks.py:106-128", rationale: "allowed_routes=[llm_api_routes] grants all LLM endpoints"} - {id: other.auth.passthrough.model_allowlist_enforced, module: other, tier: P1, area: auth, assertions: [model_allowlist_enforced], source: "route_checks.py:135-151", rationale: "Passthrough enforces per-key model allow-lists"} diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py index 3be339d28a0..5e50c380d69 100644 --- a/tests/e2e/e2e_config.py +++ b/tests/e2e/e2e_config.py @@ -72,6 +72,15 @@ POLL_TIMEOUT = float(os.environ.get("E2E_POLL_TIMEOUT", "120")) POLL_INTERVAL = float(os.environ.get("E2E_POLL_INTERVAL", "5")) REQUEST_TIMEOUT = float(os.environ.get("E2E_REQUEST_TIMEOUT", "60")) +# Keycloak, the real OIDC identity provider the JWT-auth suite validates tokens +# against. The proxy's litellm_jwtauth.issuers must trust the realm below (issuer + +# JWKS), so KEYCLOAK_URL/KEYCLOAK_REALM here must match the proxy config. Admin +# creds are the realm-provisioning bootstrap (default is Keycloak's dev bootstrap). +KEYCLOAK_URL = os.environ.get("KEYCLOAK_URL", "http://localhost:8080").rstrip("/") +KEYCLOAK_REALM = os.environ.get("KEYCLOAK_REALM", "litellm-e2e") +KEYCLOAK_ADMIN_USER = os.environ.get("KEYCLOAK_ADMIN_USER", "admin") +KEYCLOAK_ADMIN_PASSWORD = os.environ.get("KEYCLOAK_ADMIN_PASSWORD", "admin") + LOAD_USERS = int(os.environ.get("E2E_LOAD_USERS", "750")) LOAD_SPAWN_RATE = float(os.environ.get("E2E_LOAD_SPAWN_RATE", "50")) LOAD_DURATION_SECONDS = float(os.environ.get("E2E_LOAD_DURATION_SECONDS", "60")) diff --git a/tests/e2e/management/conftest.py b/tests/e2e/management/conftest.py index 6056d704f9a..65b372cb6c9 100644 --- a/tests/e2e/management/conftest.py +++ b/tests/e2e/management/conftest.py @@ -7,6 +7,9 @@ users, and orgs this suite creates. import pytest +from e2e_config import KEYCLOAK_ADMIN_PASSWORD, KEYCLOAK_ADMIN_USER, KEYCLOAK_REALM, KEYCLOAK_URL +from jwt_auth_client import JWTAuthClient, build_jwt_client +from keycloak import KeycloakAdmin, KeycloakEnv from management_client import ManagementClient, build_client from proxy_client import ProxyClient from scim_provisioning_client import SCIMProvisioningClient, build_scim_client @@ -33,3 +36,22 @@ def sso_client(proxy: ProxyClient) -> SSOManagementClient: @pytest.fixture(scope="session") def scim_client(proxy: ProxyClient) -> SCIMProvisioningClient: return build_scim_client(proxy) + + +@pytest.fixture(scope="session") +def jwt_client(proxy: ProxyClient) -> JWTAuthClient: + return build_jwt_client(proxy) + + +@pytest.fixture(scope="session") +def keycloak_env() -> KeycloakEnv: + """Provision (idempotently) the realm/clients the JWT-auth suite needs and + return a handle that mints real tokens. Hard-fails if Keycloak is unreachable - + a live e2e never skips for missing infrastructure.""" + admin = KeycloakAdmin( + base_url=KEYCLOAK_URL, + realm=KEYCLOAK_REALM, + admin_user=KEYCLOAK_ADMIN_USER, + admin_password=KEYCLOAK_ADMIN_PASSWORD, + ) + return admin.provision() diff --git a/tests/e2e/management/jwt_auth_client.py b/tests/e2e/management/jwt_auth_client.py new file mode 100644 index 00000000000..49edc706d0e --- /dev/null +++ b/tests/e2e/management/jwt_auth_client.py @@ -0,0 +1,42 @@ +"""Client for the JWT-auth e2e suite: send a bearer JWT (issued by the real IdP, +or a deliberately untrusted one) at the gateway and read the raw outcome. + +A JWT is just a bearer credential to the transport, so this is a thin wrapper that +routes a token at an arbitrary route (to prove admin access) or at chat (to prove +team-scoped model access). Team/model setup and verification go through the master +key via ManagementClient, so the suite injects both. +""" + +from __future__ import annotations + +from dataclasses import dataclass + +from e2e_http import NoBody, Result, StreamingResponse +from models import ChatBody, ChatMessage +from proxy_client import ProxyClient + + +@dataclass(frozen=True, slots=True) +class JWTAuthClient: + proxy: ProxyClient + + def get_route(self, path: str, token: str) -> Result[NoBody]: + """GET `path` authenticated with `token`. A 200 -> Success, a 401 -> + UnauthorizedError, a 403 -> UnknownApiError(403); the body is not modelled.""" + return self.proxy.transport.get( + path, + headers=self.proxy.transport.bearer(token), + params=NoBody(), + response_type=NoBody, + ) + + def chat(self, token: str, model: str, content: str) -> StreamingResponse: + return self.proxy.transport.send( + "/chat/completions", + headers=self.proxy.transport.bearer(token), + json=ChatBody(model=model, messages=[ChatMessage(role="user", content=content)], max_tokens=16), + ) + + +def build_jwt_client(proxy: ProxyClient) -> JWTAuthClient: + return JWTAuthClient(proxy=proxy) diff --git a/tests/e2e/management/keycloak.py b/tests/e2e/management/keycloak.py new file mode 100644 index 00000000000..b1aae69207b --- /dev/null +++ b/tests/e2e/management/keycloak.py @@ -0,0 +1,299 @@ +"""Keycloak driver for the JWT-auth e2e suite: idempotent provisioning of the +realm/clients the gateway trusts, plus minting real RS256 access tokens from them. + +The gateway validates JWTs against a real OIDC identity provider, so the suite +uses the Keycloak already in the e2e stack rather than a mock. This module talks +to Keycloak's admin + token endpoints over urllib (never requests - that is +reserved for e2e_http and enforced in CI), validating every response body through +pydantic so no untyped dict crosses the boundary. + +Provisioning is idempotent (create, ignore "already exists"): a realm with an +admin client (a service account whose access token carries the +``litellm_proxy_admin`` scope), a team client (a service account whose token +carries a hardcoded ``team_id`` claim), and a short-lived client (tiny access-token +lifespan, admin scope) used to exercise the expiry-denied path. +""" + +from __future__ import annotations + +import http.client +import json +import time +import urllib.error +import urllib.parse +import urllib.request +from dataclasses import dataclass +from typing import cast + +import jwt +from cryptography.hazmat.primitives.asymmetric import rsa +from pydantic import BaseModel, TypeAdapter + +ADMIN_SCOPE = "litellm_proxy_admin" +ADMIN_CLIENT = "litellm-admin" +TEAM_CLIENT = "litellm-team" +SHORTLIVED_CLIENT = "litellm-shortlived" +TEAM_CLAIM = "team_id" +JWT_TEAM_ID = "litellm-e2e-jwt-team" +SHORTLIVED_TOKEN_SECONDS = 10 + +_HTTP_TIMEOUT = 15 + + +class _TokenResponse(BaseModel): + access_token: str + + +class _SecretResponse(BaseModel): + value: str + + +class _ClientEntry(BaseModel): + id: str + clientId: str # noqa: N815 - Keycloak wire field + + +class _ScopeEntry(BaseModel): + id: str + name: str + + +_CLIENTS_ADAPTER: TypeAdapter[list[_ClientEntry]] = TypeAdapter(list[_ClientEntry]) +_SCOPES_ADAPTER: TypeAdapter[list[_ScopeEntry]] = TypeAdapter(list[_ScopeEntry]) + + +def mint_untrusted_jwt(issuer: str) -> str: + """An RS256 JWT signed by a key the IdP's JWKS does not contain, carrying an + otherwise-valid issuer, a future expiry, and admin-looking claims. The proxy + routes it to the issuer's JWKS by `iss`, finds no key matching the token's + `kid`, and must reject it - proving the gateway verifies the signature against + the IdP rather than trusting the claims.""" + key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + now = int(time.time()) + return jwt.encode( + { + "iss": issuer, + "sub": "untrusted-e2e", + "team_id": JWT_TEAM_ID, + "scope": ADMIN_SCOPE, + "iat": now, + "exp": now + 3600, + }, + key, + algorithm="RS256", + headers={"kid": "untrusted-e2e-kid"}, + ) + + +def _request(method: str, url: str, *, bearer: str | None = None, data: bytes | None = None, form: bool = False) -> bytes | None: + """One urllib round-trip returning the raw response body (None for an empty body + or a 409, so idempotent provisioning can ignore "already exists"). Any other + 4xx/5xx is a hard failure.""" + headers = {"Accept": "application/json"} + if bearer is not None: + headers["Authorization"] = f"Bearer {bearer}" + if data is not None: + headers["Content-Type"] = "application/x-www-form-urlencoded" if form else "application/json" + request = urllib.request.Request(url, data=data, method=method, headers=headers) # noqa: S310 - trusted local IdP + try: + opened = urllib.request.urlopen(request, timeout=_HTTP_TIMEOUT) # noqa: S310 # pyright: ignore[reportAny] # typeshed types urlopen() as Any + response = cast("http.client.HTTPResponse", opened) + with response: + return response.read() or None + except urllib.error.HTTPError as exc: + if exc.code == 409: + return None + detail = exc.read().decode(errors="replace")[:300] + raise AssertionError(f"keycloak {method} {url} failed {exc.code}: {detail}") from exc + + +def _expect(body: bytes | None, context: str) -> bytes: + if body is None: + raise AssertionError(f"keycloak returned an empty body for {context}") + return body + + +def _form(fields: dict[str, str]) -> bytes: + return urllib.parse.urlencode(fields).encode() + + +def _token(base_url: str, realm: str, fields: dict[str, str]) -> str: + body = _request( + "POST", f"{base_url}/realms/{realm}/protocol/openid-connect/token", data=_form(fields), form=True + ) + return _TokenResponse.model_validate_json(_expect(body, "token")).access_token + + +def _find_client_uuid(base_url: str, realm: str, token: str, client_id: str) -> str: + query = urllib.parse.urlencode({"clientId": client_id}) + body = _request("GET", f"{base_url}/admin/realms/{realm}/clients?{query}", bearer=token) + entries = _CLIENTS_ADAPTER.validate_json(_expect(body, f"client lookup {client_id}")) + match = next((entry for entry in entries if entry.clientId == client_id), None) + if match is None: + raise AssertionError(f"keycloak client {client_id!r} not found after provisioning") + return match.id + + +def _find_scope_id(base_url: str, realm: str, token: str, scope_name: str) -> str: + body = _request("GET", f"{base_url}/admin/realms/{realm}/client-scopes", bearer=token) + entries = _SCOPES_ADAPTER.validate_json(_expect(body, "client-scopes")) + match = next((entry for entry in entries if entry.name == scope_name), None) + if match is None: + raise AssertionError(f"keycloak client scope {scope_name!r} not found after provisioning") + return match.id + + +@dataclass(frozen=True, slots=True) +class KeycloakEnv: + """What the suite needs: where the IdP is and how to mint tokens from each + client. `issuer`/`jwks_url` are what the proxy's litellm_jwtauth.issuers must be + configured with.""" + + base_url: str + realm: str + admin_bootstrap_token: str + + @property + def issuer(self) -> str: + return f"{self.base_url}/realms/{self.realm}" + + @property + def jwks_url(self) -> str: + return f"{self.issuer}/protocol/openid-connect/certs" + + def admin_token(self) -> str: + return self._client_credentials_token(ADMIN_CLIENT) + + def team_token(self) -> str: + return self._client_credentials_token(TEAM_CLIENT) + + def shortlived_token(self) -> str: + return self._client_credentials_token(SHORTLIVED_CLIENT) + + def _client_credentials_token(self, client_id: str) -> str: + secret = self._client_secret(client_id) + return _token( + self.base_url, + self.realm, + {"client_id": client_id, "client_secret": secret, "grant_type": "client_credentials"}, + ) + + def _client_secret(self, client_id: str) -> str: + uuid = _find_client_uuid(self.base_url, self.realm, self.admin_bootstrap_token, client_id) + body = _request( + "GET", + f"{self.base_url}/admin/realms/{self.realm}/clients/{uuid}/client-secret", + bearer=self.admin_bootstrap_token, + ) + return _SecretResponse.model_validate_json(_expect(body, f"client-secret {client_id}")).value + + +@dataclass(frozen=True, slots=True) +class KeycloakAdmin: + base_url: str + realm: str + admin_user: str + admin_password: str + + def provision(self) -> KeycloakEnv: + token = _token( + self.base_url, + "master", + { + "client_id": "admin-cli", + "username": self.admin_user, + "password": self.admin_password, + "grant_type": "password", + }, + ) + self._ensure_realm(token) + self._ensure_client_scope(token, ADMIN_SCOPE) + self._ensure_service_client(token, ADMIN_CLIENT, default_scopes=(ADMIN_SCOPE,)) + self._ensure_service_client(token, TEAM_CLIENT) + self._ensure_hardcoded_claim(token, TEAM_CLIENT, TEAM_CLAIM, JWT_TEAM_ID) + self._ensure_service_client( + token, SHORTLIVED_CLIENT, default_scopes=(ADMIN_SCOPE,), access_token_lifespan=SHORTLIVED_TOKEN_SECONDS + ) + return KeycloakEnv(base_url=self.base_url, realm=self.realm, admin_bootstrap_token=token) + + def _ensure_realm(self, token: str) -> None: + _ = _request( + "POST", + f"{self.base_url}/admin/realms", + bearer=token, + data=json.dumps({"realm": self.realm, "enabled": True}).encode(), + ) + + def _ensure_client_scope(self, token: str, name: str) -> None: + _ = _request( + "POST", + f"{self.base_url}/admin/realms/{self.realm}/client-scopes", + bearer=token, + data=json.dumps( + { + "name": name, + "protocol": "openid-connect", + "attributes": {"include.in.token.scope": "true", "display.on.consent.screen": "false"}, + } + ).encode(), + ) + + def _ensure_service_client( + self, + token: str, + client_id: str, + *, + default_scopes: tuple[str, ...] = (), + access_token_lifespan: int | None = None, + ) -> None: + attributes = ( + {"access.token.lifespan": str(access_token_lifespan)} if access_token_lifespan is not None else {} + ) + _ = _request( + "POST", + f"{self.base_url}/admin/realms/{self.realm}/clients", + bearer=token, + data=json.dumps( + { + "clientId": client_id, + "enabled": True, + "protocol": "openid-connect", + "publicClient": False, + "serviceAccountsEnabled": True, + "standardFlowEnabled": False, + "directAccessGrantsEnabled": False, + "attributes": attributes, + } + ).encode(), + ) + for scope in default_scopes: + uuid = _find_client_uuid(self.base_url, self.realm, token, client_id) + scope_id = _find_scope_id(self.base_url, self.realm, token, scope) + _ = _request( + "PUT", + f"{self.base_url}/admin/realms/{self.realm}/clients/{uuid}/default-client-scopes/{scope_id}", + bearer=token, + ) + + def _ensure_hardcoded_claim(self, token: str, client_id: str, claim: str, value: str) -> None: + uuid = _find_client_uuid(self.base_url, self.realm, token, client_id) + _ = _request( + "POST", + f"{self.base_url}/admin/realms/{self.realm}/clients/{uuid}/protocol-mappers/models", + bearer=token, + data=json.dumps( + { + "name": claim, + "protocol": "openid-connect", + "protocolMapper": "oidc-hardcoded-claim-mapper", + "config": { + "claim.name": claim, + "claim.value": value, + "jsonType.label": "String", + "access.token.claim": "true", + "id.token.claim": "false", + "userinfo.token.claim": "false", + }, + } + ).encode(), + ) diff --git a/tests/e2e/management/test_jwt_auth_e2e.py b/tests/e2e/management/test_jwt_auth_e2e.py new file mode 100644 index 00000000000..297a7c7c4d3 --- /dev/null +++ b/tests/e2e/management/test_jwt_auth_e2e.py @@ -0,0 +1,163 @@ +"""Live e2e: authenticating to the gateway with a JWT from a real OIDC identity +provider (Keycloak), the way an enterprise fronts the proxy with its IdP instead +of virtual keys. + +The proxy runs with enable_jwt_auth and a litellm_jwtauth.issuers entry that +trusts the Keycloak realm this suite provisions (see keycloak.py / the module +docstring below for the required config). Tokens are minted from real Keycloak +clients: an admin client whose access token carries the litellm_proxy_admin scope, +and a team client whose token carries a hardcoded team_id claim. A token signed by +a key outside the IdP's JWKS, and an expired one, exercise the rejection paths. + +Proxy config this suite requires (general_settings), issuer/JWKS matching +KEYCLOAK_URL + KEYCLOAK_REALM: + + enable_jwt_auth: true + litellm_jwtauth: + admin_jwt_scope: "litellm_proxy_admin" + team_id_jwt_field: "team_id" + issuers: + - issuer: "http://localhost:8080/realms/litellm-e2e" + jwks_url: "http://localhost:8080/realms/litellm-e2e/protocol/openid-connect/certs" + disable_audience_validation: true + +Enterprise-gated: JWT auth needs a licensed proxy (LITELLM_LICENSE). +""" + +from __future__ import annotations + +import time +from collections.abc import Callable + +import pytest + +from e2e_config import unique_marker +from e2e_http import Success, UnauthorizedError +from jwt_auth_client import JWTAuthClient +from keycloak import JWT_TEAM_ID, KeycloakEnv, mint_untrusted_jwt +from lifecycle import ResourceManager +from management_client import ManagementClient +from models import LiteLLMParamsBody, TeamNewBody +from proxy_client import ProxyClient + +pytestmark = pytest.mark.e2e + +_TEAM_DENIAL_MARKER = "team_model_access_denied" + + +def _poll[T](proxy: ProxyClient, attempt: Callable[[], T | None], failure: str) -> T: + deadline = time.monotonic() + proxy.poll_timeout + while time.monotonic() < deadline: + found = attempt() + if found is not None: + return found + time.sleep(proxy.poll_interval) + pytest.fail(failure) + + +def _setup_jwt_team(client: ManagementClient, resources: ResourceManager, marker: str) -> tuple[str, str]: + """Register two mock deployments and bind a team (with the fixed id the Keycloak + team token claims) to only the first, so the second proves the deny path. The + team is deleted first to clear a leaked prior run, since its id is fixed.""" + team_model = f"jwt-team-model-{marker}" + other_model = f"jwt-other-model-{marker}" + team_model_id = client.proxy.create_model( + team_model, LiteLLMParamsBody(model="openai/gpt-4o-mini", mock_response="jwt ok") + ) + resources.defer(lambda: client.proxy.delete_model(team_model_id)) + other_model_id = client.proxy.create_model( + other_model, LiteLLMParamsBody(model="openai/gpt-4o-mini", mock_response="jwt ok") + ) + resources.defer(lambda: client.proxy.delete_model(other_model_id)) + + client.delete_team(JWT_TEAM_ID) + _ = client.create_team(TeamNewBody(team_alias=f"jwt-e2e-{marker}", team_id=JWT_TEAM_ID, models=[team_model])) + resources.defer(lambda: client.delete_team(JWT_TEAM_ID)) + return team_model, other_model + + +class TestJWTValidToken: + @pytest.mark.covers("other.auth.jwt.valid_token_allows") + def test_admin_token_allows_management_route( + self, jwt_client: JWTAuthClient, keycloak_env: KeycloakEnv + ) -> None: + token = keycloak_env.admin_token() + outcome = jwt_client.get_route("/user/list", token) + match outcome: + case Success(): + return + case _: + pytest.fail( + f"a valid IdP admin JWT (litellm_proxy_admin scope) must be allowed on /user/list, got {outcome}" + ) + + @pytest.mark.covers("other.auth.jwt.team_model_allowed") + def test_team_token_allows_its_model( + self, + client: ManagementClient, + jwt_client: JWTAuthClient, + keycloak_env: KeycloakEnv, + resources: ResourceManager, + ) -> None: + team_model, _ = _setup_jwt_team(client, resources, unique_marker()) + token = keycloak_env.team_token() + + _ = _poll( + client.proxy, + lambda: True if jwt_client.chat(token, team_model, f"hi {unique_marker()}").ok else None, + f"team JWT was never allowed to call its own model {team_model}", + ) + + @pytest.mark.covers("other.auth.jwt.team_model_denied") + def test_team_token_denied_untrusted_model( + self, + client: ManagementClient, + jwt_client: JWTAuthClient, + keycloak_env: KeycloakEnv, + resources: ResourceManager, + ) -> None: + _, other_model = _setup_jwt_team(client, resources, unique_marker()) + token = keycloak_env.team_token() + + outcome = jwt_client.chat(token, other_model, f"hi {unique_marker()}") + assert outcome.status_code == 403, ( + f"team JWT calling a model outside its team must be denied 403, got {outcome.status_code}: " + f"{outcome.body[:300]}" + ) + assert _TEAM_DENIAL_MARKER in outcome.body, ( + f"the 403 must be a team model-access denial, got: {outcome.body[:300]}" + ) + + +class TestJWTRejection: + @pytest.mark.covers("other.auth.jwt.invalid_signature_denied") + def test_untrusted_signature_denied(self, jwt_client: JWTAuthClient, keycloak_env: KeycloakEnv) -> None: + token = mint_untrusted_jwt(keycloak_env.issuer) + outcome = jwt_client.get_route("/user/list", token) + match outcome: + case UnauthorizedError(): + return + case _: + pytest.fail( + f"a JWT signed by a key outside the IdP's JWKS must be rejected 401, got {outcome}" + ) + + @pytest.mark.covers("other.auth.jwt.expired_denied") + def test_expired_token_denied(self, jwt_client: JWTAuthClient, keycloak_env: KeycloakEnv) -> None: + """The short-lived client's token carries the admin scope, so it is accepted + before expiry and rejected after - isolating expiry as the only thing that + changed, with a valid signature throughout.""" + token = keycloak_env.shortlived_token() + + before = jwt_client.get_route("/user/list", token) + match before: + case Success(): + pass + case _: + pytest.fail(f"the short-lived admin JWT should be accepted before expiry, got {before}") + + _ = _poll( + jwt_client.proxy, + lambda: True if isinstance(jwt_client.get_route("/user/list", token), UnauthorizedError) else None, + "the short-lived JWT was never rejected (401) after its expiry", + )