"""Deterministic, physical-axis RCI detection mechanics.

This module deliberately contains no production signal definitions or
scientific thresholds. It turns caller-supplied states and explicit policies
into candidates and qualified ``ConditionSignal`` records. D1 integration is
therefore deferred until an accepted definition supplies scientific rules.

For distance/time accounting, a sampled state applies forward from its route
point to the next sample. A transition is therefore bounded by the next sample
unless explicit numeric interpolation or an exact caller-provided boundary is
available. This is a deterministic sampling convention, not a claim of
physical transition precision.
"""

from __future__ import annotations

from dataclasses import dataclass, replace
from typing import Any, Iterable, Mapping, Sequence

from mountain_twin.rci.contracts import (
    BoundaryLocation,
    BoundaryReason,
    ConditionBoundary,
    ConditionSignal,
    ConditionSpan,
    ConditionState,
    ConditionStateRecord,
    CoverageState,
    EvidenceQuality,
    HysteresisDirection,
    QualificationMode,
    SignalDefinition,
    SpanKind,
    derive_condition_signal_id,
)


@dataclass(frozen=True)
class DetectionCandidate:
    """One detected active candidate, whether or not it qualifies as a signal."""

    definition_id: str
    start_route_index: int
    end_route_index: int
    start_route_distance_m: float
    end_route_distance_m: float
    start_planned_time: str | None
    end_planned_time: str | None
    total_duration_s: float | None
    active_duration_s: float | None
    active_distance_m: float
    active_spans: tuple[ConditionSpan, ...]
    bridged_false_gaps: tuple[ConditionSpan, ...]
    boundaries: tuple[ConditionBoundary, ...]
    coverage: CoverageState
    quality: EvidenceQuality
    qualified: bool
    relevant_values: Mapping[str, Any]

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


@dataclass(frozen=True)
class DetectionResult:
    """Stable ordered output; candidates include failed qualification attempts."""

    condition_states: tuple[ConditionStateRecord, ...]
    candidates: tuple[DetectionCandidate, ...]
    signals: tuple[ConditionSignal, ...]


def detect_condition_signals(
    *,
    definition: SignalDefinition,
    samples: Sequence[ConditionStateRecord],
    analysis_reference: str | None,
    catalog_version: str,
    policy_version: str,
) -> DetectionResult:
    """Detect and qualify deterministic RCI candidates along ordered route samples."""
    original = tuple(samples)
    _validate_samples(definition.signal_id, original)
    states = _apply_hysteresis(definition, original)
    bridgeable = _bridgeable_false_gaps(states, definition.gap_bridging_policy)
    candidates = []
    position = 0
    while position < len(states):
        if states[position].state is not ConditionState.TRUE:
            position += 1
            continue
        candidate, next_position = _candidate_at(
            definition=definition,
            samples=states,
            start=position,
            bridgeable=bridgeable,
        )
        candidates.append(candidate)
        position = next_position

    qualified = tuple(candidate for candidate in candidates if candidate.qualified)
    signals = tuple(
        _signal_from_candidate(
            definition=definition,
            candidate=candidate,
            analysis_reference=analysis_reference,
            catalog_version=catalog_version,
            policy_version=policy_version,
        )
        for candidate in qualified
    )
    return DetectionResult(states, tuple(candidates), signals)


def _validate_samples(definition_id: str, samples: Sequence[ConditionStateRecord]) -> None:
    for previous, current in zip(samples, samples[1:]):
        if current.definition_id != definition_id or previous.definition_id != definition_id:
            raise ValueError("all state records must match the signal definition")
        if current.route_point_index <= previous.route_point_index:
            raise ValueError("route point indices must be strictly increasing")
        if current.route_distance_m < previous.route_distance_m:
            raise ValueError("route distances must be monotonic")
    if samples and samples[0].definition_id != definition_id:
        raise ValueError("all state records must match the signal definition")
    known_times = [
        sample.elapsed_planned_seconds
        for sample in samples
        if sample.elapsed_planned_seconds is not None
    ]
    if any(current < previous for previous, current in zip(known_times, known_times[1:])):
        raise ValueError("elapsed planned time must be monotonic where available")


def _apply_hysteresis(
    definition: SignalDefinition,
    samples: Sequence[ConditionStateRecord],
) -> tuple[ConditionStateRecord, ...]:
    policy = definition.hysteresis_policy
    if not policy.is_enabled:
        return tuple(samples)
    active = False
    result = []
    for sample in samples:
        if sample.state in {ConditionState.UNKNOWN, ConditionState.NOT_APPLICABLE}:
            active = False
            result.append(sample)
            continue
        if sample.value is None:
            active = False
            result.append(replace(sample, state=ConditionState.UNKNOWN))
            continue
        if active:
            active = _matches_exit(policy, sample.value)
        else:
            active = _matches_enter(policy, sample.value)
        result.append(replace(sample, state=ConditionState.TRUE if active else ConditionState.FALSE))
    return tuple(result)


def _matches_enter(policy, value: float) -> bool:
    if policy.direction is HysteresisDirection.ABOVE:
        return value >= policy.enter_threshold if policy.enter_inclusive else value > policy.enter_threshold
    return value <= policy.enter_threshold if policy.enter_inclusive else value < policy.enter_threshold


def _matches_exit(policy, value: float) -> bool:
    if policy.direction is HysteresisDirection.ABOVE:
        return value >= policy.exit_threshold if policy.exit_inclusive else value > policy.exit_threshold
    return value <= policy.exit_threshold if policy.exit_inclusive else value < policy.exit_threshold


def _bridgeable_false_gaps(samples, policy) -> dict[int, int]:
    if policy is None:
        return {}
    bridgeable = {}
    index = 0
    while index < len(samples):
        if samples[index].state is not ConditionState.FALSE:
            index += 1
            continue
        start = index
        while index < len(samples) and samples[index].state is ConditionState.FALSE:
            index += 1
        end = index
        if start == 0 or end == len(samples):
            continue
        if (
            samples[start - 1].state is not ConditionState.TRUE
            or samples[end].state is not ConditionState.TRUE
            or not _continuous(samples[start - 1 : end + 1])
        ):
            continue
        distance = samples[end].route_distance_m - samples[start].route_distance_m
        elapsed = _duration(samples[start], samples[end])
        if policy.maximum_route_distance_m is not None and distance > policy.maximum_route_distance_m:
            continue
        if (
            policy.maximum_elapsed_seconds is not None
            and (elapsed is None or elapsed > policy.maximum_elapsed_seconds)
        ):
            continue
        bridgeable[start] = end
    return bridgeable


def _candidate_at(*, definition, samples, start, bridgeable):
    start_sample = samples[start]
    active_start = start
    cursor = start
    active_spans = []
    gaps = []
    boundaries = [_start_boundary(samples, start, definition)]
    end = start

    while True:
        if cursor == len(samples) - 1:
            active_spans.append(_span(SpanKind.ACTIVE, samples[active_start], samples[cursor]))
            end = cursor
            boundaries.append(_sample_boundary(samples[cursor], BoundaryReason.ANALYSIS_WINDOW_BOUNDARY))
            next_position = cursor + 1
            break
        following = cursor + 1
        if not _compatible(samples[cursor], samples[following]):
            active_spans.append(_span(SpanKind.ACTIVE, samples[active_start], samples[cursor]))
            end = cursor
            boundaries.append(
                ConditionBoundary(BoundaryLocation.UNKNOWN, BoundaryReason.PROVENANCE_DISCONTINUITY)
            )
            next_position = following
            break
        following_state = samples[following].state
        if following_state is ConditionState.TRUE:
            cursor = following
            continue
        if following_state is ConditionState.FALSE and following in bridgeable:
            gap_end = bridgeable[following]
            active_spans.append(_span(SpanKind.ACTIVE, samples[active_start], samples[following]))
            gaps.append(_span(SpanKind.BRIDGED_FALSE_GAP, samples[following], samples[gap_end]))
            cursor = gap_end
            active_start = cursor
            continue
        if following_state is ConditionState.FALSE:
            active_spans.append(_span(SpanKind.ACTIVE, samples[active_start], samples[following]))
            end = following
            boundaries.append(_transition_boundary(samples[cursor], samples[following], definition))
            next_position = following + 1
            break
        active_spans.append(_span(SpanKind.ACTIVE, samples[active_start], samples[cursor]))
        end = cursor
        reason = (
            BoundaryReason.DATA_GAP
            if following_state is ConditionState.UNKNOWN
            else BoundaryReason.NOT_APPLICABLE_DISCONTINUITY
        )
        boundaries.append(ConditionBoundary(BoundaryLocation.UNKNOWN, reason))
        next_position = following
        break

    active_distance = sum(span.distance_m for span in active_spans)
    durations = [span.duration_s for span in active_spans]
    active_duration = sum(durations) if all(value is not None for value in durations) else None
    total_duration = _duration(start_sample, samples[end])
    quality = _aggregate_quality(samples[start : end + 1])
    candidate = DetectionCandidate(
        definition_id=definition.signal_id,
        start_route_index=start_sample.route_point_index,
        end_route_index=samples[end].route_point_index,
        start_route_distance_m=start_sample.route_distance_m,
        end_route_distance_m=samples[end].route_distance_m,
        start_planned_time=start_sample.planned_time,
        end_planned_time=samples[end].planned_time,
        total_duration_s=total_duration,
        active_duration_s=active_duration,
        active_distance_m=active_distance,
        active_spans=tuple(active_spans),
        bridged_false_gaps=tuple(gaps),
        boundaries=tuple(boundaries),
        coverage=quality.coverage,
        quality=quality,
        qualified=_qualifies(
            definition.qualification_policy,
            active_duration_s=active_duration,
            active_distance_m=active_distance,
            has_active_state=bool(active_spans),
        ),
        relevant_values={
            "active_values": tuple(
                sample.value
                for sample in samples[start : end + 1]
                if sample.state is ConditionState.TRUE and sample.value is not None
            )
        },
    )
    return candidate, next_position


def _start_boundary(samples, start, definition):
    current = samples[start]
    if start == 0:
        return _sample_boundary(current, BoundaryReason.ANALYSIS_WINDOW_BOUNDARY)
    previous = samples[start - 1]
    if not _compatible(previous, current):
        return ConditionBoundary(BoundaryLocation.UNKNOWN, BoundaryReason.PROVENANCE_DISCONTINUITY)
    if previous.state is ConditionState.UNKNOWN:
        return ConditionBoundary(BoundaryLocation.UNKNOWN, BoundaryReason.DATA_GAP)
    if previous.state is ConditionState.NOT_APPLICABLE:
        return ConditionBoundary(BoundaryLocation.UNKNOWN, BoundaryReason.NOT_APPLICABLE_DISCONTINUITY)
    return _transition_boundary(previous, current, definition)


def _transition_boundary(before, after, definition):
    hint = after.boundary_hint if hasattr(after, "boundary_hint") else None
    if hint is not None and hint.location is BoundaryLocation.EXACT:
        return hint
    policy = definition.hysteresis_policy
    if (
        policy.is_enabled
        and before.value is not None
        and after.value is not None
        and before.value != after.value
    ):
        threshold = (
            policy.enter_threshold
            if before.state is ConditionState.FALSE and after.state is ConditionState.TRUE
            else policy.exit_threshold
        )
        fraction = (threshold - before.value) / (after.value - before.value)
        if 0.0 <= fraction <= 1.0:
            elapsed = None
            if (
                before.elapsed_planned_seconds is not None
                and after.elapsed_planned_seconds is not None
            ):
                elapsed = before.elapsed_planned_seconds + fraction * (
                    after.elapsed_planned_seconds - before.elapsed_planned_seconds
                )
            return ConditionBoundary(
                BoundaryLocation.INTERPOLATED,
                BoundaryReason.THRESHOLD_CROSSING,
                route_distance_m=before.route_distance_m
                + fraction * (after.route_distance_m - before.route_distance_m),
                elapsed_planned_seconds=elapsed,
            )
    return _sample_boundary(after, BoundaryReason.STATE_TRANSITION)


def _sample_boundary(sample, reason):
    return ConditionBoundary(
        BoundaryLocation.SAMPLE,
        reason,
        route_point_index=sample.route_point_index,
        route_distance_m=sample.route_distance_m,
        planned_time=sample.planned_time,
        elapsed_planned_seconds=sample.elapsed_planned_seconds,
    )


def _span(kind, start, end):
    return ConditionSpan(
        kind=kind,
        start_route_index=start.route_point_index,
        end_route_index=end.route_point_index,
        start_route_distance_m=start.route_distance_m,
        end_route_distance_m=end.route_distance_m,
        start_elapsed_planned_seconds=start.elapsed_planned_seconds,
        end_elapsed_planned_seconds=end.elapsed_planned_seconds,
    )


def _duration(start, end):
    if start.elapsed_planned_seconds is None or end.elapsed_planned_seconds is None:
        return None
    return end.elapsed_planned_seconds - start.elapsed_planned_seconds


def _compatible(left, right):
    return _continuity_key(left) == _continuity_key(right)


def _continuity_key(sample):
    if sample.source_continuity_key is not None:
        return ("EXPLICIT", sample.source_continuity_key)
    if sample.quality is not None:
        return ("QUALITY", sample.quality.source, sample.quality.model_run)
    return None


def _continuous(samples: Iterable[ConditionStateRecord]) -> bool:
    items = tuple(samples)
    return all(_compatible(left, right) for left, right in zip(items, items[1:]))


def _aggregate_quality(samples):
    coverage_states = {sample.coverage for sample in samples}
    if coverage_states == {CoverageState.NOT_APPLICABLE}:
        coverage = CoverageState.NOT_APPLICABLE
    elif CoverageState.UNAVAILABLE in coverage_states:
        coverage = CoverageState.UNAVAILABLE
    elif CoverageState.UNKNOWN in coverage_states:
        coverage = CoverageState.UNKNOWN
    elif coverage_states == {CoverageState.FULL}:
        coverage = CoverageState.FULL
    else:
        coverage = CoverageState.PARTIAL
    reasons = tuple(
        sorted(
            {
                reason
                for sample in samples
                if sample.quality is not None
                for reason in sample.quality.reason_codes
            }
        )
    )
    limitations = tuple(
        sorted(
            {
                limitation
                for sample in samples
                if sample.quality is not None
                for limitation in sample.quality.limitations
            }
        )
    )
    return EvidenceQuality(coverage=coverage, reason_codes=reasons, limitations=limitations)


def _qualifies(policy, *, active_duration_s, active_distance_m, has_active_state):
    if policy.mode is QualificationMode.EVENT:
        return has_active_state
    if policy.mode is QualificationMode.TIME:
        return (
            active_duration_s is not None
            and active_duration_s >= policy.minimum_active_duration_s
        )
    if policy.mode is QualificationMode.DISTANCE:
        return active_distance_m >= policy.minimum_active_distance_m
    time_passes = (
        active_duration_s is not None and active_duration_s >= policy.minimum_active_duration_s
    )
    distance_passes = active_distance_m >= policy.minimum_active_distance_m
    if policy.mode is QualificationMode.AND:
        return time_passes and distance_passes
    if policy.mode is QualificationMode.OR:
        return time_passes or distance_passes
    raise ValueError(f"unsupported qualification mode: {policy.mode}")


def _signal_from_candidate(*, definition, candidate, analysis_reference, catalog_version, policy_version):
    return ConditionSignal(
        signal_instance_id=derive_condition_signal_id(
            analysis_reference=analysis_reference,
            definition_id=definition.signal_id,
            start_route_index=candidate.start_route_index,
            end_route_index=candidate.end_route_index,
            catalog_version=catalog_version,
            policy_version=policy_version,
        ),
        definition_id=definition.signal_id,
        category=definition.category,
        start_route_index=candidate.start_route_index,
        end_route_index=candidate.end_route_index,
        start_route_distance_m=candidate.start_route_distance_m,
        end_route_distance_m=candidate.end_route_distance_m,
        active_spans=candidate.active_spans,
        bridged_false_gaps=candidate.bridged_false_gaps,
        applied_rule=definition.rule_specification,
        threshold_provenance=definition.threshold_provenance,
        coverage=candidate.coverage,
        quality=candidate.quality,
        boundaries=candidate.boundaries,
        explanation_key=definition.explanation_key,
        start_planned_time=candidate.start_planned_time,
        end_planned_time=candidate.end_planned_time,
        active_duration_s=candidate.active_duration_s,
        active_distance_m=candidate.active_distance_m,
        total_duration_s=candidate.total_duration_s,
        relevant_values=candidate.relevant_values,
        evidence_payload={"bridged_false_gap_count": len(candidate.bridged_false_gaps)},
    )
