from __future__ import annotations

import base64
import hashlib
import hmac
import secrets
from dataclasses import dataclass

PASSWORD_SCHEME = "pbkdf2_sha256"
DEFAULT_ITERATIONS = 600_000


@dataclass(frozen=True)
class PasswordHash:
    scheme: str
    iterations: int
    salt: bytes
    digest: bytes


def hash_password(password: str, *, iterations: int = DEFAULT_ITERATIONS) -> str:
    """Return a portable PBKDF2-SHA256 password hash for .env storage."""
    if len(password) < 12:
        raise ValueError("Le mot de passe doit contenir au moins 12 caractères.")
    salt = secrets.token_bytes(18)
    digest = hashlib.pbkdf2_hmac("sha256", password.encode("utf-8"), salt, iterations)
    return "$".join(
        (
            PASSWORD_SCHEME,
            str(iterations),
            base64.urlsafe_b64encode(salt).decode("ascii").rstrip("="),
            base64.urlsafe_b64encode(digest).decode("ascii").rstrip("="),
        )
    )


def _decode_base64(value: str) -> bytes:
    padding = "=" * (-len(value) % 4)
    return base64.urlsafe_b64decode(value + padding)


def parse_password_hash(encoded: str) -> PasswordHash | None:
    try:
        scheme, iterations_text, salt_text, digest_text = encoded.split("$", 3)
        if scheme != PASSWORD_SCHEME:
            return None
        iterations = int(iterations_text)
        if iterations < 100_000:
            return None
        return PasswordHash(
            scheme=scheme,
            iterations=iterations,
            salt=_decode_base64(salt_text),
            digest=_decode_base64(digest_text),
        )
    except (ValueError, TypeError):
        return None


def verify_password(password: str, encoded: str) -> bool:
    parsed = parse_password_hash(encoded)
    if not parsed:
        # Constant-ish work even when the stored value is invalid.
        hashlib.pbkdf2_hmac("sha256", password.encode("utf-8"), b"invalid-salt-value", 100_000)
        return False
    candidate = hashlib.pbkdf2_hmac(
        "sha256",
        password.encode("utf-8"),
        parsed.salt,
        parsed.iterations,
    )
    return hmac.compare_digest(candidate, parsed.digest)


def generate_secret_key() -> str:
    return secrets.token_urlsafe(48)


def generate_csrf_token() -> str:
    return secrets.token_urlsafe(32)
