"""Primitives de sécurité : Argon2id, TOTP, chiffrement des secrets MFA, jetons."""
from __future__ import annotations

import base64
import hashlib
import hmac
import secrets
import time
import unicodedata

import pyotp
from argon2 import PasswordHasher
from argon2.exceptions import InvalidHashError, VerificationError, VerifyMismatchError
from cryptography.hazmat.primitives import hashes
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
from cryptography.hazmat.primitives.kdf.hkdf import HKDF

from .config import get_config

# RFC 9106, seconde recommandation (64 Mio, 3 passes, 4 voies) : adaptée à un serveur modeste.
_hasher = PasswordHasher(time_cost=3, memory_cost=65536, parallelism=4, hash_len=32, salt_len=16)
_DUMMY_HASH = _hasher.hash("mot-de-passe-factice-pour-egaliser-les-temps")

COMMON_PASSWORDS = {
    "motdepasse", "motdepasse1", "motdepasse123", "password", "password1", "password123", "azerty",
    "azertyuiop", "azerty123", "qwerty", "qwertyuiop", "123456", "1234567890", "123456789", "000000",
    "gendarmerie", "gendarmerie52", "edcf52", "edcf", "hautemarne", "chaumont", "bienvenue", "soleil",
    "admin", "administrateur", "changeme", "jesuisla", "loulou", "doudou", "marseille", "football",
}


def hash_password(password: str) -> str:
    return _hasher.hash(password)


def verify_password(stored_hash: str | None, password: str) -> bool:
    """Vérification à temps comparable même si le compte n'existe pas (anti-énumération)."""
    try:
        return _hasher.verify(stored_hash or _DUMMY_HASH, password) and stored_hash is not None
    except (VerifyMismatchError, VerificationError, InvalidHashError):
        return False


def needs_rehash(stored_hash: str) -> bool:
    try:
        return _hasher.check_needs_rehash(stored_hash)
    except InvalidHashError:
        return True


def _fold(s: str) -> str:
    s = unicodedata.normalize("NFKD", s.lower())
    return "".join(c for c in s if c.isalnum() and not unicodedata.combining(c))


def password_problems(password: str, username: str, min_length: int) -> list[str]:
    """Politique configurable : longueur avant complexité (recommandation ANSSI)."""
    problems: list[str] = []
    if len(password) < min_length:
        problems.append(f"Le mot de passe doit contenir au moins {min_length} caractères.")
    if len(password) > 256:
        problems.append("Le mot de passe ne doit pas dépasser 256 caractères.")
    folded = _fold(password)
    if folded in COMMON_PASSWORDS or len(set(password)) < 5:
        problems.append("Ce mot de passe est trop courant ou trop répétitif.")
    if username and len(username) >= 3 and _fold(username) in folded:
        problems.append("Le mot de passe ne doit pas contenir l'identifiant.")
    classes = sum([
        any(c.islower() for c in password), any(c.isupper() for c in password),
        any(c.isdigit() for c in password), any(not c.isalnum() for c in password),
    ])
    if len(password) < 16 and classes < 3:
        problems.append("En dessous de 16 caractères, mélangez au moins 3 types de caractères "
                        "(minuscules, majuscules, chiffres, symboles).")
    return problems


def generate_password(length: int = 20) -> str:
    alphabet = "abcdefghijkmnopqrstuvwxyzABCDEFGHJKLMNPQRSTUVWXYZ23456789-_.!"
    while True:
        pw = "".join(secrets.choice(alphabet) for _ in range(length))
        if not password_problems(pw, "", 12):
            return pw


# --- Jetons -------------------------------------------------------------------------------------

def new_token() -> str:
    return secrets.token_urlsafe(32)  # 256 bits


def token_digest(token: str) -> str:
    return hashlib.sha256(token.encode("ascii", "ignore")).hexdigest()


def constant_eq(a: str, b: str) -> bool:
    return hmac.compare_digest(a.encode(), b.encode())


# --- Chiffrement des secrets MFA (AES-256-GCM, clé dérivée par HKDF) ----------------------------

def _key(purpose: bytes) -> bytes:
    return HKDF(algorithm=hashes.SHA256(), length=32, salt=b"edcf52", info=purpose).derive(
        get_config().secret_key.encode()
    )


def encrypt_secret(plain: str) -> bytes:
    nonce = secrets.token_bytes(12)
    return nonce + AESGCM(_key(b"mfa-secret")).encrypt(nonce, plain.encode(), b"edcf-mfa")


def decrypt_secret(blob: bytes) -> str:
    return AESGCM(_key(b"mfa-secret")).decrypt(blob[:12], blob[12:], b"edcf-mfa").decode()


# --- TOTP (RFC 6238) ----------------------------------------------------------------------------

def new_totp_secret() -> str:
    return pyotp.random_base32(length=32)


def totp_uri(secret: str, username: str) -> str:
    return pyotp.TOTP(secret).provisioning_uri(name=username, issuer_name="EDCF 52 Bilans")


def verify_totp(secret: str, code: str, last_step: int | None, at: float | None = None) -> int | None:
    """Renvoie le pas de temps accepté (anti-rejeu : un code ne sert qu'une fois), sinon None."""
    code = "".join(c for c in code if c.isdigit())
    if len(code) != 6:
        return None
    now = at if at is not None else time.time()
    totp = pyotp.TOTP(secret)
    current = int(now // 30)
    for step in (current - 1, current, current + 1):
        if last_step is not None and step <= last_step:
            continue
        if hmac.compare_digest(totp.at(step * 30), code):
            return step
    return None


# --- Codes de récupération ----------------------------------------------------------------------

def new_recovery_codes(n: int = 10) -> list[str]:
    alphabet = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"
    return ["-".join("".join(secrets.choice(alphabet) for _ in range(4)) for _ in range(3)) for _ in range(n)]


def recovery_digest(code: str) -> str:
    norm = "".join(c for c in code.upper() if c.isalnum())
    return hmac.new(_key(b"recovery"), norm.encode(), hashlib.sha256).hexdigest()


def b64(data: bytes) -> str:
    return base64.b64encode(data).decode()
