"""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 generated randomly once and persisted.
Encryption never falls back to plaintext.
"""

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


class CryptoUnavailableError(RuntimeError):
    """Raised when the encryption key cannot be loaded or created."""


def _get_key_file() -> Path:
    settings = get_settings()
    config_dir = settings.base_dir / "config"
    return config_dir / ".dns-token-key"


def _derive_master_key() -> bytes | None:
    """Derive a Fernet key from stable admin secret material when available."""
    settings = get_settings()
    secret_material = (
        settings.panel_login_csrf_secret
        or settings.api_admin_pass_hash
        or settings.panel_admin_pass_hash
    )
    if not secret_material:
        return None
    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)
        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)
        # Group-readable so the service user (limristem-mail) can decrypt secrets
        # when the key is owned by root:limristem-mail (install/CLI may create it as root).
        try:
            os.chmod(key_file, 0o640)
        except OSError:
            os.chmod(key_file, 0o600)
        try:
            import grp

            os.chown(key_file, 0, grp.getgrnam("limristem-mail").gr_gid)
        except (KeyError, OSError, PermissionError):
            pass
    except OSError as exc:
        raise CryptoUnavailableError(f"Unable to persist encryption key file {key_file}: {exc}") from exc


def _get_fernet() -> Fernet:
    global _fernet, _fernet_loaded
    if _fernet_loaded and _fernet is not None:
        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.error("Unable to read encryption key file %s: %s", key_file, exc)
            raise CryptoUnavailableError(f"Invalid encryption key file: {key_file}") from exc

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

    try:
        generated = Fernet.generate_key()
        _persist_key(generated)
        _fernet = Fernet(generated)
        logger.info("Generated new encryption key at %s", key_file)
        return _fernet
    except CryptoUnavailableError:
        raise
    except Exception as exc:  # pragma: no cover
        raise CryptoUnavailableError(f"Unable to generate encryption key: {exc}") from exc


def encrypt_token(plaintext: str | None) -> str | None:
    """Encrypt a token string. Returns None for None/empty input. Never stores plaintext."""
    if not plaintext:
        return plaintext
    f = _get_fernet()
    return f.encrypt(plaintext.encode("utf-8")).decode("utf-8")


def decrypt_token(ciphertext: str | None) -> str | None:
    """Decrypt a token string. Raises CryptoUnavailableError / InvalidToken on failure.

    Empty/None pass through. Plaintext legacy values are rejected (must re-enter secrets).
    """
    if not ciphertext:
        return ciphertext

    f = _get_fernet()
    try:
        return f.decrypt(ciphertext.encode("utf-8")).decode("utf-8")
    except (InvalidToken, ValueError) as exc:
        # Do not return raw secrets: force re-configuration of legacy plaintext rows.
        raise CryptoUnavailableError("Unable to decrypt secret; value may be legacy plaintext or corrupt") from exc


def decrypt_secret_flexible(value: str | None) -> str | None:
    """Decrypt a Fernet secret, accepting legacy plaintext for one-time migration.

    Fernet-looking values that fail to decrypt raise CryptoUnavailableError.
    Plain legacy strings are returned as-is so callers can re-encrypt on save.
    """
    if not value:
        return value
    try:
        return decrypt_token(value)
    except CryptoUnavailableError:
        if value.startswith("gAAAA"):
            raise
        return value


def encrypt_secret_if_needed(value: str | None) -> str | None:
    """Encrypt a secret unless empty or already a Fernet token that decrypts cleanly."""
    if not value:
        return value
    if value.startswith("gAAAA"):
        try:
            decrypt_token(value)
            return value
        except CryptoUnavailableError:
            pass
    return encrypt_token(value)


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