"""Shared complete validation for weekly-operations durable history."""

from __future__ import annotations

from collections.abc import Sequence
from dataclasses import dataclass
from datetime import date, timedelta

from pydantic import ValidationError

from checkin_cli.weekly_operations import (
    CustomerIdentityDigest,
    DayState,
    WeeklyOperationRow,
    WeeklyOperationsConflict,
    WeeklyOperationsCorruption,
    ZERO_DIGEST,
    customer_storage_digest,
    operation_logical_key,
    weekly_row_digest,
)


@dataclass(frozen=True, slots=True)
class HistoryValidation:
    rows: tuple[WeeklyOperationRow, ...]
    complete_end: int
    size: int


def _transition_failure(reason: str, *, candidate: bool) -> None:
    if candidate:
        raise WeeklyOperationsConflict(reason)
    raise WeeklyOperationsCorruption(reason)


def _validate_transition(previous: WeeklyOperationRow | None, row: WeeklyOperationRow, *, candidate: bool) -> None:
    if previous is None:
        if row.state not in {DayState.SUBMITTED, DayState.MISSED}:
            _transition_failure("late_submitted requires a prior missed row", candidate=candidate)
        return
    transition = previous.state, row.state
    correction = transition in {(DayState.SUBMITTED, DayState.SUBMITTED), (DayState.MISSED, DayState.MISSED), (DayState.LATE_SUBMITTED, DayState.LATE_SUBMITTED)}
    if transition != (DayState.MISSED, DayState.LATE_SUBMITTED) and not correction:
        _transition_failure("illegal or reversed timeliness transition", candidate=candidate)
    if correction and row.source_event_id is None:
        _transition_failure("same-state append requires source lineage", candidate=candidate)
    if row.canonical_sequence <= previous.canonical_sequence:
        _transition_failure("source lineage must advance the canonical pin", candidate=candidate)
    if correction and row.source_event_id == previous.source_event_id:
        _transition_failure("correction source lineage did not advance", candidate=candidate)


def _logical_key(customer_identity: CustomerIdentityDigest, previous: WeeklyOperationRow | None, row: WeeklyOperationRow) -> str:
    if previous is None:
        identity = "terminal"
    elif previous.state is DayState.MISSED and row.state is DayState.LATE_SUBMITTED:
        identity = "late"
    elif row.source_event_id is not None:
        identity = f"source:{row.source_event_id}"
    else:
        raise WeeklyOperationsCorruption("durable logical identity is incomplete")
    return operation_logical_key(customer_identity, row.kst_day, identity)


def validate_weekly_rows(rows: Sequence[WeeklyOperationRow], customer_identity: CustomerIdentityDigest, *, candidate: bool = False) -> None:
    """Validate every durable hash, identity, pin, logical key, and state invariant."""
    predecessor = ZERO_DIGEST
    canonical: dict[int, str] = {}
    timelines: dict[date, WeeklyOperationRow] = {}
    logical_keys: set[str] = set()
    last_sequence: int | None = None
    for row in rows:
        if row.customer_identity_digest != customer_identity or row.occurred_at_kst.utcoffset() != timedelta(hours=9) or row.occurred_at_kst.date() != row.kst_day:
            raise WeeklyOperationsCorruption("row identity or KST day mismatch")
        if (row.source_event_id is None) != (row.source_event_digest is None):
            raise WeeklyOperationsCorruption("partial source lineage")
        if row.predecessor_row_digest != predecessor or row.row_digest != weekly_row_digest(row):
            raise WeeklyOperationsCorruption("sidecar hash chain disagreement")
        if last_sequence is not None and row.canonical_sequence < last_sequence:
            _transition_failure("canonical sequence reversed", candidate=candidate)
        pinned = canonical.get(row.canonical_sequence)
        if pinned is not None and pinned != row.canonical_digest:
            _transition_failure("canonical pin drift", candidate=candidate)
        previous = timelines.get(row.kst_day)
        _validate_transition(previous, row, candidate=candidate)
        expected_key = _logical_key(customer_identity, previous, row)
        if row.logical_key != expected_key or row.logical_key in logical_keys:
            raise WeeklyOperationsCorruption("durable logical key disagreement")
        canonical[row.canonical_sequence] = row.canonical_digest
        timelines[row.kst_day] = row
        logical_keys.add(row.logical_key)
        predecessor, last_sequence = row.row_digest, row.canonical_sequence


def validate_history_bytes(payload: bytes, filename_digest: str, *, expected_customer: CustomerIdentityDigest | None = None, allow_torn_tail: bool = False) -> HistoryValidation:
    """Parse and completely validate one JSONL history and its filename binding."""
    rows: list[WeeklyOperationRow] = []
    complete_end = 0
    offset = 0
    for index, line in enumerate(payload.splitlines(keepends=True)):
        offset += len(line)
        if not line.endswith(b"\n"):
            if allow_torn_tail and offset == len(payload):
                break
            raise WeeklyOperationsCorruption("sidecar has a torn tail")
        if not line[:-1]:
            raise WeeklyOperationsCorruption("sidecar contains a blank row")
        try:
            row = WeeklyOperationRow.model_validate_json(line[:-1])
        except ValidationError as error:
            raise WeeklyOperationsCorruption(f"invalid sidecar row {index}") from error
        if customer_storage_digest(row.customer_identity_digest) != filename_digest:
            raise WeeklyOperationsCorruption("sidecar filename and history identity disagree")
        if expected_customer is not None and row.customer_identity_digest != expected_customer:
            raise WeeklyOperationsCorruption("row identity or filename mismatch")
        rows.append(row)
        complete_end = offset
    identity = expected_customer if expected_customer is not None else (rows[0].customer_identity_digest if rows else None)
    if identity is not None:
        validate_weekly_rows(rows, identity)
    return HistoryValidation(tuple(rows), complete_end, len(payload))
