diff --git a/litellm/secret_managers/aws_secret_manager_v2.py b/litellm/secret_managers/aws_secret_manager_v2.py index e86c8e7c919..d75375a01cc 100644 --- a/litellm/secret_managers/aws_secret_manager_v2.py +++ b/litellm/secret_managers/aws_secret_manager_v2.py @@ -47,6 +47,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): aws_web_identity_token: str | None = None, aws_sts_endpoint: str | None = None, replica_regions: list[str] | None = None, + kms_key_id: str | None = None, **kwargs, ): BaseSecretManager.__init__(self, **kwargs) @@ -61,6 +62,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): self.aws_web_identity_token = aws_web_identity_token self.aws_sts_endpoint = aws_sts_endpoint self.replica_regions: list[str] = replica_regions or [] + self.kms_key_id = kms_key_id @classmethod def validate_environment(cls): @@ -106,7 +108,8 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): # Remove None values aws_kwargs = {k: v for k, v in aws_kwargs.items() if v is not None} - litellm.secret_manager_client = cls(**aws_kwargs) + kms_key_id: Final = key_management_settings.kms_key_id if key_management_settings is not None else None + litellm.secret_manager_client = cls(kms_key_id=kms_key_id, **aws_kwargs) litellm._key_management_system = KeyManagementSystem.AWS_SECRET_MANAGER except Exception as e: @@ -275,6 +278,9 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): if description: data["Description"] = description + if self.kms_key_id: + data["KmsKeyId"] = self.kms_key_id + # ✅ Normalize tags to AWS format if tags: if isinstance(tags, dict): diff --git a/litellm/types/secret_managers/main.py b/litellm/types/secret_managers/main.py index 599e5746dfb..148e680a236 100644 --- a/litellm/types/secret_managers/main.py +++ b/litellm/types/secret_managers/main.py @@ -45,6 +45,9 @@ class KeyManagementSettings(LiteLLMPydanticObjectBase): tags: dict[str, str] | None = None """Optional tags to attach when creating secrets (e.g. {"Environment": "Prod", "Owner": "AI-Platform"}).""" + kms_key_id: str | None = None + """Optional customer-managed KMS key (ID, alias or ARN) used to encrypt secrets created in AWS Secrets Manager.""" + custom_secret_manager: str | None = None """ Path to custom secret manager class (e.g. "my_secret_manager.InMemorySecretManager") diff --git a/tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py b/tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py index 7e655b70756..03422d5433c 100644 --- a/tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py +++ b/tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py @@ -5,11 +5,68 @@ Tests the write/read/delete cycle for JSON and simple string secrets. """ import json -from unittest.mock import AsyncMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest +import respx +import litellm from litellm.secret_managers.aws_secret_manager_v2 import AWSSecretsManagerV2 +from litellm.types.secret_managers.main import KeyManagementSettings + +_STATIC_CREDENTIALS = {"aws_access_key_id": "test-key", "aws_secret_access_key": "test-secret"} +_CMK_ARN = "arn:aws:kms:us-east-1:123456789012:key/11111111-2222-3333-4444-555555555555" + + +async def _create_secret_body_for_settings( + monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter, settings: KeyManagementSettings +) -> dict[str, object]: + """Boot the manager from settings the way the proxy does and return the CreateSecret body it posts to AWS.""" + monkeypatch.setattr(litellm, "secret_manager_client", None) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", None) + monkeypatch.setenv("AWS_REGION_NAME", "us-east-1") + AWSSecretsManagerV2.load_aws_secret_manager(use_aws_secret_manager=True, key_management_settings=settings) + manager = litellm.secret_manager_client + assert isinstance(manager, AWSSecretsManagerV2) + + route = respx_mock.post("https://secretsmanager.us-east-1.amazonaws.com/").respond( + json={"ARN": "arn", "Name": "litellm/test-key"} + ) + await manager.async_write_secret( + secret_name="litellm/test-key", + secret_value="sk-test-value", + optional_params=dict(_STATIC_CREDENTIALS), + ) + assert route.call_count == 1 + request = route.calls.last.request + assert request.headers["X-Amz-Target"] == "secretsmanager.CreateSecret" + return json.loads(request.content) + + +@pytest.mark.asyncio +async def test_create_secret_uses_customer_managed_kms_key_from_settings( + monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter +) -> None: + body = await _create_secret_body_for_settings( + monkeypatch, + respx_mock, + KeyManagementSettings(store_virtual_keys=True, aws_region_name="us-east-1", kms_key_id=_CMK_ARN), + ) + assert body["KmsKeyId"] == _CMK_ARN + assert body["Name"] == "litellm/test-key" + assert body["SecretString"] == "sk-test-value" + + +@pytest.mark.asyncio +async def test_create_secret_omits_kms_key_id_when_not_configured( + monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter +) -> None: + body = await _create_secret_body_for_settings( + monkeypatch, respx_mock, KeyManagementSettings(store_virtual_keys=True, aws_region_name="us-east-1") + ) + assert "KmsKeyId" not in body + assert body["Name"] == "litellm/test-key" @pytest.mark.asyncio