"""Deterministic domain contracts for Route Conditions Intelligence (RCI).

These contracts describe evidence, condition state and detected signals. They
do not define scientific rules or thresholds: callers provide those through a
signal definition and its policy metadata.
"""

from __future__ import annotations

import hashlib
import json
from dataclasses import asdict, dataclass, field
from enum import Enum
from typing import Any, Mapping

from mountain_twin.analysis_contract import FreshnessState


class ConditionState(str, Enum):
    """Truth state of one condition; UNKNOWN and FALSE are intentionally distinct."""

    TRUE = "TRUE"
    FALSE = "FALSE"
    UNKNOWN = "UNKNOWN"
    NOT_APPLICABLE = "NOT_APPLICABLE"


class CoverageState(str, Enum):
    """Evidence coverage, separate from a condition's truth state."""

    FULL = "FULL"
    PARTIAL = "PARTIAL"
    UNAVAILABLE = "UNAVAILABLE"
    UNKNOWN = "UNKNOWN"
    NOT_APPLICABLE = "NOT_APPLICABLE"


class ThresholdOrigin(str, Enum):
    """Origin of a rule threshold; no origin implies no fabricated reference."""

    REFERENCE_STANDARD = "REFERENCE_STANDARD"
    REFERENCE_SCALE = "REFERENCE_SCALE"
    PROVIDER_SPECIFIC = "PROVIDER_SPECIFIC"
    MT_RND_POLICY = "MT_RND_POLICY"


class BoundaryLocation(str, Enum):
    """How a route boundary location was established, not its measurement precision."""

    EXACT = "EXACT"
    INTERPOLATED = "INTERPOLATED"
    SAMPLE = "SAMPLE"
    UNKNOWN = "UNKNOWN"


class BoundaryReason(str, Enum):
    THRESHOLD_CROSSING = "THRESHOLD_CROSSING"
    STATE_TRANSITION = "STATE_TRANSITION"
    ANALYSIS_WINDOW_BOUNDARY = "ANALYSIS_WINDOW_BOUNDARY"
    DATA_GAP = "DATA_GAP"
    NOT_APPLICABLE_DISCONTINUITY = "NOT_APPLICABLE_DISCONTINUITY"
    PROVENANCE_DISCONTINUITY = "PROVENANCE_DISCONTINUITY"


class SpanKind(str, Enum):
    ACTIVE = "ACTIVE"
    BRIDGED_FALSE_GAP = "BRIDGED_FALSE_GAP"


class QualificationMode(str, Enum):
    TIME = "TIME"
    DISTANCE = "DISTANCE"
    AND = "AND"
    OR = "OR"
    EVENT = "EVENT"


class HysteresisDirection(str, Enum):
    ABOVE = "ABOVE"
    BELOW = "BELOW"


@dataclass(frozen=True)
class ThresholdProvenance:
    origin: ThresholdOrigin
    reference_id: str | None = None
    version: str | None = None
    note: str | None = None

    def to_dict(self) -> dict[str, Any]:
        return _json_value(asdict(self))


@dataclass(frozen=True)
class EvidenceQuality:
    """Multidimensional evidence metadata without a synthetic confidence score."""

    coverage: CoverageState
    freshness: FreshnessState | None = None
    source: str | None = None
    model_run: str | None = None
    spatial_characteristics: Mapping[str, Any] = field(default_factory=dict)
    temporal_characteristics: Mapping[str, Any] = field(default_factory=dict)
    interpolation: str | None = None
    derived_status: str | None = None
    reason_codes: tuple[str, ...] = ()
    limitations: tuple[str, ...] = ()

    def to_dict(self) -> dict[str, Any]:
        return _json_value(asdict(self))


@dataclass(frozen=True)
class ConditionBoundary:
    """A boundary with explicit location semantics and reason.

    An UNKNOWN boundary deliberately carries no route point, distance or time:
    evidence discontinuity is not a fabricated threshold crossing.
    """

    location: BoundaryLocation
    reason: BoundaryReason
    route_point_index: int | None = None
    route_distance_m: float | None = None
    planned_time: str | None = None
    elapsed_planned_seconds: float | None = None
    note: str | None = None

    def __post_init__(self) -> None:
        if self.location is BoundaryLocation.UNKNOWN and any(
            value is not None
            for value in (
                self.route_point_index,
                self.route_distance_m,
                self.planned_time,
                self.elapsed_planned_seconds,
            )
        ):
            raise ValueError("UNKNOWN boundary cannot carry a fabricated route location")

    def to_dict(self) -> dict[str, Any]:
        return _json_value(asdict(self))


@dataclass(frozen=True)
class ConditionSpan:
    """A route span that is either active evidence or a bridged FALSE gap."""

    kind: SpanKind
    start_route_index: int
    end_route_index: int
    start_route_distance_m: float
    end_route_distance_m: float
    start_elapsed_planned_seconds: float | None = None
    end_elapsed_planned_seconds: float | None = None

    def __post_init__(self) -> None:
        if self.end_route_index < self.start_route_index:
            raise ValueError("span route indices must be ordered")
        if self.end_route_distance_m < self.start_route_distance_m:
            raise ValueError("span route distances must be monotonic")
        if (
            self.start_elapsed_planned_seconds is not None
            and self.end_elapsed_planned_seconds is not None
            and self.end_elapsed_planned_seconds < self.start_elapsed_planned_seconds
        ):
            raise ValueError("span elapsed planned time must be monotonic")

    @property
    def distance_m(self) -> float:
        return self.end_route_distance_m - self.start_route_distance_m

    @property
    def duration_s(self) -> float | None:
        if (
            self.start_elapsed_planned_seconds is None
            or self.end_elapsed_planned_seconds is None
        ):
            return None
        return self.end_elapsed_planned_seconds - self.start_elapsed_planned_seconds

    def to_dict(self) -> dict[str, Any]:
        document = _json_value(asdict(self))
        document["distance_m"] = self.distance_m
        document["duration_s"] = self.duration_s
        return document


@dataclass(frozen=True)
class HysteresisPolicy:
    """Optional numeric enter/exit policy; NONE preserves supplied condition states."""

    policy_id: str
    direction: HysteresisDirection | None = None
    enter_threshold: float | None = None
    exit_threshold: float | None = None
    enter_inclusive: bool = True
    exit_inclusive: bool = True

    def __post_init__(self) -> None:
        numeric_fields = (self.direction, self.enter_threshold, self.exit_threshold)
        if any(value is not None for value in numeric_fields) and any(
            value is None for value in numeric_fields
        ):
            raise ValueError("numeric hysteresis requires direction, enter and exit thresholds")

    @property
    def is_enabled(self) -> bool:
        return self.direction is not None

    @classmethod
    def none(cls, policy_id: str = "NONE") -> "HysteresisPolicy":
        return cls(policy_id=policy_id)

    def to_dict(self) -> dict[str, Any]:
        return _json_value(asdict(self))


@dataclass(frozen=True)
class GapBridgingPolicy:
    """Optional limits for bridging only FALSE gaps; configured limits all apply."""

    policy_id: str
    maximum_elapsed_seconds: float | None = None
    maximum_route_distance_m: float | None = None

    def __post_init__(self) -> None:
        if self.maximum_elapsed_seconds is None and self.maximum_route_distance_m is None:
            raise ValueError("gap bridging needs at least one physical limit")
        if self.maximum_elapsed_seconds is not None and self.maximum_elapsed_seconds < 0:
            raise ValueError("maximum elapsed duration must be non-negative")
        if self.maximum_route_distance_m is not None and self.maximum_route_distance_m < 0:
            raise ValueError("maximum route distance must be non-negative")

    def to_dict(self) -> dict[str, Any]:
        return _json_value(asdict(self))


@dataclass(frozen=True)
class QualificationPolicy:
    """Physical active-span qualification, with no global minimums implied."""

    policy_id: str
    mode: QualificationMode
    minimum_active_duration_s: float | None = None
    minimum_active_distance_m: float | None = None

    def __post_init__(self) -> None:
        for value in (self.minimum_active_duration_s, self.minimum_active_distance_m):
            if value is not None and value < 0:
                raise ValueError("qualification limits must be non-negative")
        if self.mode is QualificationMode.TIME and self.minimum_active_duration_s is None:
            raise ValueError("TIME qualification needs a duration limit")
        if self.mode is QualificationMode.DISTANCE and self.minimum_active_distance_m is None:
            raise ValueError("DISTANCE qualification needs a distance limit")
        if self.mode in {QualificationMode.AND, QualificationMode.OR} and (
            self.minimum_active_duration_s is None or self.minimum_active_distance_m is None
        ):
            raise ValueError("AND/OR qualification needs duration and distance limits")
        if self.mode is QualificationMode.EVENT and any(
            value is not None
            for value in (self.minimum_active_duration_s, self.minimum_active_distance_m)
        ):
            raise ValueError("EVENT qualification does not use duration or distance limits")

    def to_dict(self) -> dict[str, Any]:
        return _json_value(asdict(self))


@dataclass(frozen=True)
class SignalDefinition:
    """One catalogued RCI rule; rule content may remain explicitly TBD."""

    signal_id: str
    category: str
    name: str
    description: str
    input_variables: tuple[str, ...]
    rule_specification: Mapping[str, Any]
    threshold_provenance: tuple[ThresholdProvenance, ...]
    applicability_policy: Mapping[str, Any]
    hysteresis_policy: HysteresisPolicy
    qualification_policy: QualificationPolicy
    explanation_key: str
    definition_version: str | None = None
    gap_bridging_policy: GapBridgingPolicy | None = None

    def __post_init__(self) -> None:
        if not self.signal_id:
            raise ValueError("signal definition requires a stable signal ID")

    def to_dict(self) -> dict[str, Any]:
        return _json_value(asdict(self))


@dataclass(frozen=True)
class SignalDefinitionRegistry:
    """A deterministic catalog ordered by stable signal ID."""

    catalog_version: str
    definitions: tuple[SignalDefinition, ...] = ()

    def __post_init__(self) -> None:
        identifiers = tuple(definition.signal_id for definition in self.definitions)
        duplicates = tuple(sorted({identifier for identifier in identifiers if identifiers.count(identifier) > 1}))
        if duplicates:
            raise ValueError(f"duplicate signal definition IDs: {', '.join(duplicates)}")
        object.__setattr__(
            self,
            "definitions",
            tuple(sorted(self.definitions, key=lambda definition: definition.signal_id)),
        )

    def get(self, signal_id: str) -> SignalDefinition:
        for definition in self.definitions:
            if definition.signal_id == signal_id:
                return definition
        raise KeyError(signal_id)

    def to_dict(self) -> dict[str, Any]:
        return _json_value(asdict(self))


@dataclass(frozen=True)
class ConditionStateRecord:
    """One ordered state sample supplied to or retained by RCI."""

    definition_id: str
    route_point_index: int
    route_distance_m: float
    state: ConditionState
    elapsed_planned_seconds: float | None = None
    planned_time: str | None = None
    value: float | None = None
    source_continuity_key: str | None = None
    coverage: CoverageState = CoverageState.UNKNOWN
    quality: EvidenceQuality | None = None
    boundary_hint: ConditionBoundary | None = None

    def to_dict(self) -> dict[str, Any]:
        return _json_value(asdict(self))


@dataclass(frozen=True)
class ConditionSignal:
    """A future-facing qualified RCI signal, never a safety verdict."""

    signal_instance_id: str
    definition_id: str
    category: str
    start_route_index: int
    end_route_index: int
    start_route_distance_m: float
    end_route_distance_m: float
    active_spans: tuple[ConditionSpan, ...]
    bridged_false_gaps: tuple[ConditionSpan, ...]
    applied_rule: Mapping[str, Any]
    threshold_provenance: tuple[ThresholdProvenance, ...]
    coverage: CoverageState
    quality: EvidenceQuality | None
    boundaries: tuple[ConditionBoundary, ...]
    explanation_key: str
    start_planned_time: str | None = None
    end_planned_time: str | None = None
    active_duration_s: float | None = None
    active_distance_m: float = 0.0
    total_duration_s: float | None = None
    relevant_values: Mapping[str, Any] = field(default_factory=dict)
    evidence_payload: Mapping[str, Any] = field(default_factory=dict)
    why: "StructuredWhy | None" = None

    def __post_init__(self) -> None:
        if self.end_route_index < self.start_route_index:
            raise ValueError("signal route indices must be ordered")
        if self.end_route_distance_m < self.start_route_distance_m:
            raise ValueError("signal route distances must be monotonic")
        if self.active_distance_m < 0:
            raise ValueError("active route distance must be non-negative")
        if self.active_distance_m > self.total_distance_m:
            raise ValueError("active route distance cannot exceed total span")
        if self.active_duration_s is not None and self.active_duration_s < 0:
            raise ValueError("active duration must be non-negative")
        if (
            self.active_duration_s is not None
            and self.total_duration_s is not None
            and self.active_duration_s > self.total_duration_s
        ):
            raise ValueError("active duration cannot exceed total span")

    @property
    def total_distance_m(self) -> float:
        return self.end_route_distance_m - self.start_route_distance_m

    def to_dict(self) -> dict[str, Any]:
        document = _json_value(asdict(self))
        document["total_distance_m"] = self.total_distance_m
        return document


@dataclass(frozen=True)
class RouteConditionCluster:
    cluster_id: str
    signal_instance_ids: tuple[str, ...]
    definition_ids: tuple[str, ...]
    start_route_index: int | None
    end_route_index: int | None
    start_route_distance_m: float | None
    end_route_distance_m: float | None
    start_planned_time: str | None
    end_planned_time: str | None
    overlap_axis: str
    quality: EvidenceQuality | None = None
    evidence_payload: Mapping[str, Any] = field(default_factory=dict)
    explanation_key: str = "rci.cluster.overlap"
    why: "StructuredWhy | None" = None

    def __post_init__(self) -> None:
        if len(self.definition_ids) < 2:
            raise ValueError("a cluster requires at least two different signal definitions")
        if tuple(sorted(set(self.signal_instance_ids))) != self.signal_instance_ids:
            raise ValueError("cluster signal IDs must be unique and ordered")
        if tuple(sorted(set(self.definition_ids))) != self.definition_ids:
            raise ValueError("cluster definition IDs must be unique and ordered")

    def to_dict(self) -> dict[str, Any]:
        return _json_value(asdict(self))


@dataclass(frozen=True)
class RouteConditionEvent:
    event_id: str
    event_type: str
    route_point_index: int | None
    route_distance_m: float | None
    planned_time: str | None
    value: float | None = None
    unit: str | None = None
    source_reference: str | None = None
    quality: EvidenceQuality | None = None
    explanation_key: str = "rci.event"
    evidence_payload: Mapping[str, Any] = field(default_factory=dict)
    why: "StructuredWhy | None" = None

    def to_dict(self) -> dict[str, Any]:
        return _json_value(asdict(self))


@dataclass(frozen=True)
class StructuredWhy:
    """Machine-readable evidence for later UI/AI explanation, never advice."""

    what: Mapping[str, Any]
    where: Mapping[str, Any]
    when: Mapping[str, Any]
    input_values: Mapping[str, Any]
    rule: Mapping[str, Any]
    result: Mapping[str, Any]
    source: Mapping[str, Any]
    quality: Mapping[str, Any]
    boundaries: tuple[ConditionBoundary, ...] = ()

    def to_dict(self) -> dict[str, Any]:
        return _json_value(asdict(self))


@dataclass(frozen=True)
class RouteConditionsResult:
    """Versioned RCI result skeleton composed from, not duplicating, Unified Analysis."""

    analysis_metadata: Mapping[str, Any]
    route_reference: Mapping[str, Any]
    planning_reference: Mapping[str, Any]
    catalog_version: str
    policy_version: str
    condition_states: tuple[ConditionStateRecord, ...] = ()
    signals: tuple[ConditionSignal, ...] = ()
    clusters: tuple[RouteConditionCluster, ...] = ()
    events: tuple[RouteConditionEvent, ...] = ()
    quality: EvidenceQuality | None = None
    provenance: Mapping[str, Any] = field(default_factory=dict)
    diagnostics: tuple[str, ...] = ()
    state_summary: Mapping[str, Any] = field(default_factory=dict)
    candidate_diagnostics: tuple[Mapping[str, Any], ...] = ()
    contract_version: str = "0.2"

    def to_dict(self) -> dict[str, Any]:
        return _json_value(asdict(self))


def derive_condition_signal_id(
    *,
    analysis_reference: str | None,
    definition_id: str,
    start_route_index: int,
    end_route_index: int,
    catalog_version: str,
    policy_version: str,
) -> str:
    """Derive a stable signal ID from stable semantic inputs only."""
    identity = {
        "analysis_reference": analysis_reference,
        "catalog_version": catalog_version,
        "definition_id": definition_id,
        "end_route_index": end_route_index,
        "policy_version": policy_version,
        "start_route_index": start_route_index,
    }
    encoded = json.dumps(identity, sort_keys=True, separators=(",", ":"), allow_nan=False)
    return f"rci-{hashlib.sha256(encoded.encode('utf-8')).hexdigest()[:24]}"


def serialize_route_conditions_result(result: RouteConditionsResult) -> str:
    """Canonical serialization without runtime clocks or process-specific state."""
    return json.dumps(result.to_dict(), sort_keys=True, separators=(",", ":"), allow_nan=False)


def _json_value(value: Any) -> Any:
    if isinstance(value, Enum):
        return value.value
    if isinstance(value, Mapping):
        return {str(key): _json_value(value[key]) for key in sorted(value, key=str)}
    if isinstance(value, tuple):
        return [_json_value(item) for item in value]
    if isinstance(value, list):
        return [_json_value(item) for item in value]
    if hasattr(value, "__dataclass_fields__"):
        return _json_value(asdict(value))
    return value
