"""Sealed canonical file bindings and descriptor-only snapshot parsing."""

from __future__ import annotations

import hashlib
import json
import os
import stat
from dataclasses import dataclass
from typing import ClassVar, Literal

from pydantic import BaseModel, ConfigDict, Field, ValidationError

from .models import Event
from .store import CanonicalEventSnapshot
from .weekly_operations import CustomerIdentityDigest, WeeklyOperationsConflict


@dataclass(frozen=True, slots=True)
class CanonicalFilePin:
    device: int
    inode: int
    mode: int
    owner: int
    links: int
    kind: Literal["directory", "file"]
    content_size: int | None = None
    content_digest: str | None = None

    @classmethod
    def from_descriptor(
        cls, descriptor: int, kind: Literal["directory", "file"], *, invariant: bool = False
    ) -> CanonicalFilePin:
        info = os.fstat(descriptor)
        safe_kind = stat.S_ISDIR(info.st_mode) if kind == "directory" else stat.S_ISREG(info.st_mode)
        if not safe_kind or info.st_uid != os.geteuid():
            raise WeeklyOperationsConflict("canonical authority inode is unsafe")
        if kind == "file" and (stat.S_IMODE(info.st_mode) != 0o600 or info.st_nlink != 1):
            raise WeeklyOperationsConflict("canonical authority file is unsafe")
        content = read_descriptor(descriptor) if invariant else None
        return cls(
            info.st_dev,
            info.st_ino,
            stat.S_IMODE(info.st_mode),
            info.st_uid,
            info.st_nlink,
            kind,
            None if content is None else len(content),
            None if content is None else hashlib.sha256(content).hexdigest(),
        )

    def verify(self, info: os.stat_result, content: bytes | None = None) -> None:
        observed = (
            info.st_dev,
            info.st_ino,
            stat.S_IMODE(info.st_mode),
            info.st_uid,
            info.st_nlink,
        )
        expected = (self.device, self.inode, self.mode, self.owner, self.links)
        safe_kind = stat.S_ISDIR(info.st_mode) if self.kind == "directory" else stat.S_ISREG(info.st_mode)
        if observed != expected or not safe_kind:
            raise WeeklyOperationsConflict("canonical authority identity drift")
        if self.content_size is not None:
            digest = None if content is None else hashlib.sha256(content).hexdigest()
            if content is None or len(content) != self.content_size or digest != self.content_digest:
                raise WeeklyOperationsConflict("canonical lock content drift")


@dataclass(frozen=True, slots=True)
class CanonicalPinSet:
    root: CanonicalFilePin
    wizard: CanonicalFilePin
    plans: CanonicalFilePin
    events: CanonicalFilePin
    sequence: CanonicalFilePin
    lock: CanonicalFilePin

    def values(self) -> tuple[CanonicalFilePin, ...]:
        return (self.root, self.wizard, self.plans, self.events, self.sequence, self.lock)


@dataclass(frozen=True, slots=True)
class CanonicalDescriptorSet:
    root: int
    wizard: int
    plans: int
    events: int
    sequence: int
    lock: int

    def values(self) -> tuple[int, ...]:
        return (self.root, self.wizard, self.plans, self.events, self.sequence, self.lock)


@dataclass(frozen=True, slots=True)
class CanonicalBindingMaterial:
    customer_identity_digest: CustomerIdentityDigest
    registered_binding_digest: str
    schema_digest: str
    pins: CanonicalPinSet


@dataclass(frozen=True, slots=True)
class CanonicalCheckinCustomerBinding:
    customer_identity_digest: CustomerIdentityDigest
    registered_binding_digest: str
    schema_digest: str
    pins: CanonicalPinSet
    binding_digest: str

    def __post_init__(self) -> None:
        digests = (
            self.customer_identity_digest,
            self.registered_binding_digest,
            self.schema_digest,
            self.binding_digest,
        )
        if any(len(value) != 64 or any(character not in "0123456789abcdef" for character in value) for value in digests):
            raise WeeklyOperationsConflict("canonical customer binding digest is invalid")
        if self.binding_digest != binding_digest(self):
            raise WeeklyOperationsConflict("canonical customer binding digest drift")


def _material_digest(material: CanonicalBindingMaterial) -> str:
    payload = {
        "customer_identity_digest": material.customer_identity_digest,
        "registered_binding_digest": material.registered_binding_digest,
        "schema_digest": material.schema_digest,
        "pins": [
            [
                pin.device,
                pin.inode,
                pin.mode,
                pin.owner,
                pin.links,
                pin.kind,
                pin.content_size,
                pin.content_digest,
            ]
            for pin in material.pins.values()
        ],
    }
    encoded = json.dumps(payload, sort_keys=True, separators=(",", ":")).encode()
    return hashlib.sha256(encoded).hexdigest()


def binding_digest(binding: CanonicalCheckinCustomerBinding) -> str:
    return _material_digest(
        CanonicalBindingMaterial(
            binding.customer_identity_digest,
            binding.registered_binding_digest,
            binding.schema_digest,
            binding.pins,
        )
    )


def create_binding(material: CanonicalBindingMaterial) -> CanonicalCheckinCustomerBinding:
    return CanonicalCheckinCustomerBinding(
        material.customer_identity_digest,
        material.registered_binding_digest,
        material.schema_digest,
        material.pins,
        _material_digest(material),
    )


def read_descriptor(descriptor: int) -> bytes:
    info = os.fstat(descriptor)
    first = os.pread(descriptor, info.st_size, 0)
    second = os.pread(descriptor, info.st_size, 0)
    if first != second or len(first) != info.st_size:
        raise WeeklyOperationsConflict("canonical descriptor changed during read")
    return first


class _SequenceRow(BaseModel):
    model_config: ClassVar[ConfigDict] = ConfigDict(frozen=True, extra="forbid")
    schema_version: Literal["canonical_sequence_v1"]
    sequence: int = Field(ge=1)
    event_id: str
    event_digest: str = Field(pattern=r"^[0-9a-f]{64}$")
    intent_id: str | None = None
    row_digest: str = Field(pattern=r"^[0-9a-f]{64}$")


def _complete_lines(payload: bytes, label: str) -> tuple[bytes, ...]:
    if payload and not payload.endswith(b"\n"):
        raise WeeklyOperationsConflict(f"canonical {label} has a torn tail")
    return tuple(line for line in payload.splitlines() if line)


def read_canonical_snapshot(descriptors: CanonicalDescriptorSet) -> CanonicalEventSnapshot:
    """Parse and cross-check canonical bytes only from pinned descriptors."""
    event_lines = _complete_lines(read_descriptor(descriptors.events), "events")
    sequence_lines = _complete_lines(read_descriptor(descriptors.sequence), "sequence")
    if len(event_lines) != len(sequence_lines):
        raise WeeklyOperationsConflict("canonical event and sequence counts differ")
    try:
        events = tuple(Event.model_validate_json(line) for line in event_lines)
        sequence_models = tuple(
            _SequenceRow.model_validate_json(line) for line in sequence_lines
        )
    except ValidationError as error:
        raise WeeklyOperationsConflict("canonical descriptor schema is invalid") from error
    for expected, (event, row, event_line) in enumerate(
        zip(events, sequence_models, event_lines, strict=True), start=1
    ):
        body = row.model_dump(mode="json", exclude={"row_digest"}, exclude_none=True)
        encoded = json.dumps(
            body, ensure_ascii=False, sort_keys=True, separators=(",", ":")
        ).encode()
        valid = (
            row.sequence == expected
            and row.event_id == event.event_id
            and row.event_digest == hashlib.sha256(event_line + b"\n").hexdigest()
            and row.row_digest == hashlib.sha256(encoded).hexdigest()
        )
        if not valid:
            raise WeeklyOperationsConflict("canonical event and sequence relation drift")
    return CanonicalEventSnapshot(
        events,
        tuple(row.model_dump(mode="json", exclude_none=True) for row in sequence_models),
    )
