diff --git a/litellm/proxy/client/cli/commands/encryption.py b/litellm/proxy/client/cli/commands/encryption.py index 4c6ab94191e..48d71075a53 100644 --- a/litellm/proxy/client/cli/commands/encryption.py +++ b/litellm/proxy/client/cli/commands/encryption.py @@ -1,5 +1,6 @@ """CLI commands for the at-rest credential encryption migration.""" +from types import MappingProxyType from typing import Final import click @@ -27,28 +28,51 @@ def encryption(): default=False, help="Run the full migration walkers without writing any changes.", ) +@click.option( + "--mode", + type=click.Choice(("algorithm", "salt-key")), + default="algorithm", + show_default=True, + help="'algorithm' re-encrypts into AES-256-GCM; 'salt-key' re-encrypts under the active salt key.", +) @click.pass_context -def migrate(ctx: click.Context, check_only: bool, dry_run: bool): - """Re-encrypt at-rest credentials into the AES-256-GCM (v2:gcm:) format. +def migrate(ctx: click.Context, check_only: bool, dry_run: bool, mode: str): + """Re-encrypt at-rest credentials, by algorithm or under a rotated salt key. - Requires the proxy to be started with - ``general_settings.encryption_algorithm: aes-256-gcm``. Idempotent and - resumable; safe to re-run after an interruption. + ``--mode algorithm`` (the default) moves values into the AES-256-GCM + (v2:gcm:) format and requires the proxy to be started with + ``general_settings.encryption_algorithm: aes-256-gcm``. + + ``--mode salt-key`` re-encrypts values that still decrypt only under a + retired salt key. Restart the proxy with the new key in ``LITELLM_SALT_KEY`` + and the retired one(s) in ``LITELLM_SALT_KEY_PREVIOUS`` (comma-separated), run + this, then drop ``LITELLM_SALT_KEY_PREVIOUS`` once ``--check`` reports + ``residual_legacy: 0`` and an empty ``unreadable_locations``. Virtual keys are + SHA-256 hashes rather than salt-key ciphertext, so they keep working + throughout and never need regenerating. + + Both modes are idempotent and resumable; safe to re-run after an + interruption. Examples: - litellm-proxy encryption migrate --check # attestation scan, no writes - litellm-proxy encryption migrate # perform the migration + litellm-proxy encryption migrate --check # attestation scan, no writes + litellm-proxy encryption migrate # perform the migration + litellm-proxy encryption migrate --mode salt-key # rotate the salt key """ client: Final = HTTPClient(ctx.obj["base_url"], ctx.obj["api_key"]) if check_only: - response = client.request("GET", "/credentials/migrate-encryption/check") + response = client.request( + "GET", + "/credentials/migrate-encryption/check", + params=MappingProxyType({"mode": mode}), + ) else: response = client.request( "POST", "/credentials/migrate-encryption", json={}, - params={"dry_run": "true"} if dry_run else None, + params=MappingProxyType({"mode": mode, "dry_run": str(dry_run).lower()}), ) rich.print_json(data=response) diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index 9379a8577a3..5d0784a5189 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -663,7 +663,7 @@ def encrypt_callback_vars(metadata: Any) -> Any: Idempotent: a value that already decrypts cleanly is left unchanged so round-trips through edit forms don't double-encrypt. """ - return _transform_callback_vars(metadata, _encrypt_if_plaintext) + return transform_callback_vars(metadata, _encrypt_if_plaintext) def decrypt_callback_vars(metadata: Any) -> Any: @@ -671,10 +671,10 @@ def decrypt_callback_vars(metadata: Any) -> Any: Legacy plaintext rows pass through unchanged (decrypt failure → original). """ - return _transform_callback_vars(metadata, _decrypt_or_passthrough) + return transform_callback_vars(metadata, _decrypt_or_passthrough) -def _transform_callback_vars(metadata: Any, transform: Callable[[str, Any], Any]) -> Any: +def transform_callback_vars(metadata: Any, transform: Callable[[str, Any], Any]) -> Any: if not isinstance(metadata, dict): return metadata out: Final = copy.deepcopy(metadata) diff --git a/litellm/proxy/common_utils/encrypt_decrypt_utils.py b/litellm/proxy/common_utils/encrypt_decrypt_utils.py index 836a4a778bb..f395a654e3a 100644 --- a/litellm/proxy/common_utils/encrypt_decrypt_utils.py +++ b/litellm/proxy/common_utils/encrypt_decrypt_utils.py @@ -1,5 +1,6 @@ import base64 import os +from dataclasses import dataclass from typing import Final, Literal, cast from litellm._logging import verbose_proxy_logger @@ -19,6 +20,12 @@ _ENCRYPTION_ALGORITHM_SETTING: Final = "encryption_algorithm" _ALGO_AES_GCM: Final = "aes-256-gcm" _ALGO_XSALSA20: Final = "xsalsa20-poly1305" +# Comma-separated retired salt keys, accepted on read only. New writes always use +# the active LITELLM_SALT_KEY, so setting this makes a salt-key rotation +# zero-downtime: values encrypted under a retired key stay readable until the +# rotation walkers have re-encrypted them under the active key. +_PREVIOUS_SALT_KEYS_ENV: Final = "LITELLM_SALT_KEY_PREVIOUS" + def _get_salt_key(): from litellm.proxy.proxy_server import master_key @@ -31,6 +38,19 @@ def _get_salt_key(): return salt_key +def get_previous_salt_keys() -> tuple[str, ...]: + """Retired salt keys from ``LITELLM_SALT_KEY_PREVIOUS``, in configured order.""" + raw: Final = os.getenv(_PREVIOUS_SALT_KEYS_ENV) or "" + return tuple(dict.fromkeys(k.strip() for k in raw.split(",") if k.strip())) + + +def get_decryption_keys() -> tuple[str, ...]: + """Keys tried on read: the active salt key first, then any retired ones.""" + primary: Final = _get_salt_key() + previous: Final = tuple(k for k in get_previous_salt_keys() if k != primary) + return (() if primary is None else (primary,)) + previous + + def _get_encryption_algorithm() -> str: """ Resolve the configured at-rest encryption algorithm for *new writes*. @@ -79,6 +99,11 @@ def _encrypt_aes_gcm(value: str, signing_key: str) -> str: return _V2_GCM_PREFIX + base64.urlsafe_b64encode(nonce + blob).decode("utf-8") +def _encrypt_xsalsa20(value: str, signing_key: str) -> str: + """Encrypt under the legacy XSalsa20-Poly1305 (nacl) format, url-safe base64.""" + return base64.urlsafe_b64encode(encrypt_value(value=value, signing_key=signing_key)).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`.""" from cryptography.hazmat.primitives.ciphers.aead import AESGCM @@ -91,6 +116,66 @@ def _decrypt_aes_gcm(value: str, signing_key: str) -> str: return AESGCM(_derive_key(signing_key)).decrypt(nonce, blob, None).decode("utf-8") +@dataclass(frozen=True, slots=True) +class _DecryptAttempt: + plaintext: str | None + error: BaseException | None + + +def _decrypt_with_key(value: str, signing_key: str) -> str: + """Decrypt one stored value under ``signing_key``, detecting its format. + + Versioned AES-256-GCM values are detected before any base64 decode: the + prefix is the algorithm tag the legacy nacl format never carried. Legacy + values are url-safe base64 (new) or standard base64 (old). + """ + if value.startswith(_V2_GCM_PREFIX): + return _decrypt_aes_gcm(value=value, signing_key=signing_key) + + try: + decoded_b64: Final = base64.urlsafe_b64decode(value) + except (ValueError, TypeError): + return decrypt_value(value=base64.b64decode(value), signing_key=signing_key) + return decrypt_value(value=decoded_b64, signing_key=signing_key) + + +def _attempt_decrypt(value: str, signing_key: str) -> _DecryptAttempt: + try: + return _DecryptAttempt(plaintext=_decrypt_with_key(value=value, signing_key=signing_key), error=None) + except Exception as e: # noqa: BLE001 # a wrong key fails differently per algorithm (nacl CryptoError, InvalidTag, decode errors) + return _DecryptAttempt(plaintext=None, error=e) + + +def try_decrypt_with_key(value: str, signing_key: str) -> str | None: + """``_decrypt_with_key`` as a value: ``None`` when the key does not fit.""" + return _attempt_decrypt(value=value, signing_key=signing_key).plaintext + + +def _decrypt_with_any_key(value: str) -> _DecryptAttempt: + """Decrypt under the first salt key that fits (active, then retired).""" + attempts: Final = tuple(_attempt_decrypt(value=value, signing_key=k) for k in get_decryption_keys()) + fallback: Final = ( + attempts[-1] if attempts else _DecryptAttempt(plaintext=None, error=ValueError("No salt key is set")) + ) + return next((a for a in attempts if a.plaintext is not None), fallback) + + +def encrypt_value_in_format_of(value: str, reference_ciphertext: str) -> str: + """Encrypt ``value`` under the active salt key in ``reference_ciphertext``'s format. + + The algorithm and the salt key are independent axes: a key-only rotation must + leave each value's algorithm alone. Rewriting through ``encrypt_value_helper`` + would instead apply the configured *write* algorithm, silently downgrading an + AES-256-GCM value to the legacy format whenever that setting is the default. + """ + signing_key: Final = _get_salt_key() + if signing_key is None: + raise RuntimeError("Cannot re-encrypt: neither LITELLM_SALT_KEY nor a master key is configured.") + if reference_ciphertext.startswith(_V2_GCM_PREFIX): + return _encrypt_aes_gcm(value=value, signing_key=signing_key) + return _encrypt_xsalsa20(value=value, signing_key=signing_key) + + def encrypt_value_helper(value: str, new_encryption_key: str | None = None): signing_key: Final = new_encryption_key or _get_salt_key() @@ -101,11 +186,7 @@ def encrypt_value_helper(value: str, new_encryption_key: str | None = None): # is returned directly with no extra base64 wrapper. return _encrypt_aes_gcm(value=value, signing_key=cast(str, signing_key)) - encrypted_value = encrypt_value(value=value, signing_key=signing_key) - # Use urlsafe_b64encode for URL-safe base64 encoding (replaces + with - and / with _) - encrypted_value = base64.urlsafe_b64encode(encrypted_value).decode("utf-8") - - return encrypted_value + return _encrypt_xsalsa20(value=value, signing_key=signing_key) verbose_proxy_logger.debug( "Invalid value type passed to encrypt_value: %s for Value: %s\n Value must be a string", type(value), value @@ -122,25 +203,15 @@ def decrypt_value_helper( exception_type: Literal["debug", "error"] = "error", return_original_value: bool = False, ): - signing_key: Final = _get_salt_key() - try: if isinstance(value, str): - # Versioned AES-256-GCM values are detected before any base64 decode. - # The prefix is the algorithm tag the legacy nacl format never carried. - if value.startswith(_V2_GCM_PREFIX): - return _decrypt_aes_gcm(value=value, signing_key=cast(str, signing_key)) - - # Try URL-safe base64 decoding first (new format) - # Fall back to standard base64 decoding for backwards compatibility (old format) - try: - decoded_b64 = base64.urlsafe_b64decode(value) - except Exception: - # If URL-safe decoding fails, try standard base64 decoding for backwards compatibility - decoded_b64 = base64.b64decode(value) - - value = decrypt_value(value=decoded_b64, signing_key=signing_key) - return value + # Read under the active salt key, then any retired keys listed in + # LITELLM_SALT_KEY_PREVIOUS, so a salt-key rotation stays readable + # while the re-encryption walkers catch up. + attempt: Final = _decrypt_with_any_key(value) + if attempt.plaintext is not None: + return attempt.plaintext + raise attempt.error or ValueError("Unable to decrypt value") # if it's not str - do not decrypt it, return the value return value diff --git a/litellm/proxy/management_endpoints/credential_migration.py b/litellm/proxy/management_endpoints/credential_migration.py index e0725119576..337b12aab86 100644 --- a/litellm/proxy/management_endpoints/credential_migration.py +++ b/litellm/proxy/management_endpoints/credential_migration.py @@ -1,10 +1,24 @@ """ -At-rest credential re-encryption migration. +At-rest credential re-encryption. -Switches every encrypted-at-rest value from the legacy XSalsa20-Poly1305 (nacl) -format to the versioned AES-256-GCM (``v2:gcm:``) format produced by -``encrypt_decrypt_utils`` when ``general_settings.encryption_algorithm`` is set to -``aes-256-gcm``. +Two passes share one set of walkers, selected by a :class:`ReencryptPolicy`: + +* **algorithm** (:data:`ALGORITHM_POLICY`) switches every encrypted-at-rest value + from the legacy XSalsa20-Poly1305 (nacl) format to the versioned AES-256-GCM + (``v2:gcm:``) format produced by ``encrypt_decrypt_utils`` when + ``general_settings.encryption_algorithm`` is set to ``aes-256-gcm``. The key is + unchanged. +* **salt key** (:data:`SALT_KEY_POLICY`) re-encrypts every value that still + decrypts only under a retired salt key (``LITELLM_SALT_KEY_PREVIOUS``) under the + active ``LITELLM_SALT_KEY``. This is what makes a leaked salt key recoverable + without regenerating virtual keys: those are SHA-256 hashes, never salt-key + ciphertext. + +The two are independent axes, and a key-only rotation must not move the algorithm +one. Each salt-key rewrite therefore reproduces the algorithm of the value it +replaces (``encrypt_value_in_format_of``) rather than writing through the +configured write algorithm, which would downgrade an AES-256-GCM value whenever +that setting sits at its legacy default. Design properties (see case 2026-06-24 fix plan): @@ -20,19 +34,29 @@ Design properties (see case 2026-06-24 fix plan): overwritten — corrupt rows are preserved and reported, never destroyed. * **Attestable.** :func:`check_encryption` is a read-only scan that classifies every value as ``migrated`` / ``legacy`` / ``plaintext`` / ``undecryptable``. - A residual ``legacy == 0`` is the compliance attestation. + A residual ``legacy == 0`` with an empty ``unreadable_locations`` is the + compliance attestation. A store the scan could not open reports zero of + everything, which means unknown rather than clean, so it is named there + instead of quietly passing. Coverage. The covered tables (model table, credentials table, MCP credential/env -tables, config ``environment_variables``) already have a re-encryption path in -``_rotate_master_key``; this module delegates to it in *same-key* mode and adds -walkers for the locations that had no rotation path: team / verification-token -``callback_vars`` metadata, the ``vantage_settings`` / ``cloudzero_settings`` -config rows, and the SSO config table. +tables, SSO identity assertions, config ``environment_variables``) already have a +re-encryption path in ``_rotate_master_key``; this module delegates to it in +*same-key* mode and adds walkers for the locations that had no rotation path: team +/ verification-token ``callback_vars`` metadata, the ``vantage_settings`` / ``cloudzero_settings`` +config rows, and the SSO / cache / config-override settings rows. Every one of them is +scanned by :func:`check_encryption`, so a store whose rotation failed (that path +logs and carries on rather than aborting) surfaces as residual legacy instead of +being attested clean. """ import json +from collections.abc import AsyncIterator, Callable, Mapping from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Final, Literal, cast +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, Literal, assert_never, cast + +from pydantic import TypeAdapter from litellm._logging import verbose_proxy_logger @@ -46,24 +70,98 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( _get_salt_key, decrypt_value_helper, encrypt_value_helper, + encrypt_value_in_format_of, + get_previous_salt_keys, + try_decrypt_with_key, ) +_MAYBE_STR: Final = TypeAdapter(str | None) + ValueClass = Literal["migrated", "legacy", "plaintext", "undecryptable", "not-a-string"] +@dataclass(frozen=True, slots=True) +class ReencryptPolicy: + """What "already re-encrypted" means for one pass over the stored values. + + ``is_current`` decides whether a decryptable value is left alone (counted as + ``already_v2``) or rewritten (counted as ``migrated``). Everything a pass + would still rewrite is residual ``legacy``. ``encrypt`` takes the recovered + plaintext and the value as stored, and produces the replacement ciphertext: + the algorithm pass writes through the configured algorithm, while the + salt-key pass keeps the stored value's own algorithm. + """ + + name: Literal["algorithm", "salt_key"] + is_current: Callable[[str], bool] + encrypt: Callable[[str, str], str] + + +def _is_target_algorithm(value: str) -> bool: + return value.startswith(_V2_GCM_PREFIX) + + +def _is_under_active_salt_key(value: str) -> bool: + primary: Final = _get_salt_key() + return primary is not None and try_decrypt_with_key(value=value, signing_key=primary) is not None + + +def _encrypt_under_configured_algorithm(plaintext: str, _stored: str) -> str: + return encrypt_value_helper(plaintext) + + +ALGORITHM_POLICY: Final = ReencryptPolicy( + name="algorithm", + is_current=_is_target_algorithm, + encrypt=_encrypt_under_configured_algorithm, +) +SALT_KEY_POLICY: Final = ReencryptPolicy( + name="salt_key", + is_current=_is_under_active_salt_key, + encrypt=encrypt_value_in_format_of, +) + +ReencryptMode = Literal["algorithm", "salt-key"] + + +def policy_for_mode(mode: ReencryptMode) -> ReencryptPolicy: + """Map the wire-level ``mode`` of the migration endpoints onto its policy.""" + match mode: + case "algorithm": + return ALGORITHM_POLICY + case "salt-key": + return SALT_KEY_POLICY + case _: + assert_never(mode) + + @dataclass class LocationReport: """Per-location counters for one migration / check pass.""" location: str scanned: int = 0 - migrated: int = 0 # values rewritten to v2 this run - already_v2: int = 0 # values already migrated (skipped) + migrated: int = 0 # values rewritten this run + already_v2: int = 0 # values already in the target shape (skipped) plaintext: int = 0 # legacy-plaintext values (no ciphertext to migrate) undecryptable: int = 0 # could not decrypt — preserved, not overwritten # Used by --check (read-only classification): - legacy: int = 0 # nacl ciphertext still awaiting migration + legacy: int = 0 # ciphertext still awaiting re-encryption + + # The store could not be read at all, so its zero counts mean "unknown", + # never "clean". Kept out of the counters so it can never be mistaken for one. + unreadable: bool = False + + def absorb(self, other: "LocationReport") -> None: + """Fold another report's counters into this one, for per-row accumulation.""" + self.scanned += other.scanned + self.migrated += other.migrated + self.already_v2 += other.already_v2 + self.plaintext += other.plaintext + self.undecryptable += other.undecryptable + self.legacy += other.legacy + self.unreadable = self.unreadable or other.unreadable def as_dict(self) -> dict[str, int]: return { @@ -94,10 +192,21 @@ class MigrationReport: def total_undecryptable(self) -> int: return sum(loc.undecryptable for loc in self.locations) + @property + def unreadable_locations(self) -> tuple[str, ...]: + """Stores that could not be read, so their zero counts prove nothing. + + ``residual_legacy == 0`` only attests a clean rotation while this is + empty: a store the scan could not open may still hold values under the + retired key, and dropping that key would then make them unreadable. + """ + return tuple(loc.location for loc in self.locations if loc.unreadable) + def as_dict(self) -> dict[str, object]: return { "residual_legacy": self.residual_legacy, "total_undecryptable": self.total_undecryptable, + "unreadable_locations": list(self.unreadable_locations), "locations": {loc.location: loc.as_dict() for loc in self.locations}, } @@ -112,14 +221,16 @@ def is_migrated(value: object) -> bool: return isinstance(value, str) and value.startswith(_V2_GCM_PREFIX) -def classify_value(value: object, key: str = "scan") -> ValueClass: +def classify_value(value: object, key: str = "scan", policy: ReencryptPolicy = ALGORITHM_POLICY) -> ValueClass: """Classify a stored value for the residual scanner. * ``not-a-string`` — not a string (numbers/bools/None left as-is on disk). - * ``migrated`` — carries the ``v2:gcm:`` prefix. - * ``legacy`` — decrypts under the legacy nacl reader (still needs migrating). - * ``plaintext`` — a non-empty string that does not decrypt and is not v2; - treated as legacy plaintext (nothing to migrate). + * ``migrated`` — already in the shape ``policy`` targets (``v2:gcm:`` format, + or readable under the active salt key). + * ``legacy`` — decrypts (under any configured salt key) but not in the target + shape, so this pass would rewrite it. + * ``plaintext`` — a non-empty string that does not decrypt at all; treated as + legacy plaintext (nothing to migrate). * ``undecryptable`` — reserved for callers that already know a value is ciphertext but cannot decrypt it; ``classify_value`` itself cannot tell a corrupt ciphertext from plaintext, so it returns ``plaintext`` for both. @@ -128,36 +239,50 @@ def classify_value(value: object, key: str = "scan") -> ValueClass: return "not-a-string" if value == "": return "plaintext" - if value.startswith(_V2_GCM_PREFIX): + if policy.is_current(value): return "migrated" decrypted: Final = decrypt_value_helper(value=value, key=key, exception_type="debug", return_original_value=False) if decrypted is None: - # Did not decrypt under nacl and has no v2 marker: legacy plaintext. + # Did not decrypt under any configured salt key: legacy plaintext. return "plaintext" return "legacy" -def reencrypt_value(value: object, key: str = "migrate") -> object: - """Re-encrypt a single stored string into the configured (AES) format. +def reencrypt_string(value: str, key: str = "migrate", policy: ReencryptPolicy = ALGORITHM_POLICY) -> str: + """Re-encrypt one stored string into the shape ``policy`` targets. - Returns the value unchanged if it is not a string, is already ``v2:``, or - cannot be decrypted (skip-on-undecryptable). Otherwise decrypts under the - format-detecting reader and re-encrypts through ``encrypt_value_helper`` - (which writes AES when the gate is on). + Returns the value unchanged if it already matches the policy or cannot be + decrypted (skip-on-undecryptable). Otherwise decrypts under the format- and + key-detecting reader and re-encrypts through the policy's writer, always + under the active salt key. """ - if not isinstance(value, str) or value == "": - return value - if value.startswith(_V2_GCM_PREFIX): + if value == "" or policy.is_current(value): return value # idempotent: already migrated - decrypted: Final = decrypt_value_helper(value=value, key=key, exception_type="debug", return_original_value=False) + decrypted: Final = _MAYBE_STR.validate_python( + decrypt_value_helper(value=value, key=key, exception_type="debug", return_original_value=False) + ) if decrypted is None: # Either legacy plaintext (no ciphertext to migrate) or corrupt. Either # way, do not overwrite — preserve the value as stored. return value - return encrypt_value_helper(decrypted) + return policy.encrypt(decrypted, value) -def reencrypt_selective_dict(data: dict[str, object], sensitive_keys: list[str]) -> dict[str, object]: +def reencrypt_value(value: object, key: str = "migrate", policy: ReencryptPolicy = ALGORITHM_POLICY) -> object: + """:func:`reencrypt_string` over a stored value of unknown type. + + Non-strings (numbers/bools/None) are left on disk exactly as they are. + """ + if not isinstance(value, str): + return value + return reencrypt_string(value, key=key, policy=policy) + + +def reencrypt_selective_dict( + data: dict[str, object], + sensitive_keys: list[str], + policy: ReencryptPolicy = ALGORITHM_POLICY, +) -> dict[str, object]: """Return a copy of ``data`` with only ``sensitive_keys`` re-encrypted. Non-sensitive fields (e.g. ``base_url``, ``connection_id``) are left as-is. @@ -168,24 +293,33 @@ def reencrypt_selective_dict(data: dict[str, object], sensitive_keys: list[str]) v = out.get(k) if v is None: continue - out[k] = reencrypt_value(v, key=k) + out[k] = reencrypt_value(v, key=k, policy=policy) return out +def _configured_algorithm() -> object: + from litellm.proxy.proxy_server import general_settings + + return general_settings.get(_ENCRYPTION_ALGORITHM_SETTING) + + +def _aes_gate_enabled() -> bool: + """Whether new writes land in AES-256-GCM (``general_settings``, read live).""" + algo: Final = _configured_algorithm() + return isinstance(algo, str) and algo.lower() == _ALGO_AES_GCM + + def _assert_aes_gate_enabled() -> None: """Fail fast if the AES algorithm gate is not enabled. Running the migration with the gate off would decrypt then re-encrypt right back into the legacy format — a no-op that silently fails the migration. """ - from litellm.proxy.proxy_server import general_settings - - algo: Final = general_settings.get(_ENCRYPTION_ALGORITHM_SETTING) - if not (isinstance(algo, str) and algo.lower() == _ALGO_AES_GCM): + if not _aes_gate_enabled(): raise RuntimeError( "Encryption migration requires general_settings.encryption_algorithm: " - f"'{_ALGO_AES_GCM}'. Current value: {algo!r}. Set it before migrating " - "so re-encrypted values are written in the AES-256-GCM format." + f"'{_ALGO_AES_GCM}'. Current value: {_configured_algorithm()!r}. Set it before " + "migrating so re-encrypted values are written in the AES-256-GCM format." ) @@ -196,11 +330,65 @@ def _assert_aes_gate_enabled() -> None: # --------------------------------------------------------------------------- +def _reencrypt_settings_fields( + settings: Mapping[str, object], + fields: tuple[str, ...], + dry_run: bool, + policy: ReencryptPolicy, + location: str, +) -> tuple[Mapping[str, object], LocationReport]: + """Re-encrypt ``fields`` of a settings dict, with the counters for what it did. + + A dry run counts what it would rewrite as residual ``legacy`` and returns the + dict unchanged, so ``--check`` never reports something as migrated that was + never written. + """ + report: Final = LocationReport(location=location) + out: Final = dict(settings) + for fld in fields: + v = out.get(fld) + if v is None: + continue + report.scanned += 1 + cls = classify_value(v, key=fld, policy=policy) + if cls == "migrated": + report.already_v2 += 1 + elif cls != "legacy": + report.plaintext += 1 + elif dry_run: + report.legacy += 1 + else: + new_v = reencrypt_value(v, key=fld, policy=policy) + if new_v == v: + report.legacy += 1 # did not re-encrypt, so it is still residual + else: + out[fld] = new_v + report.migrated += 1 + return out, report + + +def _encrypted_string_fields(settings: Mapping[str, object]) -> tuple[str, ...]: + """Field names holding a non-empty string, the only shape that can be ciphertext.""" + return tuple(k for k, v in settings.items() if isinstance(v, str) and v != "") + + +def _parse_settings_column(raw: object) -> Mapping[str, object] | None: + """A settings column as a dict, or ``None`` when it is absent, unparseable, or not one.""" + if not isinstance(raw, str): + return raw if isinstance(raw, dict) else None + try: + parsed: Final = json.loads(raw) + except (ValueError, TypeError): + return None + return parsed if isinstance(parsed, dict) else None + + async def _migrate_config_settings_row( prisma_client: object, param_name: str, sensitive_fields: list[str], dry_run: bool, + policy: ReencryptPolicy = ALGORITHM_POLICY, ) -> LocationReport: """Migrate a single ``LiteLLM_Config`` row whose ``param_value`` is a JSON dict with selected sensitive fields (vantage_settings / cloudzero_settings). @@ -210,117 +398,94 @@ async def _migrate_config_settings_row( if record is None or record.param_value is None: return report - settings = record.param_value - if isinstance(settings, str): - settings = json.loads(settings) - if not isinstance(settings, dict): + settings: Final = _parse_settings_column(record.param_value) + if settings is None: return report - changed = False - for fld in sensitive_fields: - v = settings.get(fld) - if v is None: - continue - report.scanned += 1 - cls = classify_value(v, key=fld) - if cls == "migrated": - report.already_v2 += 1 - continue - if cls == "legacy": - if dry_run: - # Residual: would migrate, but a dry run writes nothing, so it - # stays legacy for the attestation (never counted as migrated). - report.legacy += 1 - continue - new_v = reencrypt_value(v, key=fld) - if new_v != v: - settings[fld] = new_v - report.migrated += 1 - changed = True - else: - # Defensive: a legacy value that did not re-encrypt is still - # residual, not migrated. - report.legacy += 1 - else: # plaintext / not-a-string — nothing to migrate - report.plaintext += 1 - - if changed and not dry_run: + rewritten, fields_report = _reencrypt_settings_fields( + settings, tuple(sensitive_fields), dry_run, policy, param_name + ) + report.absorb(fields_report) + if rewritten != settings: await prisma_client.db.litellm_config.update( where={"param_name": param_name}, - data={"param_value": json.dumps(settings)}, + data={"param_value": json.dumps(rewritten)}, ) return report -async def _migrate_sso_config(prisma_client: object, dry_run: bool) -> LocationReport: - """Migrate the ``LiteLLM_SSOConfig`` row. All non-null fields are encrypted - (via the same ``_encrypt_env_variables`` path used on save), so we re-encrypt - every present string field. - """ - report: Final = LocationReport(location="sso_config") - record: Final = await prisma_client.db.litellm_ssoconfig.find_unique(where={"id": "sso_config"}) - if record is None or record.sso_settings is None: +# (location, prisma db attribute, primary-key field, JSON column of encrypted values). +# Every value in these rows is written through ``_encrypt_env_variables`` (or the +# SSO save path, which is the same code), so each field is salt-key ciphertext +# with no marker. None of them has a master-key rotation path, so a salt-key +# rotation that skipped them would leave live credentials (an IdP client secret, a +# Redis password, a Vault token) readable only under the retired key while the +# attestation reported a clean run. +_SETTINGS_ROW_SPECS: Final = ( + ("sso_config", "litellm_ssoconfig", "id", "sso_settings"), + ("cache_config", "litellm_cacheconfig", "id", "cache_settings"), + ("config_overrides", "litellm_configoverrides", "config_type", "config_value"), +) + + +async def _migrate_settings_rows( + prisma_client: object, + location: str, + db_attr: str, + pk: str, + column: str, + dry_run: bool, + policy: ReencryptPolicy = ALGORITHM_POLICY, +) -> LocationReport: + """Re-encrypt every encrypted field of one settings table's rows.""" + report: Final = LocationReport(location=location) + table: Final = getattr(prisma_client.db, db_attr, None) + if table is None: + return report + try: + rows: Final = await table.find_many() + except Exception as e: # noqa: BLE001 # any driver failure means this store is unknown, not clean + verbose_proxy_logger.warning("migrate: %s could not be read: %s", location, str(e)) + report.unreadable = True return report - settings = record.sso_settings - if isinstance(settings, str): - settings = json.loads(settings) - if not isinstance(settings, dict): - return report - - new_settings: Final = dict(settings) - changed = False - for fld, v in settings.items(): - if not isinstance(v, str) or v == "": + for row in rows or []: + settings = _parse_settings_column(getattr(row, column, None)) + if settings is None: continue - report.scanned += 1 - cls = classify_value(v, key=fld) - if cls == "migrated": - report.already_v2 += 1 - continue - if cls == "legacy": - if dry_run: - # Residual: would migrate, but a dry run writes nothing, so it - # stays legacy for the attestation (never counted as migrated). - report.legacy += 1 - continue - new_v = reencrypt_value(v, key=fld) - if new_v != v: - new_settings[fld] = new_v - report.migrated += 1 - changed = True - else: - # Defensive: a legacy value that did not re-encrypt is still - # residual, not migrated. - report.legacy += 1 - else: - report.plaintext += 1 - - if changed and not dry_run: - await prisma_client.db.litellm_ssoconfig.update( - where={"id": "sso_config"}, - data={"sso_settings": json.dumps(new_settings)}, + rewritten, row_report = _reencrypt_settings_fields( + settings, _encrypted_string_fields(settings), dry_run, policy, location ) + report.absorb(row_report) + if rewritten != settings: + await table.update(where={pk: getattr(row, pk)}, data={column: json.dumps(rewritten)}) return report +async def _migrate_settings_tables( + prisma_client: object, + dry_run: bool, + policy: ReencryptPolicy = ALGORITHM_POLICY, +) -> AsyncIterator[LocationReport]: + """One report per settings table, walked in order (a table at a time, never all in memory).""" + for location, db_attr, pk, column in _SETTINGS_ROW_SPECS: + yield await _migrate_settings_rows(prisma_client, location, db_attr, pk, column, dry_run, policy=policy) + + async def _migrate_callback_vars_table( prisma_client: object, table_name: Literal["team", "verification_token"], dry_run: bool, + policy: ReencryptPolicy = ALGORITHM_POLICY, ) -> LocationReport: """Migrate callback-var credentials on the team or verification-token table. - Covers both shapes the ``decrypt_callback_vars`` / ``encrypt_callback_vars`` - transforms understand: ``metadata.logging[*].callback_vars.`` and - the top-level ``metadata.callback_settings.callback_vars.``. Reuses - those proven transforms (selective, prefix-marked; legacy plaintext is left - alone until re-encrypted). + Covers both shapes ``transform_callback_vars`` understands: + ``metadata.logging[*].callback_vars.`` and the top-level + ``metadata.callback_settings.callback_vars.``. Rewrites are + per value and prefix-marked; legacy plaintext is left alone. """ - from litellm.proxy.common_utils.callback_utils import ( - decrypt_callback_vars, - encrypt_callback_vars, - ) + from litellm.proxy.common_utils.callback_utils import transform_callback_vars report: Final = LocationReport(location=f"{table_name}.callback_vars") @@ -347,7 +512,7 @@ async def _migrate_callback_vars_table( for cvs in _iter_callback_var_dicts(metadata): for v in cvs.values(): report.scanned += 1 - cls = _classify_callback_value(v) + cls = _classify_callback_value(v, policy=policy) if cls == "migrated": report.already_v2 += 1 elif cls == "legacy": @@ -363,10 +528,15 @@ async def _migrate_callback_vars_table( report.legacy += row_legacy continue - # Real run: re-encrypt the legacy ciphertext to AES via the proven - # selective transforms and persist. Never drop a row on failure. + # Real run: rewrite value by value, so a row that mixes a retired-key + # value with an already-current one only touches the former. A whole-row + # decrypt/re-encrypt would push every value through the configured write + # algorithm and downgrade active AES-GCM ciphertext during a key-only + # rotation. Never drop a row on failure. try: - re_encrypted = encrypt_callback_vars(decrypt_callback_vars(metadata)) + re_encrypted = transform_callback_vars( + metadata, lambda k, v: _reencrypt_callback_value(v, key=k, policy=policy) + ) except Exception as e: # pragma: no cover - defensive; never drop a row verbose_proxy_logger.warning( "Skipping %s row %s callback_vars (transform failed): %s", @@ -388,7 +558,7 @@ async def _migrate_callback_vars_table( def _iter_callback_var_dicts(metadata: dict[str, object]): """Yield each ``callback_vars`` dict in a metadata structure. - Mirrors ``_transform_callback_vars``: credentials live both under + Mirrors ``transform_callback_vars``: credentials live both under ``logging[*].callback_vars`` and under the top-level ``callback_settings.callback_vars``. Counting only the former would let the walker report success while leaving ``callback_settings`` secrets in legacy @@ -406,7 +576,29 @@ def _iter_callback_var_dicts(metadata: dict[str, object]): yield cvs -def _classify_callback_value(value: object) -> ValueClass: +def _callback_marker() -> str: + """The marker that fronts an encrypted callback var (``litellm_enc::``).""" + from litellm.proxy.common_utils.callback_utils import ( + _CALLBACK_VAR_ENCRYPTED_PREFIX, + ) + + return _CALLBACK_VAR_ENCRYPTED_PREFIX + + +def _reencrypt_callback_value(value: object, key: str, policy: ReencryptPolicy) -> object: + """Re-encrypt one stored callback-var value, keeping the ``litellm_enc::`` marker. + + Only marked ciphertext is rewritten: plaintext callback vars are the write + path's business, and this pass reports them as ``plaintext`` rather than + silently encrypting them on rows that happen to also hold legacy ciphertext. + """ + marker: Final = _callback_marker() + if not isinstance(value, str) or not value.startswith(marker): + return value + return marker + reencrypt_string(value.removeprefix(marker), key=key, policy=policy) + + +def _classify_callback_value(value: object, policy: ReencryptPolicy = ALGORITHM_POLICY) -> ValueClass: """Classify one stored callback-var value, independent of the AES gate. Encrypted callback vars carry the ``litellm_enc::`` marker in front of the @@ -416,15 +608,9 @@ def _classify_callback_value(value: object) -> ValueClass: re-encrypt delta is what makes the ``check_encryption`` attestation correct even when run with the AES write gate off. """ - from litellm.proxy.common_utils.callback_utils import ( - _CALLBACK_VAR_ENCRYPTED_PREFIX, - ) - if not isinstance(value, str): return "not-a-string" - inner = value - inner = inner.removeprefix(_CALLBACK_VAR_ENCRYPTED_PREFIX) - return classify_value(inner, key="callback") + return classify_value(value.removeprefix(_callback_marker()), key="callback", policy=policy) # --------------------------------------------------------------------------- @@ -444,6 +630,8 @@ _COVERED_TABLE_SPECS: Final = [ ("mcp_server", "litellm_mcpservertable", ("credentials", "env_vars"), ()), ("mcp_user_credentials", "litellm_mcpusercredentials", (), ("credential_b64",)), ("mcp_user_env_vars", "litellm_mcpuserenvvars", (), ("values_b64",)), + ("sso_identity_assertions", "litellm_ssoidentityassertion", (), ("assertion_b64",)), + ("mcp_oauth_clients", "litellm_mcpserveroauthclient", ("credentials",), ()), ] @@ -465,7 +653,7 @@ def _iter_encrypted_strings(obj: object): stack.extend(cur) -def _classify_into_report(report: LocationReport, value: str) -> None: +def _classify_into_report(report: LocationReport, value: str, policy: ReencryptPolicy = ALGORITHM_POLICY) -> None: """Classify one stored string and bump the matching read-only counter. Only genuine nacl ciphertext lands in ``legacy``; non-secret strings (model @@ -473,7 +661,7 @@ def _classify_into_report(report: LocationReport, value: str) -> None: over-scanning a column is harmless to the residual count. """ report.scanned += 1 - cls: Final = classify_value(value, key="scan") + cls: Final = classify_value(value, key="scan", policy=policy) if cls == "migrated": report.already_v2 += 1 elif cls == "legacy": @@ -488,6 +676,7 @@ async def _scan_one_table( db_attr: str, json_columns: tuple, scalar_columns: tuple, + policy: ReencryptPolicy = ALGORITHM_POLICY, ) -> LocationReport: report: Final = LocationReport(location=location) table: Final = getattr(prisma_client.db, db_attr, None) @@ -495,8 +684,9 @@ async def _scan_one_table( return report try: rows: Final = await table.find_many() - except Exception as e: # pragma: no cover - table absent / not migrated - verbose_proxy_logger.debug("scan: %s unavailable: %s", location, str(e)) + except Exception as e: # noqa: BLE001 # any driver failure means this store is unknown, not clean + verbose_proxy_logger.warning("scan: %s could not be read: %s", location, str(e)) + report.unreadable = True return report for row in rows or []: for col in json_columns: @@ -509,21 +699,22 @@ async def _scan_one_table( except (ValueError, TypeError): pass for s in _iter_encrypted_strings(raw): - _classify_into_report(report, s) + _classify_into_report(report, s, policy=policy) for col in scalar_columns: v = getattr(row, col, None) if isinstance(v, str): - _classify_into_report(report, v) + _classify_into_report(report, v, policy=policy) return report -async def _scan_config_env_vars(prisma_client: object) -> LocationReport: +async def _scan_config_env_vars(prisma_client: object, policy: ReencryptPolicy = ALGORITHM_POLICY) -> LocationReport: """Scan the ``environment_variables`` config row (``param_value`` dict).""" report: Final = LocationReport(location="config_environment_variables") try: record: Final = await prisma_client.db.litellm_config.find_unique(where={"param_name": "environment_variables"}) - except Exception as e: # pragma: no cover - defensive - verbose_proxy_logger.debug("scan: config env vars unavailable: %s", str(e)) + except Exception as e: # noqa: BLE001 # any driver failure means this store is unknown, not clean + verbose_proxy_logger.warning("scan: config env vars could not be read: %s", str(e)) + report.unreadable = True return report if record is None or record.param_value is None: return report @@ -534,16 +725,19 @@ async def _scan_config_env_vars(prisma_client: object) -> LocationReport: except (ValueError, TypeError): value = {} for s in _iter_encrypted_strings(value): - _classify_into_report(report, s) + _classify_into_report(report, s, policy=policy) return report -async def _scan_covered_tables(prisma_client: object) -> list[LocationReport]: +async def _scan_covered_tables( + prisma_client: object, + policy: ReencryptPolicy = ALGORITHM_POLICY, +) -> list[LocationReport]: """Read-only classification of every rotation-covered table. No writes.""" reports: Final[list[LocationReport]] = [] for location, db_attr, json_cols, scalar_cols in _COVERED_TABLE_SPECS: - reports.append(await _scan_one_table(prisma_client, location, db_attr, json_cols, scalar_cols)) - reports.append(await _scan_config_env_vars(prisma_client)) + reports.append(await _scan_one_table(prisma_client, location, db_attr, json_cols, scalar_cols, policy=policy)) + reports.append(await _scan_config_env_vars(prisma_client, policy=policy)) return reports @@ -556,7 +750,11 @@ _VANTAGE_SENSITIVE: Final = ["api_key", "integration_token"] _CLOUDZERO_SENSITIVE: Final = ["api_key"] -async def _migrate_covered_tables(prisma_client: object, user_api_key_dict: object) -> list[LocationReport]: +async def _migrate_covered_tables( + prisma_client: object, + user_api_key_dict: object, + policy: ReencryptPolicy = ALGORITHM_POLICY, +) -> list[LocationReport]: """Re-encrypt the tables already covered by ``_rotate_master_key`` (model table, credentials, MCP credential/env tables, config environment_variables) by running that orchestrator in *same-key* mode. With the AES gate on, the @@ -571,7 +769,7 @@ async def _migrate_covered_tables(prisma_client: object, user_api_key_dict: obje _rotate_master_key, ) - pre: Final = {r.location: r for r in await _scan_covered_tables(prisma_client)} + pre: Final = MappingProxyType({r.location: r for r in await _scan_covered_tables(prisma_client, policy=policy)}) current_key: Final = _get_salt_key() if current_key is None: @@ -582,10 +780,10 @@ async def _migrate_covered_tables(prisma_client: object, user_api_key_dict: obje prisma_client=cast("PrismaClient", prisma_client), user_api_key_dict=cast("UserAPIKeyAuth", user_api_key_dict), current_master_key=current_key, - new_master_key=current_key, # same key, algorithm-only switch + new_master_key=current_key, # writes land under the active salt key ) - post: Final = await _scan_covered_tables(prisma_client) + post: Final = await _scan_covered_tables(prisma_client, policy=policy) for post_report in post: pre_report = pre.get(post_report.location) pre_legacy = pre_report.legacy if pre_report else 0 @@ -599,18 +797,25 @@ async def migrate_encryption( prisma_client: object, user_api_key_dict: object, dry_run: bool = False, + policy: ReencryptPolicy = ALGORITHM_POLICY, ) -> MigrationReport: - """Run the full at-rest re-encryption migration. + """Run the full at-rest re-encryption pass for ``policy``. - Requires ``general_settings.encryption_algorithm == 'aes-256-gcm'`` so writes - are produced in the AES format. Idempotent and resumable: re-running skips + Under :data:`ALGORITHM_POLICY` this requires + ``general_settings.encryption_algorithm == 'aes-256-gcm'`` so writes are + produced in the AES format. Idempotent and resumable: re-running skips already-migrated values and finishes any partial run. A ``dry_run`` performs no writes: the covered tables are scanned read-only (so their residual legacy still counts toward the attestation) and the net-new walkers run in dry-run mode. """ - _assert_aes_gate_enabled() + if policy is ALGORITHM_POLICY: + _assert_aes_gate_enabled() + else: + _assert_previous_salt_keys_configured() + if not dry_run: + await _assert_no_algorithm_downgrade(prisma_client) report: Final = MigrationReport() @@ -618,43 +823,94 @@ async def migrate_encryption( # delegate to the rotation path (with bracketing scans for counts); on a dry # run only classify them read-only. if dry_run: - for covered in await _scan_covered_tables(prisma_client): + for covered in await _scan_covered_tables(prisma_client, policy=policy): report.add(covered) else: - for covered in await _migrate_covered_tables(prisma_client, user_api_key_dict): + for covered in await _migrate_covered_tables(prisma_client, user_api_key_dict, policy=policy): report.add(covered) # Net-new walkers (items 3, 4, 11, 12, 13). - report.add(await _migrate_callback_vars_table(prisma_client, "team", dry_run)) - report.add(await _migrate_callback_vars_table(prisma_client, "verification_token", dry_run)) - report.add(await _migrate_config_settings_row(prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run)) - report.add(await _migrate_config_settings_row(prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run)) - report.add(await _migrate_sso_config(prisma_client, dry_run)) + report.add(await _migrate_callback_vars_table(prisma_client, "team", dry_run, policy=policy)) + report.add(await _migrate_callback_vars_table(prisma_client, "verification_token", dry_run, policy=policy)) + report.add( + await _migrate_config_settings_row( + prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run, policy=policy + ) + ) + report.add( + await _migrate_config_settings_row( + prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run, policy=policy + ) + ) + async for settings_report in _migrate_settings_tables(prisma_client, dry_run, policy=policy): + report.add(settings_report) return report -async def check_encryption(prisma_client: object) -> MigrationReport: +async def check_encryption(prisma_client: object, policy: ReencryptPolicy = ALGORITHM_POLICY) -> MigrationReport: """Read-only residual scan across **every** at-rest location. No writes. Covers both the rotation-managed tables (model / credentials / MCP credential and env-var tables / config ``environment_variables``) and the net-new walker locations (team and verification-token ``callback_vars``, vantage / cloudzero - config rows, SSO config). Reports how many values are still ``legacy``; + config rows, SSO / cache / config-override settings). Reports how many values + are still ``legacy``; ``residual_legacy == 0`` across this full scan is the compliance attestation. """ report: Final = MigrationReport() # Rotation-covered tables (read-only classification). - for covered in await _scan_covered_tables(prisma_client): + for covered in await _scan_covered_tables(prisma_client, policy=policy): report.add(covered) # Net-new walker locations, in dry-run (read-only) mode. - report.add(await _migrate_callback_vars_table(prisma_client, "team", dry_run=True)) - report.add(await _migrate_callback_vars_table(prisma_client, "verification_token", dry_run=True)) - report.add(await _migrate_config_settings_row(prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run=True)) + report.add(await _migrate_callback_vars_table(prisma_client, "team", dry_run=True, policy=policy)) + report.add(await _migrate_callback_vars_table(prisma_client, "verification_token", dry_run=True, policy=policy)) report.add( - await _migrate_config_settings_row(prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run=True) + await _migrate_config_settings_row( + prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run=True, policy=policy + ) ) - report.add(await _migrate_sso_config(prisma_client, dry_run=True)) + report.add( + await _migrate_config_settings_row( + prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run=True, policy=policy + ) + ) + async for settings_report in _migrate_settings_tables(prisma_client, dry_run=True, policy=policy): + report.add(settings_report) return report + + +async def _assert_no_algorithm_downgrade(prisma_client: object) -> None: + """Fail fast when a salt-key rotation would rewrite AES values in legacy format. + + The rotation-covered tables are re-encrypted by ``_rotate_master_key``, which + writes through ``general_settings.encryption_algorithm`` and so cannot + preserve a value's own algorithm. Downgrading AES-256-GCM ciphertext is not + something a key-only rotation may do silently, so refuse the run instead. + """ + if _aes_gate_enabled(): + return + aes_values: Final = sum(r.already_v2 for r in await _scan_covered_tables(prisma_client, policy=ALGORITHM_POLICY)) + if aes_values: + raise RuntimeError( + f"Salt-key rotation would rewrite {aes_values} AES-256-GCM value(s) in the legacy " + f"format, because general_settings.encryption_algorithm is not '{_ALGO_AES_GCM}'. " + "Set it to that, restart the proxy, then re-run the rotation." + ) + + +def _assert_previous_salt_keys_configured() -> None: + """Fail fast when no retired salt key is configured. + + Without ``LITELLM_SALT_KEY_PREVIOUS``, values written under the retired key + cannot be read at all, so a rotation pass would classify them as plaintext + and leave them behind while reporting a clean run. + """ + if not get_previous_salt_keys(): + raise RuntimeError( + "Salt key rotation requires LITELLM_SALT_KEY_PREVIOUS to list the retired " + "salt key(s), with LITELLM_SALT_KEY set to the new one. Restart the proxy " + "with both set, then re-run the rotation." + ) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 97999cbb6d7..237dbb3c396 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -86,6 +86,7 @@ from litellm.proxy.management_endpoints.common_utils import ( validate_budget_duration, validate_finite_spend, ) +from litellm.proxy.management_endpoints.credential_migration import ReencryptMode from litellm.proxy.management_endpoints.model_management_endpoints import ( _add_model_to_db, ) @@ -4568,17 +4569,24 @@ async def migrate_encryption_endpoint( False, description="If true, scan and report without writing any changes.", ), + mode: ReencryptMode = "algorithm", ): """ - Re-encrypt all at-rest credentials into the AES-256-GCM (``v2:gcm:``) format. + Re-encrypt all at-rest credentials, either into the AES-256-GCM (``v2:gcm:``) + format (``mode=algorithm``, the default) or under the active salt key + (``mode=salt-key``). - Admin only. Requires ``general_settings.encryption_algorithm: aes-256-gcm``. - Idempotent and resumable — re-running skips already-migrated values. Pass - ``dry_run=true`` for a non-mutating scan (equivalent to ``--check``). + Admin only. ``mode=algorithm`` requires + ``general_settings.encryption_algorithm: aes-256-gcm``; ``mode=salt-key`` + requires the retired key(s) in ``LITELLM_SALT_KEY_PREVIOUS`` and the new one + in ``LITELLM_SALT_KEY``. Idempotent and resumable, so re-running skips values + that are already current. Pass ``dry_run=true`` for a non-mutating scan + (equivalent to ``--check``). """ from litellm.proxy._types import CommonProxyErrors from litellm.proxy.management_endpoints.credential_migration import ( migrate_encryption, + policy_for_mode, ) from litellm.proxy.proxy_server import prisma_client @@ -4589,12 +4597,17 @@ async def migrate_encryption_endpoint( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - report: Final = await migrate_encryption( - prisma_client=prisma_client, - user_api_key_dict=user_api_key_dict, - dry_run=dry_run, - ) - return {"status": "success", "dry_run": dry_run, "report": report.as_dict()} + try: + report: Final = await migrate_encryption( + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + dry_run=dry_run, + policy=policy_for_mode(mode), + ) + except RuntimeError as e: + # Unmet precondition (AES gate off, or no retired salt key configured). + raise HTTPException(status_code=400, detail={"error": str(e)}) from e + return {"status": "success", "dry_run": dry_run, "mode": mode, "report": report.as_dict()} @router.get( @@ -4604,15 +4617,22 @@ async def migrate_encryption_endpoint( ) async def check_encryption_endpoint( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + mode: ReencryptMode = "algorithm", ): """ - Read-only residual scan for compliance attestation. Reports how many at-rest - values are still in the legacy format. ``residual_legacy == 0`` attests no - legacy ciphertext remains. Admin only; performs no writes. + Read-only residual scan for compliance attestation. Admin only; performs no + writes. With ``mode=algorithm`` it reports how many at-rest values are still + in the legacy format; with ``mode=salt-key`` it reports how many still + decrypt only under a retired salt key. ``residual_legacy == 0`` with an empty + ``unreadable_locations`` is the attestation for the selected mode, and for + ``salt-key`` it means ``LITELLM_SALT_KEY_PREVIOUS`` can be dropped. A store + the scan could not read is named in ``unreadable_locations``, because zero + counts from a store nobody could open mean unknown rather than clean. """ from litellm.proxy._types import CommonProxyErrors from litellm.proxy.management_endpoints.credential_migration import ( check_encryption, + policy_for_mode, ) from litellm.proxy.proxy_server import prisma_client @@ -4623,8 +4643,8 @@ async def check_encryption_endpoint( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - report: Final = await check_encryption(prisma_client=prisma_client) - return {"status": "success", "report": report.as_dict()} + report: Final = await check_encryption(prisma_client=prisma_client, policy=policy_for_mode(mode)) + return {"status": "success", "mode": mode, "report": report.as_dict()} async def get_new_token(data: RegenerateKeyRequest | None) -> str: diff --git a/tests/test_litellm/proxy/client/cli/test_encryption_commands.py b/tests/test_litellm/proxy/client/cli/test_encryption_commands.py index 43e53cf5be2..4935bdcfb75 100644 --- a/tests/test_litellm/proxy/client/cli/test_encryption_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_encryption_commands.py @@ -51,14 +51,41 @@ def test_migrate_default_posts_without_dry_run(runner): assert result.exit_code == 0, result.output assert _FakeHTTPClient.last["method"] == "POST" assert _FakeHTTPClient.last["path"] == "/credentials/migrate-encryption" - assert _FakeHTTPClient.last["params"] is None + assert dict(_FakeHTTPClient.last["params"]) == { + "mode": "algorithm", + "dry_run": "false", + } def test_migrate_dry_run_sets_param(runner): result = runner.invoke(cli_main.cli, ["encryption", "migrate", "--dry-run"]) assert result.exit_code == 0, result.output assert _FakeHTTPClient.last["method"] == "POST" - assert _FakeHTTPClient.last["params"] == {"dry_run": "true"} + assert dict(_FakeHTTPClient.last["params"]) == { + "mode": "algorithm", + "dry_run": "true", + } + + +def test_migrate_salt_key_mode_is_forwarded(runner): + result = runner.invoke( + cli_main.cli, ["encryption", "migrate", "--mode", "salt-key"] + ) + assert result.exit_code == 0, result.output + assert _FakeHTTPClient.last["method"] == "POST" + assert dict(_FakeHTTPClient.last["params"]) == { + "mode": "salt-key", + "dry_run": "false", + } + + +def test_migrate_salt_key_check_is_forwarded(runner): + result = runner.invoke( + cli_main.cli, ["encryption", "migrate", "--mode", "salt-key", "--check"] + ) + assert result.exit_code == 0, result.output + assert _FakeHTTPClient.last["method"] == "GET" + assert dict(_FakeHTTPClient.last["params"]) == {"mode": "salt-key"} def test_migrate_reports_residual_legacy(runner): diff --git a/tests/test_litellm/proxy/common_utils/test_encrypt_decrypt_utils.py b/tests/test_litellm/proxy/common_utils/test_encrypt_decrypt_utils.py index 08cf1e45812..d7d52b98e3f 100644 --- a/tests/test_litellm/proxy/common_utils/test_encrypt_decrypt_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_encrypt_decrypt_utils.py @@ -13,6 +13,8 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( _V2_GCM_PREFIX, decrypt_value_helper, encrypt_value_helper, + get_decryption_keys, + get_previous_salt_keys, ) @@ -185,3 +187,60 @@ def test_decrypt_failure_debug_log_omits_raw_value(monkeypatch): "the failing key should still be named in the breadcrumb" ) assert result == secret + + +# ---------------------------- salt key rotation ---------------------------- + + +def test_retired_salt_key_value_still_decrypts(monkeypatch): + """The rotation guarantee: old ciphertext is readable under the new key. + + Values written before the rotation stay readable as long as the retired key + is listed in LITELLM_SALT_KEY_PREVIOUS, which is what makes rotating the salt + key safe without regenerating virtual keys. + """ + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-old-salt") + old_ct = encrypt_value_helper("provider-api-key") + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-new-salt") + assert decrypt_value_helper(old_ct, key="t") is None + + monkeypatch.setenv("LITELLM_SALT_KEY_PREVIOUS", "sk-old-salt") + assert decrypt_value_helper(old_ct, key="t") == "provider-api-key" + + +def test_retired_salt_key_works_for_aes_values(monkeypatch): + """Fallback is format-agnostic: it also covers v2:gcm: ciphertext.""" + _use_aes(monkeypatch) + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-old-salt") + old_ct = encrypt_value_helper("aes-secret") + assert old_ct.startswith(_V2_GCM_PREFIX) + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-new-salt") + monkeypatch.setenv("LITELLM_SALT_KEY_PREVIOUS", "sk-unrelated,sk-old-salt") + assert decrypt_value_helper(old_ct, key="t") == "aes-secret" + + +def test_new_writes_use_the_active_salt_key_only(monkeypatch): + """A retired key must never be used for writes, or rotation never converges.""" + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-new-salt") + monkeypatch.setenv("LITELLM_SALT_KEY_PREVIOUS", "sk-old-salt") + + ct = encrypt_value_helper("fresh-secret") + + monkeypatch.delenv("LITELLM_SALT_KEY_PREVIOUS") + assert decrypt_value_helper(ct, key="t") == "fresh-secret" + + +def test_decryption_keys_order_and_dedup(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-active") + monkeypatch.setenv( + "LITELLM_SALT_KEY_PREVIOUS", " sk-active , sk-a ,, sk-b , sk-a " + ) + + assert get_decryption_keys() == ("sk-active", "sk-a", "sk-b") + + +def test_no_previous_salt_keys_configured(monkeypatch): + monkeypatch.delenv("LITELLM_SALT_KEY_PREVIOUS", raising=False) + assert get_previous_salt_keys() == () diff --git a/tests/test_litellm/proxy/management_endpoints/test_credential_migration.py b/tests/test_litellm/proxy/management_endpoints/test_credential_migration.py index 81226981089..a62b2a73757 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_credential_migration.py +++ b/tests/test_litellm/proxy/management_endpoints/test_credential_migration.py @@ -15,6 +15,7 @@ import pytest from litellm.proxy import proxy_server from litellm.proxy.common_utils.encrypt_decrypt_utils import ( _V2_GCM_PREFIX, + decrypt_value_helper, encrypt_value_helper, ) from litellm.proxy.management_endpoints import credential_migration as cm @@ -40,13 +41,15 @@ def _enable_aes(monkeypatch): def _empty_covered_tables(client): - """Wire every rotation-covered table on `client` to return no rows. + """Wire every table the scanner walks on `client` to return no rows. Lets a `check_encryption` / scanner test isolate the location under test - without the other covered tables raising on an unconfigured mock. + without the other tables raising on an unconfigured mock. """ for _, db_attr, _, _ in cm._COVERED_TABLE_SPECS: getattr(client.db, db_attr).find_many = AsyncMock(return_value=[]) + for _, db_attr, _, _ in cm._SETTINGS_ROW_SPECS: + getattr(client.db, db_attr).find_many = AsyncMock(return_value=[]) # --------------------------- pure engine --------------------------- @@ -129,6 +132,13 @@ async def test_migrate_requires_aes_gate(salt_key, monkeypatch): # --------------------------- config-row walker --------------------------- +async def _migrate_sso(client, dry_run, policy=cm.ALGORITHM_POLICY): + """Run the settings-row walker over the SSO config table.""" + return await cm._migrate_settings_rows( + client, "sso_config", "litellm_ssoconfig", "id", "sso_settings", dry_run, policy=policy + ) + + def _config_prisma(record): """Build an AsyncMock prisma client whose litellm_config returns `record`.""" client = MagicMock() @@ -215,12 +225,12 @@ async def test_sso_walker_real_run_migrates_and_clears_residual(salt_key, monkey """SSO real run: a migrated field is counted as migrated, not residual legacy.""" legacy = _legacy_ct("client-secret", monkeypatch) _enable_aes(monkeypatch) - record = SimpleNamespace(sso_settings={"client_secret": legacy, "client_id": "id"}) + record = SimpleNamespace(id="sso_config", sso_settings={"client_secret": legacy, "client_id": "id"}) client = MagicMock() - client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=record) + client.db.litellm_ssoconfig.find_many = AsyncMock(return_value=[record]) client.db.litellm_ssoconfig.update = AsyncMock() - report = await cm._migrate_sso_config(client, dry_run=False) + report = await _migrate_sso(client, dry_run=False) assert report.migrated == 1 assert report.legacy == 0 # migrated -> no longer residual @@ -232,12 +242,12 @@ async def test_sso_walker_dry_run_reports_residual_not_migrated(salt_key, monkey """SSO dry run: residual legacy only; migrated stays 0 (never contradictory).""" legacy = _legacy_ct("client-secret", monkeypatch) _enable_aes(monkeypatch) - record = SimpleNamespace(sso_settings={"client_secret": legacy}) + record = SimpleNamespace(id="sso_config", sso_settings={"client_secret": legacy}) client = MagicMock() - client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=record) + client.db.litellm_ssoconfig.find_many = AsyncMock(return_value=[record]) client.db.litellm_ssoconfig.update = AsyncMock() - report = await cm._migrate_sso_config(client, dry_run=True) + report = await _migrate_sso(client, dry_run=True) assert report.legacy == 1 assert report.migrated == 0 @@ -514,3 +524,441 @@ async def test_migrate_covered_tables_reports_real_counts(salt_key, monkeypatch) assert by_loc["model_table"].migrated == 1 # was legacy pre, v2 post assert by_loc["model_table"].legacy == 0 # residual zero after rotation assert by_loc["model_table"].already_v2 == 1 + + +# --------------------------- salt-key rotation --------------------------- + + +@pytest.fixture +def rotated_salt_key(monkeypatch): + """Simulate a completed key swap: new key active, old one retired.""" + monkeypatch.setattr(proxy_server, "general_settings", {}) + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-old") + old_ct = encrypt_value_helper("provider-api-key") + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-new") + monkeypatch.setenv("LITELLM_SALT_KEY_PREVIOUS", "sk-salt-old") + return old_ct + + +def test_salt_key_policy_classifies_retired_ciphertext_as_legacy(rotated_salt_key): + assert cm.classify_value(rotated_salt_key, policy=cm.SALT_KEY_POLICY) == "legacy" + # The algorithm pass has nothing to do here: the format is already correct. + assert cm.classify_value(rotated_salt_key, policy=cm.ALGORITHM_POLICY) == "legacy" + + +def test_salt_key_policy_reencrypts_under_active_key(rotated_salt_key, monkeypatch): + out = cm.reencrypt_value(rotated_salt_key, policy=cm.SALT_KEY_POLICY) + + assert out != rotated_salt_key + # Readable with the retired key removed from the environment: the whole point. + monkeypatch.delenv("LITELLM_SALT_KEY_PREVIOUS") + assert decrypt_value_helper(out, key="t") == "provider-api-key" + assert cm.classify_value(out, policy=cm.SALT_KEY_POLICY) == "migrated" + + +def test_salt_key_policy_preserves_the_legacy_format(rotated_salt_key, monkeypatch): + """A salt-only rotation must not silently switch algorithms. + + The configured write algorithm is AES here, so a rewrite that went through + ``encrypt_value_helper`` would upgrade the format as a side effect of a + key-only rotation. The two axes are independent, so the format must survive. + """ + _enable_aes(monkeypatch) + + out = cm.reencrypt_value(rotated_salt_key, policy=cm.SALT_KEY_POLICY) + + assert out != rotated_salt_key + assert not out.startswith(_V2_GCM_PREFIX) + + +def test_salt_key_policy_preserves_the_aes_format(monkeypatch): + """Regression: a key-only rotation must never downgrade AES-GCM ciphertext. + + A value written in AES under the retired key, rotated while the configured + write algorithm sits at its legacy default, used to come back as nacl: the + rewrite went through ``encrypt_value_helper``, which reads that setting. + """ + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-old") + _enable_aes(monkeypatch) + old_aes = encrypt_value_helper("provider-api-key") + assert old_aes.startswith(_V2_GCM_PREFIX) + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-new") + monkeypatch.setenv("LITELLM_SALT_KEY_PREVIOUS", "sk-salt-old") + monkeypatch.setattr(proxy_server, "general_settings", {}) # legacy write algo + + out = cm.reencrypt_value(old_aes, policy=cm.SALT_KEY_POLICY) + + assert out != old_aes + assert out.startswith(_V2_GCM_PREFIX) + monkeypatch.delenv("LITELLM_SALT_KEY_PREVIOUS") + assert decrypt_value_helper(out, key="t") == "provider-api-key" + + +@pytest.mark.asyncio +async def test_salt_key_rotation_preserves_each_algorithm_in_a_mixed_row(monkeypatch): + """Regression: a callback row mixing a retired-key value with an active-key + AES value must have only the former rewritten. + + The walker used to round-trip the whole row through + ``encrypt_callback_vars(decrypt_callback_vars(...))``, which pushed every + value through the configured write algorithm and downgraded the AES one. + """ + from litellm.proxy.common_utils.callback_utils import ( + _CALLBACK_VAR_ENCRYPTED_PREFIX, + ) + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-old") + monkeypatch.setattr(proxy_server, "general_settings", {}) + retired = _CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper("sa-secret") + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-new") + _enable_aes(monkeypatch) + active_aes = _CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper("lf-secret") + + monkeypatch.setenv("LITELLM_SALT_KEY_PREVIOUS", "sk-salt-old") + monkeypatch.setattr(proxy_server, "general_settings", {}) # legacy write algo + + row = SimpleNamespace( + team_id="t1", + metadata={ + "logging": [ + { + "callback_vars": { + "gcs_path_service_account": retired, + "langfuse_secret_key": active_aes, + } + } + ] + }, + ) + client = MagicMock() + client.db.litellm_teamtable.find_many = AsyncMock(return_value=[row]) + client.db.litellm_teamtable.update = AsyncMock() + + report = await cm._migrate_callback_vars_table( + client, "team", dry_run=False, policy=cm.SALT_KEY_POLICY + ) + + assert report.migrated == 1 + assert report.already_v2 == 1 + written = json.loads( + client.db.litellm_teamtable.update.call_args.kwargs["data"]["metadata"] + ) + cvs = written["logging"][0]["callback_vars"] + + # Already under the active key: left byte-for-byte alone, so no downgrade. + assert cvs["langfuse_secret_key"] == active_aes + + # Under the retired key: moved to the active key, keeping its own format. + rotated = cvs["gcs_path_service_account"] + assert rotated != retired + inner = rotated.removeprefix(_CALLBACK_VAR_ENCRYPTED_PREFIX) + assert not inner.startswith(_V2_GCM_PREFIX) + monkeypatch.delenv("LITELLM_SALT_KEY_PREVIOUS") + assert decrypt_value_helper(inner, key="t") == "sa-secret" + + +@pytest.mark.asyncio +async def test_check_counts_sso_identity_assertion_residual(rotated_salt_key): + """Regression: the SSO identity assertion store is part of the attestation. + + ``_rotate_master_key`` rotates it but logs and carries on when that store + fails, so a scan that skipped the table could report ``residual_legacy == 0`` + while an assertion still depended on the retired key. Dropping + ``LITELLM_SALT_KEY_PREVIOUS`` then reads that assertion as absent. + """ + client = MagicMock() + _empty_covered_tables(client) + client.db.litellm_teamtable.find_many = AsyncMock(return_value=[]) + client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) + client.db.litellm_config.find_unique = AsyncMock(return_value=None) + client.db.litellm_ssoidentityassertion.find_many = AsyncMock( + return_value=[SimpleNamespace(user_id="u1", assertion_b64=rotated_salt_key)] + ) + + report = await cm.check_encryption(client, policy=cm.SALT_KEY_POLICY) + + assert report.as_dict()["locations"]["sso_identity_assertions"]["legacy"] == 1 + assert report.residual_legacy == 1 + + +@pytest.mark.asyncio +async def test_check_counts_mcp_oauth_client_residual(rotated_salt_key): + """Regression: the MCP OAuth client table is part of the attestation. + + The shared rotation path rewrites those DCR client credentials, so leaving + the table out of the scan let the check report ``residual_legacy == 0`` + while a client secret still decrypted only under the retired key. + """ + client = MagicMock() + _empty_covered_tables(client) + client.db.litellm_teamtable.find_many = AsyncMock(return_value=[]) + client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + client.db.litellm_config.find_unique = AsyncMock(return_value=None) + client.db.litellm_mcpserveroauthclient.find_many = AsyncMock( + return_value=[ + SimpleNamespace( + server_id="srv-1", + credentials={"client_id": "public-id", "client_secret": rotated_salt_key}, + ) + ] + ) + + report = await cm.check_encryption(client, policy=cm.SALT_KEY_POLICY) + + assert report.as_dict()["locations"]["mcp_oauth_clients"]["legacy"] == 1 + assert report.residual_legacy == 1 + + +@pytest.mark.asyncio +async def test_salt_key_migration_refuses_to_downgrade_oauth_client_credentials(monkeypatch): + """Regression: the downgrade preflight covers the MCP OAuth client table. + + Its credentials go through the same rotation path, which writes via the + configured algorithm, so an AES client secret would be silently rewritten in + the legacy format by a key-only rotation. + """ + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-old") + _enable_aes(monkeypatch) + old_aes = encrypt_value_helper("dcr-client-secret") + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-new") + monkeypatch.setenv("LITELLM_SALT_KEY_PREVIOUS", "sk-salt-old") + monkeypatch.setattr(proxy_server, "general_settings", {}) # legacy write algo + + client = MagicMock() + _empty_covered_tables(client) + client.db.litellm_mcpserveroauthclient.find_many = AsyncMock( + return_value=[SimpleNamespace(server_id="srv-1", credentials={"client_secret": old_aes})] + ) + client.db.litellm_config.find_unique = AsyncMock(return_value=None) + + with pytest.raises(RuntimeError, match="aes-256-gcm"): + await cm.migrate_encryption( + prisma_client=client, + user_api_key_dict=MagicMock(), + policy=cm.SALT_KEY_POLICY, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "location,db_attr,pk,pk_value,column,field", + [ + ("cache_config", "litellm_cacheconfig", "id", "cache_config", "cache_settings", "redis_password"), + ( + "config_overrides", + "litellm_configoverrides", + "config_type", + "hashicorp_vault", + "config_value", + "HCP_VAULT_TOKEN", + ), + ], +) +async def test_salt_key_pass_rotates_settings_rows( + rotated_salt_key, monkeypatch, location, db_attr, pk, pk_value, column, field +): + """Regression: the cache-config and config-override rows are rotated too. + + Both hold live secrets written through the same encryption path and neither + has a master-key rotation path, so skipping them would strand a Redis + password or a Vault token under the retired key while the check read clean. + """ + row = SimpleNamespace(**{pk: pk_value, column: {field: rotated_salt_key, "host": "localhost"}}) + client = MagicMock() + table = getattr(client.db, db_attr) + table.find_many = AsyncMock(return_value=[row]) + table.update = AsyncMock() + + report = await cm._migrate_settings_rows( + client, location, db_attr, pk, column, dry_run=False, policy=cm.SALT_KEY_POLICY + ) + + assert report.migrated == 1 + written = json.loads(table.update.call_args.kwargs["data"][column]) + assert written["host"] == "localhost" # non-ciphertext left alone + monkeypatch.delenv("LITELLM_SALT_KEY_PREVIOUS") + assert decrypt_value_helper(written[field], key="t") == "provider-api-key" + + +@pytest.mark.asyncio +async def test_check_counts_settings_row_residual(rotated_salt_key): + """A secret left under the retired key in either settings row blocks the attestation.""" + client = MagicMock() + _empty_covered_tables(client) + client.db.litellm_teamtable.find_many = AsyncMock(return_value=[]) + client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + client.db.litellm_config.find_unique = AsyncMock(return_value=None) + client.db.litellm_cacheconfig.find_many = AsyncMock( + return_value=[SimpleNamespace(id="cache_config", cache_settings={"redis_password": rotated_salt_key})] + ) + client.db.litellm_configoverrides.find_many = AsyncMock( + return_value=[ + SimpleNamespace(config_type="hashicorp_vault", config_value={"HCP_VAULT_TOKEN": rotated_salt_key}) + ] + ) + + report = await cm.check_encryption(client, policy=cm.SALT_KEY_POLICY) + locations = report.as_dict()["locations"] + + assert locations["cache_config"]["legacy"] == 1 + assert locations["config_overrides"]["legacy"] == 1 + assert report.residual_legacy == 2 + client.db.litellm_cacheconfig.update.assert_not_called() # read-only + + +@pytest.mark.asyncio +@pytest.mark.parametrize("db_attr,location", [("litellm_proxymodeltable", "model_table"), ("litellm_cacheconfig", "cache_config")]) +async def test_check_names_a_store_it_could_not_read(rotated_salt_key, db_attr, location): + """Regression: a store the scan cannot open is unknown, never clean. + + Both scanners swallow a driver failure and return zero of every counter, so + without this the check answered `residual_legacy: 0` for a table nobody + read, and the operator dropped the retired key on that word. + """ + client = MagicMock() + _empty_covered_tables(client) + client.db.litellm_teamtable.find_many = AsyncMock(return_value=[]) + client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + client.db.litellm_config.find_unique = AsyncMock(return_value=None) + getattr(client.db, db_attr).find_many = AsyncMock(side_effect=RuntimeError("connection refused")) + + report = await cm.check_encryption(client, policy=cm.SALT_KEY_POLICY) + + assert report.residual_legacy == 0 # nothing was read, so nothing counted + assert location in report.unreadable_locations + assert report.as_dict()["unreadable_locations"] == [location] + + +@pytest.mark.asyncio +async def test_check_reports_no_unreadable_stores_on_a_healthy_scan(rotated_salt_key): + """The unreadable list stays empty when every store answers, so it can only + ever mean a real read failure.""" + client = MagicMock() + _empty_covered_tables(client) + client.db.litellm_teamtable.find_many = AsyncMock(return_value=[]) + client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + client.db.litellm_config.find_unique = AsyncMock(return_value=None) + + report = await cm.check_encryption(client, policy=cm.SALT_KEY_POLICY) + + assert report.unreadable_locations == () + assert report.residual_legacy == 0 + + +@pytest.mark.asyncio +async def test_salt_key_migration_refuses_to_downgrade_covered_tables(monkeypatch): + """The rotation-covered tables are re-encrypted by ``_rotate_master_key``, + which writes through the configured algorithm and so cannot preserve a + value's own. Refuse the run rather than downgrade AES ciphertext silently. + """ + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-old") + _enable_aes(monkeypatch) + old_aes = encrypt_value_helper("model-secret") + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-new") + monkeypatch.setenv("LITELLM_SALT_KEY_PREVIOUS", "sk-salt-old") + monkeypatch.setattr(proxy_server, "general_settings", {}) # legacy write algo + + client = MagicMock() + _empty_covered_tables(client) + client.db.litellm_proxymodeltable.find_many = AsyncMock( + return_value=[SimpleNamespace(litellm_params={"api_key": old_aes})] + ) + client.db.litellm_config.find_unique = AsyncMock(return_value=None) + + with pytest.raises(RuntimeError, match="aes-256-gcm"): + await cm.migrate_encryption( + prisma_client=client, + user_api_key_dict=MagicMock(), + policy=cm.SALT_KEY_POLICY, + ) + + +def test_salt_key_policy_is_idempotent(rotated_salt_key): + once = cm.reencrypt_value(rotated_salt_key, policy=cm.SALT_KEY_POLICY) + assert cm.reencrypt_value(once, policy=cm.SALT_KEY_POLICY) == once + + +def test_salt_key_policy_does_not_require_the_aes_gate(rotated_salt_key): + """The AES gate guards the algorithm pass only, never the salt-key pass.""" + assert cm.reencrypt_value(rotated_salt_key, policy=cm.SALT_KEY_POLICY) != rotated_salt_key + + +@pytest.mark.asyncio +async def test_salt_key_migration_requires_previous_keys(monkeypatch): + """Without the retired key, old ciphertext reads as plaintext and is skipped, + so the pass would report a clean run while leaving values behind. + """ + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-new") + monkeypatch.delenv("LITELLM_SALT_KEY_PREVIOUS", raising=False) + + with pytest.raises(RuntimeError, match="LITELLM_SALT_KEY_PREVIOUS"): + await cm.migrate_encryption( + prisma_client=MagicMock(), + user_api_key_dict=MagicMock(), + policy=cm.SALT_KEY_POLICY, + ) + + +def test_policy_for_mode(): + assert cm.policy_for_mode("algorithm") is cm.ALGORITHM_POLICY + assert cm.policy_for_mode("salt-key") is cm.SALT_KEY_POLICY + + +@pytest.mark.asyncio +async def test_salt_key_check_reports_residual(rotated_salt_key): + client = MagicMock() + _empty_covered_tables(client) + client.db.litellm_teamtable.find_many = AsyncMock(return_value=[]) + client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None) + client.db.litellm_config.find_unique = AsyncMock(return_value=None) + client.db.litellm_config.update = AsyncMock() + client.db.litellm_proxymodeltable.find_many = AsyncMock( + return_value=[SimpleNamespace(litellm_params={"api_key": rotated_salt_key})] + ) + + report = await cm.check_encryption(client, policy=cm.SALT_KEY_POLICY) + + assert report.residual_legacy == 1 + assert report.as_dict()["locations"]["model_table"]["legacy"] == 1 + client.db.litellm_config.update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_salt_key_pass_walks_callback_vars(monkeypatch): + """Locations with no master-key rotation path are covered by the salt pass too.""" + from litellm.proxy.common_utils.callback_utils import ( + decrypt_callback_vars, + encrypt_callback_vars, + ) + + monkeypatch.setattr(proxy_server, "general_settings", {}) + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-old") + old_meta = encrypt_callback_vars( + {"logging": [{"callback_vars": {"gcs_path_service_account": "sa-secret"}}]} + ) + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-new") + monkeypatch.setenv("LITELLM_SALT_KEY_PREVIOUS", "sk-salt-old") + + row = SimpleNamespace(team_id="t1", metadata=old_meta) + client = MagicMock() + client.db.litellm_teamtable.find_many = AsyncMock(return_value=[row]) + client.db.litellm_teamtable.update = AsyncMock() + + report = await cm._migrate_callback_vars_table( + client, "team", dry_run=False, policy=cm.SALT_KEY_POLICY + ) + + assert report.migrated == 1 + written = json.loads( + client.db.litellm_teamtable.update.call_args.kwargs["data"]["metadata"] + ) + # Readable once the retired key is gone: the rewrite used the active key. + monkeypatch.delenv("LITELLM_SALT_KEY_PREVIOUS") + rotated = decrypt_callback_vars(written) + assert rotated["logging"][0]["callback_vars"]["gcs_path_service_account"] == "sa-secret" diff --git a/tests/test_litellm/proxy/management_endpoints/test_encryption_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_encryption_endpoints.py index e92cb3b13a7..6c06e603bb4 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_encryption_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_encryption_endpoints.py @@ -98,3 +98,47 @@ async def test_migrate_endpoint_db_not_connected(monkeypatch): with pytest.raises(HTTPException) as exc: await migrate_encryption_endpoint(user_api_key_dict=ADMIN) assert exc.value.status_code == 500 + + +# ------------------------------ salt-key mode ------------------------------ + + +@pytest.mark.asyncio +async def test_salt_key_mode_selects_the_salt_key_policy(monkeypatch): + monkeypatch.setattr(proxy_server, "prisma_client", object()) + fake = AsyncMock(return_value=_sample_report()) + monkeypatch.setattr(cm, "migrate_encryption", fake) + + out = await migrate_encryption_endpoint(user_api_key_dict=ADMIN, mode="salt-key") + + assert out["mode"] == "salt-key" + assert fake.await_args.kwargs["policy"] is cm.SALT_KEY_POLICY + + +@pytest.mark.asyncio +async def test_check_endpoint_salt_key_mode_selects_the_salt_key_policy(monkeypatch): + monkeypatch.setattr(proxy_server, "prisma_client", object()) + fake = AsyncMock(return_value=_sample_report()) + monkeypatch.setattr(cm, "check_encryption", fake) + + out = await check_encryption_endpoint(user_api_key_dict=ADMIN, mode="salt-key") + + assert out["mode"] == "salt-key" + assert fake.await_args.kwargs["policy"] is cm.SALT_KEY_POLICY + + +@pytest.mark.asyncio +async def test_unmet_precondition_is_400_not_500(monkeypatch): + """A missing LITELLM_SALT_KEY_PREVIOUS is operator error, not a server fault.""" + monkeypatch.setattr(proxy_server, "prisma_client", object()) + monkeypatch.setattr( + cm, + "migrate_encryption", + AsyncMock(side_effect=RuntimeError("LITELLM_SALT_KEY_PREVIOUS")), + ) + + with pytest.raises(HTTPException) as exc: + await migrate_encryption_endpoint(user_api_key_dict=ADMIN, mode="salt-key") + + assert exc.value.status_code == 400 + assert "LITELLM_SALT_KEY_PREVIOUS" in exc.value.detail["error"] diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 9df7d8baba7..d0946b8ae92 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -3240,11 +3240,16 @@ export interface paths { put?: never; /** * Migrate Encryption Endpoint - * @description Re-encrypt all at-rest credentials into the AES-256-GCM (``v2:gcm:``) format. + * @description Re-encrypt all at-rest credentials, either into the AES-256-GCM (``v2:gcm:``) + * format (``mode=algorithm``, the default) or under the active salt key + * (``mode=salt-key``). * - * Admin only. Requires ``general_settings.encryption_algorithm: aes-256-gcm``. - * Idempotent and resumable — re-running skips already-migrated values. Pass - * ``dry_run=true`` for a non-mutating scan (equivalent to ``--check``). + * Admin only. ``mode=algorithm`` requires + * ``general_settings.encryption_algorithm: aes-256-gcm``; ``mode=salt-key`` + * requires the retired key(s) in ``LITELLM_SALT_KEY_PREVIOUS`` and the new one + * in ``LITELLM_SALT_KEY``. Idempotent and resumable, so re-running skips values + * that are already current. Pass ``dry_run=true`` for a non-mutating scan + * (equivalent to ``--check``). */ post: operations["migrate_encryption_endpoint_credentials_migrate_encryption_post"]; delete?: never; @@ -3262,9 +3267,14 @@ export interface paths { }; /** * Check Encryption Endpoint - * @description Read-only residual scan for compliance attestation. Reports how many at-rest - * values are still in the legacy format. ``residual_legacy == 0`` attests no - * legacy ciphertext remains. Admin only; performs no writes. + * @description Read-only residual scan for compliance attestation. Admin only; performs no + * writes. With ``mode=algorithm`` it reports how many at-rest values are still + * in the legacy format; with ``mode=salt-key`` it reports how many still + * decrypt only under a retired salt key. ``residual_legacy == 0`` with an empty + * ``unreadable_locations`` is the attestation for the selected mode, and for + * ``salt-key`` it means ``LITELLM_SALT_KEY_PREVIOUS`` can be dropped. A store + * the scan could not read is named in ``unreadable_locations``, because zero + * counts from a store nobody could open mean unknown rather than clean. */ get: operations["check_encryption_endpoint_credentials_migrate_encryption_check_get"]; put?: never; @@ -43147,6 +43157,7 @@ export interface operations { query?: { /** @description If true, scan and report without writing any changes. */ dry_run?: boolean; + mode?: "algorithm" | "salt-key"; }; header?: never; path?: never; @@ -43176,7 +43187,9 @@ export interface operations { }; check_encryption_endpoint_credentials_migrate_encryption_check_get: { parameters: { - query?: never; + query?: { + mode?: "algorithm" | "salt-key"; + }; header?: never; path?: never; cookie?: never; @@ -43192,6 +43205,15 @@ export interface operations { "application/json": unknown; }; }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; }; }; delete_credential_credentials__credential_name__delete: {