"""P-256 device possession proofs for connected licensing API requests."""

from __future__ import annotations

import base64
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
import hashlib
from typing import Any, Mapping, Optional

from cryptography.exceptions import InvalidSignature
from cryptography.hazmat.primitives import hashes, serialization
from cryptography.hazmat.primitives.asymmetric import ec

from .canonical_json import canonicalize_json
from .constants import validate_identifier
from .errors import LicenseErrorCode, LicensingError
from .models import format_rfc3339, parse_rfc3339
from .offline import device_public_key_thumbprint


DEVICE_PROOF_SCHEMA = "apolon.licensing.device-proof"
DEVICE_PROOF_SCHEMA_VERSION = 1
DEVICE_PROOF_ALGORITHM = "ECDSA_P256_SHA256"
DEVICE_PROOF_DOMAIN = b"APOLON-DEVICE-PROOF-V1\x00"
MAX_DEVICE_PROOF_LIFETIME = timedelta(minutes=15)
MAX_DEVICE_PROOF_CLOCK_SKEW = timedelta(minutes=5)
MAX_DEVICE_PROOF_SIGNATURE_CHARACTERS = 256


def device_action_payload_digest(payload: Mapping[str, Any]) -> str:
    if not isinstance(payload, Mapping):
        raise TypeError("payload must be a mapping")
    return hashlib.sha256(canonicalize_json(payload)).hexdigest()


@dataclass(frozen=True)
class DeviceActionProof:
    action: str
    installation_id: str
    device_public_key_base64: str
    device_key_thumbprint: str
    payload_digest: str
    created_at: datetime
    expires_at: datetime
    nonce: str
    signature_base64: str
    algorithm: str = DEVICE_PROOF_ALGORITHM
    schema: str = DEVICE_PROOF_SCHEMA
    schema_version: int = DEVICE_PROOF_SCHEMA_VERSION

    def __post_init__(self) -> None:
        if self.schema != DEVICE_PROOF_SCHEMA or self.schema_version != DEVICE_PROOF_SCHEMA_VERSION:
            raise ValueError("unsupported device-proof schema")
        if self.algorithm != DEVICE_PROOF_ALGORITHM:
            raise ValueError("unsupported device-proof algorithm")
        object.__setattr__(self, "action", validate_identifier(self.action, "action"))
        object.__setattr__(
            self,
            "installation_id",
            validate_identifier(self.installation_id, "installation_id"),
        )
        if not isinstance(self.device_public_key_base64, str):
            raise TypeError("device_public_key_base64 must be a string")
        normalized_key = self.device_public_key_base64.strip()
        try:
            public_key_der = base64.b64decode(normalized_key, validate=True)
        except (TypeError, ValueError) as exc:
            raise ValueError("device_public_key_base64 is not valid Base64") from exc
        expected_thumbprint = device_public_key_thumbprint(public_key_der)
        normalized_thumbprint = validate_identifier(
            self.device_key_thumbprint,
            "device_key_thumbprint",
        )
        if normalized_thumbprint != expected_thumbprint:
            raise ValueError("device_key_thumbprint does not match device public key")
        object.__setattr__(self, "device_public_key_base64", normalized_key)
        object.__setattr__(self, "device_key_thumbprint", normalized_thumbprint)
        if not isinstance(self.payload_digest, str):
            raise TypeError("payload_digest must be a string")
        normalized_digest = self.payload_digest.strip().lower()
        if len(normalized_digest) != 64 or any(
            character not in "0123456789abcdef" for character in normalized_digest
        ):
            raise ValueError("payload_digest must be a SHA-256 hexadecimal digest")
        object.__setattr__(self, "payload_digest", normalized_digest)
        created = parse_rfc3339(format_rfc3339(self.created_at, "created_at"), "created_at")
        expires = parse_rfc3339(format_rfc3339(self.expires_at, "expires_at"), "expires_at")
        if expires <= created:
            raise ValueError("expires_at must be later than created_at")
        if expires - created > MAX_DEVICE_PROOF_LIFETIME:
            raise ValueError("device proof lifetime exceeds the allowed maximum")
        object.__setattr__(self, "created_at", created)
        object.__setattr__(self, "expires_at", expires)
        if not isinstance(self.nonce, str) or not self.nonce.strip():
            raise ValueError("nonce must be a non-empty string")
        if len(self.nonce.strip()) > 256:
            raise ValueError("nonce is too long")
        object.__setattr__(self, "nonce", self.nonce.strip())
        if not isinstance(self.signature_base64, str):
            raise TypeError("signature_base64 must be a string")
        signature = self.signature_base64.strip()
        if not signature or len(signature) > MAX_DEVICE_PROOF_SIGNATURE_CHARACTERS:
            raise ValueError("device proof signature is empty or too long")
        object.__setattr__(self, "signature_base64", signature)

    def unsigned_mapping(self) -> dict[str, Any]:
        return {
            "schema": self.schema,
            "schemaVersion": self.schema_version,
            "algorithm": self.algorithm,
            "action": self.action,
            "installationId": self.installation_id,
            "devicePublicKey": self.device_public_key_base64,
            "deviceKeyThumbprint": self.device_key_thumbprint,
            "payloadDigest": self.payload_digest,
            "createdAt": format_rfc3339(self.created_at, "created_at"),
            "expiresAt": format_rfc3339(self.expires_at, "expires_at"),
            "nonce": self.nonce,
        }

    def to_mapping(self) -> dict[str, Any]:
        result = self.unsigned_mapping()
        result["signature"] = self.signature_base64
        return result

    @classmethod
    def from_mapping(cls, raw: Mapping[str, Any]) -> "DeviceActionProof":
        if not isinstance(raw, Mapping):
            raise TypeError("device proof must be an object")
        try:
            return cls(
                action=raw.get("action"),
                installation_id=raw.get("installationId"),
                device_public_key_base64=raw.get("devicePublicKey"),
                device_key_thumbprint=raw.get("deviceKeyThumbprint"),
                payload_digest=raw.get("payloadDigest"),
                created_at=parse_rfc3339(raw.get("createdAt"), "createdAt"),
                expires_at=parse_rfc3339(raw.get("expiresAt"), "expiresAt"),
                nonce=raw.get("nonce"),
                signature_base64=raw.get("signature"),
                algorithm=raw.get("algorithm"),
                schema=raw.get("schema"),
                schema_version=raw.get("schemaVersion"),
            )
        except LicensingError:
            raise
        except (TypeError, ValueError) as exc:
            raise LicensingError(
                LicenseErrorCode.INVALID_DOCUMENT,
                f"invalid device proof: {exc}",
            ) from exc


def device_proof_signing_bytes(proof: DeviceActionProof) -> bytes:
    if not isinstance(proof, DeviceActionProof):
        raise TypeError("proof must be a DeviceActionProof")
    return DEVICE_PROOF_DOMAIN + canonicalize_json(proof.unsigned_mapping())


def verify_device_action_proof(
    proof: DeviceActionProof | Mapping[str, Any],
    expected_action: str,
    payload: Mapping[str, Any],
    *,
    at: Optional[datetime] = None,
) -> DeviceActionProof:
    parsed = proof if isinstance(proof, DeviceActionProof) else DeviceActionProof.from_mapping(proof)
    normalized_action = validate_identifier(expected_action, "expected_action")
    if parsed.action != normalized_action:
        raise LicensingError(
            LicenseErrorCode.INVALID_DOCUMENT,
            "device proof action does not match the request",
        )
    expected_digest = device_action_payload_digest(payload)
    if parsed.payload_digest != expected_digest:
        raise LicensingError(
            LicenseErrorCode.INVALID_SIGNATURE,
            "device proof payload digest does not match the request",
        )
    checked_at = datetime.now(timezone.utc) if at is None else at
    if not isinstance(checked_at, datetime):
        raise TypeError("at must be a datetime or None")
    normalized_at = parse_rfc3339(format_rfc3339(checked_at, "at"), "at")
    if normalized_at < parsed.created_at - MAX_DEVICE_PROOF_CLOCK_SKEW:
        raise LicensingError(
            LicenseErrorCode.CLAIM_NOT_YET_VALID,
            "device proof creation time is in the future",
        )
    if normalized_at >= parsed.expires_at:
        raise LicensingError(
            LicenseErrorCode.INVALID_DOCUMENT,
            "device proof has expired",
        )
    try:
        public_key_der = base64.b64decode(parsed.device_public_key_base64, validate=True)
        public_key = serialization.load_der_public_key(public_key_der)
        signature = base64.b64decode(parsed.signature_base64, validate=True)
        if not isinstance(public_key, ec.EllipticCurvePublicKey):
            raise ValueError("device public key is not an EC key")
        public_key.verify(
            signature,
            device_proof_signing_bytes(parsed),
            ec.ECDSA(hashes.SHA256()),
        )
    except (InvalidSignature, TypeError, ValueError) as exc:
        raise LicensingError(
            LicenseErrorCode.INVALID_SIGNATURE,
            "device proof signature verification failed",
        ) from exc
    return parsed


__all__ = [
    "DEVICE_PROOF_ALGORITHM",
    "DEVICE_PROOF_SCHEMA",
    "DEVICE_PROOF_SCHEMA_VERSION",
    "DeviceActionProof",
    "device_action_payload_digest",
    "device_proof_signing_bytes",
    "verify_device_action_proof",
]

