From 3f90e5489498debe8a57d81a2c3133d1bfea334b Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 24 Sep 2026 08:00:36 +0000 Subject: [PATCH] fix(proxy): replace Fernet plugin session claims with AES-256-GCM Fernet (AES-128-CBC + HMAC-SHA256) is not an approved construction for the FIPS 140-3 image. Plugin session claims now use AES-256-GCM through a key-taking pair (encrypt_aes_gcm_with_key / decrypt_aes_gcm_with_key) that the existing _encrypt_aes_gcm / _decrypt_aes_gcm wrappers delegate to, keyed by the unchanged HMAC-SHA256(LITELLM_SALT_KEY, plugin_name) per-plugin key. Claims live 30 seconds so no Fernet read-compatibility window is kept. Resolves LIT-8428 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../common_utils/encrypt_decrypt_utils.py | 27 +++- litellm/proxy/plugin_routes.py | 33 +++-- .../test_plugin_session_claim.py | 132 ++++++++++++++++++ .../test_litellm/proxy/test_plugin_routes.py | 95 ++++++++++++- 4 files changed, 261 insertions(+), 26 deletions(-) create mode 100644 tests/integration/authorization/test_plugin_session_claim.py diff --git a/litellm/proxy/common_utils/encrypt_decrypt_utils.py b/litellm/proxy/common_utils/encrypt_decrypt_utils.py index 3584aaaf833..cc7327c7ccb 100644 --- a/litellm/proxy/common_utils/encrypt_decrypt_utils.py +++ b/litellm/proxy/common_utils/encrypt_decrypt_utils.py @@ -72,26 +72,41 @@ def _derive_key(signing_key: str) -> bytes: return hashlib.sha256(signing_key.encode()).digest() -def _encrypt_aes_gcm(value: str, signing_key: str) -> str: - """Encrypt under AES-256-GCM and return the versioned ``v2:gcm:`` string.""" +def encrypt_aes_gcm_with_key(value: str, key: bytes) -> str: + """Encrypt under AES-256-GCM with a caller-derived 32-byte key; returns the versioned ``v2:gcm:`` string.""" from cryptography.hazmat.primitives.ciphers.aead import AESGCM nonce: Final = os.urandom(12) # AESGCM.encrypt returns ciphertext || tag(16); wire format is nonce || that. - blob: Final = AESGCM(_derive_key(signing_key)).encrypt(nonce, value.encode("utf-8"), None) + blob: Final = AESGCM(key).encrypt(nonce, value.encode("utf-8"), None) return _V2_GCM_PREFIX + base64.urlsafe_b64encode(nonce + blob).decode("utf-8") -def _decrypt_aes_gcm(value: str, signing_key: str) -> str: - """Decrypt a versioned ``v2:gcm:`` string produced by :func:`_encrypt_aes_gcm`.""" +def decrypt_aes_gcm_with_key(value: str, key: bytes) -> str: + """Decrypt a ``v2:gcm:`` string produced by :func:`encrypt_aes_gcm_with_key` under the same key. + + Raises ``ValueError`` when the prefix is missing and ``InvalidTag`` when the key or ciphertext is wrong. + """ from cryptography.hazmat.primitives.ciphers.aead import AESGCM + if not value.startswith(_V2_GCM_PREFIX): + raise ValueError("not a v2:gcm ciphertext") raw: Final = base64.urlsafe_b64decode(value[len(_V2_GCM_PREFIX) :]) # An empty plaintext still serializes to nonce(12) || tag(16) = 28 bytes, so a # short/empty buffer here is a corrupt value: let AESGCM.decrypt raise and be # swallowed by decrypt_value_helper (returns None/original), same as legacy. nonce, blob = raw[:12], raw[12:] - return AESGCM(_derive_key(signing_key)).decrypt(nonce, blob, None).decode("utf-8") + return AESGCM(key).decrypt(nonce, blob, None).decode("utf-8") + + +def _encrypt_aes_gcm(value: str, signing_key: str) -> str: + """Encrypt under AES-256-GCM and return the versioned ``v2:gcm:`` string.""" + return encrypt_aes_gcm_with_key(value, _derive_key(signing_key)) + + +def _decrypt_aes_gcm(value: str, signing_key: str) -> str: + """Decrypt a versioned ``v2:gcm:`` string produced by :func:`_encrypt_aes_gcm`.""" + return decrypt_aes_gcm_with_key(value, _derive_key(signing_key)) def encrypt_value_helper(value: str, new_encryption_key: str | None = None): diff --git a/litellm/proxy/plugin_routes.py b/litellm/proxy/plugin_routes.py index eb6fe7dd177..9402940a32e 100644 --- a/litellm/proxy/plugin_routes.py +++ b/litellm/proxy/plugin_routes.py @@ -13,14 +13,13 @@ Config (in litellm config.yaml general_settings): Plugin iframe auth: The UI calls GET /api/plugins/auth-token to receive a short-lived identity - claim ({user_id, user_role, plugin, exp}) encrypted with a per-plugin key - derived as HMAC-SHA256(LITELLM_SALT_KEY, plugin_name). The claim carries no + claim ({user_id, user_role, plugin, exp}) encrypted with AES-256-GCM under a + per-plugin key derived as HMAC-SHA256(LITELLM_SALT_KEY, plugin_name). The claim carries no litellm bearer token, so a compromised plugin learns only the caller's identity, never their credential. LITELLM_SALT_KEY itself is never shared with plugins — each plugin holds only its own derived key. """ -import base64 import hashlib import hmac as _hmac import json @@ -29,12 +28,14 @@ import time from collections.abc import Mapping from typing import Final -from cryptography.fernet import Fernet, InvalidToken +from cryptography.exceptions import InvalidTag from fastapi import APIRouter, Depends, HTTPException, Request, Response +from pydantic import JsonValue, TypeAdapter, ValidationError from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.proxy._types import PluginConfig, SpecialHeaders, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_aes_gcm_with_key, encrypt_aes_gcm_with_key from litellm.types.llms.custom_http import httpxSpecialProvider router: Final = APIRouter() @@ -124,16 +125,15 @@ _plugin_registry: Final[dict[str, PluginConfig]] = {} # shared with plugins; each plugin only receives a key derived from # HMAC(LITELLM_SALT_KEY, plugin_name) which reveals nothing about the master. # --------------------------------------------------------------------------- -def _plugin_fernet(plugin_name: str) -> Fernet: - """Return a Fernet cipher whose key is scoped to a specific plugin. +def _plugin_key(plugin_name: str) -> bytes: + """Return the 32-byte AES-256-GCM key scoped to a specific plugin. Key material: HMAC-SHA256(LITELLM_SALT_KEY, plugin_name). A plugin possessing its own key cannot derive the master salt or forge claims intended for a different plugin. """ salt: Final = os.getenv("LITELLM_SALT_KEY", "").encode() - derived: Final = _hmac.new(salt, plugin_name.encode(), hashlib.sha256).digest() - return Fernet(base64.urlsafe_b64encode(derived)) + return _hmac.new(salt, plugin_name.encode(), hashlib.sha256).digest() _CLAIM_TTL_SECONDS: Final = 30 # identity claims expire after 30 s @@ -152,24 +152,27 @@ def issue_plugin_session_claim(plugin_name: str, user_id: str | None, user_role: "user_role": user_role or "", "exp": int(time.time()) + _CLAIM_TTL_SECONDS, } - return _plugin_fernet(plugin_name).encrypt(json.dumps(claim).encode()).decode() + return encrypt_aes_gcm_with_key(json.dumps(claim), _plugin_key(plugin_name)) -def verify_plugin_session_claim(plugin_name: str, ciphertext: str) -> dict: +_CLAIM_ADAPTER: Final = TypeAdapter(dict[str, JsonValue]) + + +def verify_plugin_session_claim(plugin_name: str, ciphertext: str) -> dict[str, JsonValue]: """Verify and decode a plugin session claim. - Raises ValueError if the HMAC is invalid, the audience is wrong, or + Raises ValueError if the GCM tag is invalid, the audience is wrong, or the claim is expired. Returns the decoded claim dict on success. """ try: - raw: Final = _plugin_fernet(plugin_name).decrypt(ciphertext.encode(), ttl=_CLAIM_TTL_SECONDS) - claim: Final = json.loads(raw) - except (InvalidToken, Exception) as exc: + claim: Final = _CLAIM_ADAPTER.validate_json(decrypt_aes_gcm_with_key(ciphertext, _plugin_key(plugin_name))) + except (ValueError, ValidationError, InvalidTag) as exc: raise ValueError("Invalid, tampered, or expired plugin session claim") from exc if claim.get("plugin") != plugin_name: raise ValueError("Plugin claim audience mismatch") - if int(claim.get("exp", 0)) < int(time.time()): + expiry: Final = claim.get("exp") + if not isinstance(expiry, int) or expiry < int(time.time()): raise ValueError("Plugin session claim expired") return claim diff --git a/tests/integration/authorization/test_plugin_session_claim.py b/tests/integration/authorization/test_plugin_session_claim.py new file mode 100644 index 00000000000..7ffb3cf4528 --- /dev/null +++ b/tests/integration/authorization/test_plugin_session_claim.py @@ -0,0 +1,132 @@ +import base64 +import hashlib +import hmac +import os +import signal +import time +import uuid +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from typing import Final + +import psutil +from cryptography.exceptions import InvalidTag +from cryptography.hazmat.primitives.ciphers.aead import AESGCM +from integration._support.client import Gateway, eventually, string_value +from integration._support.process import owned_proxy, owned_proxy_process +from pydantic import JsonValue, TypeAdapter + +_SALT: Final = "sk-integration-plugin-salt" +_PLUGIN: Final = "integration-plugin" +_OTHER_PLUGIN: Final = "integration-other-plugin" +_GCM_PREFIX: Final = "v2:gcm:" +_TTL_SECONDS: Final = 30 +_CLAIM_ADAPTER: Final = TypeAdapter(dict[str, JsonValue]) + + +def _plugin_config(directory: Path) -> Path: + config: Final = directory / "proxy_config.yaml" + config.write_text( + "model_list: []\n" + "general_settings:\n" + " master_key: os.environ/LITELLM_MASTER_KEY\n" + " database_url: os.environ/DATABASE_URL\n" + " store_model_in_db: true\n" + " plugins:\n" + f" - name: {_PLUGIN}\n" + " url: http://127.0.0.1:9\n" + f" - name: {_OTHER_PLUGIN}\n" + " url: http://127.0.0.1:9\n" + ) + return config + + +def _plugin_key(plugin_name: str, salt: str = _SALT) -> bytes: + return hmac.new(salt.encode(), plugin_name.encode(), hashlib.sha256).digest() + + +def _decrypt_claim(claim: str, key: bytes) -> dict[str, JsonValue]: + assert claim.startswith(_GCM_PREFIX), claim + raw: Final = base64.urlsafe_b64decode(claim[len(_GCM_PREFIX) :]) + return _CLAIM_ADAPTER.validate_json(AESGCM(key).decrypt(raw[:12], raw[12:], None)) + + +def _issue(candidate: Gateway, key: str, plugin_name: str = _PLUGIN) -> str: + response: Final = candidate.request("GET", "/api/plugins/auth-token", key=key, params={"plugin_name": plugin_name}) + assert response.status_code == 200, response.text + return string_value(_CLAIM_ADAPTER.validate_json(response.text)["session_claim"]) + + +def test_plugin_claim_is_aes_gcm_under_the_hmac_derived_plugin_key(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy(gateway, tmp_path, {"LITELLM_SALT_KEY": _SALT}, config=_plugin_config(tmp_path)) as candidate: + with candidate.scenario() as scenario: + user_id: Final = scenario.user(user_id=f"plugin-admin-{uuid.uuid4().hex}", user_role="proxy_admin") + key: Final = scenario.key(user_id=user_id) + issued_at: Final = int(time.time()) + claim: Final = _issue(candidate, key) + decrypted: Final = _decrypt_claim(claim, _plugin_key(_PLUGIN)) + assert decrypted["plugin"] == _PLUGIN, decrypted + assert decrypted["user_id"] == user_id, decrypted + assert isinstance(decrypted["user_role"], str), decrypted + expiry: Final = decrypted["exp"] + assert isinstance(expiry, int), decrypted + assert issued_at + _TTL_SECONDS <= expiry <= issued_at + _TTL_SECONDS + 5, decrypted + + +def test_plugin_claim_rejects_other_plugin_key_wrong_salt_and_tampering(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy(gateway, tmp_path, {"LITELLM_SALT_KEY": _SALT}, config=_plugin_config(tmp_path)) as candidate: + claim: Final = _issue(candidate, candidate.key) + raw: Final = base64.urlsafe_b64decode(claim[len(_GCM_PREFIX) :]) + flipped: Final = raw[:-1] + bytes([raw[-1] ^ 0x01]) + tampered: Final = _GCM_PREFIX + base64.urlsafe_b64encode(flipped).decode() + for wrong_key in (_plugin_key(_OTHER_PLUGIN), _plugin_key(_PLUGIN, salt="sk-not-the-proxy-salt")): + try: + _decrypt_claim(claim, wrong_key) + except InvalidTag: + continue + raise AssertionError(f"claim for {_PLUGIN} decrypted under a foreign key: {claim}") + try: + _decrypt_claim(tampered, _plugin_key(_PLUGIN)) + except InvalidTag: + pass + else: + raise AssertionError(f"tampered claim decrypted: {tampered}") + assert _decrypt_claim(claim, _plugin_key(_PLUGIN))["plugin"] == _PLUGIN + + +def test_plugin_auth_token_still_gates_registration_and_bearer(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy(gateway, tmp_path, {"LITELLM_SALT_KEY": _SALT}, config=_plugin_config(tmp_path)) as candidate: + missing: Final = candidate.request("GET", "/api/plugins/auth-token", params={"plugin_name": "not-registered"}) + assert missing.status_code == 404, missing.text + assert missing.json()["detail"] == "Plugin 'not-registered' is not registered." + anonymous: Final = candidate.client.get("/api/plugins/auth-token", params={"plugin_name": _PLUGIN}) + assert anonymous.status_code == 401, anonymous.text + assert _decrypt_claim(_issue(candidate, candidate.key), _plugin_key(_PLUGIN))["plugin"] == _PLUGIN + + +def test_plugin_claims_survive_a_worker_kill_during_a_burst(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy_process( + gateway, tmp_path, {"LITELLM_SALT_KEY": _SALT}, config=_plugin_config(tmp_path), workers=2 + ) as owned: + candidate: Final = owned.gateway + root: Final = psutil.Process(owned.process.pid) + workers: Final = eventually( + lambda: tuple(root.children(recursive=True)), lambda found: len(found) >= 2, seconds=30 + ) + with ThreadPoolExecutor(max_workers=8) as pool: + before: Final = tuple( + future.result() for future in [pool.submit(_issue, candidate, candidate.key) for _ in range(12)] + ) + os.kill(workers[0].pid, signal.SIGKILL) + probe: Final = eventually( + lambda: candidate.request("GET", "/api/plugins/auth-token", params={"plugin_name": _PLUGIN}), + lambda response: response.status_code == 200, + seconds=30, + ) + after: Final = tuple( + future.result() for future in [pool.submit(_issue, candidate, candidate.key) for _ in range(12)] + ) + claims: Final = before + (string_value(_CLAIM_ADAPTER.validate_json(probe.text)["session_claim"]),) + after + assert len(set(claims)) == len(claims), "nonce reuse: two claims shared ciphertext" + decrypted: Final = tuple(_decrypt_claim(claim, _plugin_key(_PLUGIN)) for claim in claims) + assert all(value["plugin"] == _PLUGIN for value in decrypted), decrypted diff --git a/tests/test_litellm/proxy/test_plugin_routes.py b/tests/test_litellm/proxy/test_plugin_routes.py index c8d2939385d..77a50022bac 100644 --- a/tests/test_litellm/proxy/test_plugin_routes.py +++ b/tests/test_litellm/proxy/test_plugin_routes.py @@ -12,9 +12,16 @@ Covers three bugs: """ import asyncio -from unittest.mock import MagicMock +import base64 +import hashlib +import hmac +import json +import time import pytest +from cryptography.exceptions import InvalidTag +from cryptography.hazmat.primitives.ciphers.aead import AESGCM +from pydantic import JsonValue, TypeAdapter from litellm.proxy._types import ( ConfigGeneralSettings, @@ -22,7 +29,13 @@ from litellm.proxy._types import ( PluginConfig, UserAPIKeyAuth, ) -from litellm.proxy.plugin_routes import list_plugins, register_plugins_from_config +from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_aes_gcm_with_key +from litellm.proxy.plugin_routes import ( + issue_plugin_session_claim, + list_plugins, + register_plugins_from_config, + verify_plugin_session_claim, +) def _admin() -> UserAPIKeyAuth: @@ -97,9 +110,7 @@ def test_registered_plugins_appear_in_list_without_restart() -> None: ] } ) - names = sorted( - p["name"] for p in asyncio.run(list_plugins(user_api_key_dict=_admin())) - ) + names = sorted(p["name"] for p in asyncio.run(list_plugins(user_api_key_dict=_admin()))) assert names == ["agent-builder", "chat-ui"] # Removing a plugin from config drops it from the live list. @@ -230,3 +241,77 @@ def test_configured_custom_key_header_is_stripped() -> None: assert "x-my-tenant-key" in _request_strip_headers() finally: proxy_server.general_settings = original + + +_CLAIM_SALT = "sk-unit-test-salt" + + +def _claim_key(plugin_name: str, salt: str = _CLAIM_SALT) -> bytes: + return hmac.new(salt.encode(), plugin_name.encode(), hashlib.sha256).digest() + + +def _decrypt_raw(claim: str, key: bytes) -> dict[str, JsonValue]: + assert claim.startswith("v2:gcm:"), claim + raw = base64.urlsafe_b64decode(claim[len("v2:gcm:") :]) + return TypeAdapter(dict[str, JsonValue]).validate_json(AESGCM(key).decrypt(raw[:12], raw[12:], None)) + + +def test_plugin_claim_round_trip_is_aes_gcm_under_hmac_plugin_key(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", _CLAIM_SALT) + before = int(time.time()) + claim = issue_plugin_session_claim("chat-ui", "user-1", "proxy_admin") + + verified = verify_plugin_session_claim("chat-ui", claim) + assert verified["plugin"] == "chat-ui" + assert verified["user_id"] == "user-1" + assert verified["user_role"] == "proxy_admin" + expiry = verified["exp"] + assert isinstance(expiry, int) and before + 30 <= expiry <= before + 31, verified + + plugin_side = _decrypt_raw(claim, _claim_key("chat-ui")) + assert plugin_side == verified + assert issue_plugin_session_claim("chat-ui", "user-1", "proxy_admin") != claim + + +def test_plugin_claim_for_other_plugin_is_rejected(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", _CLAIM_SALT) + claim = issue_plugin_session_claim("chat-ui", "user-1", "proxy_admin") + + with pytest.raises(ValueError, match="Invalid, tampered, or expired"): + verify_plugin_session_claim("agent-builder", claim) + with pytest.raises(InvalidTag): + _decrypt_raw(claim, _claim_key("agent-builder")) + + +def test_plugin_claim_with_forged_audience_is_rejected(monkeypatch: pytest.MonkeyPatch) -> None: + """A claim encrypted under plugin A's key but naming plugin B fails the audience check.""" + monkeypatch.setenv("LITELLM_SALT_KEY", _CLAIM_SALT) + payload = {"plugin": "agent-builder", "user_id": "user-1", "user_role": "proxy_admin", "exp": int(time.time()) + 30} + forged = encrypt_aes_gcm_with_key(json.dumps(payload), _claim_key("chat-ui")) + + with pytest.raises(ValueError, match="audience mismatch"): + verify_plugin_session_claim("chat-ui", forged) + + +def test_expired_plugin_claim_is_rejected(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", _CLAIM_SALT) + payload = {"plugin": "chat-ui", "user_id": "user-1", "user_role": "proxy_admin", "exp": int(time.time()) - 1} + expired = encrypt_aes_gcm_with_key(json.dumps(payload), _claim_key("chat-ui")) + + with pytest.raises(ValueError, match="expired"): + verify_plugin_session_claim("chat-ui", expired) + + +def test_tampered_plugin_claim_is_rejected(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", _CLAIM_SALT) + claim = issue_plugin_session_claim("chat-ui", "user-1", "proxy_admin") + raw = bytearray(base64.urlsafe_b64decode(claim[len("v2:gcm:") :])) + raw[-1] ^= 0x01 + tampered = "v2:gcm:" + base64.urlsafe_b64encode(bytes(raw)).decode() + + with pytest.raises(ValueError, match="Invalid, tampered, or expired"): + verify_plugin_session_claim("chat-ui", tampered) + with pytest.raises(ValueError, match="Invalid, tampered, or expired"): + verify_plugin_session_claim("chat-ui", "gAAAAABlegacyFernetToken==") + with pytest.raises(ValueError, match="Invalid, tampered, or expired"): + verify_plugin_session_claim("chat-ui", encrypt_aes_gcm_with_key("[1, 2]", _claim_key("chat-ui")))