import abc import base64 import binascii import boto3 import botocore.exceptions from fansifter_common.utils.cache import make_cache_key from dmp.core.types import EncryptedToken, PlainToken class KMSInvalidCiphertextException(Exception): pass class KMSRetryableException(Exception): pass class BaseKMS(abc.ABC): @abc.abstractmethod def encrypt( self, plain_token: PlainToken, *, context: dict[str, str] | None = None ) -> EncryptedToken: pass @abc.abstractmethod def decrypt( self, encrypted_token: EncryptedToken, *, context: dict[str, str] | None = None ) -> PlainToken: pass class KMS(BaseKMS): retryable_error_codes = { "DependencyTimeoutException", "KeyUnavailableException", "KMSInternalException", } def __init__(self, region_name: str, key_id: str, session: boto3.Session) -> None: self.client = session.client("kms", region_name=region_name) self.key_id = key_id self._cache: dict[str, PlainToken] = {} def encrypt( self, plain_token: PlainToken, *, context: dict[str, str] | None = None, ) -> EncryptedToken: context = context or {} try: encrypt_response = self.client.encrypt( KeyId=self.key_id, Plaintext=plain_token, EncryptionContext=context, ) except botocore.exceptions.ClientError as error: if "Error" in error.response and "Code" in error.response["Error"]: error_code = error.response["Error"]["Code"] if error_code in self.retryable_error_codes: raise KMSRetryableException from error 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 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 = self.client.decrypt( KeyId=self.key_id, CiphertextBlob=base64.b64decode(encrypted_token), EncryptionContext=context, ) except botocore.exceptions.ClientError as exc: if "Error" in exc.response and "Code" in exc.response["Error"]: error_code = exc.response["Error"]["Code"] if error_code == "InvalidCiphertextException": raise KMSInvalidCiphertextException from exc elif error_code in self.retryable_error_codes: raise KMSRetryableException from exc 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_cache_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() } def encrypt( self, plain_token: PlainToken, *, context: dict[str, str] | None = None, # noqa: ARG002 ) -> EncryptedToken: return EncryptedToken(self.decrypted.get(plain_token, plain_token)) def decrypt( self, encrypted_token: EncryptedToken, *, context: dict[str, str] | None = None, # noqa: ARG002 ) -> PlainToken: return self.encrypted.get(encrypted_token, PlainToken(encrypted_token))