diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index c644ecc3dae..26988609140 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -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) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 71cf2db3dfb..3547daa859a 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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 diff --git a/tests/test_litellm/proxy/common_utils/test_callback_utils.py b/tests/test_litellm/proxy/common_utils/test_callback_utils.py index 36ff3f3c399..fd0210a2e4c 100644 --- a/tests/test_litellm/proxy/common_utils/test_callback_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_callback_utils.py @@ -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" + ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 4fb3df52cf6..4f96a76bb69 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -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