from __future__ import annotations

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


PBKDF2_ALGORITHM = "pbkdf2_sha256"
PBKDF2_ITERATIONS = 210_000
LICENCE_KEY_GROUPS = 5
LICENCE_KEY_GROUP_LENGTH = 5
LICENCE_KEY_ALPHABET = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"


@dataclass(frozen=True)
class GeneratedSecret:
    plaintext: str
    hashed: str
    public_identifier: str
    suffix: str


def generate_licence_key() -> GeneratedSecret:
    groups = [
        "".join(secrets.choice(LICENCE_KEY_ALPHABET) for _ in range(LICENCE_KEY_GROUP_LENGTH))
        for _ in range(LICENCE_KEY_GROUPS)
    ]
    plaintext = "TD-" + "-".join(groups)
    return GeneratedSecret(
        plaintext=plaintext,
        hashed=hash_secret(plaintext),
        public_identifier=public_identifier_for_secret(plaintext),
        suffix=secret_suffix(plaintext),
    )


def generate_activation_token() -> GeneratedSecret:
    plaintext = secrets.token_urlsafe(32)
    return GeneratedSecret(
        plaintext=plaintext,
        hashed=hash_secret(plaintext),
        public_identifier=public_identifier_for_secret(plaintext),
        suffix=secret_suffix(plaintext),
    )


def hash_secret(secret: str) -> str:
    salt = secrets.token_bytes(16)
    digest = hashlib.pbkdf2_hmac(
        "sha256",
        secret.encode("utf-8"),
        salt,
        PBKDF2_ITERATIONS,
    )
    return "$".join(
        [
            PBKDF2_ALGORITHM,
            str(PBKDF2_ITERATIONS),
            base64.urlsafe_b64encode(salt).decode("ascii"),
            base64.urlsafe_b64encode(digest).decode("ascii"),
        ]
    )


def verify_secret(secret: str, stored_hash: str) -> bool:
    try:
        algorithm, iterations_text, salt_text, digest_text = stored_hash.split("$", 3)
        if algorithm != PBKDF2_ALGORITHM:
            return False
        salt = base64.urlsafe_b64decode(salt_text.encode("ascii"))
        expected = base64.urlsafe_b64decode(digest_text.encode("ascii"))
        actual = hashlib.pbkdf2_hmac(
            "sha256",
            secret.encode("utf-8"),
            salt,
            int(iterations_text),
        )
        return hmac.compare_digest(actual, expected)
    except Exception:
        return False


def normalise_licence_key(licence_key: str) -> str:
    return "".join(ch for ch in licence_key.upper() if ch in string.ascii_uppercase + string.digits)


def public_identifier_for_secret(secret: str) -> str:
    digest = hashlib.sha256(secret.encode("utf-8")).hexdigest().upper()
    return digest[:12]


def secret_suffix(secret: str, length: int = 4) -> str:
    compact = "".join(ch for ch in secret if ch.isalnum())
    return compact[-length:].upper() if compact else ""


def redacted_secret(secret: str) -> str:
    suffix = secret_suffix(secret)
    return f"****-{suffix}" if suffix else "****"
