"""Sweep-line clustering over qualified signal ACTIVE spans only."""

from __future__ import annotations

import hashlib
import json
from dataclasses import replace
from typing import Any, Iterable, Mapping, Sequence

from mountain_twin.rci.contracts import (
    ConditionSignal,
    CoverageState,
    EvidenceQuality,
    RouteConditionCluster,
)
from mountain_twin.rci.why import cluster_why


def build_condition_clusters(
    signals: Sequence[ConditionSignal],
    locations: Mapping[int, Mapping[str, Any]],
) -> tuple[RouteConditionCluster, ...]:
    """Return maximal stable-member overlap intervals for different definitions.

    Planned elapsed time is used only when every active span has it. Otherwise
    route distance is the deterministic fallback. Bridged FALSE gaps are absent
    because only ``active_spans`` are considered.
    """
    intervals = []
    for signal in signals:
        for span in signal.active_spans:
            intervals.append((signal, span))
    if not intervals:
        return ()
    use_time = all(
        span.start_elapsed_planned_seconds is not None
        and span.end_elapsed_planned_seconds is not None
        for _, span in intervals
    )
    axis = "PLANNED_TIME" if use_time else "ROUTE_DISTANCE"
    events = []
    for signal, span in intervals:
        start = span.start_elapsed_planned_seconds if use_time else span.start_route_distance_m
        end = span.end_elapsed_planned_seconds if use_time else span.end_route_distance_m
        if end <= start:
            continue
        events.append((start, 1, signal, span))
        events.append((end, 0, signal, span))
    if not events:
        return ()
    positions = sorted({event[0] for event in events})
    by_position: dict[float, list[tuple[float, int, ConditionSignal, Any]]] = {}
    for event in events:
        by_position.setdefault(event[0], []).append(event)
    active: dict[str, tuple[ConditionSignal, Any]] = {}
    segments = []
    for position, following in zip(positions, positions[1:]):
        # End events precede starts at the same position: touching spans do not overlap.
        for _, kind, signal, span in sorted(by_position[position], key=lambda item: item[1]):
            if kind == 0:
                active.pop(signal.signal_instance_id, None)
            else:
                active[signal.signal_instance_id] = (signal, span)
        members = tuple(sorted(active))
        definitions = tuple(sorted({active[item][0].definition_id for item in members}))
        if following > position and len(definitions) >= 2:
            segments.append((position, following, members, definitions, tuple(active.values())))
    merged = []
    for segment in segments:
        if merged and merged[-1][1] == segment[0] and merged[-1][2:4] == segment[2:4]:
            merged[-1] = (merged[-1][0], segment[1], *merged[-1][2:])
        else:
            merged.append(segment)
    clusters = []
    for start, end, members, definitions, supporting in merged:
        start_index, end_index, start_distance, end_distance, start_time, end_time = _location_bounds(
            start, end, axis, supporting, locations
        )
        identity = {
            "axis": axis,
            "definitions": definitions,
            "end": end,
            "members": members,
            "start": start,
        }
        cluster = RouteConditionCluster(
            cluster_id="rci-cluster-"
            + hashlib.sha256(
                json.dumps(identity, sort_keys=True, separators=(",", ":")).encode("utf-8")
            ).hexdigest()[:24],
            signal_instance_ids=members,
            definition_ids=definitions,
            start_route_index=start_index,
            end_route_index=end_index,
            start_route_distance_m=start_distance,
            end_route_distance_m=end_distance,
            start_planned_time=start_time,
            end_planned_time=end_time,
            overlap_axis=axis,
            quality=_cluster_quality(signal for signal, _ in supporting),
            evidence_payload={
                "supporting_active_span_count": len(supporting),
                "supporting_active_spans": tuple(
                    {
                        "definition_id": signal.definition_id,
                        "signal_instance_id": signal.signal_instance_id,
                        "span": span.to_dict(),
                    }
                    for signal, span in sorted(
                        supporting,
                        key=lambda item: (item[0].definition_id, item[0].signal_instance_id),
                    )
                ),
            },
        )
        clusters.append(replace(cluster, why=cluster_why(cluster)))
    return tuple(clusters)


def _location_bounds(
    start: float,
    end: float,
    axis: str,
    supporting: Sequence[tuple[ConditionSignal, Any]],
    locations: Mapping[int, Mapping[str, Any]],
) -> tuple[int | None, int | None, float | None, float | None, str | None, str | None]:
    axis_key = "elapsed_planned_seconds" if axis == "PLANNED_TIME" else "route_distance_m"
    known_start_indices = [
        index for index, location in locations.items() if location.get(axis_key) == start
    ]
    known_end_indices = [
        index for index, location in locations.items() if location.get(axis_key) == end
    ]
    start_indices = []
    end_indices = []
    for _, span in supporting:
        if axis == "PLANNED_TIME":
            if span.start_elapsed_planned_seconds == start:
                start_indices.append(span.start_route_index)
            if span.end_elapsed_planned_seconds == start:
                start_indices.append(span.end_route_index)
            if span.start_elapsed_planned_seconds == end:
                end_indices.append(span.start_route_index)
            if span.end_elapsed_planned_seconds == end:
                end_indices.append(span.end_route_index)
        else:
            if span.start_route_distance_m == start:
                start_indices.append(span.start_route_index)
            if span.end_route_distance_m == start:
                start_indices.append(span.end_route_index)
            if span.start_route_distance_m == end:
                end_indices.append(span.start_route_index)
            if span.end_route_distance_m == end:
                end_indices.append(span.end_route_index)
    start_candidates = known_start_indices or [
        index for index in start_indices if index in locations
    ]
    end_candidates = known_end_indices or [index for index in end_indices if index in locations]
    start_index = min(start_candidates) if start_candidates else None
    end_index = max(end_candidates) if end_candidates else None
    start_location = locations.get(start_index, {})
    end_location = locations.get(end_index, {})
    return (
        start_index,
        end_index,
        start_location.get("route_distance_m"),
        end_location.get("route_distance_m"),
        start_location.get("planned_time"),
        end_location.get("planned_time"),
    )


def _cluster_quality(signals: Iterable[ConditionSignal]) -> EvidenceQuality:
    items = tuple(signals)
    coverage_states = {signal.coverage for signal in items}
    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
    return EvidenceQuality(coverage=coverage)
