mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
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:
parent
b8248a21d2
commit
35d423398d
4 changed files with 239 additions and 4 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue