diff --git a/litellm/proxy/hooks/key_management_event_hooks.py b/litellm/proxy/hooks/key_management_event_hooks.py index e8d5c2abd42..f75197532b4 100644 --- a/litellm/proxy/hooks/key_management_event_hooks.py +++ b/litellm/proxy/hooks/key_management_event_hooks.py @@ -21,6 +21,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.utils import _hash_token_if_needed +from litellm.secret_managers.base_secret_manager import BaseSecretManager # NOTE: This is the prefix for all virtual keys stored in AWS Secrets Manager LITELLM_PREFIX_STORED_VIRTUAL_KEYS: Final = "litellm/" @@ -100,6 +101,7 @@ class KeyManagementEventHooks: Post /key/update processing hook Handles the following: + - Renaming the key's secret in the secret manager when the alias changes - Storing Audit Logs for key update """ from litellm.proxy.management_helpers.audit_logs import ( @@ -109,6 +111,16 @@ class KeyManagementEventHooks: ) from litellm.proxy.proxy_server import litellm_proxy_admin_name + if data.key_alias is not None and data.key_alias != existing_key_row.key_alias: + try: + await KeyManagementEventHooks._rename_virtual_key_in_secret_manager( + current_secret_name=existing_key_row.key_alias or f"virtual-key-{existing_key_row.token}", + new_secret_name=data.key_alias, + team_id=existing_key_row.team_id, + ) + except Exception as e: + verbose_proxy_logger.warning("Failed to rename virtual key in secret manager: %s", e) + if is_audit_logging_enabled(): updated_fields: Final = { **data.model_dump(exclude_none=True), @@ -306,21 +318,66 @@ class KeyManagementEventHooks: new_secret_value: New value of the virtual key (example: sk-1234) team_id: Optional team ID to get team-specific secret manager settings """ - if litellm._key_management_settings is not None: - if litellm._key_management_settings.store_virtual_keys is True: - from litellm.secret_managers.base_secret_manager import ( - BaseSecretManager, - ) + secret_manager: Final = KeyManagementEventHooks._stored_virtual_key_secret_manager() + if secret_manager is None: + return + optional_params: Final = await KeyManagementEventHooks._get_secret_manager_optional_params(team_id) + await secret_manager.async_rotate_secret( + current_secret_name=KeyManagementEventHooks._get_secret_name(current_secret_name), + new_secret_name=KeyManagementEventHooks._get_secret_name(new_secret_name), + new_secret_value=new_secret_value, + optional_params=optional_params, + ) - # store the key in the secret manager - if isinstance(litellm.secret_manager_client, BaseSecretManager): - optional_params: Final = await KeyManagementEventHooks._get_secret_manager_optional_params(team_id) - await litellm.secret_manager_client.async_rotate_secret( - current_secret_name=KeyManagementEventHooks._get_secret_name(current_secret_name), - new_secret_name=KeyManagementEventHooks._get_secret_name(new_secret_name), - new_secret_value=new_secret_value, - optional_params=optional_params, - ) + @staticmethod + def _stored_virtual_key_secret_manager() -> BaseSecretManager | None: + """ + The secret manager client that stores virtual keys, or None when virtual keys are not stored in one + """ + if litellm._key_management_settings is None or litellm._key_management_settings.store_virtual_keys is not True: + return None + if not isinstance(litellm.secret_manager_client, BaseSecretManager): + return None + return litellm.secret_manager_client + + @staticmethod + async def _rename_virtual_key_in_secret_manager( + current_secret_name: str, + new_secret_name: str, + team_id: str | None = None, + ) -> None: + """ + Move a virtual key to a new secret name, keeping its current value + + Args: + current_secret_name: Current name of the virtual key + new_secret_name: New name of the virtual key + team_id: Optional team ID to get team-specific secret manager settings + """ + secret_manager: Final = KeyManagementEventHooks._stored_virtual_key_secret_manager() + if secret_manager is None: + return + optional_params: Final = await KeyManagementEventHooks._get_secret_manager_optional_params(team_id) + current_secret_value: Final = await secret_manager.async_read_secret( + secret_name=KeyManagementEventHooks._get_secret_name(current_secret_name), + optional_params=optional_params, + ) + if current_secret_value is None: + verbose_proxy_logger.warning( + "Secret %s not found in secret manager, skipping rename to %s", current_secret_name, new_secret_name + ) + return + verbose_proxy_logger.info( + "Renaming secret in secret manager: current_secret_name=%s new_secret_name=%s", + current_secret_name, + new_secret_name, + ) + await secret_manager.async_rotate_secret( + current_secret_name=KeyManagementEventHooks._get_secret_name(current_secret_name), + new_secret_name=KeyManagementEventHooks._get_secret_name(new_secret_name), + new_secret_value=current_secret_value, + optional_params=optional_params, + ) @staticmethod def _get_secret_name(secret_name: str) -> str: diff --git a/tests/test_litellm/proxy/hooks/test_key_management_event_hooks.py b/tests/test_litellm/proxy/hooks/test_key_management_event_hooks.py index e9aa3c4c701..ede8c89a2ef 100644 --- a/tests/test_litellm/proxy/hooks/test_key_management_event_hooks.py +++ b/tests/test_litellm/proxy/hooks/test_key_management_event_hooks.py @@ -514,6 +514,117 @@ class TestRotateVirtualKeyInSecretManager: mock_secret_manager.async_rotate_secret.assert_not_called() +class TestKeyUpdatedSecretManagerSync: + """Tests that /key/update moves the stored secret when the key alias changes.""" + + @staticmethod + def _configure_secret_manager( + monkeypatch: pytest.MonkeyPatch, stored_value: str | None, store_virtual_keys: bool = True + ) -> MagicMock: + import litellm + from litellm.secret_managers.base_secret_manager import BaseSecretManager + from litellm.types.secret_managers.main import KeyManagementSettings, KeyManagementSystem + + mock_secret_manager: Final = MagicMock(spec=BaseSecretManager) + mock_secret_manager.async_read_secret = AsyncMock(return_value=stored_value) + mock_secret_manager.async_rotate_secret = AsyncMock(return_value={"status": "success"}) + monkeypatch.setattr(litellm, "secret_manager_client", mock_secret_manager) + monkeypatch.setattr(litellm, "_key_management_system", KeyManagementSystem.AWS_SECRET_MANAGER) + monkeypatch.setattr( + litellm, + "_key_management_settings", + KeyManagementSettings(store_virtual_keys=store_virtual_keys, prefix_for_stored_virtual_keys="litellm/"), + ) + monkeypatch.setattr(litellm, "store_audit_logs", False) + return mock_secret_manager + + @pytest.mark.parametrize("existing_alias", ["old-alias", None]) + @pytest.mark.asyncio + async def test_updated_hook_renames_secret_when_alias_changes( + self, monkeypatch: pytest.MonkeyPatch, existing_alias: str | None + ): + """A new alias on /key/update must move the secret to the new name, keeping the stored key value.""" + from litellm.proxy._types import LiteLLM_VerificationToken, UpdateKeyRequest + + mock_secret_manager: Final = self._configure_secret_manager(monkeypatch, stored_value="sk-stored-key") + existing_key_row: Final = LiteLLM_VerificationToken(token="hashed-token", key_alias=existing_alias) + + await KeyManagementEventHooks.async_key_updated_hook( + data=UpdateKeyRequest(key="hashed-token", key_alias="new-alias"), + existing_key_row=existing_key_row, + response=MagicMock(), + user_api_key_dict=MagicMock(), + ) + + current_secret_name: Final = f"litellm/{existing_alias or 'virtual-key-hashed-token'}" + mock_secret_manager.async_read_secret.assert_awaited_once_with( + secret_name=current_secret_name, optional_params=None + ) + mock_secret_manager.async_rotate_secret.assert_awaited_once_with( + current_secret_name=current_secret_name, + new_secret_name="litellm/new-alias", + new_secret_value="sk-stored-key", + optional_params=None, + ) + + @pytest.mark.parametrize("requested_alias", ["same-alias", None]) + @pytest.mark.asyncio + async def test_updated_hook_leaves_secret_alone_when_alias_unchanged( + self, monkeypatch: pytest.MonkeyPatch, requested_alias: str | None + ): + """An update that keeps or omits the alias must not touch the secret manager.""" + from litellm.proxy._types import LiteLLM_VerificationToken, UpdateKeyRequest + + mock_secret_manager: Final = self._configure_secret_manager(monkeypatch, stored_value="sk-stored-key") + + await KeyManagementEventHooks.async_key_updated_hook( + data=UpdateKeyRequest(key="hashed-token", key_alias=requested_alias, max_budget=10.0), + existing_key_row=LiteLLM_VerificationToken(token="hashed-token", key_alias="same-alias"), + response=MagicMock(), + user_api_key_dict=MagicMock(), + ) + + mock_secret_manager.async_read_secret.assert_not_awaited() + mock_secret_manager.async_rotate_secret.assert_not_awaited() + + @pytest.mark.asyncio + async def test_updated_hook_skips_rename_when_secret_missing(self, monkeypatch: pytest.MonkeyPatch): + """If the key was never stored under its current name there is nothing to move, so no secret is created.""" + from litellm.proxy._types import LiteLLM_VerificationToken, UpdateKeyRequest + + mock_secret_manager: Final = self._configure_secret_manager(monkeypatch, stored_value=None) + + await KeyManagementEventHooks.async_key_updated_hook( + data=UpdateKeyRequest(key="hashed-token", key_alias="new-alias"), + existing_key_row=LiteLLM_VerificationToken(token="hashed-token", key_alias="old-alias"), + response=MagicMock(), + user_api_key_dict=MagicMock(), + ) + + mock_secret_manager.async_rotate_secret.assert_not_awaited() + + @pytest.mark.asyncio + async def test_updated_hook_ignores_alias_change_when_store_virtual_keys_disabled( + self, monkeypatch: pytest.MonkeyPatch + ): + """With store_virtual_keys off, an alias change must not read or write any secret.""" + from litellm.proxy._types import LiteLLM_VerificationToken, UpdateKeyRequest + + mock_secret_manager: Final = self._configure_secret_manager( + monkeypatch, stored_value="sk-stored-key", store_virtual_keys=False + ) + + await KeyManagementEventHooks.async_key_updated_hook( + data=UpdateKeyRequest(key="hashed-token", key_alias="new-alias"), + existing_key_row=LiteLLM_VerificationToken(token="hashed-token", key_alias="old-alias"), + response=MagicMock(), + user_api_key_dict=MagicMock(), + ) + + mock_secret_manager.async_read_secret.assert_not_awaited() + mock_secret_manager.async_rotate_secret.assert_not_awaited() + + class TestKeyUpdatedAuditLogObjectId: """Tests that /key/update audit logs never store the raw virtual key (issue #31620)."""