"""Deterministic RCI composition over existing Unified Route Analysis output."""

from __future__ import annotations

from collections import Counter
from dataclasses import replace
from typing import Any, Mapping, Sequence

from mountain_twin.rci.catalog import POLICY_VERSION, initial_signal_registry
from mountain_twin.rci.clustering import build_condition_clusters
from mountain_twin.rci.contracts import (
    ConditionState,
    ConditionStateRecord,
    CoverageState,
    EvidenceQuality,
    RouteConditionsResult,
    SignalDefinitionRegistry,
    serialize_route_conditions_result,
)
from mountain_twin.rci.detection import DetectionResult, detect_condition_signals
from mountain_twin.rci.events import extract_route_condition_events
from mountain_twin.rci.why import signal_why


def compose_rci_route_conditions(
    unified_analysis: Any,
    *,
    registry: SignalDefinitionRegistry | None = None,
    policy_version: str = POLICY_VERSION,
) -> RouteConditionsResult:
    """Compose compact RCI facts without recalculating upstream domains.

    The public result retains state summaries rather than Unified's complete
    point matrix. Callers that need route points retain the upstream Unified
    result identified by ``analysis_metadata``.
    """
    registry = registry or initial_signal_registry()
    points = tuple(_value(unified_analysis, "points", ()))
    analysis_reference = _analysis_reference(unified_analysis)
    locations = _route_locations(points)

    detection_results = tuple(
        _detect_definition(
            unified_analysis,
            definition_id=definition.signal_id,
            analysis_reference=analysis_reference,
            catalog_version=registry.catalog_version,
            policy_version=policy_version,
            registry=registry,
        )
        for definition in registry.definitions
    )
    signals = tuple(
        sorted(
            (
                _enrich_signal(signal, analysis_reference, points)
                for detection in detection_results
                for signal in detection.signals
            ),
            key=lambda item: (item.start_route_index, item.definition_id, item.signal_instance_id),
        )
    )
    result = RouteConditionsResult(
        analysis_metadata={
            "analysis_reference": analysis_reference,
            "upstream_analysis_type": _value(_value(unified_analysis, "identity", {}), "analysis_type"),
            "upstream_analysis_version": _value(_value(unified_analysis, "identity", {}), "version"),
            "state_series_retained_publicly": False,
        },
        route_reference=dict(_value(unified_analysis, "route", {})),
        planning_reference=dict(_value(unified_analysis, "scenario", {})),
        catalog_version=registry.catalog_version,
        policy_version=policy_version,
        signals=signals,
        clusters=build_condition_clusters(signals, locations),
        events=extract_route_condition_events(unified_analysis),
        quality=EvidenceQuality(coverage=_domain_coverage(unified_analysis, "solar")),
        provenance={
            "signal_catalog": registry.to_dict(),
            "upstream_analysis_reference": analysis_reference,
            "upstream_provenance_reference": _value(unified_analysis, "provenance", {}),
        },
        diagnostics=("DATA_QUALITY_REMAINS_QUALITY_METADATA_NOT_ENVIRONMENTAL_SIGNAL",),
        state_summary=_state_summary(registry, detection_results),
        candidate_diagnostics=_failed_candidate_diagnostics(detection_results),
        contract_version="0.2",
    )
    _validate_result(result, registry)
    return result


def build_condition_state_series(
    unified_analysis: Any, definition_id: str
) -> tuple[ConditionStateRecord, ...]:
    """Build one lightweight, deterministic state series for an active definition.

    The initial catalog only contains terrain-solar state equality definitions.
    Missing or unresolved solar state becomes ``UNKNOWN``; it is never treated
    as an inactive/``FALSE`` sample.
    """
    if definition_id not in {"solar.terrain_direct", "solar.terrain_shadow"}:
        raise KeyError(f"no state composer for signal definition: {definition_id}")
    target = "DIRECT" if definition_id == "solar.terrain_direct" else "SHADOW"
    coverage = _domain_coverage(unified_analysis, "solar")
    terrain_context = _value(
        _value(_value(unified_analysis, "provenance", {}), "solar", {}), "terrain_context_id"
    )
    records = []
    for point in _value(unified_analysis, "points", ()):
        solar = _value(point, "solar", {})
        status = _enum_or_value(_value(solar, "status"))
        state = _terrain_visibility_state(status, target)
        solar_quality = _value(point, "solar_quality", {})
        continuity_key = _value(solar, "terrain_context_id", terrain_context)
        records.append(
            ConditionStateRecord(
                definition_id=definition_id,
                route_point_index=_value(point, "point_index"),
                route_distance_m=_value(point, "route_distance_m"),
                state=state,
                elapsed_planned_seconds=_value(point, "elapsed_route_seconds"),
                planned_time=_value(point, "planned_arrival_time"),
                value=None,
                source_continuity_key=continuity_key,
                coverage=coverage,
                quality=EvidenceQuality(
                    coverage=coverage,
                    source=_value(point, "solar_provenance_ref"),
                    derived_status="DETERMINISTIC_TERRAIN_SOLAR_STATE",
                    reason_codes=tuple(_value(solar_quality, "reason_codes", ())),
                    limitations=(
                        "TERRAIN_VISIBILITY_IS_DISTINCT_FROM_ASTRONOMICAL_SOLAR_STATE",
                    ),
                ),
            )
        )
    return tuple(records)


def serialize_rci_route_conditions(result: RouteConditionsResult) -> str:
    """Canonical serialization for RCI output without runtime clock fields."""
    return serialize_route_conditions_result(result)


def _detect_definition(
    analysis: Any,
    *,
    definition_id: str,
    analysis_reference: str,
    catalog_version: str,
    policy_version: str,
    registry: SignalDefinitionRegistry,
) -> DetectionResult:
    return detect_condition_signals(
        definition=registry.get(definition_id),
        samples=build_condition_state_series(analysis, definition_id),
        analysis_reference=analysis_reference,
        catalog_version=catalog_version,
        policy_version=policy_version,
    )


def _enrich_signal(signal: Any, analysis_reference: str, points: Sequence[Any]):
    visible_state = signal.applied_rule["terrain_solar_visibility_equals"]
    source_references = tuple(
        sorted(
            {
                reference
                for point in points
                if signal.start_route_index <= _value(point, "point_index") <= signal.end_route_index
                for reference in (_value(point, "solar_provenance_ref"),)
                if reference is not None
            }
        )
    )
    terrain_context_ids = tuple(
        sorted(
            {
                context_id
                for point in points
                if signal.start_route_index <= _value(point, "point_index") <= signal.end_route_index
                for context_id in (_value(_value(point, "solar", {}), "terrain_context_id"),)
                if context_id is not None
            }
        )
    )
    enriched = replace(
        signal,
        relevant_values={
            **dict(signal.relevant_values),
            "terrain_solar_visibility": visible_state,
        },
        evidence_payload={
            **dict(signal.evidence_payload),
            "analysis_reference": analysis_reference,
            "solar_provenance_references": source_references,
            "terrain_context_ids": terrain_context_ids,
            "upstream_domain": "solar",
            "upstream_state": visible_state,
        },
    )
    return replace(enriched, why=signal_why(enriched))


def _terrain_visibility_state(status: str | None, target: str) -> ConditionState:
    if status == target:
        return ConditionState.TRUE
    if status in {"DIRECT", "SHADOW"}:
        return ConditionState.FALSE
    if status == "NOT_APPLICABLE":
        return ConditionState.NOT_APPLICABLE
    return ConditionState.UNKNOWN


def _state_summary(
    registry: SignalDefinitionRegistry, results: Sequence[DetectionResult]
) -> Mapping[str, Any]:
    return {
        definition.signal_id: {
            "counts": dict(
                sorted(Counter(record.state.value for record in result.condition_states).items())
            ),
            "sample_count": len(result.condition_states),
        }
        for definition, result in zip(registry.definitions, results)
    }


def _failed_candidate_diagnostics(results: Sequence[DetectionResult]) -> tuple[Mapping[str, Any], ...]:
    return tuple(
        {
            "active_distance_m": candidate.active_distance_m,
            "active_duration_s": candidate.active_duration_s,
            "definition_id": candidate.definition_id,
            "end_route_index": candidate.end_route_index,
            "qualified": candidate.qualified,
            "start_route_index": candidate.start_route_index,
        }
        for result in results
        for candidate in result.candidates
        if not candidate.qualified
    )


def _route_locations(points: Sequence[Any]) -> Mapping[int, Mapping[str, Any]]:
    return {
        _value(point, "point_index"): {
            "elapsed_planned_seconds": _value(point, "elapsed_route_seconds"),
            "planned_time": _value(point, "planned_arrival_time"),
            "route_distance_m": _value(point, "route_distance_m"),
        }
        for point in points
    }


def _domain_coverage(analysis: Any, domain: str) -> CoverageState:
    mode = _enum_or_value(_value(_value(_value(analysis, "coverage", {}), domain, {}), "mode"))
    return {
        "FULL": CoverageState.FULL,
        "PARTIAL": CoverageState.PARTIAL,
        "SAMPLED": CoverageState.PARTIAL,
        "UNAVAILABLE": CoverageState.UNAVAILABLE,
    }.get(mode, CoverageState.UNKNOWN)


def _analysis_reference(analysis: Any) -> str:
    identity = _value(analysis, "identity", {})
    return ":".join(
        str(value)
        for value in (
            _value(identity, "analysis_type", "unified_route_analysis"),
            _value(identity, "version", "unknown"),
            _value(identity, "subject_id", _value(_value(analysis, "route", {}), "route_id", "unknown")),
            _value(_value(analysis, "scenario", {}), "name", "unknown"),
        )
    )


def _validate_result(result: RouteConditionsResult, registry: SignalDefinitionRegistry) -> None:
    definition_ids = {definition.signal_id for definition in registry.definitions}
    signal_ids = {signal.signal_instance_id for signal in result.signals}
    if any(signal.definition_id not in definition_ids for signal in result.signals):
        raise ValueError("result signal definition is absent from catalog")
    if any(
        any(member not in signal_ids for member in cluster.signal_instance_ids)
        for cluster in result.clusters
    ):
        raise ValueError("cluster references a missing signal")
    if any(
        any(definition_id not in definition_ids for definition_id in cluster.definition_ids)
        for cluster in result.clusters
    ):
        raise ValueError("cluster definition is absent from catalog")


def _enum_or_value(value: Any) -> Any:
    return getattr(value, "value", value)


def _value(item: Any, key: str, default: Any = None) -> Any:
    return item.get(key, default) if isinstance(item, Mapping) else getattr(item, key, default)
