diff --git a/.github/demo-assets/TEST_RESULTS.md b/.github/demo-assets/TEST_RESULTS.md deleted file mode 100644 index 9e34862f9e4..00000000000 --- a/.github/demo-assets/TEST_RESULTS.md +++ /dev/null @@ -1,72 +0,0 @@ -## JWT Auth Fixes — Integration Test Results - -All three fixes verified end-to-end against a live LiteLLM proxy with mock JWKS/OIDC server, fake LLM backend, and RSA-signed JWTs. - -### Setup - -| Component | Port | Description | -|---|---|---| -| Mock JWKS + OIDC | :19900 | Serves `/.well-known/openid-configuration` → `jwks_uri` | -| Fake LLM | :19901 | Returns canned chat completions | -| LiteLLM Proxy | :19902 | JWT auth enabled, `premium_user` patched, `team_id_upsert: true` | - ---- - -### Fix 1: OIDC Discovery URL Resolution - -`JWT_PUBLIC_KEY_URL` set to `.well-known/openid-configuration` (not a direct JWKS URL). - -| Test | Token | Expected | Got | -|---|---|---|---| -| 1a | Valid JWT (RSA-signed, correct key) | **200** | ✅ **200** — chat completion returned | -| 1b | Tampered JWT (last 5 chars replaced) | **401** | ✅ **401** — "Signature verification failed" | - -The proxy correctly fetches the OIDC discovery doc, resolves `jwks_uri`, and validates the JWT signature. - ---- - -### Fix 2: Roles as Array - -Config: `team_id_jwt_field: "roles"` — JWT `roles` claim is a JSON array. - -| Test | Token Claims | Expected | Got | -|---|---|---|---| -| 2a | `"roles": ["team-beta", "team-gamma"]` | **200** (first element used) | ✅ **200** — `team-beta` extracted, completion returned | -| 2b | `"roles": []` | **401** | ✅ **401** — "No team found in token" | - -![Fix 1+2 results](https://raw.githubusercontent.com/BerriAI/litellm/68c86d22396ea4dce14e646dc9298c7b8e5f18ee/.github/demo-assets/fix1_fix2_results.webp) - ---- - -### Fix 3: Dot-notation Error Hint - -Config: `team_id_jwt_field: "roles.0"` — common Azure AD misconfiguration. - -| Test | Token Claims | Expected | Got | -|---|---|---|---| -| 3 | `"roles": ["team-alpha"]` | **401** with hint | ✅ **401** — hint says: *"Use 'roles' instead — LiteLLM automatically uses the first element when the field value is a list."* | - -![Fix 3 results](https://raw.githubusercontent.com/BerriAI/litellm/68c86d22396ea4dce14e646dc9298c7b8e5f18ee/.github/demo-assets/fix3_results.webp) - ---- - -### Unit Tests - -13/13 new tests pass: - -``` -tests/test_litellm/proxy/auth/test_handle_jwt.py::test_resolve_jwks_url_passthrough_for_direct_jwks_url PASSED -tests/test_litellm/proxy/auth/test_handle_jwt.py::test_resolve_jwks_url_resolves_oidc_discovery_document PASSED -tests/test_litellm/proxy/auth/test_handle_jwt.py::test_resolve_jwks_url_caches_resolved_jwks_uri PASSED -tests/test_litellm/proxy/auth/test_handle_jwt.py::test_resolve_jwks_url_raises_if_no_jwks_uri_in_discovery_doc PASSED -tests/test_litellm/proxy/auth/test_handle_jwt.py::test_get_team_id_returns_first_element_when_roles_is_list PASSED -tests/test_litellm/proxy/auth/test_handle_jwt.py::test_get_team_id_returns_first_element_from_multi_value_roles_list PASSED -tests/test_litellm/proxy/auth/test_handle_jwt.py::test_get_team_id_returns_default_when_roles_list_is_empty PASSED -tests/test_litellm/proxy/auth/test_handle_jwt.py::test_get_team_id_still_works_with_string_value PASSED -tests/test_litellm/proxy/auth/test_handle_jwt.py::test_get_team_id_list_result_is_hashable PASSED -tests/test_litellm/proxy/auth/test_handle_jwt.py::test_find_and_validate_specific_team_id_hints_bracket_notation PASSED -tests/test_litellm/proxy/auth/test_handle_jwt.py::test_find_and_validate_specific_team_id_hints_bracket_index_notation PASSED -tests/test_litellm/proxy/auth/test_handle_jwt.py::test_find_and_validate_specific_team_id_no_hint_for_valid_field PASSED -``` - -All 13 passed ✅ diff --git a/.github/demo-assets/fix1_fix2_results.webp b/.github/demo-assets/fix1_fix2_results.webp deleted file mode 100644 index 389d8d5873d..00000000000 Binary files a/.github/demo-assets/fix1_fix2_results.webp and /dev/null differ diff --git a/.github/demo-assets/fix3_results.webp b/.github/demo-assets/fix3_results.webp deleted file mode 100644 index b41785bb894..00000000000 Binary files a/.github/demo-assets/fix3_results.webp and /dev/null differ diff --git a/demo_jwt_fixes.py b/demo_jwt_fixes.py deleted file mode 100644 index a2edaa8663a..00000000000 --- a/demo_jwt_fixes.py +++ /dev/null @@ -1,329 +0,0 @@ -#!/usr/bin/env python3 -""" -Integration demo for JWT auth fixes on branch fix/jwt-auth-oidc-array-roles. - -Spins up: - - Mock JWKS + OIDC discovery server on port 19900 - - Fake LLM backend on port 19901 - - LiteLLM proxy on port 19902 (premium_user patched to bypass license) - -Tests three scenarios: - Fix 1 OIDC discovery URL resolution (valid JWT → 200, tampered → 401) - Fix 2 roles claim as JSON array (["team-beta","team-gamma"] → 200, [] → 401) - Fix 3 Helpful error for dot-notation (roles.0 → 401 with hint) -""" - -import asyncio, base64, json, os, signal, subprocess, sys, textwrap, time, threading -from http.server import HTTPServer, BaseHTTPRequestHandler - -from cryptography.hazmat.primitives.asymmetric import rsa -from cryptography.hazmat.primitives import serialization -import jwt as pyjwt - -# ── 1. RSA key pair ────────────────────────────────────────────────────────── - -RSA_KEY = rsa.generate_private_key(public_exponent=65537, key_size=2048) -RSA_PUB = RSA_KEY.public_key() -PUB_NUMBERS = RSA_PUB.public_numbers() - - -def _b64url(n: int, length: int) -> str: - return base64.urlsafe_b64encode(n.to_bytes(length, "big")).rstrip(b"=").decode() - - -KID = "demo-kid-1" -JWKS = { - "keys": [{ - "kty": "RSA", "kid": KID, "use": "sig", "alg": "RS256", - "n": _b64url(PUB_NUMBERS.n, 256), - "e": _b64url(PUB_NUMBERS.e, 3), - }] -} - -PEM_PRIVATE = RSA_KEY.private_bytes( - serialization.Encoding.PEM, - serialization.PrivateFormat.PKCS8, - serialization.NoEncryption(), -) - -ISSUER = "https://demo-idp.example.com" -AUDIENCE = "litellm-proxy-demo" - - -def sign_token(claims: dict) -> str: - payload = {"iss": ISSUER, "aud": AUDIENCE, "exp": int(time.time()) + 3600, **claims} - return pyjwt.encode(payload, PEM_PRIVATE, algorithm="RS256", headers={"kid": KID}) - - -# ── 2. Mock JWKS + OIDC discovery server (port 19900) ──────────────────────── - -class JWKSHandler(BaseHTTPRequestHandler): - def do_GET(self): - if self.path == "/.well-known/openid-configuration": - body = json.dumps({ - "issuer": ISSUER, - "jwks_uri": "http://127.0.0.1:19900/jwks", - "authorization_endpoint": "https://example.com/authorize", - "token_endpoint": "https://example.com/token", - }).encode() - elif self.path == "/jwks": - body = json.dumps(JWKS).encode() - else: - self.send_response(404); self.end_headers(); return - self.send_response(200) - self.send_header("Content-Type", "application/json") - self.end_headers() - self.wfile.write(body) - - def log_message(self, *_): - pass - - -# ── 3. Fake LLM backend (port 19901) ───────────────────────────────────────── - -class FakeLLMHandler(BaseHTTPRequestHandler): - def do_POST(self): - body = json.dumps({ - "id": "chatcmpl-fake", "object": "chat.completion", "created": 1700000000, - "model": "fake-model", - "choices": [{"index": 0, "message": {"role": "assistant", "content": "Hello from fake LLM!"}, "finish_reason": "stop"}], - "usage": {"prompt_tokens": 5, "completion_tokens": 7, "total_tokens": 12}, - }).encode() - self.send_response(200) - self.send_header("Content-Type", "application/json") - self.end_headers() - self.wfile.write(body) - - def do_GET(self): - if "/models" in self.path: - body = json.dumps({"data": [{"id": "fake-model", "object": "model"}], "object": "list"}).encode() - self.send_response(200) - self.send_header("Content-Type", "application/json") - self.end_headers() - self.wfile.write(body) - else: - self.send_response(404); self.end_headers() - - def log_message(self, *_): - pass - - -# ── 4. Proxy startup that patches premium_user ─────────────────────────────── - -PROXY_PORT = 19902 - -CONFIG_FIX12 = textwrap.dedent("""\ - model_list: - - model_name: fake-model - litellm_params: - model: openai/fake-model - api_key: fake-key - api_base: http://127.0.0.1:19901/ - general_settings: - enable_jwt_auth: true - litellm_jwtauth: - team_id_jwt_field: "roles" - team_id_default: "" - litellm_settings: - drop_params: true -""") - -CONFIG_FIX3 = textwrap.dedent("""\ - model_list: - - model_name: fake-model - litellm_params: - model: openai/fake-model - api_key: fake-key - api_base: http://127.0.0.1:19901/ - general_settings: - enable_jwt_auth: true - litellm_jwtauth: - team_id_jwt_field: "roles.0" - team_id_default: "" - litellm_settings: - drop_params: true -""") - - -def _write_config(content: str, path: str): - with open(path, "w") as f: - f.write(content) - - -PROXY_WRAPPER = textwrap.dedent("""\ - import sys, os, unittest.mock - os.environ["JWT_PUBLIC_KEY_URL"] = "http://127.0.0.1:19900/.well-known/openid-configuration" - os.environ["JWT_AUDIENCE"] = "{audience}" - os.environ["DISABLE_SCHEMA_UPDATE"] = "true" - # Patch premium_user BEFORE proxy imports the variable - import litellm.proxy.proxy_server as _ps - _ps.premium_user = True - from litellm.proxy.proxy_server import app, initialize - import asyncio, uvicorn - asyncio.get_event_loop().run_until_complete(initialize(config="{config_path}")) - uvicorn.run(app, host="127.0.0.1", port={port}, log_level="warning") -""") - - -def start_proxy(config_path: str) -> subprocess.Popen: - wrapper_code = PROXY_WRAPPER.format( - audience=AUDIENCE, - config_path=config_path, - port=PROXY_PORT, - ) - wrapper_path = "/tmp/_proxy_wrapper.py" - with open(wrapper_path, "w") as f: - f.write(wrapper_code) - proc = subprocess.Popen( - [sys.executable, wrapper_path], - stdout=subprocess.PIPE, stderr=subprocess.STDOUT, - ) - return proc - - -def wait_for_proxy(timeout: int = 90) -> bool: - import urllib.request - deadline = time.time() + timeout - while time.time() < deadline: - try: - r = urllib.request.urlopen(f"http://127.0.0.1:{PROXY_PORT}/health/readiness", timeout=2) - if r.status == 200: - return True - except Exception: - pass - time.sleep(1) - return False - - -def stop_proxy(proc: subprocess.Popen): - proc.send_signal(signal.SIGTERM) - try: - proc.wait(timeout=10) - except subprocess.TimeoutExpired: - proc.kill() - - -def proxy_request(token: str, model: str = "fake-model") -> tuple: - import urllib.request, urllib.error - data = json.dumps({"model": model, "messages": [{"role": "user", "content": "Hi"}]}).encode() - req = urllib.request.Request( - f"http://127.0.0.1:{PROXY_PORT}/v1/chat/completions", - data=data, - headers={"Content-Type": "application/json", "Authorization": f"Bearer {token}"}, - ) - try: - r = urllib.request.urlopen(req, timeout=15) - return r.status, r.read().decode() - except urllib.error.HTTPError as e: - return e.code, e.read().decode() - - -# ── 5. Pretty output ───────────────────────────────────────────────────────── - -GREEN = "\033[92m" -RED = "\033[91m" -BOLD = "\033[1m" -RESET = "\033[0m" -results = [] - - -def header(text: str): - print(f"\n{BOLD}{'='*72}{RESET}") - print(f"{BOLD} {text}{RESET}") - print(f"{BOLD}{'='*72}{RESET}") - - -def check(label: str, condition: bool, detail: str = ""): - mark = f"{GREEN}PASS{RESET}" if condition else f"{RED}FAIL{RESET}" - print(f" [{mark}] {label}") - if detail: - for line in detail.strip().split("\n"): - print(f" {line}") - results.append((label, condition)) - - -# ── 6. Main ────────────────────────────────────────────────────────────────── - -def main(): - _write_config(CONFIG_FIX12, "/tmp/config_fix12.yaml") - _write_config(CONFIG_FIX3, "/tmp/config_fix3.yaml") - - print(f"{BOLD}Starting mock JWKS + OIDC discovery server on :19900 ...{RESET}") - jwks_srv = HTTPServer(("127.0.0.1", 19900), JWKSHandler) - threading.Thread(target=jwks_srv.serve_forever, daemon=True).start() - - print(f"{BOLD}Starting fake LLM backend on :19901 ...{RESET}") - llm_srv = HTTPServer(("127.0.0.1", 19901), FakeLLMHandler) - threading.Thread(target=llm_srv.serve_forever, daemon=True).start() - - # ── Fix 1 + 2 ──────────────────────────────────────────────────────── - header("Fix 1 + Fix 2: OIDC Discovery & Array Roles") - proxy = start_proxy("/tmp/config_fix12.yaml") - try: - print(f" Waiting for proxy on :{PROXY_PORT} (may take ~30s for Prisma) ...") - if not wait_for_proxy(): - print(f"{RED}Proxy failed to start!{RESET}") - out = proxy.stdout.read().decode() if proxy.stdout else "" - print(out[-3000:]) - raise SystemExit(1) - print(f" {GREEN}Proxy is up.{RESET}\n") - - # Fix 1a - token_ok = sign_token({"sub": "user-1", "roles": ["team-beta", "team-gamma"]}) - status, body = proxy_request(token_ok) - check("Fix 1 — OIDC discovery: valid JWT → HTTP 200", status == 200, f"status={status}") - - # Fix 1b - tampered = token_ok[:-5] + "XXXXX" - status, body = proxy_request(tampered) - check("Fix 1 — OIDC discovery: tampered JWT → HTTP 401", status == 401, f"status={status}") - - # Fix 2a - token_arr = sign_token({"sub": "user-2", "roles": ["team-beta", "team-gamma"]}) - status, body = proxy_request(token_arr) - check("Fix 2 — roles array ['team-beta','team-gamma'] → HTTP 200", status == 200, f"status={status}") - - # Fix 2b - token_empty = sign_token({"sub": "user-3", "roles": []}) - status, body = proxy_request(token_empty) - check("Fix 2 — empty roles [] → HTTP 401", status == 401, f"status={status}") - - finally: - stop_proxy(proxy) - time.sleep(2) - - # ── Fix 3 ───────────────────────────────────────────────────────────── - header("Fix 3: Dot-notation hint (roles.0)") - proxy = start_proxy("/tmp/config_fix3.yaml") - try: - print(f" Waiting for proxy on :{PROXY_PORT} ...") - if not wait_for_proxy(): - print(f"{RED}Proxy failed to start!{RESET}") - out = proxy.stdout.read().decode() if proxy.stdout else "" - print(out[-3000:]) - raise SystemExit(1) - print(f" {GREEN}Proxy is up.{RESET}\n") - - token_dot = sign_token({"sub": "user-4", "roles": ["team-alpha"]}) - status, body = proxy_request(token_dot) - hint_text = "Use 'roles' instead" - has_hint = hint_text in body - check( - "Fix 3 — roles.0 config → HTTP 401 with helpful hint", - status == 401 and has_hint, - f"status={status}, hint_present={has_hint}\n" - f"body excerpt: {body[:400]}", - ) - - finally: - stop_proxy(proxy) - - # ── Summary ─────────────────────────────────────────────────────────── - all_pass = all(ok for _, ok in results) - header("ALL TESTS PASSED" if all_pass else "SOME TESTS FAILED") - print() - raise SystemExit(0 if all_pass else 1) - - -if __name__ == "__main__": - main() diff --git a/demo_servers.py b/demo_servers.py deleted file mode 100644 index 3e941d3b690..00000000000 --- a/demo_servers.py +++ /dev/null @@ -1,103 +0,0 @@ -#!/usr/bin/env python3 -""" -Start mock JWKS/OIDC server (:19900) and fake LLM backend (:19901). -Also writes signed JWTs to /tmp/tokens.env for use in shell demos. -Runs forever — Ctrl-C to stop. -""" - -import base64, json, os, time, threading -from http.server import HTTPServer, BaseHTTPRequestHandler -from cryptography.hazmat.primitives.asymmetric import rsa -from cryptography.hazmat.primitives import serialization -import jwt as pyjwt - -# ── RSA key pair ────────────────────────────────────────────────────────────── -RSA_KEY = rsa.generate_private_key(public_exponent=65537, key_size=2048) -RSA_PUB = RSA_KEY.public_key() -PUB_NUMBERS = RSA_PUB.public_numbers() - -def _b64url(n, length): - return base64.urlsafe_b64encode(n.to_bytes(length, "big")).rstrip(b"=").decode() - -KID = "demo-kid-1" -JWKS = {"keys": [{"kty": "RSA", "kid": KID, "use": "sig", "alg": "RS256", - "n": _b64url(PUB_NUMBERS.n, 256), "e": _b64url(PUB_NUMBERS.e, 3)}]} - -PEM_PRIVATE = RSA_KEY.private_bytes( - serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption()) - -ISSUER = "https://demo-idp.example.com" -AUDIENCE = "litellm-proxy-demo" - -def sign_token(claims): - payload = {"iss": ISSUER, "aud": AUDIENCE, "exp": int(time.time()) + 3600, **claims} - return pyjwt.encode(payload, PEM_PRIVATE, algorithm="RS256", headers={"kid": KID}) - -# ── Write tokens to file ───────────────────────────────────────────────────── -valid_token = sign_token({"sub": "user-1", "roles": ["team-beta", "team-gamma"]}) -tampered_token = valid_token[:-5] + "XXXXX" -empty_roles_token = sign_token({"sub": "user-2", "roles": []}) -single_role_token = sign_token({"sub": "user-3", "roles": ["team-alpha"]}) - -with open("/tmp/tokens.env", "w") as f: - f.write(f'VALID_TOKEN="{valid_token}"\n') - f.write(f'TAMPERED_TOKEN="{tampered_token}"\n') - f.write(f'EMPTY_ROLES_TOKEN="{empty_roles_token}"\n') - f.write(f'SINGLE_ROLE_TOKEN="{single_role_token}"\n') - -print("[tokens] Written to /tmp/tokens.env") - -# ── Mock JWKS + OIDC discovery server ──────────────────────────────────────── -class JWKSHandler(BaseHTTPRequestHandler): - def do_GET(self): - if self.path == "/.well-known/openid-configuration": - body = json.dumps({"issuer": ISSUER, - "jwks_uri": "http://127.0.0.1:19900/jwks", - "authorization_endpoint": "https://example.com/authorize", - "token_endpoint": "https://example.com/token"}).encode() - elif self.path == "/jwks": - body = json.dumps(JWKS).encode() - else: - self.send_response(404); self.end_headers(); return - self.send_response(200) - self.send_header("Content-Type", "application/json") - self.end_headers() - self.wfile.write(body) - def log_message(self, *_): pass - -# ── Fake LLM backend ───────────────────────────────────────────────────────── -class FakeLLMHandler(BaseHTTPRequestHandler): - def do_POST(self): - body = json.dumps({"id": "chatcmpl-fake", "object": "chat.completion", - "created": 1700000000, "model": "fake-model", - "choices": [{"index": 0, "message": {"role": "assistant", - "content": "Hello from fake LLM!"}, "finish_reason": "stop"}], - "usage": {"prompt_tokens": 5, "completion_tokens": 7, "total_tokens": 12}}).encode() - self.send_response(200) - self.send_header("Content-Type", "application/json") - self.end_headers() - self.wfile.write(body) - def do_GET(self): - if "/models" in self.path: - body = json.dumps({"data": [{"id": "fake-model", "object": "model"}], "object": "list"}).encode() - self.send_response(200); self.send_header("Content-Type", "application/json") - self.end_headers(); self.wfile.write(body) - else: - self.send_response(404); self.end_headers() - def log_message(self, *_): pass - -# ── Start servers ───────────────────────────────────────────────────────────── -print("[JWKS] Starting on :19900 (OIDC discovery + JWKS)") -jwks_srv = HTTPServer(("127.0.0.1", 19900), JWKSHandler) -threading.Thread(target=jwks_srv.serve_forever, daemon=True).start() - -print("[LLM] Starting fake LLM backend on :19901") -llm_srv = HTTPServer(("127.0.0.1", 19901), FakeLLMHandler) -threading.Thread(target=llm_srv.serve_forever, daemon=True).start() - -print("[READY] Mock servers running. Press Ctrl-C to stop.") -try: - while True: - time.sleep(60) -except KeyboardInterrupt: - print("\nShutting down.")