"""Device-signed activation and release requests for air-gapped systems."""

from __future__ import annotations

import base64
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from enum import Enum
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, parse_bounded_json
from .constants import MAX_SIGNED_DOCUMENT_BYTES, PRODUCT_ID, validate_identifier
from .errors import LicenseErrorCode, LicensingError
from .models import format_rfc3339, parse_rfc3339


OFFLINE_REQUEST_SCHEMA = "apolon.licensing.offline-request"
OFFLINE_REQUEST_SCHEMA_VERSION = 2
SUPPORTED_OFFLINE_REQUEST_SCHEMA_VERSIONS = (1, OFFLINE_REQUEST_SCHEMA_VERSION)
DEVICE_SIGNED_DOCUMENT_SCHEMA = "apolon.licensing.device-signed-document"
DEVICE_SIGNED_DOCUMENT_SCHEMA_VERSION = 1
DEVICE_SIGNATURE_ALGORITHM = "ECDSA_P256_SHA256"
DEVICE_SIGNATURE_DOMAIN = b"APOLON-OFFLINE-REQUEST-V1\x00"
MAX_OFFLINE_REQUEST_LIFETIME = timedelta(days=7)
MAX_DEVICE_PUBLIC_KEY_BYTES = 1024
MAX_EVIDENCE_COMPONENTS = 16
MAX_EVIDENCE_DIGEST_CHARACTERS = 128
MAX_FRIENDLY_DEVICE_NAME_CHARACTERS = 128


class OfflineRequestType(str, Enum):
    ACTIVATE = "activate"
    TRIAL = "trial"
    RENEW = "renew"
    RELEASE = "release"


def device_public_key_thumbprint(public_key_der: bytes) -> str:
    if not isinstance(public_key_der, bytes):
        raise TypeError("public_key_der must be bytes")
    if not public_key_der or len(public_key_der) > MAX_DEVICE_PUBLIC_KEY_BYTES:
        raise ValueError("device public key is empty or too large")
    try:
        public_key = serialization.load_der_public_key(public_key_der)
    except ValueError as exc:
        raise ValueError("device public key is not valid DER") from exc
    if not isinstance(public_key, ec.EllipticCurvePublicKey) or not isinstance(
        public_key.curve,
        ec.SECP256R1,
    ):
        raise ValueError("device public key must use ECDSA P-256")
    digest = hashlib.sha256(public_key_der).hexdigest()
    return f"sha256.{digest}"


def _optional_identifier(value: object, field_name: str) -> Optional[str]:
    return None if value is None else validate_identifier(value, field_name)


def _normalize_evidence(raw: Mapping[str, str]) -> tuple[tuple[str, str], ...]:
    if not isinstance(raw, Mapping):
        raise TypeError("evidence_digests must be a mapping")
    if len(raw) > MAX_EVIDENCE_COMPONENTS:
        raise ValueError("too many device-evidence components")
    normalized: list[tuple[str, str]] = []
    for key, value in raw.items():
        component_id = validate_identifier(key, "evidence component ID")
        if not isinstance(value, str):
            raise TypeError("evidence digest must be a string")
        digest = value.strip().lower()
        if not digest or len(digest) > MAX_EVIDENCE_DIGEST_CHARACTERS:
            raise ValueError("evidence digest is empty or too long")
        if any(character not in "0123456789abcdef" for character in digest):
            raise ValueError("evidence digest must be lowercase hexadecimal")
        normalized.append((component_id, digest))
    return tuple(sorted(normalized))


@dataclass(frozen=True)
class OfflineRequestPayload:
    request_type: OfflineRequestType
    request_id: str
    product_id: str
    application_major_version: int
    installation_id: str
    device_id: Optional[str]
    device_public_key_base64: str
    device_key_thumbprint: str
    evidence_digests: tuple[tuple[str, str], ...]
    key_provider: str
    friendly_name: str
    requested_sku_id: Optional[str]
    created_at: datetime
    expires_at: datetime
    nonce: str
    target_license_id: Optional[str] = None
    schema: str = OFFLINE_REQUEST_SCHEMA
    schema_version: int = OFFLINE_REQUEST_SCHEMA_VERSION

    def __post_init__(self) -> None:
        if (
            self.schema != OFFLINE_REQUEST_SCHEMA
            or self.schema_version not in SUPPORTED_OFFLINE_REQUEST_SCHEMA_VERSIONS
        ):
            raise ValueError("unsupported offline-request schema")
        if not isinstance(self.request_type, OfflineRequestType):
            raise TypeError("request_type must be an OfflineRequestType")
        object.__setattr__(self, "request_id", validate_identifier(self.request_id, "request_id"))
        object.__setattr__(self, "product_id", validate_identifier(self.product_id, "product_id"))
        if isinstance(self.application_major_version, bool) or not isinstance(
            self.application_major_version,
            int,
        ):
            raise TypeError("application_major_version must be an integer")
        if self.application_major_version < 1:
            raise ValueError("application_major_version must be at least one")
        object.__setattr__(
            self,
            "installation_id",
            validate_identifier(self.installation_id, "installation_id"),
        )
        object.__setattr__(self, "device_id", _optional_identifier(self.device_id, "device_id"))
        if (
            self.request_type in (OfflineRequestType.RENEW, OfflineRequestType.RELEASE)
            and self.device_id is None
        ):
            raise ValueError("renewal and release requests require device_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_bytes = 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_bytes)
        object.__setattr__(self, "device_public_key_base64", normalized_key)
        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_key_thumbprint",
            normalized_thumbprint,
        )
        if not isinstance(self.evidence_digests, tuple):
            raise TypeError("evidence_digests must be a tuple")
        normalized_evidence = _normalize_evidence(dict(self.evidence_digests))
        if normalized_evidence != self.evidence_digests:
            raise ValueError("evidence_digests must be sorted and unique")
        object.__setattr__(
            self,
            "key_provider",
            validate_identifier(self.key_provider, "key_provider"),
        )
        if not isinstance(self.friendly_name, str):
            raise TypeError("friendly_name must be a string")
        friendly_name = self.friendly_name.strip()
        if not friendly_name or len(friendly_name) > MAX_FRIENDLY_DEVICE_NAME_CHARACTERS:
            raise ValueError("friendly_name is empty or too long")
        object.__setattr__(self, "friendly_name", friendly_name)
        object.__setattr__(
            self,
            "requested_sku_id",
            _optional_identifier(self.requested_sku_id, "requested_sku_id"),
        )
        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_OFFLINE_REQUEST_LIFETIME:
            raise ValueError("offline request 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())
        normalized_target = _optional_identifier(
            self.target_license_id,
            "target_license_id",
        )
        if self.schema_version == 1 and normalized_target is not None:
            raise ValueError("offline-request schema version 1 cannot bind a target license")
        if self.request_type is OfflineRequestType.RENEW:
            if normalized_target is None:
                raise ValueError("renewal requests require target_license_id")
            if self.requested_sku_id is not None:
                raise ValueError("renewal requests must not request a purchase SKU")
        if self.request_type is OfflineRequestType.TRIAL:
            if normalized_target is not None or self.device_id is not None:
                raise ValueError("trial requests must not bind an existing license or device")
            if self.requested_sku_id is not None:
                raise ValueError("trial requests must not request a purchase SKU")
        object.__setattr__(self, "target_license_id", normalized_target)

    def to_mapping(self) -> dict[str, Any]:
        result = {
            "schema": self.schema,
            "schemaVersion": self.schema_version,
            "requestType": self.request_type.value,
            "requestId": self.request_id,
            "productId": self.product_id,
            "applicationMajorVersion": self.application_major_version,
            "installationId": self.installation_id,
            "deviceId": self.device_id,
            "devicePublicKey": self.device_public_key_base64,
            "deviceKeyThumbprint": self.device_key_thumbprint,
            "evidenceDigests": dict(self.evidence_digests),
            "keyProvider": self.key_provider,
            "friendlyName": self.friendly_name,
            "requestedSkuId": self.requested_sku_id,
            "createdAt": format_rfc3339(self.created_at, "created_at"),
            "expiresAt": format_rfc3339(self.expires_at, "expires_at"),
            "nonce": self.nonce,
        }
        if self.schema_version >= 2:
            result["targetLicenseId"] = self.target_license_id
        return result

    @classmethod
    def from_mapping(cls, raw: Mapping[str, Any]) -> "OfflineRequestPayload":
        if not isinstance(raw, Mapping):
            raise TypeError("offline request payload must be an object")
        if (
            raw.get("schema") != OFFLINE_REQUEST_SCHEMA
            or raw.get("schemaVersion") not in SUPPORTED_OFFLINE_REQUEST_SCHEMA_VERSIONS
        ):
            raise LicensingError(
                LicenseErrorCode.UNSUPPORTED_SCHEMA,
                "unsupported offline-request schema",
            )
        evidence = raw.get("evidenceDigests")
        if not isinstance(evidence, Mapping):
            raise TypeError("evidenceDigests must be an object")
        try:
            return cls(
                request_type=OfflineRequestType(raw.get("requestType")),
                request_id=raw.get("requestId"),
                product_id=raw.get("productId"),
                application_major_version=raw.get("applicationMajorVersion"),
                installation_id=raw.get("installationId"),
                device_id=raw.get("deviceId"),
                device_public_key_base64=raw.get("devicePublicKey"),
                device_key_thumbprint=raw.get("deviceKeyThumbprint"),
                evidence_digests=_normalize_evidence(evidence),
                key_provider=raw.get("keyProvider"),
                friendly_name=raw.get("friendlyName"),
                requested_sku_id=raw.get("requestedSkuId"),
                created_at=parse_rfc3339(raw.get("createdAt"), "createdAt"),
                expires_at=parse_rfc3339(raw.get("expiresAt"), "expiresAt"),
                nonce=raw.get("nonce"),
                target_license_id=raw.get("targetLicenseId"),
                schema=raw.get("schema"),
                schema_version=raw.get("schemaVersion"),
            )
        except LicensingError:
            raise
        except (TypeError, ValueError) as exc:
            raise LicensingError(
                LicenseErrorCode.INVALID_DOCUMENT,
                f"invalid offline request: {exc}",
            ) from exc


@dataclass(frozen=True)
class DeviceSignedOfflineRequest:
    payload: OfflineRequestPayload
    signature_base64: str
    algorithm: str = DEVICE_SIGNATURE_ALGORITHM
    schema: str = DEVICE_SIGNED_DOCUMENT_SCHEMA
    schema_version: int = DEVICE_SIGNED_DOCUMENT_SCHEMA_VERSION

    def __post_init__(self) -> None:
        if not isinstance(self.payload, OfflineRequestPayload):
            raise TypeError("payload must be an OfflineRequestPayload")
        if self.schema != DEVICE_SIGNED_DOCUMENT_SCHEMA or self.schema_version != DEVICE_SIGNED_DOCUMENT_SCHEMA_VERSION:
            raise ValueError("unsupported device-signed document schema")
        if self.algorithm != DEVICE_SIGNATURE_ALGORITHM:
            raise ValueError("unsupported device signature algorithm")
        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) > 256:
            raise ValueError("device signature is empty or too long")
        object.__setattr__(self, "signature_base64", signature)

    def to_mapping(self) -> dict[str, Any]:
        return {
            "schema": self.schema,
            "schemaVersion": self.schema_version,
            "algorithm": self.algorithm,
            "payload": self.payload.to_mapping(),
            "signature": self.signature_base64,
        }

    @classmethod
    def from_mapping(cls, raw: Mapping[str, Any]) -> "DeviceSignedOfflineRequest":
        if not isinstance(raw, Mapping):
            raise TypeError("device-signed request must be an object")
        if (
            raw.get("schema") != DEVICE_SIGNED_DOCUMENT_SCHEMA
            or raw.get("schemaVersion") != DEVICE_SIGNED_DOCUMENT_SCHEMA_VERSION
        ):
            raise LicensingError(
                LicenseErrorCode.UNSUPPORTED_SCHEMA,
                "unsupported device-signed request schema",
            )
        try:
            return cls(
                payload=OfflineRequestPayload.from_mapping(raw.get("payload")),
                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-signed request: {exc}",
            ) from exc


def offline_request_signing_bytes(payload: OfflineRequestPayload) -> bytes:
    if not isinstance(payload, OfflineRequestPayload):
        raise TypeError("payload must be an OfflineRequestPayload")
    return DEVICE_SIGNATURE_DOMAIN + canonicalize_json(payload.to_mapping())


def parse_and_verify_offline_request(
    document: str | bytes | bytearray | Mapping[str, Any],
    *,
    at: Optional[datetime] = None,
    expected_product_id: str = PRODUCT_ID,
) -> DeviceSignedOfflineRequest:
    if isinstance(document, Mapping):
        raw = document
    else:
        raw = parse_bounded_json(document, maximum_bytes=MAX_SIGNED_DOCUMENT_BYTES)
    if not isinstance(raw, Mapping):
        raise LicensingError(
            LicenseErrorCode.INVALID_DOCUMENT,
            "offline request root must be an object",
        )
    request = DeviceSignedOfflineRequest.from_mapping(raw)
    normalized_product = validate_identifier(expected_product_id, "expected_product_id")
    if request.payload.product_id != normalized_product:
        raise LicensingError(
            LicenseErrorCode.WRONG_PRODUCT,
            "offline request is for a different product",
        )
    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 < request.payload.created_at - timedelta(minutes=5):
        raise LicensingError(
            LicenseErrorCode.CLAIM_NOT_YET_VALID,
            "offline request creation time is in the future",
        )
    if normalized_at >= request.payload.expires_at:
        raise LicensingError(
            LicenseErrorCode.INVALID_DOCUMENT,
            "offline request has expired",
        )
    try:
        signature = base64.b64decode(request.signature_base64, validate=True)
        public_key_der = base64.b64decode(
            request.payload.device_public_key_base64,
            validate=True,
        )
        public_key = serialization.load_der_public_key(public_key_der)
        if not isinstance(public_key, ec.EllipticCurvePublicKey):
            raise ValueError("device public key is not an EC key")
        public_key.verify(
            signature,
            offline_request_signing_bytes(request.payload),
            ec.ECDSA(hashes.SHA256()),
        )
    except (InvalidSignature, TypeError, ValueError) as exc:
        raise LicensingError(
            LicenseErrorCode.INVALID_SIGNATURE,
            "offline request device signature verification failed",
        ) from exc
    return request


__all__ = [
    "DEVICE_SIGNATURE_ALGORITHM",
    "DeviceSignedOfflineRequest",
    "OfflineRequestPayload",
    "OfflineRequestType",
    "SUPPORTED_OFFLINE_REQUEST_SCHEMA_VERSIONS",
    "device_public_key_thumbprint",
    "offline_request_signing_bytes",
    "parse_and_verify_offline_request",
]
