"""Server-only signing, secret-digest, and identifier primitives."""

from __future__ import annotations

import base64
from dataclasses import dataclass
import hashlib
import hmac
import os
from pathlib import Path
import secrets
from typing import Mapping, Protocol
import uuid

from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric.ed25519 import (
    Ed25519PrivateKey,
    Ed25519PublicKey,
)

from licensing_shared.constants import validate_identifier
from licensing_shared.models import SignedLicenseDocument
from licensing_shared.verifier import LicenseVerifier

from .constants import FINGERPRINT_PEPPER_MINIMUM_BYTES, SERIAL_PEPPER_MINIMUM_BYTES


SERIAL_DIGEST_DOMAIN = b"APOLON-SERIAL-DIGEST-V1\x00"
FINGERPRINT_DIGEST_DOMAIN = b"APOLON-DEVICE-EVIDENCE-V1\x00"
ANONYMOUS_SUBJECT_DOMAIN = b"APOLON-ANONYMOUS-SUBJECT-V1\x00"
MAX_SECRET_INPUT_BYTES = 4096
MAX_SIGNING_PRIVATE_KEY_FILE_BYTES = 64 * 1024


def new_identifier(prefix: str) -> str:
    normalized_prefix = validate_identifier(prefix, "identifier prefix")
    return validate_identifier(
        f"{normalized_prefix}.{uuid.uuid4().hex}",
        "generated identifier",
    )


def random_token(character_count: int, alphabet: str) -> str:
    if isinstance(character_count, bool) or not isinstance(character_count, int):
        raise TypeError("character_count must be an integer")
    if character_count < 1:
        raise ValueError("character_count must be at least one")
    if not isinstance(alphabet, str):
        raise TypeError("alphabet must be a string")
    if len(alphabet) < 2 or len(set(alphabet)) != len(alphabet):
        raise ValueError("alphabet must contain at least two unique characters")
    return "".join(secrets.choice(alphabet) for _ in range(character_count))


def _hmac_digest(key: bytes, domain: bytes, value: bytes) -> bytes:
    if not isinstance(key, bytes):
        raise TypeError("key must be bytes")
    if not isinstance(domain, bytes) or not domain:
        raise ValueError("domain must be non-empty bytes")
    if not isinstance(value, bytes):
        raise TypeError("value must be bytes")
    if not value or len(value) > MAX_SECRET_INPUT_BYTES:
        raise ValueError("digest input is empty or too large")
    return hmac.new(key, domain + value, hashlib.sha256).digest()


def serial_secret_digest(serial_pepper: bytes, compact_serial: str) -> bytes:
    if len(serial_pepper) < SERIAL_PEPPER_MINIMUM_BYTES:
        raise ValueError("serial_pepper is too short")
    if not isinstance(compact_serial, str) or not compact_serial:
        raise ValueError("compact_serial must be a non-empty string")
    return _hmac_digest(
        serial_pepper,
        SERIAL_DIGEST_DOMAIN,
        compact_serial.encode("ascii"),
    )


def anonymous_subject_digest(fingerprint_pepper: bytes, subject: str) -> bytes:
    if len(fingerprint_pepper) < FINGERPRINT_PEPPER_MINIMUM_BYTES:
        raise ValueError("fingerprint_pepper is too short")
    if not isinstance(subject, str) or not subject.strip():
        raise ValueError("subject must be a non-empty string")
    return _hmac_digest(
        fingerprint_pepper,
        ANONYMOUS_SUBJECT_DOMAIN,
        subject.strip().encode("utf-8"),
    )


def digest_device_evidence(
    fingerprint_pepper: bytes,
    evidence: Mapping[str, str],
) -> dict[str, str]:
    if len(fingerprint_pepper) < FINGERPRINT_PEPPER_MINIMUM_BYTES:
        raise ValueError("fingerprint_pepper is too short")
    if not isinstance(evidence, Mapping):
        raise TypeError("evidence must be a mapping")
    if len(evidence) > 16:
        raise ValueError("too many device-evidence components")
    result: dict[str, str] = {}
    for component, raw_value in evidence.items():
        normalized_component = validate_identifier(component, "evidence component")
        if not isinstance(raw_value, str):
            raise TypeError("device-evidence values must be strings")
        normalized_value = raw_value.strip().lower()
        if not normalized_value or len(normalized_value) > 512:
            raise ValueError("device-evidence value is empty or too long")
        digest = _hmac_digest(
            fingerprint_pepper,
            FINGERPRINT_DIGEST_DOMAIN + normalized_component.encode("ascii") + b"\x00",
            normalized_value.encode("utf-8"),
        )
        result[normalized_component] = digest.hex()
    return dict(sorted(result.items()))


class SnapshotSigner(Protocol):
    @property
    def key_id(self) -> str:
        ...

    def public_key_bytes(self) -> bytes:
        ...

    def sign_payload(self, payload: Mapping[str, object]) -> SignedLicenseDocument:
        ...


@dataclass(frozen=True)
class Ed25519SnapshotSigner:
    _key_id: str
    _private_key: Ed25519PrivateKey

    def __post_init__(self) -> None:
        object.__setattr__(self, "_key_id", validate_identifier(self._key_id, "key_id"))
        if not isinstance(self._private_key, Ed25519PrivateKey):
            raise TypeError("_private_key must be an Ed25519PrivateKey")

    @property
    def key_id(self) -> str:
        return self._key_id

    def public_key_bytes(self) -> bytes:
        return self._private_key.public_key().public_bytes(
            encoding=serialization.Encoding.Raw,
            format=serialization.PublicFormat.Raw,
        )

    def sign_payload(self, payload: Mapping[str, object]) -> SignedLicenseDocument:
        if not isinstance(payload, Mapping):
            raise TypeError("payload must be a mapping")
        signature = self._private_key.sign(LicenseVerifier.signing_bytes(payload))
        return SignedLicenseDocument(
            key_id=self.key_id,
            payload=dict(payload),
            signature_base64=base64.b64encode(signature).decode("ascii"),
        )

    @classmethod
    def from_private_key_file(
        cls,
        key_id: str,
        path: Path,
        *,
        require_owner_only: bool = False,
    ) -> "Ed25519SnapshotSigner":
        normalized_key_id = validate_identifier(key_id, "key_id")
        if not isinstance(path, Path):
            raise TypeError("path must be a Path")
        if not isinstance(require_owner_only, bool):
            raise TypeError("require_owner_only must be a Boolean")
        if not path.is_file() or path.is_symlink():
            raise ValueError("signing key path must be a regular non-symlink file")
        file_status = path.stat()
        if file_status.st_size < 1 or file_status.st_size > MAX_SIGNING_PRIVATE_KEY_FILE_BYTES:
            raise ValueError("signing key file size is invalid")
        if require_owner_only and os.name != "nt" and file_status.st_mode & 0o077:
            raise PermissionError("production signing key file must be owner-only")
        document = path.read_bytes()
        try:
            loaded = serialization.load_pem_private_key(document, password=None)
        except ValueError:
            try:
                raw = base64.b64decode(document.strip(), validate=True)
                loaded = Ed25519PrivateKey.from_private_bytes(raw)
            except (TypeError, ValueError) as exc:
                raise ValueError("signing key file is not unencrypted PEM or Base64 raw Ed25519") from exc
        if not isinstance(loaded, Ed25519PrivateKey):
            raise ValueError("signing key file does not contain an Ed25519 private key")
        return cls(normalized_key_id, loaded)


def decode_ed25519_public_key(value: str) -> Ed25519PublicKey:
    if not isinstance(value, str) or not value.strip():
        raise ValueError("value must be a non-empty Base64 string")
    try:
        raw = base64.b64decode(value.strip(), validate=True)
        return Ed25519PublicKey.from_public_bytes(raw)
    except (TypeError, ValueError) as exc:
        raise ValueError("value is not a valid raw Ed25519 public key") from exc


__all__ = [
    "Ed25519SnapshotSigner",
    "SnapshotSigner",
    "anonymous_subject_digest",
    "decode_ed25519_public_key",
    "digest_device_evidence",
    "new_identifier",
    "random_token",
    "serial_secret_digest",
]

