mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
feat(secret_managers): support customer-managed KMS key for virtual keys stored in AWS Secrets Manager (#40475)
Co-authored-by: yassin <yassin@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
47bba14336
commit
de79310954
3 changed files with 68 additions and 2 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue