"""Token encryption at rest using Fernet (AES-128-CBC + HMAC-SHA256).

A dedicated key is stored in config/.dns-token-key with mode 0600.
If the key file is missing, it is derived once from a master secret (or
generated randomly) and persisted so restarts and password rotations do
not invalidate previously encrypted tokens.
"""

import base64
import hashlib
import logging
import os
from pathlib import Path

from cryptography.fernet import Fernet, InvalidToken

from .settings import get_settings

logger = logging.getLogger(__name__)

_fernet: Fernet | None = None
_fernet_loaded = False


def _get_key_file() -> Path:
    """Return the path to the stored encryption key."""
    settings = get_settings()
    config_dir = settings.base_dir / "config"
    return config_dir / ".dns-token-key"


def _derive_master_key() -> bytes:
    """Derive a Fernet-compatible 32-byte key from the best available secret."""
    settings = get_settings()

    secret_material = (
        settings.panel_login_csrf_secret
        or settings.api_admin_pass_hash
        or settings.panel_admin_pass_hash
        or settings.api_admin_pass
        or settings.panel_admin_pass
    )
    if not secret_material:
        return b""

    raw = hashlib.pbkdf2_hmac(
        "sha256",
        secret_material.encode("utf-8"),
        b"limristem-mail-dns-token-encryption-v1",
        iterations=600_000,
        dklen=32,
    )
    return base64.urlsafe_b64encode(raw)


def _persist_key(key_data: bytes) -> None:
    key_file = _get_key_file()
    try:
        key_file.parent.mkdir(parents=True, exist_ok=True)
        # Write privately, then replace atomically.
        tmp_path = key_file.with_suffix(".key.tmp")
        flags = os.O_WRONLY | os.O_CREAT | os.O_TRUNC
        fd = os.open(tmp_path, flags, 0o600)
        try:
            os.write(fd, key_data + (b"" if key_data.endswith(b"\n") else b"\n"))
        finally:
            os.close(fd)
        os.replace(tmp_path, key_file)
        os.chmod(key_file, 0o600)
    except OSError as exc:
        logger.warning("Unable to persist encryption key file %s: %s", key_file, exc)


def _get_fernet() -> Fernet | None:
    global _fernet, _fernet_loaded
    if _fernet_loaded:
        return _fernet

    _fernet_loaded = True
    key_file = _get_key_file()

    if key_file.exists():
        try:
            key_data = key_file.read_bytes().strip()
            _fernet = Fernet(key_data)
            return _fernet
        except (ValueError, OSError) as exc:
            logger.warning("Unable to read encryption key file %s: %s", key_file, exc)

    master_key = _derive_master_key()
    if master_key:
        _persist_key(master_key)
        _fernet = Fernet(master_key)
        return _fernet

    # Last resort: generate a random key and persist it so future restarts work.
    try:
        generated = Fernet.generate_key()
        _persist_key(generated)
        _fernet = Fernet(generated)
        logger.info("Generated new DNS token encryption key at %s", key_file)
        return _fernet
    except Exception as exc:  # pragma: no cover - extremely unlikely
        logger.warning("Unable to generate encryption key: %s", exc)
        return None


def encrypt_token(plaintext: str | None) -> str | None:
    """Encrypt a token string. Returns None for None/empty input."""
    if not plaintext:
        return plaintext
    f = _get_fernet()
    if f is None:
        logger.warning("Encryption not available, storing token in plaintext")
        return plaintext
    return f.encrypt(plaintext.encode("utf-8")).decode("utf-8")


def decrypt_token(ciphertext: str | None) -> str | None:
    """Decrypt a token string. Returns None for None/empty input.

    Falls back to returning the raw value if it doesn't look like encrypted data,
    for backwards compatibility with existing plaintext tokens.
    """
    if not ciphertext:
        return ciphertext

    f = _get_fernet()
    if f is None:
        return ciphertext

    try:
        return f.decrypt(ciphertext.encode("utf-8")).decode("utf-8")
    except (InvalidToken, ValueError):
        logger.debug("Token is not encrypted or decryption failed, returning raw value")
        return ciphertext


def reset_crypto_state_for_tests() -> None:
    """Clear cached Fernet state (tests only)."""
    global _fernet, _fernet_loaded
    _fernet = None
    _fernet_loaded = False
