import base64
import hashlib

from cryptography.fernet import Fernet

from config import settings


def _get_fernet() -> Fernet:
    """Derive a Fernet key from SECRET_KEY.

    Fernet requires a 32-byte URL-safe base64-encoded key.  We derive one
    deterministically from the configured SECRET_KEY by SHA-256 hashing it
    and then base64-encoding the result.
    """
    raw = settings.SECRET_KEY.encode()
    digest = hashlib.sha256(raw).digest()  # always 32 bytes
    fernet_key = base64.urlsafe_b64encode(digest)
    return Fernet(fernet_key)


_fernet: Fernet | None = None


def _fernet_instance() -> Fernet:
    """Return a cached Fernet instance (re-used across calls for efficiency)."""
    global _fernet  # noqa: PLW0603
    if _fernet is None:
        _fernet = _get_fernet()
    return _fernet


def encrypt_string(value: str) -> str:
    """Encrypt *value* with AES-256 (via Fernet) and return a base64 string.

    Args:
        value: Plaintext string to encrypt.

    Returns:
        URL-safe base64-encoded ciphertext string.
    """
    if not value:
        return value
    token: bytes = _fernet_instance().encrypt(value.encode("utf-8"))
    return token.decode("utf-8")


def decrypt_string(value: str) -> str:
    """Decrypt a previously-encrypted string.

    Args:
        value: URL-safe base64-encoded ciphertext returned by :func:`encrypt_string`.

    Returns:
        Original plaintext string.

    Raises:
        cryptography.fernet.InvalidToken: If the token is invalid or tampered.
    """
    if not value:
        return value
    plaintext: bytes = _fernet_instance().decrypt(value.encode("utf-8"))
    return plaintext.decode("utf-8")
