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>
This commit is contained in:
yucheng 2026-09-24 08:00:36 +00:00
parent 09ebb28473
commit 3f90e54894
4 changed files with 261 additions and 26 deletions

View file

@ -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):

View file

@ -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

View file

@ -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

View file

@ -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")))