from __future__ import annotations import secrets from abc import ABC, abstractmethod from cryptography.fernet import Fernet from fansifter_common.utils.functional import lazy_proxy from pydantic import SecretStr from resonance_engine.adapters import aws_secretsmanager from resonance_engine.config import settings class Encrypter(ABC): @abstractmethod def encrypt(self, plaintext: str) -> str: """Encrypt plaintext and return a base64-encoded ciphertext string.""" @abstractmethod def decrypt(self, ciphertext: str) -> str: """Decrypt a base64-encoded ciphertext string and return plaintext.""" @property def key_ids(self) -> list[str]: return [] @property def active_key_id(self) -> str | None: return None # --------------------------------------------------------------------------- # Fernet # --------------------------------------------------------------------------- KEY_ID_LEN = 8 class FernetEncrypter(Encrypter): def __init__(self, keys: list[SecretStr]) -> None: if not keys: raise ValueError( "FernetEncrypter requires at least one key — " "set ENCRYPTION_KEYS or ENCRYPTION_KEYS_SECRET_NAME" ) self._fernets: dict[str, Fernet] = {} for k in keys: key_id, _, fernet_key = k.get_secret_value().partition(":") self._fernets[key_id] = Fernet(fernet_key.encode()) self._active_key_id = keys[0].get_secret_value().partition(":")[0] @property def active_key_id(self) -> str: return self._active_key_id @property def key_ids(self) -> list[str]: return list(self._fernets.keys()) def encrypt(self, plaintext: str) -> str: token = self._fernets[self._active_key_id].encrypt(plaintext.encode()).decode() return f"{self._active_key_id}:{token}" def decrypt(self, ciphertext: str) -> str: key_id, _, token = ciphertext.partition(":") return self._fernets[key_id].decrypt(token.encode()).decode() @staticmethod def generate_key() -> SecretStr: key_id = secrets.token_hex(KEY_ID_LEN // 2) return SecretStr(f"{key_id}:{Fernet.generate_key().decode()}") # --------------------------------------------------------------------------- # Plaintext (no-op, for local development) # --------------------------------------------------------------------------- class PlaintextEncrypter(Encrypter): """Passes values through unchanged — no encryption. Use only in development.""" def encrypt(self, plaintext: str) -> str: return plaintext def decrypt(self, ciphertext: str) -> str: return ciphertext # --------------------------------------------------------------------------- # Factory # --------------------------------------------------------------------------- def get_encrypter() -> Encrypter: if settings.encryption_backend == "plain": return PlaintextEncrypter() keys = settings.encryption_keys if settings.encryption_keys_secret_name: keys = aws_secretsmanager.get_encryption_keys() return FernetEncrypter(keys=keys) encrypter = lazy_proxy(get_encrypter) def encrypt(plaintext: str) -> str: return encrypter.encrypt(plaintext) def decrypt(ciphertext: str) -> str: return encrypter.decrypt(ciphertext)