import abc import asyncio import base64 import binascii import boto3 import botocore.exceptions from audience_common.aws.utils import get_default_session from audience_common.cache.utils import make_key from campaigns.core.types import EncryptedToken, PlainToken from campaigns.utils.asyncio import run_async class KMSInvalidCiphertextException(Exception): pass class KMSRetryableException(Exception): pass class BaseKMS(abc.ABC): @abc.abstractmethod async def encrypt( self, plain_token: PlainToken, *, context: dict[str, str] | None = None ) -> EncryptedToken: pass async def encrypt_many( self, plain_tokens: list[PlainToken], *, context: dict[str, str] | None = None ) -> dict[PlainToken, EncryptedToken]: if not plain_tokens: return {} tasks = [] for plain_token in plain_tokens: tasks.append(self.encrypt(plain_token, context=context)) return dict(zip(plain_tokens, await asyncio.gather(*tasks), strict=True)) @abc.abstractmethod async def decrypt( self, encrypted_token: EncryptedToken, *, context: dict[str, str] | None = None ) -> PlainToken: pass async def decrypt_many( self, encrypted_tokens: list[EncryptedToken], *, context: dict[str, str] | None = None, ) -> dict[EncryptedToken, PlainToken]: if not encrypted_tokens: return {} tasks = [] for ciphertext in encrypted_tokens: tasks.append(self.decrypt(ciphertext, context=context)) return dict(zip(encrypted_tokens, await asyncio.gather(*tasks), strict=True)) class KMS(BaseKMS): retryable_error_codes = { "DependencyTimeoutException", "KeyUnavailableException", "KMSInternalException", } def __init__( self, region_name: str, key_id: str, session: boto3.Session | None = None ) -> None: session = session or get_default_session() self.client = session.client("kms", region_name=region_name) self.key_id = key_id self._cache: dict[str, PlainToken] = {} async def encrypt( self, plain_token: PlainToken, *, context: dict[str, str] | None = None, ) -> EncryptedToken: context = context or {} try: encrypt_response = await run_async( self.client.encrypt, KeyId=self.key_id, Plaintext=plain_token, EncryptionContext=context, ) except botocore.exceptions.ClientError as error: if error.response["Error"]["Code"] in self.retryable_error_codes: raise KMSRetryableException from error else: raise error decoded = base64.b64encode(encrypt_response["CiphertextBlob"]).decode() ciphertext = EncryptedToken(decoded) cache_key = self._cache_key(ciphertext, context=context) self._cache[cache_key] = plain_token return ciphertext async def decrypt( self, encrypted_token: EncryptedToken, *, context: dict[str, str] | None = None, ) -> PlainToken: context = context or {} cache_key = self._cache_key(encrypted_token, context=context) if cache_key in self._cache: return self._cache[cache_key] try: decrypt_response = await run_async( self.client.decrypt, KeyId=self.key_id, CiphertextBlob=base64.b64decode(encrypted_token), EncryptionContext=context, ) except botocore.exceptions.ClientError as exc: if exc.response["Error"]["Code"] == "InvalidCiphertextException": raise KMSInvalidCiphertextException from exc elif exc.response["Error"]["Code"] in self.retryable_error_codes: raise KMSRetryableException from exc else: raise exc except binascii.Error as exc: raise KMSInvalidCiphertextException from exc plaintext = PlainToken(decrypt_response["Plaintext"].decode()) self._cache[cache_key] = plaintext return plaintext def _cache_key( self, encrypted_token: EncryptedToken, *, context: dict[str, str], key: str = "kms:decrypt", ) -> str: return make_key(key, encrypted_token, key_id=self.key_id, context=context) class DummyKMS(BaseKMS): encrypted = { # Encrypted with QA KMS key: key/773a9ce5-29c8-4838-a801-a826ba9f0bc3 ( EncryptedToken( "AQICAHgOb4p+nJKCMOleADjroRR/AeyrFsMaCBezDETi3y" "UnDgHhQcRvsp5BayKXZHccMnBGAAAAazBpBgkqhkiG9w0B" "BwagXDBaAgEAMFUGCSqGSIb3DQEHATAeBglghkgBZQMEAS" "4wEQQMbAM9GZpJ0KfaqGopAgEQgChXmDW7kSZwysDfOvgA" "Hl2oTK9QU2hGnF8DfWgXziMPvDQ+5acuhfs+" ) ): PlainToken("test@mail.com"), EncryptedToken( "AQICAHgOb4p+nJKCMOleADjroRR/AeyrFsMaCBezDETi3yU" "nDgEfOHFwGhVokSZ0mcPs/DnjAAAA5jCB4wYJKoZIhvcNAQ" "cGoIHVMIHSAgEAMIHMBgkqhkiG9w0BBwEwHgYJYIZIAWUDB" "AEuMBEEDOcJXFCVvG94lMXGEQIBEICBnu8F8LhpkriGVxRp" "t4OJ3+ofwNERwn1+v52NcZv9Zu+JoKVLGx/weJzwJzr+fuZ" "J+rS9hLmar5C4gAEf5O+zWVu5W9KU7H4v62NHBRyZeiUeVe" "vSaslXlWohaZRsJs8pwb94AjRUx1Xi1Fozf9+xEA8hjIxxz" "DyEOUKo3fcGIKfJYwCnPHuNPeZmBoXr8XXn2hgN2knTi5dM" "gjjWhdVV" ): PlainToken( "AQByP2zB56gRzCSPs6Z6vAFBXh4zTWJomakh3w9DRFaiodIK" "yOy8CQluwv3d3dR9flBu-QK0MY-gM2mc2NPcxPkn_bXIYopF" "3gLXiiEnLaiTgy8EPesfGy_hVMk7E6la3OM" ), } decrypted = { plain_token: encrypted_token for encrypted_token, plain_token in encrypted.items() } async def encrypt( self, plain_token: PlainToken, *, context: dict[str, str] | None = None ) -> EncryptedToken: return EncryptedToken(self.decrypted.get(plain_token, plain_token)) async def decrypt( self, encrypted_token: EncryptedToken, *, context: dict[str, str] | None = None ) -> PlainToken: return self.encrypted.get(encrypted_token, PlainToken(encrypted_token))