"""Validated product catalog shared by desktop presentation and server policy."""

from __future__ import annotations

from dataclasses import dataclass
from enum import Enum
import hashlib
from importlib import resources
from types import MappingProxyType
from typing import Any, Mapping

from .canonical_json import canonicalize_json, parse_bounded_json
from .constants import (
    CATALOG_LEGACY_SCHEMA_VERSION,
    CATALOG_SCHEMA,
    CATALOG_SCHEMA_VERSION,
    CATALOG_SUPPORTED_SCHEMA_VERSIONS,
    MAX_CATALOG_DOCUMENT_BYTES,
    MAX_DISPLAY_TEXT_LENGTH,
    PRODUCT_ID,
    TRIAL_MAXIMUM_TOTAL_DURATION_HOURS,
    validate_identifier,
)
from .errors import LicenseErrorCode, LicensingError


class SkuKind(str, Enum):
    EDITION = "edition"
    ADDON = "addon"
    SUBSCRIPTION = "subscription"
    TRIAL = "trial"


class BillingMode(str, Enum):
    PERPETUAL = "perpetual"
    RECURRING = "recurring"
    FREE = "free"
    CONTRACT = "contract"


HOURS_PER_DAY = 24


def _require_exact_fields(
    raw: Mapping[str, Any],
    expected: frozenset[str],
    field_name: str,
) -> None:
    if not isinstance(raw, Mapping):
        raise TypeError(f"{field_name} must be an object")
    actual = set(raw)
    missing = sorted(expected.difference(actual))
    unknown = sorted(actual.difference(expected))
    if missing or unknown:
        raise ValueError(
            f"{field_name} fields are invalid; missing={missing}, unknown={unknown}"
        )


def _display_text(value: object, field_name: str) -> str:
    if not isinstance(value, str):
        raise TypeError(f"{field_name} must be a string")
    normalized = value.strip()
    if not normalized:
        raise ValueError(f"{field_name} must not be empty")
    if len(normalized) > MAX_DISPLAY_TEXT_LENGTH:
        raise ValueError(
            f"{field_name} exceeds the maximum length of {MAX_DISPLAY_TEXT_LENGTH}"
        )
    return normalized


def _positive_integer(value: object, field_name: str, *, allow_zero: bool = False) -> int:
    if isinstance(value, bool) or not isinstance(value, int):
        raise TypeError(f"{field_name} must be an integer")
    minimum = 0 if allow_zero else 1
    if value < minimum:
        raise ValueError(f"{field_name} must be at least {minimum}")
    return value


@dataclass(frozen=True)
class EntitlementDefinition:
    entitlement_id: str
    label: str
    description: str

    @classmethod
    def from_mapping(cls, entitlement_id: str, raw: Mapping[str, Any]) -> "EntitlementDefinition":
        if not isinstance(raw, Mapping):
            raise TypeError("entitlement definition must be an object")
        _require_exact_fields(
            raw,
            frozenset(("description", "label")),
            "entitlement definition",
        )
        return cls(
            entitlement_id=validate_identifier(entitlement_id, "entitlement_id"),
            label=_display_text(raw.get("label"), "entitlement label"),
            description=_display_text(raw.get("description"), "entitlement description"),
        )

    def to_mapping(self) -> dict[str, str]:
        return {
            "label": self.label,
            "description": self.description,
        }


@dataclass(frozen=True)
class DevicePolicyDefinition:
    policy_id: str
    maximum_devices: int
    named_seats: int
    forced_release_limit: int
    forced_release_window_days: int
    rapid_churn_device_limit: int
    rapid_churn_window_days: int
    connected_lease_days: int
    offline_certificate_days: int
    trial_days: int
    trial_extension_days: int
    trial_extension_limit: int
    subscription_payment_grace_days: int
    connected_refresh_hours: int | None = None
    offline_refresh_reminder_days: int | None = None

    @classmethod
    def from_mapping(
        cls,
        policy_id: str,
        raw: Mapping[str, Any],
        *,
        schema_version: int = CATALOG_SCHEMA_VERSION,
    ) -> "DevicePolicyDefinition":
        if not isinstance(raw, Mapping):
            raise TypeError("device policy definition must be an object")
        normalized_schema_version = _positive_integer(
            schema_version,
            "schema_version",
        )
        if normalized_schema_version not in CATALOG_SUPPORTED_SCHEMA_VERSIONS:
            raise ValueError("device policy schema version is not supported")
        expected_fields = {
            "connectedLeaseDays",
            "forcedReleaseLimit",
            "forcedReleaseWindowDays",
            "maximumDevices",
            "namedSeats",
            "offlineCertificateDays",
            "rapidChurnDeviceLimit",
            "rapidChurnWindowDays",
            "subscriptionPaymentGraceDays",
            "trialDays",
            "trialExtensionDays",
            "trialExtensionLimit",
        }
        if normalized_schema_version >= CATALOG_SCHEMA_VERSION:
            expected_fields.update(
                ("connectedRefreshHours", "offlineRefreshReminderDays")
            )
        _require_exact_fields(
            raw,
            frozenset(expected_fields),
            "device policy definition",
        )
        connected_lease_days = _positive_integer(
            raw.get("connectedLeaseDays"),
            "connectedLeaseDays",
        )
        offline_certificate_days = _positive_integer(
            raw.get("offlineCertificateDays"),
            "offlineCertificateDays",
            allow_zero=True,
        )
        connected_refresh_hours = None
        offline_refresh_reminder_days = None
        if normalized_schema_version >= CATALOG_SCHEMA_VERSION:
            connected_refresh_hours = _positive_integer(
                raw.get("connectedRefreshHours"),
                "connectedRefreshHours",
            )
            offline_refresh_reminder_days = _positive_integer(
                raw.get("offlineRefreshReminderDays"),
                "offlineRefreshReminderDays",
                allow_zero=True,
            )
            if connected_refresh_hours > connected_lease_days * HOURS_PER_DAY:
                raise ValueError(
                    "connectedRefreshHours must not exceed the connected lease"
                )
            if offline_certificate_days == 0:
                if offline_refresh_reminder_days != 0:
                    raise ValueError(
                        "offlineRefreshReminderDays must be zero when offline "
                        "certificates are disabled"
                    )
            elif not 1 <= offline_refresh_reminder_days <= offline_certificate_days:
                raise ValueError(
                    "offlineRefreshReminderDays must be within the offline "
                    "certificate validity interval"
                )
        trial_days = _positive_integer(
            raw.get("trialDays"),
            "trialDays",
            allow_zero=True,
        )
        trial_extension_days = _positive_integer(
            raw.get("trialExtensionDays", 0),
            "trialExtensionDays",
            allow_zero=True,
        )
        trial_extension_limit = _positive_integer(
            raw.get("trialExtensionLimit", 0),
            "trialExtensionLimit",
            allow_zero=True,
        )
        if (trial_extension_days == 0) != (trial_extension_limit == 0):
            raise ValueError(
                "trialExtensionDays and trialExtensionLimit must both be zero or positive"
            )
        maximum_trial_hours = (
            trial_days + trial_extension_days * trial_extension_limit
        ) * HOURS_PER_DAY
        if maximum_trial_hours > TRIAL_MAXIMUM_TOTAL_DURATION_HOURS:
            raise ValueError(
                "trial policy exceeds the maximum duration accepted by shipped clients"
            )
        return cls(
            policy_id=validate_identifier(policy_id, "policy_id"),
            maximum_devices=_positive_integer(raw.get("maximumDevices"), "maximumDevices"),
            named_seats=_positive_integer(raw.get("namedSeats"), "namedSeats"),
            forced_release_limit=_positive_integer(
                raw.get("forcedReleaseLimit"),
                "forcedReleaseLimit",
                allow_zero=True,
            ),
            forced_release_window_days=_positive_integer(
                raw.get("forcedReleaseWindowDays"),
                "forcedReleaseWindowDays",
            ),
            rapid_churn_device_limit=_positive_integer(
                raw.get("rapidChurnDeviceLimit"),
                "rapidChurnDeviceLimit",
            ),
            rapid_churn_window_days=_positive_integer(
                raw.get("rapidChurnWindowDays"),
                "rapidChurnWindowDays",
            ),
            connected_lease_days=connected_lease_days,
            offline_certificate_days=offline_certificate_days,
            trial_days=trial_days,
            trial_extension_days=trial_extension_days,
            trial_extension_limit=trial_extension_limit,
            subscription_payment_grace_days=_positive_integer(
                raw.get("subscriptionPaymentGraceDays"),
                "subscriptionPaymentGraceDays",
                allow_zero=True,
            ),
            connected_refresh_hours=connected_refresh_hours,
            offline_refresh_reminder_days=offline_refresh_reminder_days,
        )

    def to_mapping(self) -> dict[str, int]:
        result = {
            "maximumDevices": self.maximum_devices,
            "namedSeats": self.named_seats,
            "forcedReleaseLimit": self.forced_release_limit,
            "forcedReleaseWindowDays": self.forced_release_window_days,
            "rapidChurnDeviceLimit": self.rapid_churn_device_limit,
            "rapidChurnWindowDays": self.rapid_churn_window_days,
            "connectedLeaseDays": self.connected_lease_days,
            "offlineCertificateDays": self.offline_certificate_days,
            "trialDays": self.trial_days,
            "trialExtensionDays": self.trial_extension_days,
            "trialExtensionLimit": self.trial_extension_limit,
            "subscriptionPaymentGraceDays": self.subscription_payment_grace_days,
        }
        if (
            self.connected_refresh_hours is not None
            and self.offline_refresh_reminder_days is not None
        ):
            result["connectedRefreshHours"] = self.connected_refresh_hours
            result["offlineRefreshReminderDays"] = (
                self.offline_refresh_reminder_days
            )
        return result


@dataclass(frozen=True)
class SkuDefinition:
    sku_id: str
    kind: SkuKind
    edition_level: int
    billing_mode: BillingMode
    label: str
    entitlements: tuple[str, ...]
    device_policy_id: str
    account_required: bool
    active: bool
    application_major_minimum: int
    application_major_maximum: int
    compatible_editions: tuple[str, ...]

    @classmethod
    def from_mapping(cls, sku_id: str, raw: Mapping[str, Any]) -> "SkuDefinition":
        if not isinstance(raw, Mapping):
            raise TypeError("SKU definition must be an object")
        kind = SkuKind(raw.get("kind"))
        expected_fields = {
            "accountRequired",
            "active",
            "applicationMajorMaximum",
            "applicationMajorMinimum",
            "billingMode",
            "compatibleEditions",
            "devicePolicyId",
            "entitlements",
            "kind",
            "label",
        }
        if kind is SkuKind.EDITION:
            expected_fields.add("editionLevel")
        _require_exact_fields(raw, frozenset(expected_fields), "SKU definition")
        entitlement_values = raw.get("entitlements")
        compatible_values = raw.get("compatibleEditions", [])
        if not isinstance(entitlement_values, list) or not entitlement_values:
            raise ValueError("SKU entitlements must be a non-empty list")
        if not isinstance(compatible_values, list):
            raise TypeError("compatibleEditions must be a list")
        if not isinstance(raw.get("accountRequired"), bool):
            raise TypeError("accountRequired must be a Boolean")
        if not isinstance(raw.get("active"), bool):
            raise TypeError("active must be a Boolean")
        minimum = _positive_integer(raw.get("applicationMajorMinimum"), "applicationMajorMinimum")
        maximum = _positive_integer(raw.get("applicationMajorMaximum"), "applicationMajorMaximum")
        if maximum < minimum:
            raise ValueError("applicationMajorMaximum must not be below the minimum")
        edition_level = _positive_integer(
            raw.get("editionLevel", 0),
            "editionLevel",
            allow_zero=True,
        )
        if kind is SkuKind.EDITION and edition_level < 1:
            raise ValueError("edition SKUs require a positive editionLevel")
        if kind is not SkuKind.EDITION and edition_level != 0:
            raise ValueError("only edition SKUs may define editionLevel")
        normalized_entitlements = tuple(
            validate_identifier(value, "SKU entitlement")
            for value in entitlement_values
        )
        normalized_compatible_editions = tuple(
            validate_identifier(value, "compatible edition")
            for value in compatible_values
        )
        if len(normalized_entitlements) != len(set(normalized_entitlements)):
            raise ValueError("SKU entitlements must be unique")
        if len(normalized_compatible_editions) != len(
            set(normalized_compatible_editions)
        ):
            raise ValueError("compatibleEditions must be unique")
        return cls(
            sku_id=validate_identifier(sku_id, "sku_id"),
            kind=kind,
            edition_level=edition_level,
            billing_mode=BillingMode(raw.get("billingMode")),
            label=_display_text(raw.get("label"), "SKU label"),
            entitlements=normalized_entitlements,
            device_policy_id=validate_identifier(
                raw.get("devicePolicyId"),
                "devicePolicyId",
            ),
            account_required=raw["accountRequired"],
            active=raw["active"],
            application_major_minimum=minimum,
            application_major_maximum=maximum,
            compatible_editions=normalized_compatible_editions,
        )

    def to_mapping(self) -> dict[str, object]:
        result: dict[str, object] = {
            "kind": self.kind.value,
            "billingMode": self.billing_mode.value,
            "label": self.label,
            "entitlements": sorted(self.entitlements),
            "devicePolicyId": self.device_policy_id,
            "accountRequired": self.account_required,
            "active": self.active,
            "applicationMajorMinimum": self.application_major_minimum,
            "applicationMajorMaximum": self.application_major_maximum,
            "compatibleEditions": sorted(self.compatible_editions),
        }
        if self.kind is SkuKind.EDITION:
            result["editionLevel"] = self.edition_level
        return result


class LicensingCatalog:
    """Immutable, fully cross-referenced product catalog."""

    def __init__(
        self,
        *,
        product_id: str,
        revision: int,
        entitlements: Mapping[str, EntitlementDefinition],
        device_policies: Mapping[str, DevicePolicyDefinition],
        skus: Mapping[str, SkuDefinition],
        schema_version: int = CATALOG_SCHEMA_VERSION,
    ) -> None:
        normalized_product = validate_identifier(product_id, "product_id")
        normalized_revision = _positive_integer(revision, "revision")
        normalized_schema_version = _positive_integer(
            schema_version,
            "schema_version",
        )
        if normalized_schema_version not in CATALOG_SUPPORTED_SCHEMA_VERSIONS:
            raise ValueError("catalog schema version is not supported")
        entitlement_map = dict(entitlements)
        policy_map = dict(device_policies)
        sku_map = dict(skus)
        if not entitlement_map:
            raise ValueError("catalog must contain entitlements")
        if not policy_map:
            raise ValueError("catalog must contain device policies")
        if not sku_map:
            raise ValueError("catalog must contain SKUs")
        for identifier, definition in entitlement_map.items():
            if identifier != definition.entitlement_id:
                raise ValueError("entitlement mapping key does not match definition")
        for identifier, definition in policy_map.items():
            if identifier != definition.policy_id:
                raise ValueError("device policy mapping key does not match definition")
            has_any_refresh_policy = (
                definition.connected_refresh_hours is not None
                or definition.offline_refresh_reminder_days is not None
            )
            has_complete_refresh_policy = (
                definition.connected_refresh_hours is not None
                and definition.offline_refresh_reminder_days is not None
            )
            if normalized_schema_version == CATALOG_LEGACY_SCHEMA_VERSION:
                if has_any_refresh_policy:
                    raise ValueError(
                        "legacy catalog policies must not define refresh cadence"
                    )
            elif not has_complete_refresh_policy:
                raise ValueError("catalog policies must define refresh cadence")
        for identifier, definition in sku_map.items():
            if identifier != definition.sku_id:
                raise ValueError("SKU mapping key does not match definition")
            missing = set(definition.entitlements).difference(entitlement_map)
            if missing:
                raise ValueError(
                    f"SKU {identifier} references unknown entitlements: {sorted(missing)}"
                )
            if definition.device_policy_id not in policy_map:
                raise ValueError(
                    f"SKU {identifier} references unknown device policy "
                    f"{definition.device_policy_id}"
                )
        edition_ids = {
            sku_id for sku_id, definition in sku_map.items() if definition.kind is SkuKind.EDITION
        }
        edition_levels = [
            definition.edition_level
            for definition in sku_map.values()
            if definition.kind is SkuKind.EDITION
        ]
        if len(edition_levels) != len(set(edition_levels)):
            raise ValueError("editionLevel values must be unique")
        for identifier, definition in sku_map.items():
            missing_editions = set(definition.compatible_editions).difference(edition_ids)
            if missing_editions:
                raise ValueError(
                    f"SKU {identifier} references unknown compatible editions: "
                    f"{sorted(missing_editions)}"
                )
        self.product_id = normalized_product
        self.revision = normalized_revision
        self.schema_version = normalized_schema_version
        self.entitlements = MappingProxyType(entitlement_map)
        self.device_policies = MappingProxyType(policy_map)
        self.skus = MappingProxyType(sku_map)

    @classmethod
    def from_mapping(cls, raw: Mapping[str, Any]) -> "LicensingCatalog":
        if not isinstance(raw, Mapping):
            raise TypeError("catalog document must be an object")
        try:
            _require_exact_fields(
                raw,
                frozenset(
                    (
                        "devicePolicies",
                        "entitlements",
                        "productId",
                        "revision",
                        "schema",
                        "schemaVersion",
                        "skus",
                    )
                ),
                "catalog document",
            )
        except (TypeError, ValueError) as exc:
            raise LicensingError(
                LicenseErrorCode.CATALOG_INVALID,
                f"invalid licensing catalog: {exc}",
            ) from exc
        schema_version = raw.get("schemaVersion")
        if (
            raw.get("schema") != CATALOG_SCHEMA
            or isinstance(schema_version, bool)
            or not isinstance(schema_version, int)
            or schema_version not in CATALOG_SUPPORTED_SCHEMA_VERSIONS
        ):
            raise LicensingError(
                LicenseErrorCode.UNSUPPORTED_SCHEMA,
                "unsupported licensing catalog schema",
            )
        raw_entitlements = raw.get("entitlements")
        raw_policies = raw.get("devicePolicies")
        raw_skus = raw.get("skus")
        if not isinstance(raw_entitlements, Mapping):
            raise TypeError("catalog entitlements must be an object")
        if not isinstance(raw_policies, Mapping):
            raise TypeError("catalog devicePolicies must be an object")
        if not isinstance(raw_skus, Mapping):
            raise TypeError("catalog skus must be an object")
        try:
            entitlements = {
                identifier: EntitlementDefinition.from_mapping(identifier, value)
                for identifier, value in raw_entitlements.items()
            }
            policies = {
                identifier: DevicePolicyDefinition.from_mapping(
                    identifier,
                    value,
                    schema_version=schema_version,
                )
                for identifier, value in raw_policies.items()
            }
            skus = {
                identifier: SkuDefinition.from_mapping(identifier, value)
                for identifier, value in raw_skus.items()
            }
            return cls(
                product_id=raw.get("productId"),
                revision=raw.get("revision"),
                entitlements=entitlements,
                device_policies=policies,
                skus=skus,
                schema_version=schema_version,
            )
        except LicensingError:
            raise
        except (TypeError, ValueError) as exc:
            raise LicensingError(
                LicenseErrorCode.CATALOG_INVALID,
                f"invalid licensing catalog: {exc}",
            ) from exc

    def require_sku(self, sku_id: str, *, require_active: bool = True) -> SkuDefinition:
        normalized = validate_identifier(sku_id, "sku_id")
        if not isinstance(require_active, bool):
            raise TypeError("require_active must be a Boolean")
        definition = self.skus.get(normalized)
        if definition is None:
            raise LicensingError(
                LicenseErrorCode.CATALOG_INVALID,
                f"unknown SKU: {normalized}",
            )
        if require_active and not definition.active:
            raise LicensingError(
                LicenseErrorCode.CATALOG_INVALID,
                f"SKU is not active: {normalized}",
            )
        return definition

    def validate_entitlement_ids(self, values: tuple[str, ...] | list[str]) -> tuple[str, ...]:
        if not isinstance(values, (tuple, list)):
            raise TypeError("values must be a tuple or list")
        normalized = tuple(validate_identifier(value, "entitlement") for value in values)
        missing = set(normalized).difference(self.entitlements)
        if missing:
            raise LicensingError(
                LicenseErrorCode.UNKNOWN_ENTITLEMENT,
                f"unknown entitlement identifiers: {sorted(missing)}",
            )
        return tuple(sorted(set(normalized)))

    def to_mapping(self) -> dict[str, object]:
        return {
            "schema": CATALOG_SCHEMA,
            "schemaVersion": self.schema_version,
            "productId": self.product_id,
            "revision": self.revision,
            "entitlements": {
                identifier: definition.to_mapping()
                for identifier, definition in self.entitlements.items()
            },
            "devicePolicies": {
                identifier: definition.to_mapping()
                for identifier, definition in self.device_policies.items()
            },
            "skus": {
                identifier: definition.to_mapping()
                for identifier, definition in self.skus.items()
            },
        }

    def canonical_bytes(self) -> bytes:
        return canonicalize_json(self.to_mapping())

    def sha256(self) -> str:
        return hashlib.sha256(self.canonical_bytes()).hexdigest()


def load_catalog_document(document: str | bytes | bytearray) -> LicensingCatalog:
    raw = parse_bounded_json(document, maximum_bytes=MAX_CATALOG_DOCUMENT_BYTES)
    if not isinstance(raw, Mapping):
        raise LicensingError(
            LicenseErrorCode.CATALOG_INVALID,
            "licensing catalog root must be an object",
        )
    return LicensingCatalog.from_mapping(raw)


def load_builtin_catalog() -> LicensingCatalog:
    catalog_resource = resources.files(__package__).joinpath("resources/catalog.v1.json")
    document = catalog_resource.read_bytes()
    catalog = load_catalog_document(document)
    if catalog.product_id != PRODUCT_ID:
        raise LicensingError(
            LicenseErrorCode.CATALOG_INVALID,
            "built-in catalog product ID does not match the application product ID",
        )
    return catalog


__all__ = [
    "BillingMode",
    "DevicePolicyDefinition",
    "EntitlementDefinition",
    "LicensingCatalog",
    "SkuDefinition",
    "SkuKind",
    "load_builtin_catalog",
    "load_catalog_document",
]
