mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(proxy): add AES-256-GCM at-rest credential encryption with versioned format and re-encryption migration (#31215)
Some checks are pending
GitHub Actions Security Analysis / zizmor (push) Waiting to run
Some checks are pending
GitHub Actions Security Analysis / zizmor (push) Waiting to run
* feat(proxy): add AES-256-GCM at-rest credential encryption with versioned format and re-encryption migration * test(proxy): add behavior scenarios for credential migration endpoints * fix(proxy): scan covered tables in encryption check, fix CI lint and route types * fix(proxy): migrate callback_settings credentials, clear CI lint/recursion gates, add encryption endpoint+CLI tests * fix(proxy): correct dry-run/real-run migrated vs residual-legacy counters in config and SSO walkers * fix(proxy): make callback-vars residual detection gate-independent in encryption check
This commit is contained in:
parent
8e30cfbeb1
commit
d7654d07ab
11 changed files with 1935 additions and 1 deletions
60
litellm/proxy/client/cli/commands/encryption.py
Normal file
60
litellm/proxy/client/cli/commands/encryption.py
Normal file
|
|
@ -0,0 +1,60 @@
|
|||
"""CLI commands for the at-rest credential encryption migration."""
|
||||
|
||||
import click
|
||||
import rich
|
||||
|
||||
from ...http_client import HTTPClient
|
||||
|
||||
|
||||
@click.group()
|
||||
def encryption():
|
||||
"""Migrate at-rest credentials to AES-256-GCM and attest residual state."""
|
||||
pass
|
||||
|
||||
|
||||
@encryption.command(name="migrate")
|
||||
@click.option(
|
||||
"--check",
|
||||
"check_only",
|
||||
is_flag=True,
|
||||
default=False,
|
||||
help="Read-only residual scan (no writes). Reports legacy values remaining.",
|
||||
)
|
||||
@click.option(
|
||||
"--dry-run",
|
||||
is_flag=True,
|
||||
default=False,
|
||||
help="Run the full migration walkers without writing any changes.",
|
||||
)
|
||||
@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.
|
||||
|
||||
Requires the proxy to be started with
|
||||
``general_settings.encryption_algorithm: aes-256-gcm``. 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
|
||||
"""
|
||||
client = HTTPClient(ctx.obj["base_url"], ctx.obj["api_key"])
|
||||
|
||||
if check_only:
|
||||
response = client.request("GET", "/credentials/migrate-encryption/check")
|
||||
else:
|
||||
response = client.request(
|
||||
"POST",
|
||||
"/credentials/migrate-encryption",
|
||||
json={},
|
||||
params={"dry_run": "true"} if dry_run else None,
|
||||
)
|
||||
|
||||
rich.print_json(data=response)
|
||||
|
||||
report = response.get("report", {}) if isinstance(response, dict) else {}
|
||||
residual = report.get("residual_legacy")
|
||||
if residual is not None and residual > 0:
|
||||
rich.print(f"[yellow]Residual legacy values remaining: {residual}[/yellow]")
|
||||
elif residual == 0:
|
||||
rich.print("[green]No legacy values remaining (residual_legacy == 0).[/green]")
|
||||
|
|
@ -11,6 +11,7 @@ from .commands.agents import agent_commands
|
|||
from .commands.auth import get_stored_api_key, login, logout, whoami
|
||||
from .commands.chat import chat
|
||||
from .commands.credentials import credentials
|
||||
from .commands.encryption import encryption
|
||||
from .commands.http import http
|
||||
from .commands.keys import keys
|
||||
|
||||
|
|
@ -103,6 +104,8 @@ cli.add_command(whoami)
|
|||
cli.add_command(models)
|
||||
# Add the credentials command group
|
||||
cli.add_command(credentials)
|
||||
# Add the encryption migration command group
|
||||
cli.add_command(encryption)
|
||||
# Add the chat command group
|
||||
cli.add_command(chat)
|
||||
# Add the http command group
|
||||
|
|
|
|||
|
|
@ -1,9 +1,24 @@
|
|||
import base64
|
||||
import os
|
||||
from typing import Literal, Optional
|
||||
from typing import Literal, Optional, cast
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
# Versioned ciphertext marker for AES-256-GCM values.
|
||||
# Format: "v2:gcm:" + base64url(nonce(12) || ciphertext || tag(16)).
|
||||
# Legacy XSalsa20-Poly1305 (nacl) values carry no marker; the colon in the
|
||||
# prefix can never appear in base64url(nacl output), so the prefix check is an
|
||||
# unambiguous discriminator between the two formats on read.
|
||||
_V2_GCM_PREFIX = "v2:gcm:"
|
||||
|
||||
# general_settings key selecting the at-rest encryption algorithm for new writes.
|
||||
# Default preserves the legacy algorithm so existing deployments are byte-for-byte
|
||||
# unchanged until they explicitly opt in. Decrypt is always format-detecting, so
|
||||
# flipping this flag forward (or back) never strands previously-written data.
|
||||
_ENCRYPTION_ALGORITHM_SETTING = "encryption_algorithm"
|
||||
_ALGO_AES_GCM = "aes-256-gcm"
|
||||
_ALGO_XSALSA20 = "xsalsa20-poly1305"
|
||||
|
||||
|
||||
def _get_salt_key():
|
||||
from litellm.proxy.proxy_server import master_key
|
||||
|
|
@ -16,11 +31,76 @@ def _get_salt_key():
|
|||
return salt_key
|
||||
|
||||
|
||||
def _get_encryption_algorithm() -> str:
|
||||
"""
|
||||
Resolve the configured at-rest encryption algorithm for *new writes*.
|
||||
|
||||
Read from ``general_settings.encryption_algorithm`` at write time. Defaults to
|
||||
the legacy XSalsa20-Poly1305 algorithm so deployments that have not opted in
|
||||
keep producing byte-for-byte identical ciphertext.
|
||||
"""
|
||||
try:
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
algo = general_settings.get(_ENCRYPTION_ALGORITHM_SETTING, _ALGO_XSALSA20)
|
||||
except Exception:
|
||||
# general_settings may not be importable in some contexts (e.g. SDK-only
|
||||
# use of these helpers). Fall back to the legacy algorithm.
|
||||
return _ALGO_XSALSA20
|
||||
|
||||
if isinstance(algo, str) and algo.lower() == _ALGO_AES_GCM:
|
||||
return _ALGO_AES_GCM
|
||||
return _ALGO_XSALSA20
|
||||
|
||||
|
||||
def _derive_key(signing_key: str) -> bytes:
|
||||
"""Derive a 32-byte key from the salt/master key (shared by both algorithms).
|
||||
|
||||
Known limitation: this is a single-pass, unsalted ``SHA-256`` of the key, not
|
||||
a dedicated KDF (HKDF/PBKDF2). It is the *same* derivation the legacy nacl
|
||||
path already uses, so the AES path introduces no new weakness and stays
|
||||
interoperable with existing key sourcing; AES-256-GCM's per-value 12-byte
|
||||
random nonce gives the unique (key, nonce) pairs GCM requires. Moving both
|
||||
algorithms to HKDF-SHA256 would be more defensible in an audit but is a
|
||||
separate, coordinated change (it must re-derive or re-encrypt existing data).
|
||||
"""
|
||||
import hashlib
|
||||
|
||||
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."""
|
||||
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
|
||||
|
||||
nonce = os.urandom(12)
|
||||
# AESGCM.encrypt returns ciphertext || tag(16); wire format is nonce || that.
|
||||
blob = AESGCM(_derive_key(signing_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`."""
|
||||
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
|
||||
|
||||
raw = 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")
|
||||
|
||||
|
||||
def encrypt_value_helper(value: str, new_encryption_key: Optional[str] = None):
|
||||
signing_key = new_encryption_key or _get_salt_key()
|
||||
|
||||
try:
|
||||
if isinstance(value, str):
|
||||
if _get_encryption_algorithm() == _ALGO_AES_GCM:
|
||||
# AES path: the v2:gcm: output is already a base64url string, so it
|
||||
# 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) # type: ignore
|
||||
# Use urlsafe_b64encode for URL-safe base64 encoding (replaces + with - and / with _)
|
||||
encrypted_value = base64.urlsafe_b64encode(encrypted_value).decode("utf-8")
|
||||
|
|
@ -46,6 +126,11 @@ def decrypt_value_helper(
|
|||
|
||||
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:
|
||||
|
|
|
|||
702
litellm/proxy/management_endpoints/credential_migration.py
Normal file
702
litellm/proxy/management_endpoints/credential_migration.py
Normal file
|
|
@ -0,0 +1,702 @@
|
|||
"""
|
||||
At-rest credential re-encryption migration.
|
||||
|
||||
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``.
|
||||
|
||||
Design properties (see case 2026-06-24 fix plan):
|
||||
|
||||
* **Same key, new algorithm.** The migration does not change the encryption key;
|
||||
it re-encrypts existing ciphertext under the same derived key but in the new
|
||||
AES format. This is achieved by decrypting with the format-detecting reader and
|
||||
re-encrypting through ``encrypt_value_helper`` with the AES gate enabled.
|
||||
* **Idempotent.** A value already carrying the ``v2:gcm:`` prefix is recognised
|
||||
and left untouched, so re-running the migration is a no-op on migrated rows.
|
||||
* **Resumable.** Walkers commit per row (or per small table), so an interrupted
|
||||
run leaves a clean mixed state that a re-run completes.
|
||||
* **Skip-on-undecryptable.** A value that cannot be decrypted is never
|
||||
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.
|
||||
|
||||
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.
|
||||
"""
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, Literal, cast
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
||||
_ALGO_AES_GCM,
|
||||
_ENCRYPTION_ALGORITHM_SETTING,
|
||||
_V2_GCM_PREFIX,
|
||||
_get_salt_key,
|
||||
decrypt_value_helper,
|
||||
encrypt_value_helper,
|
||||
)
|
||||
|
||||
ValueClass = Literal["migrated", "legacy", "plaintext", "undecryptable", "not-a-string"]
|
||||
|
||||
|
||||
@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)
|
||||
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
|
||||
|
||||
def as_dict(self) -> dict[str, int]:
|
||||
return {
|
||||
"scanned": self.scanned,
|
||||
"migrated": self.migrated,
|
||||
"already_v2": self.already_v2,
|
||||
"plaintext": self.plaintext,
|
||||
"undecryptable": self.undecryptable,
|
||||
"legacy": self.legacy,
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class MigrationReport:
|
||||
"""Aggregate report across all locations."""
|
||||
|
||||
locations: list[LocationReport] = field(default_factory=list)
|
||||
|
||||
def add(self, report: LocationReport) -> None:
|
||||
self.locations.append(report)
|
||||
|
||||
@property
|
||||
def residual_legacy(self) -> int:
|
||||
"""Total legacy ciphertext still un-migrated (the TRO attestation number)."""
|
||||
return sum(loc.legacy for loc in self.locations)
|
||||
|
||||
@property
|
||||
def total_undecryptable(self) -> int:
|
||||
return sum(loc.undecryptable for loc in self.locations)
|
||||
|
||||
def as_dict(self) -> dict[str, object]:
|
||||
return {
|
||||
"residual_legacy": self.residual_legacy,
|
||||
"total_undecryptable": self.total_undecryptable,
|
||||
"locations": {loc.location: loc.as_dict() for loc in self.locations},
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pure engine — no DB I/O, fully unit-testable.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def is_migrated(value: object) -> bool:
|
||||
"""True if ``value`` is already an AES-256-GCM (``v2:gcm:``) ciphertext."""
|
||||
return isinstance(value, str) and value.startswith(_V2_GCM_PREFIX)
|
||||
|
||||
|
||||
def classify_value(value: object, key: str = "scan") -> 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).
|
||||
* ``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.
|
||||
"""
|
||||
if not isinstance(value, str):
|
||||
return "not-a-string"
|
||||
if value == "":
|
||||
return "plaintext"
|
||||
if value.startswith(_V2_GCM_PREFIX):
|
||||
return "migrated"
|
||||
decrypted = 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.
|
||||
return "plaintext"
|
||||
return "legacy"
|
||||
|
||||
|
||||
def reencrypt_value(value: object, key: str = "migrate") -> object:
|
||||
"""Re-encrypt a single stored string into the configured (AES) format.
|
||||
|
||||
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).
|
||||
"""
|
||||
if not isinstance(value, str) or value == "":
|
||||
return value
|
||||
if value.startswith(_V2_GCM_PREFIX):
|
||||
return value # idempotent: already migrated
|
||||
decrypted = 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)
|
||||
|
||||
|
||||
def reencrypt_selective_dict(
|
||||
data: dict[str, object], sensitive_keys: list[str]
|
||||
) -> 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.
|
||||
Null/missing fields are skipped.
|
||||
"""
|
||||
out = dict(data)
|
||||
for k in sensitive_keys:
|
||||
v = out.get(k)
|
||||
if v is None:
|
||||
continue
|
||||
out[k] = reencrypt_value(v, key=k)
|
||||
return out
|
||||
|
||||
|
||||
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 = general_settings.get(_ENCRYPTION_ALGORITHM_SETTING)
|
||||
if not (isinstance(algo, str) and algo.lower() == _ALGO_AES_GCM):
|
||||
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."
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Walkers for the locations with no pre-existing rotation path.
|
||||
# Each walker delegates the structural transform to the existing, tested helper
|
||||
# for that table and only adds the per-row re-encrypt + commit + counters.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def _migrate_config_settings_row(
|
||||
prisma_client: object,
|
||||
param_name: str,
|
||||
sensitive_fields: list[str],
|
||||
dry_run: bool,
|
||||
) -> LocationReport:
|
||||
"""Migrate a single ``LiteLLM_Config`` row whose ``param_value`` is a JSON
|
||||
dict with selected sensitive fields (vantage_settings / cloudzero_settings).
|
||||
"""
|
||||
report = LocationReport(location=param_name)
|
||||
record = await prisma_client.db.litellm_config.find_unique(
|
||||
where={"param_name": param_name}
|
||||
)
|
||||
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):
|
||||
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:
|
||||
await prisma_client.db.litellm_config.update(
|
||||
where={"param_name": param_name},
|
||||
data={"param_value": json.dumps(settings)},
|
||||
)
|
||||
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 = LocationReport(location="sso_config")
|
||||
record = await prisma_client.db.litellm_ssoconfig.find_unique(
|
||||
where={"id": "sso_config"}
|
||||
)
|
||||
if record is None or record.sso_settings is None:
|
||||
return report
|
||||
|
||||
settings = record.sso_settings
|
||||
if isinstance(settings, str):
|
||||
settings = json.loads(settings)
|
||||
if not isinstance(settings, dict):
|
||||
return report
|
||||
|
||||
new_settings = dict(settings)
|
||||
changed = False
|
||||
for fld, v in settings.items():
|
||||
if not isinstance(v, str) or v == "":
|
||||
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)},
|
||||
)
|
||||
return report
|
||||
|
||||
|
||||
async def _migrate_callback_vars_table(
|
||||
prisma_client: object,
|
||||
table_name: Literal["team", "verification_token"],
|
||||
dry_run: bool,
|
||||
) -> 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.<sensitive>`` and
|
||||
the top-level ``metadata.callback_settings.callback_vars.<sensitive>``. Reuses
|
||||
those proven transforms (selective, prefix-marked; legacy plaintext is left
|
||||
alone until re-encrypted).
|
||||
"""
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
decrypt_callback_vars,
|
||||
encrypt_callback_vars,
|
||||
)
|
||||
|
||||
report = LocationReport(location=f"{table_name}.callback_vars")
|
||||
|
||||
if table_name == "team":
|
||||
table = prisma_client.db.litellm_teamtable
|
||||
pk = "team_id"
|
||||
else:
|
||||
table = prisma_client.db.litellm_verificationtoken
|
||||
pk = "token"
|
||||
|
||||
rows = await table.find_many()
|
||||
for row in rows or []:
|
||||
metadata = getattr(row, "metadata", None)
|
||||
if not isinstance(metadata, dict) or (
|
||||
"logging" not in metadata and "callback_settings" not in metadata
|
||||
):
|
||||
continue
|
||||
|
||||
# Classify every callback-var value directly (strip the litellm_enc::
|
||||
# marker, then prefix/decrypt-classify), exactly like the covered-table
|
||||
# scanner. Detecting legacy this way is independent of the AES gate, so
|
||||
# the check_encryption (dry-run) attestation is correct even when run
|
||||
# before the gate is enabled -- a re-encrypt-delta heuristic would read
|
||||
# zero residual here with the gate off.
|
||||
row_legacy = 0
|
||||
for cvs in _iter_callback_var_dicts(metadata):
|
||||
for v in cvs.values():
|
||||
report.scanned += 1
|
||||
cls = _classify_callback_value(v)
|
||||
if cls == "migrated":
|
||||
report.already_v2 += 1
|
||||
elif cls == "legacy":
|
||||
row_legacy += 1
|
||||
else: # plaintext / not-a-string
|
||||
report.plaintext += 1
|
||||
|
||||
if row_legacy == 0:
|
||||
continue # no legacy ciphertext in this row
|
||||
|
||||
if dry_run:
|
||||
# Residual for the attestation; a dry run writes nothing.
|
||||
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.
|
||||
try:
|
||||
re_encrypted = encrypt_callback_vars(decrypt_callback_vars(metadata))
|
||||
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",
|
||||
table_name,
|
||||
getattr(row, pk, "?"),
|
||||
str(e),
|
||||
)
|
||||
report.undecryptable += row_legacy
|
||||
continue
|
||||
report.migrated += row_legacy
|
||||
await table.update(
|
||||
where={pk: getattr(row, pk)},
|
||||
data={"metadata": json.dumps(re_encrypted)},
|
||||
)
|
||||
|
||||
return report
|
||||
|
||||
|
||||
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
|
||||
``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
|
||||
format at rest.
|
||||
"""
|
||||
for entry in metadata.get("logging", []) or []:
|
||||
if isinstance(entry, dict):
|
||||
cvs = entry.get("callback_vars")
|
||||
if isinstance(cvs, dict):
|
||||
yield cvs
|
||||
callback_settings = metadata.get("callback_settings")
|
||||
if isinstance(callback_settings, dict):
|
||||
cvs = callback_settings.get("callback_vars")
|
||||
if isinstance(cvs, dict):
|
||||
yield cvs
|
||||
|
||||
|
||||
def _classify_callback_value(value: object) -> ValueClass:
|
||||
"""Classify one stored callback-var value, independent of the AES gate.
|
||||
|
||||
Encrypted callback vars carry the ``litellm_enc::`` marker in front of the
|
||||
ciphertext; strip it, then classify the inner value the same way the
|
||||
covered-table scanner does (``v2:gcm:`` prefix -> migrated, nacl-decryptable
|
||||
-> legacy, otherwise plaintext). Detecting legacy by decrypt rather than by a
|
||||
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
|
||||
if inner.startswith(_CALLBACK_VAR_ENCRYPTED_PREFIX):
|
||||
inner = inner[len(_CALLBACK_VAR_ENCRYPTED_PREFIX) :]
|
||||
return classify_value(inner, key="callback")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Read-only scanner for the rotation-covered tables.
|
||||
#
|
||||
# ``_rotate_master_key`` re-encrypts these tables but returns no counts, so on
|
||||
# its own it can neither attest residual legacy nor report how many rows it
|
||||
# migrated. This scanner reads (never writes) the same encrypted columns the
|
||||
# rotation path touches and classifies every value, giving both the attestation
|
||||
# coverage and the pre/post counts the rotation path can't supply itself.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# (location, prisma db attribute, JSON columns to walk, scalar string columns).
|
||||
_COVERED_TABLE_SPECS = [
|
||||
("model_table", "litellm_proxymodeltable", ("litellm_params",), ()),
|
||||
("credentials", "litellm_credentialstable", ("credential_values",), ()),
|
||||
("mcp_server", "litellm_mcpservertable", ("credentials", "env_vars"), ()),
|
||||
("mcp_user_credentials", "litellm_mcpusercredentials", (), ("credential_b64",)),
|
||||
("mcp_user_env_vars", "litellm_mcpuserenvvars", (), ("values_b64",)),
|
||||
]
|
||||
|
||||
|
||||
def _iter_encrypted_strings(obj: object):
|
||||
"""Yield every string leaf in a nested dict/list/scalar structure.
|
||||
|
||||
Iterative (explicit stack) on purpose: recursion here is banned by the
|
||||
code-quality recursive-function detector (unbounded nesting has caused CPU
|
||||
spikes in the past), and an explicit stack walks arbitrary depth safely.
|
||||
"""
|
||||
stack: list[object] = [obj]
|
||||
while stack:
|
||||
cur = stack.pop()
|
||||
if isinstance(cur, str):
|
||||
yield cur
|
||||
elif isinstance(cur, dict):
|
||||
stack.extend(cur.values())
|
||||
elif isinstance(cur, list):
|
||||
stack.extend(cur)
|
||||
|
||||
|
||||
def _classify_into_report(report: LocationReport, value: str) -> None:
|
||||
"""Classify one stored string and bump the matching read-only counter.
|
||||
|
||||
Only genuine nacl ciphertext lands in ``legacy``; non-secret strings (model
|
||||
names, base URLs, …) do not decrypt and fall through to ``plaintext``, so
|
||||
over-scanning a column is harmless to the residual count.
|
||||
"""
|
||||
report.scanned += 1
|
||||
cls = classify_value(value, key="scan")
|
||||
if cls == "migrated":
|
||||
report.already_v2 += 1
|
||||
elif cls == "legacy":
|
||||
report.legacy += 1
|
||||
else: # plaintext / not-a-string
|
||||
report.plaintext += 1
|
||||
|
||||
|
||||
async def _scan_one_table(
|
||||
prisma_client: object,
|
||||
location: str,
|
||||
db_attr: str,
|
||||
json_columns: tuple,
|
||||
scalar_columns: tuple,
|
||||
) -> LocationReport:
|
||||
report = LocationReport(location=location)
|
||||
table = getattr(prisma_client.db, db_attr, None)
|
||||
if table is None:
|
||||
return report
|
||||
try:
|
||||
rows = 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))
|
||||
return report
|
||||
for row in rows or []:
|
||||
for col in json_columns:
|
||||
raw = getattr(row, col, None)
|
||||
if raw is None:
|
||||
continue
|
||||
if isinstance(raw, str):
|
||||
try:
|
||||
raw = json.loads(raw)
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
for s in _iter_encrypted_strings(raw):
|
||||
_classify_into_report(report, s)
|
||||
for col in scalar_columns:
|
||||
v = getattr(row, col, None)
|
||||
if isinstance(v, str):
|
||||
_classify_into_report(report, v)
|
||||
return report
|
||||
|
||||
|
||||
async def _scan_config_env_vars(prisma_client: object) -> LocationReport:
|
||||
"""Scan the ``environment_variables`` config row (``param_value`` dict)."""
|
||||
report = LocationReport(location="config_environment_variables")
|
||||
try:
|
||||
record = 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))
|
||||
return report
|
||||
if record is None or record.param_value is None:
|
||||
return report
|
||||
value = record.param_value
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
value = json.loads(value)
|
||||
except (ValueError, TypeError):
|
||||
value = {}
|
||||
for s in _iter_encrypted_strings(value):
|
||||
_classify_into_report(report, s)
|
||||
return report
|
||||
|
||||
|
||||
async def _scan_covered_tables(prisma_client: object) -> list[LocationReport]:
|
||||
"""Read-only classification of every rotation-covered table. No writes."""
|
||||
reports: 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))
|
||||
return reports
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Orchestrator
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# vantage_settings / cloudzero_settings sensitive fields (see *_endpoints.py).
|
||||
_VANTAGE_SENSITIVE = ["api_key", "integration_token"]
|
||||
_CLOUDZERO_SENSITIVE = ["api_key"]
|
||||
|
||||
|
||||
async def _migrate_covered_tables(
|
||||
prisma_client: object, user_api_key_dict: object
|
||||
) -> 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
|
||||
re-encrypt writes land in ``v2:`` format.
|
||||
|
||||
``_rotate_master_key`` returns no counts, so we bracket it with read-only
|
||||
scans: the pre-scan's legacy total minus the post-scan's gives the number
|
||||
actually migrated per location, and the post-scan supplies the residual /
|
||||
already-v2 / scanned figures. Returns one report per covered location.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
_rotate_master_key,
|
||||
)
|
||||
|
||||
pre = {r.location: r for r in await _scan_covered_tables(prisma_client)}
|
||||
|
||||
current_key = _get_salt_key()
|
||||
if current_key is None:
|
||||
raise RuntimeError(
|
||||
"Cannot migrate covered tables: no salt key / master key is set. "
|
||||
"Set LITELLM_SALT_KEY before migrating."
|
||||
)
|
||||
await _rotate_master_key(
|
||||
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
|
||||
)
|
||||
|
||||
post = await _scan_covered_tables(prisma_client)
|
||||
for post_report in post:
|
||||
pre_report = pre.get(post_report.location)
|
||||
pre_legacy = pre_report.legacy if pre_report else 0
|
||||
# Everything that was legacy before and is no longer legacy now was
|
||||
# converted this run.
|
||||
post_report.migrated = max(0, pre_legacy - post_report.legacy)
|
||||
return post
|
||||
|
||||
|
||||
async def migrate_encryption(
|
||||
prisma_client: object,
|
||||
user_api_key_dict: object,
|
||||
dry_run: bool = False,
|
||||
) -> MigrationReport:
|
||||
"""Run the full at-rest re-encryption migration.
|
||||
|
||||
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()
|
||||
|
||||
report = MigrationReport()
|
||||
|
||||
# Tables that already have a rotation path (items 1, 2, 5-10). On a real run
|
||||
# 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):
|
||||
report.add(covered)
|
||||
else:
|
||||
for covered in await _migrate_covered_tables(prisma_client, user_api_key_dict):
|
||||
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))
|
||||
|
||||
return report
|
||||
|
||||
|
||||
async def check_encryption(prisma_client: object) -> 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``;
|
||||
``residual_legacy == 0`` across this full scan is the compliance attestation.
|
||||
"""
|
||||
report = MigrationReport()
|
||||
|
||||
# Rotation-covered tables (read-only classification).
|
||||
for covered in await _scan_covered_tables(prisma_client):
|
||||
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_config_settings_row(
|
||||
prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run=True
|
||||
)
|
||||
)
|
||||
report.add(await _migrate_sso_config(prisma_client, dry_run=True))
|
||||
return report
|
||||
|
|
@ -4091,6 +4091,86 @@ async def _rotate_master_key(
|
|||
verbose_proxy_logger.debug(f"Successfully re-encrypted {len(credentials)} credentials with new master key")
|
||||
|
||||
|
||||
def _require_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> None:
|
||||
from litellm.proxy._types import CommonProxyErrors
|
||||
|
||||
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={"error": CommonProxyErrors.not_allowed_access.value},
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/credentials/migrate-encryption",
|
||||
tags=["credential management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def migrate_encryption_endpoint(
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
dry_run: bool = Query(
|
||||
False,
|
||||
description="If true, scan and report without writing any changes.",
|
||||
),
|
||||
):
|
||||
"""
|
||||
Re-encrypt all at-rest credentials into the AES-256-GCM (``v2:gcm:``) format.
|
||||
|
||||
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``).
|
||||
"""
|
||||
from litellm.proxy._types import CommonProxyErrors
|
||||
from litellm.proxy.management_endpoints.credential_migration import (
|
||||
migrate_encryption,
|
||||
)
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
_require_proxy_admin(user_api_key_dict)
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": CommonProxyErrors.db_not_connected_error.value},
|
||||
)
|
||||
|
||||
report = 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()}
|
||||
|
||||
|
||||
@router.get(
|
||||
"/credentials/migrate-encryption/check",
|
||||
tags=["credential management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def check_encryption_endpoint(
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
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.
|
||||
"""
|
||||
from litellm.proxy._types import CommonProxyErrors
|
||||
from litellm.proxy.management_endpoints.credential_migration import (
|
||||
check_encryption,
|
||||
)
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
_require_proxy_admin(user_api_key_dict)
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": CommonProxyErrors.db_not_connected_error.value},
|
||||
)
|
||||
|
||||
report = await check_encryption(prisma_client=prisma_client)
|
||||
return {"status": "success", "report": report.as_dict()}
|
||||
|
||||
|
||||
async def get_new_token(data: Optional[RegenerateKeyRequest]) -> str:
|
||||
if data and data.new_key is not None:
|
||||
# Reject custom key values if disabled by admin
|
||||
|
|
|
|||
|
|
@ -0,0 +1,71 @@
|
|||
"""Behavior scenarios for the credential re-encryption migration endpoints.
|
||||
|
||||
These run against the live ASGI app + DB. The migration POST is *not* exercised
|
||||
end-to-end here because it mutates shared at-rest data (it delegates to the
|
||||
master-key rotation path); that full flow is covered by the unit suite and a
|
||||
live proxy run. Here we pin the HTTP-boundary contract: the read-only check is
|
||||
admin-reachable, and both routes are admin-gated.
|
||||
|
||||
Both routes are admin-only management routes, so a non-admin key is rejected by
|
||||
the ``user_api_key_auth`` layer (401) before the endpoint's own admin guard runs
|
||||
-- the negative scenarios assert that framework-level rejection.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from .conftest import MASTER_KEY
|
||||
|
||||
pytestmark = pytest.mark.asyncio(loop_scope="session")
|
||||
|
||||
|
||||
async def test_migrate_encryption_check_as_admin_is_read_only(proxy_client):
|
||||
"""GET /credentials/migrate-encryption/check returns a residual report (no writes)."""
|
||||
resp = await proxy_client.get(
|
||||
"/credentials/migrate-encryption/check",
|
||||
headers={"Authorization": f"Bearer {MASTER_KEY}"},
|
||||
)
|
||||
assert resp.status_code == 200, resp.text
|
||||
body = resp.json()
|
||||
assert body["status"] == "success"
|
||||
assert "residual_legacy" in body["report"]
|
||||
|
||||
|
||||
async def test_migrate_encryption_check_requires_admin(proxy_client, scratch):
|
||||
"""A non-admin key cannot reach the residual scan (auth layer rejects, 401)."""
|
||||
gen = await proxy_client.post(
|
||||
"/key/generate",
|
||||
headers={"Authorization": f"Bearer {MASTER_KEY}"},
|
||||
json={"key_alias": scratch.tag("check"), "user_id": scratch.tag("check-user")},
|
||||
)
|
||||
assert gen.status_code == 200, gen.text
|
||||
nonadmin_key = gen.json()["key"]
|
||||
|
||||
resp = await proxy_client.get(
|
||||
"/credentials/migrate-encryption/check",
|
||||
headers={"Authorization": f"Bearer {nonadmin_key}"},
|
||||
)
|
||||
assert resp.status_code == 401, resp.text
|
||||
|
||||
|
||||
async def test_migrate_encryption_requires_admin(proxy_client, scratch):
|
||||
"""A non-admin key cannot trigger the migration (auth layer rejects, 401).
|
||||
|
||||
Rejection happens before any write: the admin-only route check fires in
|
||||
``user_api_key_auth``, ahead of the endpoint body.
|
||||
"""
|
||||
gen = await proxy_client.post(
|
||||
"/key/generate",
|
||||
headers={"Authorization": f"Bearer {MASTER_KEY}"},
|
||||
json={
|
||||
"key_alias": scratch.tag("migrate"),
|
||||
"user_id": scratch.tag("migrate-user"),
|
||||
},
|
||||
)
|
||||
assert gen.status_code == 200, gen.text
|
||||
nonadmin_key = gen.json()["key"]
|
||||
|
||||
resp = await proxy_client.post(
|
||||
"/credentials/migrate-encryption",
|
||||
headers={"Authorization": f"Bearer {nonadmin_key}"},
|
||||
)
|
||||
assert resp.status_code == 401, resp.text
|
||||
|
|
@ -0,0 +1,71 @@
|
|||
"""CLI tests for the ``litellm-proxy encryption migrate`` command.
|
||||
|
||||
The HTTP client is mocked, so these assert the command's request routing (GET
|
||||
check vs POST migrate, dry-run param) and its residual-state messaging without a
|
||||
live proxy.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from click.testing import CliRunner
|
||||
|
||||
from litellm.proxy.client.cli import main as cli_main
|
||||
from litellm.proxy.client.cli.commands import encryption as enc_cli
|
||||
|
||||
|
||||
class _FakeHTTPClient:
|
||||
"""Stand-in for HTTPClient: records the last request and returns a canned body."""
|
||||
|
||||
last = None
|
||||
response = {"status": "success", "report": {"residual_legacy": 0, "locations": {}}}
|
||||
|
||||
def __init__(self, base_url, api_key):
|
||||
self.base_url = base_url
|
||||
self.api_key = api_key
|
||||
|
||||
def request(self, method, path, **kwargs):
|
||||
_FakeHTTPClient.last = {"method": method, "path": path, **kwargs}
|
||||
return _FakeHTTPClient.response
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def runner(monkeypatch):
|
||||
_FakeHTTPClient.last = None
|
||||
_FakeHTTPClient.response = {
|
||||
"status": "success",
|
||||
"report": {"residual_legacy": 0, "locations": {}},
|
||||
}
|
||||
monkeypatch.setattr(enc_cli, "HTTPClient", _FakeHTTPClient)
|
||||
return CliRunner()
|
||||
|
||||
|
||||
def test_migrate_check_hits_check_route(runner):
|
||||
result = runner.invoke(cli_main.cli, ["encryption", "migrate", "--check"])
|
||||
assert result.exit_code == 0, result.output
|
||||
assert _FakeHTTPClient.last["method"] == "GET"
|
||||
assert _FakeHTTPClient.last["path"] == "/credentials/migrate-encryption/check"
|
||||
assert "No legacy values remaining" in result.output
|
||||
|
||||
|
||||
def test_migrate_default_posts_without_dry_run(runner):
|
||||
result = runner.invoke(cli_main.cli, ["encryption", "migrate"])
|
||||
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
|
||||
|
||||
|
||||
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"}
|
||||
|
||||
|
||||
def test_migrate_reports_residual_legacy(runner):
|
||||
_FakeHTTPClient.response = {
|
||||
"status": "success",
|
||||
"report": {"residual_legacy": 3, "locations": {}},
|
||||
}
|
||||
result = runner.invoke(cli_main.cli, ["encryption", "migrate", "--check"])
|
||||
assert result.exit_code == 0, result.output
|
||||
assert "Residual legacy values remaining: 3" in result.output
|
||||
|
|
@ -0,0 +1,148 @@
|
|||
"""
|
||||
Tests for the at-rest credential encryption chokepoint.
|
||||
|
||||
Covers the AES-256-GCM (``v2:gcm:``) path, the ``encryption_algorithm`` config
|
||||
gate, and the backward-compatibility guarantees that let legacy XSalsa20-Poly1305
|
||||
(nacl) ciphertext and new AES values coexist and decrypt correctly.
|
||||
"""
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
def _use_aes(monkeypatch):
|
||||
"""Flip the write-time algorithm to AES-256-GCM for the duration of a test."""
|
||||
monkeypatch.setattr(
|
||||
proxy_server, "general_settings", {"encryption_algorithm": "aes-256-gcm"}
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _salt_key(monkeypatch):
|
||||
# Dominant convention in the test_litellm/ tree: set the key via env.
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-salt-aes-1234")
|
||||
# Ensure the legacy default is in force unless a test opts into AES.
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {})
|
||||
yield
|
||||
|
||||
|
||||
def test_aes_gcm_round_trip(monkeypatch):
|
||||
"""A value written under AES-256-GCM is tagged v2:gcm: and decrypts back."""
|
||||
_use_aes(monkeypatch)
|
||||
|
||||
ct = encrypt_value_helper("super-secret")
|
||||
|
||||
assert ct.startswith(_V2_GCM_PREFIX)
|
||||
assert decrypt_value_helper(ct, key="t") == "super-secret"
|
||||
|
||||
|
||||
def test_default_is_legacy_algorithm(monkeypatch):
|
||||
"""With no config, writes stay on the legacy algorithm (no v2: marker)."""
|
||||
ct = encrypt_value_helper("legacy-secret")
|
||||
|
||||
assert not ct.startswith(_V2_GCM_PREFIX)
|
||||
assert decrypt_value_helper(ct, key="t") == "legacy-secret"
|
||||
|
||||
|
||||
def test_legacy_nacl_value_still_decrypts_after_flag_flip(monkeypatch):
|
||||
"""A value written under the old algorithm decrypts unchanged once AES is on.
|
||||
|
||||
This is the mixed-format readback guarantee: decrypt is format-detecting, so
|
||||
flipping the flag forward never strands previously-written data.
|
||||
"""
|
||||
legacy = encrypt_value_helper("legacy-secret") # default = xsalsa20
|
||||
assert not legacy.startswith(_V2_GCM_PREFIX)
|
||||
|
||||
_use_aes(monkeypatch)
|
||||
# New writes are now AES, but the old value must still come back.
|
||||
assert decrypt_value_helper(legacy, key="t") == "legacy-secret"
|
||||
assert encrypt_value_helper("fresh").startswith(_V2_GCM_PREFIX)
|
||||
|
||||
|
||||
def test_v2_prefix_is_idempotent_marker(monkeypatch):
|
||||
"""The migration's skip-check: an already-v2 value is recognized by its prefix.
|
||||
|
||||
Re-encrypting an AES value yields a fresh (different nonce) AES value, but the
|
||||
prefix is what lets a migration skip already-migrated rows without decrypting.
|
||||
"""
|
||||
_use_aes(monkeypatch)
|
||||
|
||||
ct = encrypt_value_helper("secret")
|
||||
assert ct.startswith(_V2_GCM_PREFIX)
|
||||
|
||||
# Round-tripping does not change the plaintext, and the marker is stable.
|
||||
again = encrypt_value_helper(decrypt_value_helper(ct, key="t"))
|
||||
assert again.startswith(_V2_GCM_PREFIX)
|
||||
assert decrypt_value_helper(again, key="t") == "secret"
|
||||
|
||||
|
||||
def test_aes_decrypt_failure_returns_none_not_raise(monkeypatch):
|
||||
"""Decrypt contract preserved: a garbled v2 value returns None, never raises."""
|
||||
_use_aes(monkeypatch)
|
||||
|
||||
garbled = _V2_GCM_PREFIX + "not-valid-base64-or-ciphertext!!!"
|
||||
# exception_type="debug" exercises the swallow path; must not raise.
|
||||
assert decrypt_value_helper(garbled, key="t", exception_type="debug") is None
|
||||
|
||||
|
||||
def test_aes_decrypt_failure_returns_original_when_requested(monkeypatch):
|
||||
"""With return_original_value=True a bad v2 value comes back as-is, not None."""
|
||||
_use_aes(monkeypatch)
|
||||
|
||||
garbled = _V2_GCM_PREFIX + "###"
|
||||
assert (
|
||||
decrypt_value_helper(
|
||||
garbled, key="t", exception_type="debug", return_original_value=True
|
||||
)
|
||||
== garbled
|
||||
)
|
||||
|
||||
|
||||
def test_empty_string_round_trips_under_aes(monkeypatch):
|
||||
"""Empty string is preserved through the AES path (parity with legacy)."""
|
||||
_use_aes(monkeypatch)
|
||||
|
||||
ct = encrypt_value_helper("")
|
||||
assert ct.startswith(_V2_GCM_PREFIX)
|
||||
assert decrypt_value_helper(ct, key="t") == ""
|
||||
|
||||
|
||||
def test_callback_prefix_composes_with_v2(monkeypatch):
|
||||
"""litellm_enc:: + v2:gcm:... round-trips through the callback read path.
|
||||
|
||||
Callback vars are stored as ``litellm_enc::<helper output>``; the read path
|
||||
strips ``litellm_enc::`` then calls the helper, so the value handed to the
|
||||
helper is ``v2:gcm:...``. Ordering must work end to end.
|
||||
"""
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
_CALLBACK_VAR_ENCRYPTED_PREFIX,
|
||||
_decrypt_or_passthrough,
|
||||
_encrypt_if_plaintext,
|
||||
)
|
||||
|
||||
_use_aes(monkeypatch)
|
||||
|
||||
# "gcs_path_service_account" is a known-sensitive callback key.
|
||||
stored = _encrypt_if_plaintext("gcs_path_service_account", "my-sa-secret")
|
||||
|
||||
assert stored.startswith(_CALLBACK_VAR_ENCRYPTED_PREFIX)
|
||||
inner = stored[len(_CALLBACK_VAR_ENCRYPTED_PREFIX) :]
|
||||
assert inner.startswith(_V2_GCM_PREFIX)
|
||||
assert _decrypt_or_passthrough("gcs_path_service_account", stored) == "my-sa-secret"
|
||||
|
||||
|
||||
def test_unknown_algorithm_falls_back_to_legacy(monkeypatch):
|
||||
"""An unrecognized encryption_algorithm value does not produce v2 writes."""
|
||||
monkeypatch.setattr(
|
||||
proxy_server, "general_settings", {"encryption_algorithm": "rot13"}
|
||||
)
|
||||
|
||||
ct = encrypt_value_helper("secret")
|
||||
assert not ct.startswith(_V2_GCM_PREFIX)
|
||||
assert decrypt_value_helper(ct, key="t") == "secret"
|
||||
|
|
@ -0,0 +1,516 @@
|
|||
"""
|
||||
Tests for the at-rest credential re-encryption migration engine.
|
||||
|
||||
The pure engine (classify / reencrypt / selective-dict) is tested directly; the
|
||||
DB walkers are tested against an AsyncMock Prisma client. Live end-to-end
|
||||
proof-of-fix (real proxy + DB) is performed separately on the repro server.
|
||||
"""
|
||||
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
||||
_V2_GCM_PREFIX,
|
||||
encrypt_value_helper,
|
||||
)
|
||||
from litellm.proxy.management_endpoints import credential_migration as cm
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def salt_key(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-migration-salt-1234")
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {})
|
||||
return "sk-migration-salt-1234"
|
||||
|
||||
|
||||
def _legacy_ct(value: str, monkeypatch) -> str:
|
||||
"""Produce a legacy (nacl) ciphertext with the AES gate off."""
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {})
|
||||
return encrypt_value_helper(value)
|
||||
|
||||
|
||||
def _enable_aes(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
proxy_server, "general_settings", {"encryption_algorithm": "aes-256-gcm"}
|
||||
)
|
||||
|
||||
|
||||
def _empty_covered_tables(client):
|
||||
"""Wire every rotation-covered table 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.
|
||||
"""
|
||||
for _, db_attr, _, _ in cm._COVERED_TABLE_SPECS:
|
||||
getattr(client.db, db_attr).find_many = AsyncMock(return_value=[])
|
||||
|
||||
|
||||
# --------------------------- pure engine ---------------------------
|
||||
|
||||
|
||||
def test_classify_value(salt_key, monkeypatch):
|
||||
legacy = _legacy_ct("secret", monkeypatch)
|
||||
_enable_aes(monkeypatch)
|
||||
migrated = encrypt_value_helper("secret")
|
||||
|
||||
assert cm.classify_value(legacy) == "legacy"
|
||||
assert cm.classify_value(migrated) == "migrated"
|
||||
assert cm.classify_value("just-plaintext") == "plaintext"
|
||||
assert cm.classify_value("") == "plaintext"
|
||||
assert cm.classify_value(123) == "not-a-string"
|
||||
assert cm.classify_value(None) == "not-a-string"
|
||||
|
||||
|
||||
def test_is_migrated(salt_key, monkeypatch):
|
||||
_enable_aes(monkeypatch)
|
||||
assert cm.is_migrated(encrypt_value_helper("x")) is True
|
||||
assert cm.is_migrated("plaintext") is False
|
||||
assert cm.is_migrated(5) is False
|
||||
|
||||
|
||||
def test_reencrypt_value_legacy_to_v2(salt_key, monkeypatch):
|
||||
legacy = _legacy_ct("secret", monkeypatch)
|
||||
_enable_aes(monkeypatch)
|
||||
|
||||
out = cm.reencrypt_value(legacy)
|
||||
assert out != legacy
|
||||
assert out.startswith(_V2_GCM_PREFIX)
|
||||
|
||||
|
||||
def test_reencrypt_value_is_idempotent(salt_key, monkeypatch):
|
||||
_enable_aes(monkeypatch)
|
||||
v2 = encrypt_value_helper("secret")
|
||||
# Already v2 -> returned byte-for-byte unchanged (no re-wrap).
|
||||
assert cm.reencrypt_value(v2) == v2
|
||||
|
||||
|
||||
def test_reencrypt_value_preserves_non_string_and_empty(salt_key, monkeypatch):
|
||||
_enable_aes(monkeypatch)
|
||||
assert cm.reencrypt_value(42) == 42
|
||||
assert cm.reencrypt_value("") == ""
|
||||
assert cm.reencrypt_value(None) is None
|
||||
|
||||
|
||||
def test_reencrypt_value_skips_undecryptable(salt_key, monkeypatch):
|
||||
"""A value that does not decrypt (legacy plaintext or corrupt) is preserved."""
|
||||
_enable_aes(monkeypatch)
|
||||
plaintext = "not-actually-encrypted"
|
||||
assert cm.reencrypt_value(plaintext) == plaintext
|
||||
|
||||
|
||||
def test_reencrypt_selective_dict(salt_key, monkeypatch):
|
||||
legacy_key = _legacy_ct("the-api-key", monkeypatch)
|
||||
_enable_aes(monkeypatch)
|
||||
|
||||
data = {"api_key": legacy_key, "base_url": "https://x", "integration_token": None}
|
||||
out = cm.reencrypt_selective_dict(data, ["api_key", "integration_token"])
|
||||
|
||||
assert out["api_key"].startswith(_V2_GCM_PREFIX)
|
||||
assert out["base_url"] == "https://x" # untouched non-sensitive
|
||||
assert out["integration_token"] is None # null skipped
|
||||
|
||||
|
||||
# --------------------------- gate enforcement ---------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_migrate_requires_aes_gate(salt_key, monkeypatch):
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {}) # gate off
|
||||
with pytest.raises(RuntimeError, match="encryption_algorithm"):
|
||||
await cm.migrate_encryption(
|
||||
prisma_client=MagicMock(), user_api_key_dict=MagicMock()
|
||||
)
|
||||
|
||||
|
||||
# --------------------------- config-row walker ---------------------------
|
||||
|
||||
|
||||
def _config_prisma(record):
|
||||
"""Build an AsyncMock prisma client whose litellm_config returns `record`."""
|
||||
client = MagicMock()
|
||||
client.db.litellm_config.find_unique = AsyncMock(return_value=record)
|
||||
client.db.litellm_config.update = AsyncMock()
|
||||
return client
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vantage_walker_migrates_legacy_field(salt_key, monkeypatch):
|
||||
legacy_api_key = _legacy_ct("vantage-secret", monkeypatch)
|
||||
_enable_aes(monkeypatch)
|
||||
record = SimpleNamespace(
|
||||
param_value={
|
||||
"api_key": legacy_api_key,
|
||||
"integration_token": None,
|
||||
"base_url": "https://api.vantage.sh",
|
||||
}
|
||||
)
|
||||
client = _config_prisma(record)
|
||||
|
||||
report = await cm._migrate_config_settings_row(
|
||||
client, "vantage_settings", cm._VANTAGE_SENSITIVE, dry_run=False
|
||||
)
|
||||
|
||||
assert report.migrated == 1
|
||||
assert report.legacy == 0 # migrated -> no longer residual legacy
|
||||
client.db.litellm_config.update.assert_awaited_once()
|
||||
written = json.loads(
|
||||
client.db.litellm_config.update.call_args.kwargs["data"]["param_value"]
|
||||
)
|
||||
assert written["api_key"].startswith(_V2_GCM_PREFIX)
|
||||
assert written["base_url"] == "https://api.vantage.sh" # non-sensitive untouched
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vantage_walker_idempotent_no_write(salt_key, monkeypatch):
|
||||
_enable_aes(monkeypatch)
|
||||
record = SimpleNamespace(
|
||||
param_value={"api_key": encrypt_value_helper("already-v2"), "base_url": "x"}
|
||||
)
|
||||
client = _config_prisma(record)
|
||||
|
||||
report = await cm._migrate_config_settings_row(
|
||||
client, "vantage_settings", cm._VANTAGE_SENSITIVE, dry_run=False
|
||||
)
|
||||
|
||||
assert report.already_v2 == 1
|
||||
assert report.migrated == 0
|
||||
client.db.litellm_config.update.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_config_walker_dry_run_does_not_write(salt_key, monkeypatch):
|
||||
legacy_api_key = _legacy_ct("vantage-secret", monkeypatch)
|
||||
_enable_aes(monkeypatch)
|
||||
record = SimpleNamespace(param_value={"api_key": legacy_api_key})
|
||||
client = _config_prisma(record)
|
||||
|
||||
report = await cm._migrate_config_settings_row(
|
||||
client, "vantage_settings", cm._VANTAGE_SENSITIVE, dry_run=True
|
||||
)
|
||||
|
||||
# A dry run reports residual legacy only; nothing is migrated (no write), so
|
||||
# `migrated` and `residual_legacy` are never contradictory in --check output.
|
||||
assert report.legacy == 1
|
||||
assert report.migrated == 0
|
||||
client.db.litellm_config.update.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_config_walker_handles_missing_row(salt_key, monkeypatch):
|
||||
_enable_aes(monkeypatch)
|
||||
client = _config_prisma(None)
|
||||
report = await cm._migrate_config_settings_row(
|
||||
client, "cloudzero_settings", cm._CLOUDZERO_SENSITIVE, dry_run=False
|
||||
)
|
||||
assert report.scanned == 0
|
||||
client.db.litellm_config.update.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sso_walker_real_run_migrates_and_clears_residual(salt_key, monkeypatch):
|
||||
"""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"})
|
||||
client = MagicMock()
|
||||
client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=record)
|
||||
client.db.litellm_ssoconfig.update = AsyncMock()
|
||||
|
||||
report = await cm._migrate_sso_config(client, dry_run=False)
|
||||
|
||||
assert report.migrated == 1
|
||||
assert report.legacy == 0 # migrated -> no longer residual
|
||||
client.db.litellm_ssoconfig.update.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sso_walker_dry_run_reports_residual_not_migrated(salt_key, monkeypatch):
|
||||
"""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})
|
||||
client = MagicMock()
|
||||
client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=record)
|
||||
client.db.litellm_ssoconfig.update = AsyncMock()
|
||||
|
||||
report = await cm._migrate_sso_config(client, dry_run=True)
|
||||
|
||||
assert report.legacy == 1
|
||||
assert report.migrated == 0
|
||||
client.db.litellm_ssoconfig.update.assert_not_awaited()
|
||||
|
||||
|
||||
# --------------------------- --check scanner ---------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_reports_residual_legacy(salt_key, monkeypatch):
|
||||
legacy_api_key = _legacy_ct("vantage-secret", monkeypatch)
|
||||
_enable_aes(monkeypatch)
|
||||
|
||||
client = MagicMock()
|
||||
# Net-new walker tables: empty team / token / sso, one legacy vantage field.
|
||||
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.update = AsyncMock()
|
||||
_empty_covered_tables(client)
|
||||
|
||||
def _find_unique(where):
|
||||
if where.get("param_name") == "vantage_settings":
|
||||
return SimpleNamespace(param_value={"api_key": legacy_api_key})
|
||||
return None
|
||||
|
||||
client.db.litellm_config.find_unique = AsyncMock(side_effect=_find_unique)
|
||||
|
||||
report = await cm.check_encryption(client)
|
||||
|
||||
assert report.residual_legacy == 1
|
||||
client.db.litellm_config.update.assert_not_awaited() # read-only
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_reports_zero_after_migration(salt_key, monkeypatch):
|
||||
_enable_aes(monkeypatch)
|
||||
client = MagicMock()
|
||||
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.update = AsyncMock()
|
||||
_empty_covered_tables(client)
|
||||
|
||||
def _find_unique(where):
|
||||
if where.get("param_name") == "vantage_settings":
|
||||
return SimpleNamespace(
|
||||
param_value={"api_key": encrypt_value_helper("already-v2")}
|
||||
)
|
||||
return None
|
||||
|
||||
client.db.litellm_config.find_unique = AsyncMock(side_effect=_find_unique)
|
||||
|
||||
report = await cm.check_encryption(client)
|
||||
assert report.residual_legacy == 0
|
||||
|
||||
|
||||
# --------------------------- callback_vars walker ---------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_callback_vars_walker_migrates_team_metadata(salt_key, monkeypatch):
|
||||
"""A team row with a legacy-encrypted callback var is rewritten to v2."""
|
||||
from litellm.proxy.common_utils.callback_utils import encrypt_callback_vars
|
||||
|
||||
# Legacy-encrypt a callback var via the real callback path (gate off).
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {})
|
||||
legacy_meta = encrypt_callback_vars(
|
||||
{"logging": [{"callback_vars": {"gcs_path_service_account": "sa-secret"}}]}
|
||||
)
|
||||
_enable_aes(monkeypatch)
|
||||
|
||||
team_row = SimpleNamespace(team_id="team-1", metadata=legacy_meta)
|
||||
client = MagicMock()
|
||||
client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row])
|
||||
client.db.litellm_teamtable.update = AsyncMock()
|
||||
|
||||
report = await cm._migrate_callback_vars_table(client, "team", dry_run=False)
|
||||
|
||||
assert report.migrated == 1
|
||||
assert report.scanned == 1 # one field examined, not "post-v2" count
|
||||
client.db.litellm_teamtable.update.assert_awaited_once()
|
||||
written = json.loads(
|
||||
client.db.litellm_teamtable.update.call_args.kwargs["data"]["metadata"]
|
||||
)
|
||||
inner = written["logging"][0]["callback_vars"]["gcs_path_service_account"]
|
||||
assert "v2:gcm:" in inner
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_callback_vars_walker_dry_run_reports_legacy(salt_key, monkeypatch):
|
||||
"""In --check (dry-run) mode, a legacy callback var counts as residual legacy."""
|
||||
from litellm.proxy.common_utils.callback_utils import encrypt_callback_vars
|
||||
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {})
|
||||
legacy_meta = encrypt_callback_vars(
|
||||
{"logging": [{"callback_vars": {"gcs_path_service_account": "sa-secret"}}]}
|
||||
)
|
||||
_enable_aes(monkeypatch)
|
||||
|
||||
team_row = SimpleNamespace(team_id="team-1", metadata=legacy_meta)
|
||||
client = MagicMock()
|
||||
client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row])
|
||||
client.db.litellm_teamtable.update = AsyncMock()
|
||||
|
||||
report = await cm._migrate_callback_vars_table(client, "team", dry_run=True)
|
||||
|
||||
assert report.scanned == 1
|
||||
assert report.legacy == 1 # would-migrate -> residual legacy in attestation
|
||||
assert report.migrated == 0
|
||||
client.db.litellm_teamtable.update.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_callback_vars_walker_migrates_callback_settings_shape(
|
||||
salt_key, monkeypatch
|
||||
):
|
||||
"""Regression: credentials under ``metadata.callback_settings.callback_vars``
|
||||
with no top-level ``logging`` key must be migrated, not skipped.
|
||||
|
||||
The walker previously early-continued on ``"logging" not in metadata``, so
|
||||
this credential shape (which ``encrypt_callback_vars`` does encrypt) was left
|
||||
in legacy format at rest while the migration still reported success.
|
||||
"""
|
||||
from litellm.proxy.common_utils.callback_utils import encrypt_callback_vars
|
||||
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {})
|
||||
legacy_meta = encrypt_callback_vars(
|
||||
{
|
||||
"callback_settings": {
|
||||
"callback_vars": {"gcs_path_service_account": "sa-secret"}
|
||||
}
|
||||
}
|
||||
)
|
||||
_enable_aes(monkeypatch)
|
||||
assert "logging" not in legacy_meta # the shape that used to be skipped
|
||||
|
||||
team_row = SimpleNamespace(team_id="team-1", metadata=legacy_meta)
|
||||
client = MagicMock()
|
||||
client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row])
|
||||
client.db.litellm_teamtable.update = AsyncMock()
|
||||
|
||||
report = await cm._migrate_callback_vars_table(client, "team", dry_run=False)
|
||||
|
||||
assert report.migrated == 1
|
||||
assert report.scanned == 1
|
||||
client.db.litellm_teamtable.update.assert_awaited_once()
|
||||
written = json.loads(
|
||||
client.db.litellm_teamtable.update.call_args.kwargs["data"]["metadata"]
|
||||
)
|
||||
inner = written["callback_settings"]["callback_vars"]["gcs_path_service_account"]
|
||||
assert "v2:gcm:" in inner
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_reports_callback_var_legacy_with_gate_off(salt_key, monkeypatch):
|
||||
"""check_encryption must report residual legacy callback vars even when the
|
||||
AES gate is OFF.
|
||||
|
||||
Detection is decrypt-based, not a re-encrypt delta, so it does not depend on
|
||||
the write gate. A heuristic that re-encrypts and counts new v2 values would
|
||||
read zero here (gate off -> no v2 produced) and emit a false-clean
|
||||
attestation -- exactly the compliance trap this guards against.
|
||||
"""
|
||||
from litellm.proxy.common_utils.callback_utils import encrypt_callback_vars
|
||||
|
||||
# Legacy-encrypt a callback var, and leave the gate OFF for the check itself.
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {})
|
||||
legacy_meta = encrypt_callback_vars(
|
||||
{"logging": [{"callback_vars": {"gcs_path_service_account": "sa-secret"}}]}
|
||||
)
|
||||
team_row = SimpleNamespace(team_id="team-1", metadata=legacy_meta)
|
||||
|
||||
client = MagicMock()
|
||||
_empty_covered_tables(client)
|
||||
client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row])
|
||||
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_teamtable.update = AsyncMock()
|
||||
|
||||
report = await cm.check_encryption(client)
|
||||
|
||||
assert report.residual_legacy == 1
|
||||
assert report.as_dict()["locations"]["team.callback_vars"]["legacy"] == 1
|
||||
client.db.litellm_teamtable.update.assert_not_awaited() # read-only
|
||||
|
||||
|
||||
# --------------------------- covered-tables scanner ---------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scan_covered_tables_classifies_legacy_and_v2(salt_key, monkeypatch):
|
||||
"""The read-only scanner classifies the model and credentials tables."""
|
||||
legacy = _legacy_ct("model-secret", monkeypatch)
|
||||
_enable_aes(monkeypatch)
|
||||
v2 = encrypt_value_helper("cred-secret")
|
||||
|
||||
client = MagicMock()
|
||||
_empty_covered_tables(client)
|
||||
client.db.litellm_proxymodeltable.find_many = AsyncMock(
|
||||
return_value=[
|
||||
SimpleNamespace(litellm_params={"api_key": legacy, "model": "gpt-4"})
|
||||
]
|
||||
)
|
||||
client.db.litellm_credentialstable.find_many = AsyncMock(
|
||||
return_value=[SimpleNamespace(credential_values={"api_key": v2})]
|
||||
)
|
||||
client.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
by_loc = {r.location: r for r in await cm._scan_covered_tables(client)}
|
||||
|
||||
assert by_loc["model_table"].legacy == 1
|
||||
assert by_loc["model_table"].plaintext == 1 # "gpt-4" model name, not ciphertext
|
||||
assert by_loc["credentials"].already_v2 == 1
|
||||
assert by_loc["credentials"].legacy == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_counts_covered_table_residual(salt_key, monkeypatch):
|
||||
"""check_encryption now scans the rotation-covered tables (model table here),
|
||||
so a legacy value there counts toward residual_legacy (the P1 attestation gap).
|
||||
"""
|
||||
legacy = _legacy_ct("model-secret", monkeypatch)
|
||||
_enable_aes(monkeypatch)
|
||||
|
||||
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": legacy})]
|
||||
)
|
||||
|
||||
report = await cm.check_encryption(client)
|
||||
|
||||
assert report.residual_legacy == 1
|
||||
assert report.as_dict()["locations"]["model_table"]["legacy"] == 1
|
||||
client.db.litellm_config.update.assert_not_awaited() # read-only
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_migrate_covered_tables_reports_real_counts(salt_key, monkeypatch):
|
||||
"""_migrate_covered_tables derives real per-table counts from pre/post scans,
|
||||
instead of the always-zero report Greptile flagged (P1).
|
||||
"""
|
||||
legacy = _legacy_ct("model-secret", monkeypatch)
|
||||
_enable_aes(monkeypatch)
|
||||
v2 = encrypt_value_helper("model-secret")
|
||||
|
||||
row = SimpleNamespace(litellm_params={"api_key": legacy})
|
||||
client = MagicMock()
|
||||
_empty_covered_tables(client)
|
||||
client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[row])
|
||||
client.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
async def fake_rotate(**kwargs):
|
||||
# Stand in for _rotate_master_key: re-encrypt the model api_key in place.
|
||||
row.litellm_params["api_key"] = v2
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._rotate_master_key",
|
||||
fake_rotate,
|
||||
)
|
||||
|
||||
by_loc = {
|
||||
r.location: r for r in await cm._migrate_covered_tables(client, MagicMock())
|
||||
}
|
||||
|
||||
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
|
||||
|
|
@ -0,0 +1,100 @@
|
|||
"""Unit tests for the at-rest encryption-migration HTTP endpoints.
|
||||
|
||||
The endpoint bodies are exercised directly with the migration engine mocked, so
|
||||
the admin guard, db-not-connected guard, and success path are all covered
|
||||
without touching a live DB. The live ASGI/auth contract is covered separately in
|
||||
``tests/proxy_behavior/management/test_credential_migration_endpoint.py``.
|
||||
"""
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
from litellm.proxy.management_endpoints import credential_migration as cm
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
check_encryption_endpoint,
|
||||
migrate_encryption_endpoint,
|
||||
)
|
||||
|
||||
ADMIN = SimpleNamespace(user_role=LitellmUserRoles.PROXY_ADMIN.value)
|
||||
NONADMIN = SimpleNamespace(user_role=LitellmUserRoles.INTERNAL_USER.value)
|
||||
|
||||
|
||||
def _sample_report() -> cm.MigrationReport:
|
||||
report = cm.MigrationReport()
|
||||
report.add(
|
||||
cm.LocationReport(location="model_table", scanned=2, migrated=1, legacy=0)
|
||||
)
|
||||
return report
|
||||
|
||||
|
||||
# ------------------------------- check endpoint -------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_endpoint_success(monkeypatch):
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", object())
|
||||
monkeypatch.setattr(
|
||||
cm, "check_encryption", AsyncMock(return_value=_sample_report())
|
||||
)
|
||||
|
||||
out = await check_encryption_endpoint(user_api_key_dict=ADMIN)
|
||||
|
||||
assert out["status"] == "success"
|
||||
assert out["report"]["residual_legacy"] == 0
|
||||
assert out["report"]["locations"]["model_table"]["scanned"] == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_endpoint_requires_admin(monkeypatch):
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", object())
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await check_encryption_endpoint(user_api_key_dict=NONADMIN)
|
||||
assert exc.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_endpoint_db_not_connected(monkeypatch):
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await check_encryption_endpoint(user_api_key_dict=ADMIN)
|
||||
assert exc.value.status_code == 500
|
||||
|
||||
|
||||
# ------------------------------ migrate endpoint ------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("dry_run", [False, True])
|
||||
async def test_migrate_endpoint_success(monkeypatch, dry_run):
|
||||
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, dry_run=dry_run)
|
||||
|
||||
assert out["status"] == "success"
|
||||
assert out["dry_run"] is dry_run
|
||||
assert out["report"]["locations"]["model_table"]["migrated"] == 1
|
||||
# dry_run is threaded through to the engine unchanged.
|
||||
assert fake.await_args.kwargs["dry_run"] is dry_run
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_migrate_endpoint_requires_admin(monkeypatch):
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", object())
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await migrate_encryption_endpoint(user_api_key_dict=NONADMIN)
|
||||
assert exc.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_migrate_endpoint_db_not_connected(monkeypatch):
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await migrate_encryption_endpoint(user_api_key_dict=ADMIN)
|
||||
assert exc.value.status_code == 500
|
||||
98
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
98
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -2479,6 +2479,52 @@ export interface paths {
|
|||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/credentials/migrate-encryption": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
get?: never;
|
||||
put?: never;
|
||||
/**
|
||||
* Migrate Encryption Endpoint
|
||||
* @description Re-encrypt all at-rest credentials into the AES-256-GCM (``v2:gcm:``) format.
|
||||
*
|
||||
* 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``).
|
||||
*/
|
||||
post: operations["migrate_encryption_endpoint_credentials_migrate_encryption_post"];
|
||||
delete?: never;
|
||||
options?: never;
|
||||
head?: never;
|
||||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/credentials/migrate-encryption/check": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
/**
|
||||
* 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.
|
||||
*/
|
||||
get: operations["check_encryption_endpoint_credentials_migrate_encryption_check_get"];
|
||||
put?: never;
|
||||
post?: never;
|
||||
delete?: never;
|
||||
options?: never;
|
||||
head?: never;
|
||||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/credentials/{credential_name}": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
@ -37122,6 +37168,58 @@ export interface operations {
|
|||
};
|
||||
};
|
||||
};
|
||||
migrate_encryption_endpoint_credentials_migrate_encryption_post: {
|
||||
parameters: {
|
||||
query?: {
|
||||
/** @description If true, scan and report without writing any changes. */
|
||||
dry_run?: boolean;
|
||||
};
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": unknown;
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
check_encryption_endpoint_credentials_migrate_encryption_check_get: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": unknown;
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
delete_credential_credentials__credential_name__delete: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue