fix(proxy): re-encrypt team/key callback_vars on master key rotation

_rotate_master_key re-encrypted models, config env vars, MCP
credentials and the credentials table but skipped callback_vars stored
in LiteLLM_VerificationToken.metadata and LiteLLM_TeamTable.metadata, so
after a rotation those values stayed encrypted under the old key and
produced recurring decryption errors.

Thread new_encryption_key through encrypt_callback_vars and add
rotate_callback_vars_master_key, which decrypts each row's callback vars
with the current key and re-encrypts under the new key, then wire it
into _rotate_master_key.
This commit is contained in:
Devin AI 2026-07-07 20:36:39 +00:00
parent b8248a21d2
commit 35d423398d
4 changed files with 239 additions and 4 deletions

View file

@ -1,4 +1,6 @@
import copy
import json
from functools import partial
from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Literal, Optional
import litellm
@ -34,6 +36,7 @@ TRUSTED_PILLAR_RESPONSE_HEADERS_METADATA_KEY = "_pillar_response_headers_trusted
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.proxy.utils import PrismaClient
def initialize_callbacks_on_proxy(
@ -579,13 +582,17 @@ def normalize_callback_names(callbacks: Iterable[Any]) -> List[Any]:
return [c.lower() if isinstance(c, str) else c for c in callbacks]
def encrypt_callback_vars(metadata: Any) -> Any:
def encrypt_callback_vars(metadata: Any, new_encryption_key: Optional[str] = None) -> Any:
"""Return a deep copy of metadata with callback_vars values encrypted at rest.
Idempotent: a value that already decrypts cleanly is left unchanged so
round-trips through edit forms don't double-encrypt.
``new_encryption_key`` overrides the signing key used for encryption. Master
key rotation decrypts with the current key and re-encrypts under the new key
by passing it here, mirroring the model / credential / MCP rotation paths.
"""
return _transform_callback_vars(metadata, _encrypt_if_plaintext)
return _transform_callback_vars(metadata, partial(_encrypt_if_plaintext, new_encryption_key=new_encryption_key))
def decrypt_callback_vars(metadata: Any) -> Any:
@ -626,7 +633,7 @@ def is_sensitive_callback_key(
return _CALLBACK_VAR_MASKER.is_sensitive_key(key)
def _encrypt_if_plaintext(key: str, value: Any) -> Any:
def _encrypt_if_plaintext(key: str, value: Any, new_encryption_key: Optional[str] = None) -> Any:
if not isinstance(value, str) or not value:
return value
if not is_sensitive_callback_key(key):
@ -639,7 +646,7 @@ def _encrypt_if_plaintext(key: str, value: Any) -> Any:
# plaintext under K2 and wrap them a second time.
return value
try:
return _CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper(value)
return _CALLBACK_VAR_ENCRYPTED_PREFIX + encrypt_value_helper(value, new_encryption_key=new_encryption_key)
except Exception:
# No salt key / master key configured — leave the value as-is rather
# than crash the write. Dev environments without LITELLM_SALT_KEY hit
@ -656,3 +663,65 @@ def _decrypt_or_passthrough(key: str, value: Any) -> Any:
inner = value[len(_CALLBACK_VAR_ENCRYPTED_PREFIX) :]
decrypted = decrypt_value_helper(value=inner, key=key, exception_type="debug", return_original_value=False)
return decrypted if decrypted is not None else value
def _iter_callback_var_dicts(metadata: Dict[str, Any]) -> Iterable[Dict[str, Any]]:
for entry in metadata.get("logging", []) or []:
if isinstance(entry, dict) and isinstance(entry.get("callback_vars"), dict):
yield entry["callback_vars"]
callback_settings = metadata.get("callback_settings")
if isinstance(callback_settings, dict) and isinstance(callback_settings.get("callback_vars"), dict):
yield callback_settings["callback_vars"]
def _has_encrypted_callback_vars(metadata: Any) -> bool:
if not isinstance(metadata, dict):
return False
return any(
isinstance(v, str) and v.startswith(_CALLBACK_VAR_ENCRYPTED_PREFIX)
for callback_vars in _iter_callback_var_dicts(metadata)
for v in callback_vars.values()
)
async def rotate_callback_vars_master_key(prisma_client: "PrismaClient", new_master_key: str) -> None:
"""Re-encrypt team and verification-token ``callback_vars`` under ``new_master_key``.
Key-level (``LiteLLM_VerificationToken.metadata``) and team-level
(``LiteLLM_TeamTable.metadata``) logging callback credentials (e.g. Langfuse /
Langsmith secrets) are encrypted at rest but were not covered by master key
rotation, so after a rotation they stayed encrypted under the old key and
produced recurring decryption errors. This walks both tables and, for every
row that holds encrypted callback vars, decrypts with the current key and
re-encrypts under the new key via the proven ``decrypt_callback_vars`` /
``encrypt_callback_vars`` transforms.
"""
for table_name in ("team", "verification_token"):
await _rotate_callback_vars_table(prisma_client, table_name, new_master_key)
async def _rotate_callback_vars_table(
prisma_client: "PrismaClient",
table_name: Literal["team", "verification_token"],
new_master_key: str,
) -> None:
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()
rotated = 0
for row in rows or []:
metadata = getattr(row, "metadata", None)
if not _has_encrypted_callback_vars(metadata):
continue
re_encrypted = encrypt_callback_vars(decrypt_callback_vars(metadata), new_encryption_key=new_master_key)
await table.update(
where={pk: getattr(row, pk)},
data={"metadata": json.dumps(re_encrypted)},
)
rotated += 1
verbose_proxy_logger.info("rotate_callback_vars_master_key: rotated %d %s row(s)", rotated, table_name)

View file

@ -4256,6 +4256,19 @@ async def _rotate_master_key(
continue
verbose_proxy_logger.debug(f"Successfully re-encrypted {len(credentials)} credentials with new master key")
# 6. process team + verification-token callback_vars metadata
from litellm.proxy.common_utils.callback_utils import (
rotate_callback_vars_master_key,
)
try:
await rotate_callback_vars_master_key(
prisma_client=prisma_client,
new_master_key=new_master_key,
)
except prisma.errors.PrismaError as e:
verbose_proxy_logger.warning("Failed to rotate callback_vars: %s", str(e))
def _require_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> None:
from litellm.proxy._types import CommonProxyErrors

View file

@ -1,7 +1,9 @@
import copy
import json
import sys
import os
from types import ModuleType, SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -17,9 +19,11 @@ from litellm.proxy.common_utils.callback_utils import (
initialize_callbacks_on_proxy,
get_remaining_tokens_and_requests_from_request_data,
normalize_callback_names,
rotate_callback_vars_master_key,
sanitize_openai_provider_metadata,
)
import litellm
from litellm.proxy import proxy_server
from unittest.mock import patch
from litellm.proxy.common_utils.callback_utils import process_callback
@ -406,3 +410,95 @@ def test_initialize_callbacks_on_proxy_non_dict_callback_specific_params_root(
)
finally:
litellm.callbacks = original_callbacks
# ---------------------------------------------------------------------------
# master key rotation of callback_vars (LIT-1531)
# ---------------------------------------------------------------------------
def test_encrypt_callback_vars_uses_new_encryption_key(monkeypatch):
"""new_encryption_key must sign with the override, not the active salt key.
This is the primitive master key rotation relies on: values are re-encrypted
under the incoming key while the old key is still the active one.
"""
monkeypatch.setattr(proxy_server, "general_settings", {})
monkeypatch.setenv("LITELLM_SALT_KEY", "old-key-aaaaaaaaaaaaaaaaaaaaaaaa")
original = _sample_metadata()
encrypted = encrypt_callback_vars(
original, new_encryption_key="new-key-bbbbbbbbbbbbbbbbbbbbbbbb"
)
# The active (old) key cannot decrypt a value written under the new key, so
# the ciphertext falls through unchanged. If the override were ignored the
# old key would decrypt it and this assertion would fail.
stuck = decrypt_callback_vars(encrypted)["logging"][0]["callback_vars"]
assert stuck["langfuse_secret_key"].startswith("litellm_enc::")
monkeypatch.setenv("LITELLM_SALT_KEY", "new-key-bbbbbbbbbbbbbbbbbbbbbbbb")
recovered = decrypt_callback_vars(encrypted)
assert (
recovered["logging"][0]["callback_vars"]
== original["logging"][0]["callback_vars"]
)
assert (
recovered["callback_settings"]["callback_vars"]
== original["callback_settings"]["callback_vars"]
)
@pytest.mark.asyncio
async def test_rotate_callback_vars_master_key_reencrypts_under_new_key(monkeypatch):
"""Regression for LIT-1531: master key rotation must re-encrypt team and
verification-token callback_vars so they decrypt under the new key.
Rows without encrypted callback vars are left untouched.
"""
old_key = "old-key-aaaaaaaaaaaaaaaaaaaaaaaa"
new_key = "new-key-bbbbbbbbbbbbbbbbbbbbbbbb"
monkeypatch.setattr(proxy_server, "general_settings", {})
monkeypatch.setenv("LITELLM_SALT_KEY", old_key)
# Team row: real encrypted callback vars (under the current/old key).
team_meta = encrypt_callback_vars(_sample_metadata())
team_row = SimpleNamespace(team_id="team-1", metadata=team_meta)
# Token row: only a routing field (never encrypted) -> must be skipped.
token_row = SimpleNamespace(
token="tok-1",
metadata={"logging": [{"callback_vars": {"langfuse_host": "https://h"}}]},
)
client = MagicMock()
client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row])
client.db.litellm_teamtable.update = AsyncMock()
client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[token_row])
client.db.litellm_verificationtoken.update = AsyncMock()
await rotate_callback_vars_master_key(client, new_master_key=new_key)
client.db.litellm_teamtable.update.assert_awaited_once()
client.db.litellm_verificationtoken.update.assert_not_awaited()
written = json.loads(
client.db.litellm_teamtable.update.call_args.kwargs["data"]["metadata"]
)
# Still on the old key: the rewritten value no longer decrypts here.
stuck = decrypt_callback_vars(written)["logging"][0]["callback_vars"]
assert stuck["langfuse_secret_key"].startswith("litellm_enc::")
# New key recovers the plaintext for both metadata shapes.
monkeypatch.setenv("LITELLM_SALT_KEY", new_key)
recovered = decrypt_callback_vars(written)
assert (
recovered["logging"][0]["callback_vars"]["langfuse_secret_key"]
== "sk-lf-secret"
)
assert (
recovered["callback_settings"]["callback_vars"]["langsmith_api_key"]
== "ls-api-key"
)

View file

@ -7707,11 +7707,16 @@ async def test_key_does_not_override_explicit_budget_duration():
@pytest.mark.asyncio
@patch(
"litellm.proxy.common_utils.callback_utils.rotate_callback_vars_master_key",
new_callable=AsyncMock,
)
@patch(
"litellm.proxy.management_endpoints.key_management_endpoints.rotate_mcp_server_credentials_master_key"
)
async def test_rotate_master_key_model_data_valid_for_prisma(
mock_rotate_mcp,
mock_rotate_callback_vars,
):
"""
Test that _rotate_master_key produces valid data for Prisma create_many().
@ -7830,6 +7835,58 @@ async def test_rotate_master_key_model_data_valid_for_prisma(
mock_tx.litellm_proxymodeltable.delete_many.assert_called_once()
@pytest.mark.asyncio
@patch(
"litellm.proxy.common_utils.callback_utils.rotate_callback_vars_master_key",
new_callable=AsyncMock,
)
@patch(
"litellm.proxy.management_endpoints.key_management_endpoints.rotate_mcp_server_credentials_master_key"
)
async def test_rotate_master_key_rotates_callback_vars(
mock_rotate_mcp,
mock_rotate_callback_vars,
):
"""Regression for LIT-1531: _rotate_master_key must re-encrypt team and
verification-token callback_vars under the new master key.
"""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.management_endpoints.key_management_endpoints import (
_rotate_master_key,
)
mock_prisma_client = AsyncMock()
mock_prisma_client.db = MagicMock()
mock_prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
mock_prisma_client.db.litellm_config.find_many = AsyncMock(return_value=[])
mock_prisma_client.db.litellm_credentialstable.find_many = AsyncMock(
return_value=[]
)
mock_rotate_mcp.return_value = None
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-1234",
user_id="test-user",
)
with patch("litellm.proxy.proxy_server.proxy_config", MagicMock()):
await _rotate_master_key(
prisma_client=mock_prisma_client,
user_api_key_dict=user_api_key_dict,
current_master_key="sk-old-master-key",
new_master_key="sk-new-master-key",
)
mock_rotate_callback_vars.assert_awaited_once()
assert (
mock_rotate_callback_vars.await_args.kwargs["new_master_key"]
== "sk-new-master-key"
)
async def test_default_key_generate_params_duration(monkeypatch):
"""
Test that default_key_generate_params with 'duration' is applied