"""Bounded signed update-manifest contracts and public verification."""

from __future__ import annotations

import base64
from dataclasses import dataclass
from datetime import datetime, timedelta
import hashlib
import os
from pathlib import Path
import re
import stat
from typing import Any, Mapping, Optional
from urllib.parse import urlparse

from cryptography.exceptions import InvalidSignature
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PublicKey

from .canonical_json import canonicalize_json, parse_bounded_json
from .constants import MAX_SAFE_JSON_INTEGER, PRODUCT_ID, validate_identifier
from .errors import LicenseErrorCode, LicensingError
from .models import format_rfc3339, parse_rfc3339
from .verifier import ED25519_SIGNATURE_BYTES, PublicKeyRing


UPDATE_MANIFEST_SCHEMA = "apolon.update.manifest"
UPDATE_MANIFEST_SCHEMA_VERSION = 1
SIGNED_UPDATE_MANIFEST_SCHEMA = "apolon.update.signed-manifest"
SIGNED_UPDATE_MANIFEST_SCHEMA_VERSION = 1
SIGNED_UPDATE_MANIFEST_ALGORITHM = "Ed25519"
SIGNED_UPDATE_MANIFEST_DOMAIN = b"APOLON-SIGNED-UPDATE-MANIFEST-V1\x00"
UPDATE_SIGNING_KEY_ID_PREFIX = "update."
MAXIMUM_SIGNED_UPDATE_MANIFEST_BYTES = 128 * 1024
MAXIMUM_UPDATE_VERSION_LENGTH = 128
MAXIMUM_UPDATE_FILE_NAME_LENGTH = 255
MAXIMUM_UPDATE_URL_LENGTH = 4096
UPDATE_ARTIFACT_HASH_CHUNK_BYTES = 1024 * 1024
UPDATE_CLOCK_TOLERANCE = timedelta(minutes=5)
SHA256_HEX_PATTERN = re.compile(r"^[0-9a-f]{64}$")
SEMANTIC_VERSION_PATTERN = re.compile(
    r"^(0|[1-9][0-9]*)\."
    r"(0|[1-9][0-9]*)\."
    r"(0|[1-9][0-9]*)"
    r"(?:-((?:0|[1-9][0-9]*|[0-9A-Za-z-]*[A-Za-z-][0-9A-Za-z-]*)"
    r"(?:\.(?:0|[1-9][0-9]*|[0-9A-Za-z-]*[A-Za-z-][0-9A-Za-z-]*))*))?"
    r"(?:\+([0-9A-Za-z-]+(?:\.[0-9A-Za-z-]+)*))?$"
)
SIGNED_UPDATE_MANIFEST_FIELDS = frozenset(
    (
        "algorithm",
        "keyId",
        "payload",
        "schema",
        "schemaVersion",
        "signature",
    )
)
UPDATE_MANIFEST_PAYLOAD_FIELDS = frozenset(
    (
        "artifact",
        "channel",
        "minimumSupportedVersion",
        "productId",
        "releaseId",
        "releasedAt",
        "schema",
        "schemaVersion",
        "signingKeyId",
        "version",
    )
)
UPDATE_ARTIFACT_FIELDS = frozenset(
    (
        "downloadUrl",
        "fileName",
        "sha256",
        "sizeBytes",
    )
)


def validate_update_signing_key_id(value: object) -> str:
    key_id = validate_identifier(value, "update signing key ID")
    if not key_id.startswith(UPDATE_SIGNING_KEY_ID_PREFIX):
        raise ValueError(
            f"update signing key ID must start with {UPDATE_SIGNING_KEY_ID_PREFIX!r}"
        )
    return key_id


def validate_semantic_version(value: object, field_name: str) -> str:
    if not isinstance(field_name, str) or not field_name.strip():
        raise ValueError("field_name must be a non-empty string")
    label = field_name.strip()
    if not isinstance(value, str):
        raise TypeError(f"{label} must be a string")
    normalized = value.strip()
    if not normalized or len(normalized) > MAXIMUM_UPDATE_VERSION_LENGTH:
        raise ValueError(f"{label} is empty or too long")
    if SEMANTIC_VERSION_PATTERN.fullmatch(normalized) is None:
        raise ValueError(f"{label} must be a SemVer 2.0.0 version")
    return normalized


def _semantic_version_parts(
    value: object,
    field_name: str,
) -> tuple[tuple[int, int, int], tuple[str, ...]]:
    normalized = validate_semantic_version(value, field_name)
    match = SEMANTIC_VERSION_PATTERN.fullmatch(normalized)
    if match is None:
        raise AssertionError("validated semantic version did not match its parser")
    prerelease = () if match.group(4) is None else tuple(match.group(4).split("."))
    return (
        (int(match.group(1)), int(match.group(2)), int(match.group(3))),
        prerelease,
    )


def compare_semantic_versions(left: object, right: object) -> int:
    left_core, left_prerelease = _semantic_version_parts(left, "left version")
    right_core, right_prerelease = _semantic_version_parts(right, "right version")
    if left_core != right_core:
        return -1 if left_core < right_core else 1
    if not left_prerelease or not right_prerelease:
        if left_prerelease == right_prerelease:
            return 0
        return -1 if left_prerelease else 1
    for left_part, right_part in zip(left_prerelease, right_prerelease):
        if left_part == right_part:
            continue
        left_numeric = left_part.isdigit()
        right_numeric = right_part.isdigit()
        if left_numeric and right_numeric:
            return -1 if int(left_part) < int(right_part) else 1
        if left_numeric != right_numeric:
            return -1 if left_numeric else 1
        return -1 if left_part < right_part else 1
    if len(left_prerelease) == len(right_prerelease):
        return 0
    return -1 if len(left_prerelease) < len(right_prerelease) else 1


def semantic_version_major(value: object, field_name: str = "version") -> int:
    core, _prerelease = _semantic_version_parts(value, field_name)
    return core[0]


def validate_update_file_name(value: object) -> str:
    if not isinstance(value, str):
        raise TypeError("artifact fileName must be a string")
    normalized = value.strip()
    if (
        not normalized
        or len(normalized) > MAXIMUM_UPDATE_FILE_NAME_LENGTH
        or normalized in (".", "..")
        or "/" in normalized
        or "\\" in normalized
        or any(ord(character) < 32 for character in normalized)
    ):
        raise ValueError("artifact fileName must be one safe file name")
    return normalized


def validate_update_download_url(value: object) -> str:
    if not isinstance(value, str):
        raise TypeError("artifact downloadUrl must be a string")
    normalized = value.strip()
    if (
        not normalized
        or len(normalized) > MAXIMUM_UPDATE_URL_LENGTH
        or any(character.isspace() for character in normalized)
    ):
        raise ValueError("artifact downloadUrl is empty, too long, or contains whitespace")
    parsed = urlparse(normalized)
    try:
        parsed_port = parsed.port
    except ValueError as exc:
        raise ValueError("artifact downloadUrl contains an invalid port") from exc
    if (
        parsed.scheme != "https"
        or parsed.hostname is None
        or parsed.username is not None
        or parsed.password is not None
        or parsed.fragment
        or parsed.path in ("", "/")
        or parsed_port is not None and not 1 <= parsed_port <= 65535
    ):
        raise ValueError(
            "artifact downloadUrl must be an HTTPS file URL without credentials or a fragment"
        )
    return normalized


def _positive_artifact_size(value: object) -> int:
    if isinstance(value, bool) or not isinstance(value, int):
        raise TypeError("artifact sizeBytes must be an integer")
    if value < 1 or value > MAX_SAFE_JSON_INTEGER:
        raise ValueError("artifact sizeBytes is outside the supported range")
    return value


def _sha256_hex(value: object) -> str:
    if not isinstance(value, str):
        raise TypeError("artifact sha256 must be a string")
    normalized = value.strip().lower()
    if SHA256_HEX_PATTERN.fullmatch(normalized) is None:
        raise ValueError("artifact sha256 must contain 64 lowercase hexadecimal digits")
    return normalized


def hash_update_artifact(path: Path) -> tuple[int, str]:
    if not isinstance(path, Path):
        raise TypeError("path must be a Path")
    if path.is_symlink():
        raise ValueError("update artifact must be a regular non-symlink file")
    flags = os.O_RDONLY
    if hasattr(os, "O_BINARY"):
        flags |= os.O_BINARY
    if hasattr(os, "O_NOFOLLOW"):
        flags |= os.O_NOFOLLOW
    try:
        descriptor = os.open(path, flags)
    except OSError as exc:
        raise ValueError("update artifact could not be opened safely") from exc
    try:
        initial_status = os.fstat(descriptor)
        if not stat.S_ISREG(initial_status.st_mode):
            raise ValueError("update artifact must be a regular non-symlink file")
        if initial_status.st_size < 1 or initial_status.st_size > MAX_SAFE_JSON_INTEGER:
            raise ValueError("update artifact size is outside the supported range")
        digest = hashlib.sha256()
        with os.fdopen(descriptor, "rb", closefd=False) as stream:
            while True:
                block = stream.read(UPDATE_ARTIFACT_HASH_CHUNK_BYTES)
                if not block:
                    break
                digest.update(block)
        final_status = os.fstat(descriptor)
        if (
            final_status.st_size != initial_status.st_size
            or final_status.st_mtime_ns != initial_status.st_mtime_ns
        ):
            raise ValueError("update artifact changed while it was being hashed")
        return initial_status.st_size, digest.hexdigest()
    finally:
        os.close(descriptor)


@dataclass(frozen=True)
class UpdateArtifact:
    file_name: str
    size_bytes: int
    sha256: str
    download_url: str

    def __post_init__(self) -> None:
        object.__setattr__(self, "file_name", validate_update_file_name(self.file_name))
        object.__setattr__(self, "size_bytes", _positive_artifact_size(self.size_bytes))
        object.__setattr__(self, "sha256", _sha256_hex(self.sha256))
        object.__setattr__(
            self,
            "download_url",
            validate_update_download_url(self.download_url),
        )

    def to_mapping(self) -> dict[str, object]:
        return {
            "fileName": self.file_name,
            "sizeBytes": self.size_bytes,
            "sha256": self.sha256,
            "downloadUrl": self.download_url,
        }

    @classmethod
    def from_mapping(cls, raw: Mapping[str, Any]) -> "UpdateArtifact":
        if not isinstance(raw, Mapping):
            raise TypeError("artifact must be an object")
        if set(raw) != UPDATE_ARTIFACT_FIELDS:
            raise ValueError("artifact fields are missing or unknown")
        return cls(
            file_name=raw.get("fileName"),
            size_bytes=raw.get("sizeBytes"),
            sha256=raw.get("sha256"),
            download_url=raw.get("downloadUrl"),
        )


@dataclass(frozen=True)
class UpdateManifest:
    release_id: str
    version: str
    channel: str
    released_at: datetime
    signing_key_id: str
    artifact: UpdateArtifact
    minimum_supported_version: Optional[str] = None
    product_id: str = PRODUCT_ID
    schema: str = UPDATE_MANIFEST_SCHEMA
    schema_version: int = UPDATE_MANIFEST_SCHEMA_VERSION

    def __post_init__(self) -> None:
        if (
            self.schema != UPDATE_MANIFEST_SCHEMA
            or self.schema_version != UPDATE_MANIFEST_SCHEMA_VERSION
        ):
            raise ValueError("unsupported update-manifest schema")
        object.__setattr__(self, "release_id", validate_identifier(self.release_id, "release_id"))
        object.__setattr__(self, "version", validate_semantic_version(self.version, "version"))
        object.__setattr__(self, "channel", validate_identifier(self.channel, "channel"))
        object.__setattr__(
            self,
            "released_at",
            parse_rfc3339(format_rfc3339(self.released_at, "released_at"), "released_at"),
        )
        object.__setattr__(
            self,
            "signing_key_id",
            validate_update_signing_key_id(self.signing_key_id),
        )
        if not isinstance(self.artifact, UpdateArtifact):
            raise TypeError("artifact must be an UpdateArtifact")
        object.__setattr__(self, "product_id", validate_identifier(self.product_id, "product_id"))
        if self.minimum_supported_version is not None:
            object.__setattr__(
                self,
                "minimum_supported_version",
                validate_semantic_version(
                    self.minimum_supported_version,
                    "minimum_supported_version",
                ),
            )
            if (
                compare_semantic_versions(
                    self.minimum_supported_version,
                    self.version,
                )
                > 0
            ):
                raise ValueError(
                    "minimum_supported_version must not exceed the update version"
                )

    def to_mapping(self) -> dict[str, object]:
        return {
            "schema": self.schema,
            "schemaVersion": self.schema_version,
            "productId": self.product_id,
            "releaseId": self.release_id,
            "version": self.version,
            "channel": self.channel,
            "releasedAt": format_rfc3339(self.released_at, "released_at"),
            "minimumSupportedVersion": self.minimum_supported_version,
            "signingKeyId": self.signing_key_id,
            "artifact": self.artifact.to_mapping(),
        }

    @classmethod
    def from_mapping(cls, raw: Mapping[str, Any]) -> "UpdateManifest":
        if not isinstance(raw, Mapping):
            raise TypeError("update-manifest payload must be an object")
        if set(raw) != UPDATE_MANIFEST_PAYLOAD_FIELDS:
            raise ValueError("update-manifest payload fields are missing or unknown")
        minimum_version = raw.get("minimumSupportedVersion")
        return cls(
            release_id=raw.get("releaseId"),
            version=raw.get("version"),
            channel=raw.get("channel"),
            released_at=parse_rfc3339(raw.get("releasedAt"), "releasedAt"),
            signing_key_id=raw.get("signingKeyId"),
            artifact=UpdateArtifact.from_mapping(raw.get("artifact")),
            minimum_supported_version=minimum_version,
            product_id=raw.get("productId"),
            schema=raw.get("schema"),
            schema_version=raw.get("schemaVersion"),
        )


@dataclass(frozen=True)
class SignedUpdateManifest:
    key_id: str
    payload: Mapping[str, Any]
    signature_base64: str
    algorithm: str = SIGNED_UPDATE_MANIFEST_ALGORITHM
    schema: str = SIGNED_UPDATE_MANIFEST_SCHEMA
    schema_version: int = SIGNED_UPDATE_MANIFEST_SCHEMA_VERSION

    def __post_init__(self) -> None:
        if (
            self.schema != SIGNED_UPDATE_MANIFEST_SCHEMA
            or self.schema_version != SIGNED_UPDATE_MANIFEST_SCHEMA_VERSION
        ):
            raise ValueError("unsupported signed update-manifest schema")
        if self.algorithm != SIGNED_UPDATE_MANIFEST_ALGORITHM:
            raise ValueError("unsupported signed update-manifest algorithm")
        object.__setattr__(self, "key_id", validate_update_signing_key_id(self.key_id))
        if not isinstance(self.payload, Mapping):
            raise TypeError("payload must be an object")
        object.__setattr__(self, "payload", dict(self.payload))
        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("signature_base64 is empty or too long")
        object.__setattr__(self, "signature_base64", signature)

    def to_mapping(self) -> dict[str, object]:
        return {
            "schema": self.schema,
            "schemaVersion": self.schema_version,
            "algorithm": self.algorithm,
            "keyId": self.key_id,
            "payload": dict(self.payload),
            "signature": self.signature_base64,
        }

    @classmethod
    def from_mapping(cls, raw: Mapping[str, Any]) -> "SignedUpdateManifest":
        if not isinstance(raw, Mapping):
            raise TypeError("signed update manifest must be an object")
        if set(raw) != SIGNED_UPDATE_MANIFEST_FIELDS:
            raise ValueError("signed update-manifest fields are missing or unknown")
        return cls(
            key_id=raw.get("keyId"),
            payload=raw.get("payload"),
            signature_base64=raw.get("signature"),
            algorithm=raw.get("algorithm"),
            schema=raw.get("schema"),
            schema_version=raw.get("schemaVersion"),
        )


class UpdateManifestVerifier:
    """Verify an update manifest with a dedicated public update-key ring."""

    def __init__(
        self,
        key_ring: PublicKeyRing,
        *,
        expected_product_id: str = PRODUCT_ID,
    ) -> None:
        if not isinstance(key_ring, PublicKeyRing):
            raise TypeError("key_ring must be a PublicKeyRing")
        self.key_ring = key_ring
        self.expected_product_id = validate_identifier(
            expected_product_id,
            "expected_product_id",
        )

    @staticmethod
    def signing_bytes(payload: Mapping[str, Any]) -> bytes:
        if not isinstance(payload, Mapping):
            raise TypeError("payload must be an object")
        return SIGNED_UPDATE_MANIFEST_DOMAIN + canonicalize_json(payload)

    def parse_document(
        self,
        document: str | bytes | bytearray | Mapping[str, Any],
    ) -> SignedUpdateManifest:
        if isinstance(document, Mapping):
            raw = document
        else:
            raw = parse_bounded_json(
                document,
                maximum_bytes=MAXIMUM_SIGNED_UPDATE_MANIFEST_BYTES,
            )
        if not isinstance(raw, Mapping):
            raise LicensingError(
                LicenseErrorCode.INVALID_DOCUMENT,
                "signed update-manifest root must be an object",
            )
        try:
            return SignedUpdateManifest.from_mapping(raw)
        except LicensingError:
            raise
        except (TypeError, ValueError) as exc:
            raise LicensingError(
                LicenseErrorCode.INVALID_DOCUMENT,
                f"invalid signed update manifest: {exc}",
            ) from exc

    def verify(
        self,
        document: str | bytes | bytearray | Mapping[str, Any],
    ) -> UpdateManifest:
        signed = self.parse_document(document)
        record = self.key_ring.records.get(signed.key_id)
        if record is None:
            raise LicensingError(
                LicenseErrorCode.UNKNOWN_SIGNING_KEY,
                f"update manifest uses unknown key ID: {signed.key_id}",
            )
        try:
            signature = base64.b64decode(signed.signature_base64, validate=True)
        except (TypeError, ValueError) as exc:
            raise LicensingError(
                LicenseErrorCode.INVALID_DOCUMENT,
                "update-manifest signature is not valid Base64",
            ) from exc
        if len(signature) != ED25519_SIGNATURE_BYTES:
            raise LicensingError(
                LicenseErrorCode.INVALID_DOCUMENT,
                f"Ed25519 signature must contain {ED25519_SIGNATURE_BYTES} bytes",
            )
        try:
            Ed25519PublicKey.from_public_bytes(record.public_key_bytes).verify(
                signature,
                self.signing_bytes(signed.payload),
            )
        except InvalidSignature as exc:
            raise LicensingError(
                LicenseErrorCode.INVALID_SIGNATURE,
                "update-manifest signature verification failed",
            ) from exc
        try:
            manifest = UpdateManifest.from_mapping(signed.payload)
        except LicensingError:
            raise
        except (TypeError, ValueError) as exc:
            raise LicensingError(
                LicenseErrorCode.INVALID_DOCUMENT,
                f"invalid update-manifest payload: {exc}",
            ) from exc
        if manifest.signing_key_id != signed.key_id:
            raise LicensingError(
                LicenseErrorCode.INVALID_DOCUMENT,
                "payload update signing key ID does not match the envelope",
            )
        if manifest.product_id != self.expected_product_id:
            raise LicensingError(
                LicenseErrorCode.WRONG_PRODUCT,
                "update manifest is for a different product",
            )
        record.validate_issuance_time(manifest.released_at)
        return manifest

    @staticmethod
    def verify_artifact(manifest: UpdateManifest, path: Path) -> None:
        if not isinstance(manifest, UpdateManifest):
            raise TypeError("manifest must be an UpdateManifest")
        if not isinstance(path, Path):
            raise TypeError("path must be a Path")
        if path.name != manifest.artifact.file_name:
            raise ValueError("update artifact file name does not match the manifest")
        size_bytes, sha256 = hash_update_artifact(path)
        if size_bytes != manifest.artifact.size_bytes:
            raise ValueError("update artifact size does not match the manifest")
        if sha256 != manifest.artifact.sha256:
            raise ValueError("update artifact digest does not match the manifest")


__all__ = [
    "MAXIMUM_SIGNED_UPDATE_MANIFEST_BYTES",
    "SIGNED_UPDATE_MANIFEST_ALGORITHM",
    "SIGNED_UPDATE_MANIFEST_DOMAIN",
    "SIGNED_UPDATE_MANIFEST_SCHEMA",
    "SIGNED_UPDATE_MANIFEST_SCHEMA_VERSION",
    "SignedUpdateManifest",
    "UPDATE_CLOCK_TOLERANCE",
    "UPDATE_MANIFEST_SCHEMA",
    "UPDATE_MANIFEST_SCHEMA_VERSION",
    "UPDATE_SIGNING_KEY_ID_PREFIX",
    "UpdateArtifact",
    "UpdateManifest",
    "UpdateManifestVerifier",
    "compare_semantic_versions",
    "hash_update_artifact",
    "semantic_version_major",
    "validate_semantic_version",
    "validate_update_download_url",
    "validate_update_file_name",
    "validate_update_signing_key_id",
]
