"""Offline overenie WORM+TSA atestacie epoch rootu (vendored do verifikatora).

Samostatna kopia overovacej casti api/app/worm.py — verifikator je NEZAVISLY od
beziacej appky. Cross-checkne ZIVU epoch_anchor (export/RDS) proti NAJNOVSEJ WORM
atestacii (KMS+TSA), aby chytil zmazanie VLASTNEHO/CHVOSTOVEHO anchor riadku DB
superuserom (in-DB kotva sama o sebe tail-delete nedetekuje).

KRITICKE: attestation_core tvar + kanonikalizacia + ISO + attestation_digest MUSIA
byt bajt-za-bajt zhodne s api/app/worm.py (parita zabezpecuju testy v
api/tests/test_audit_chain_integration.py).
"""

from __future__ import annotations

import base64
import hashlib
import json
from collections.abc import Callable, Sequence
from dataclasses import dataclass
from datetime import UTC, datetime
from pathlib import Path
from typing import Any

ANCHOR_GENESIS = "0" * 64
ATTESTATION_KEY_PREFIX = "anchor-attestations/"


def _canonical_bytes(obj: Any) -> bytes:  # noqa: ANN401
    return json.dumps(obj, sort_keys=True, separators=(",", ":"), ensure_ascii=False).encode(
        "utf-8"
    )


def _sha256_hex(data: bytes) -> str:
    return hashlib.sha256(data).hexdigest()


def _iso(ts: Any) -> str:  # noqa: ANN401 — datetime alebo ISO string z exportu
    dt: datetime
    if isinstance(ts, str):
        dt = datetime.fromisoformat(ts)
    elif isinstance(ts, datetime):
        dt = ts
    else:
        msg = f"server_ts must be a datetime or ISO string, got {type(ts)!r}"
        raise TypeError(msg)
    if dt.tzinfo is None:
        dt = dt.replace(tzinfo=UTC)
    return dt.astimezone(UTC).isoformat().replace("+00:00", "Z")


def attestation_digest(core: dict[str, Any]) -> str:
    """SHA256 hex kanonickeho atestacneho core (parita s api/app/worm.attestation_digest).

    `core` MUSI mat presny tvar (kluce: epoch_root_hash, anchor_seq, anchor_count,
    server_ts, key_id). int polia normalizujeme (export moze prist ako str).
    """
    normalized = {
        "epoch_root_hash": str(core["epoch_root_hash"]),
        "anchor_seq": int(core["anchor_seq"]),
        "anchor_count": int(core["anchor_count"]),
        "server_ts": core["server_ts"],
        "key_id": core.get("key_id"),
    }
    return _sha256_hex(_canonical_bytes(normalized))


@dataclass(frozen=True)
class Attestation:
    """Jedna atestacia epoch rootu (obsah WORM objektu; parita s api/app/worm.Attestation)."""

    core: dict[str, Any]
    digest: str
    kms_key_id: str | None
    signature: bytes | None
    tsa_token: bytes | None

    @property
    def anchor_seq(self) -> int:
        return int(self.core["anchor_seq"])

    @property
    def epoch_root_hash(self) -> str:
        return str(self.core["epoch_root_hash"])

    @classmethod
    def from_json_bytes(cls, raw: bytes) -> Attestation:
        obj = json.loads(raw)
        sig_b64 = obj.get("signature_b64")
        tok_b64 = obj.get("tsa_token_b64")
        return cls(
            core=obj["core"],
            digest=str(obj["digest"]),
            kms_key_id=obj.get("kms_key_id"),
            signature=base64.b64decode(sig_b64) if sig_b64 else None,
            tsa_token=base64.b64decode(tok_b64) if tok_b64 else None,
        )


def _anchor_core(row: dict[str, Any]) -> dict[str, Any]:
    """Re-zostav anchor_core PRESNE ako api/app/anchor.py:_anchor_core (parita)."""
    return {
        "seq": int(row["seq"]),
        "sealed_inspection_id": str(row["sealed_inspection_id"]),
        "sealed_head_hash": row["sealed_head_hash"],
        "prev_anchor_hash": row["prev_anchor_hash"],
        "ts": _iso(row["server_ts"]),
    }


@dataclass(frozen=True)
class AttestationCheck:
    """Vysledok cross-checku zivej kotvy proti WORM atestacii (parita s api/app/worm)."""

    ok: bool
    attested: bool
    covers_seq: int = 0
    reason: str | None = None


def _live_root_at_seq(anchor_rows: Sequence[dict[str, Any]], attested_seq: int) -> str | None:
    """Prepocitaj running-hash root ZIVEJ epoch_anchor po `attested_seq` (parita s app.worm).

    Re-derivuje retazec od GENESIS cez VSETKY riadky so seq <= attested_seq. Vrati
    anchor_hash riadku so seq==attested_seq (root). None ak ziva kotva nema riadok so
    seq==attested_seq (skratena) alebo sa retazec rozbije.
    """
    rows = sorted(
        (r for r in anchor_rows if int(r["seq"]) <= attested_seq),
        key=lambda r: int(r["seq"]),
    )
    if not rows:
        return None
    prev = ANCHOR_GENESIS
    head_hash: str | None = None
    max_seq = 0
    for row in rows:
        if row["prev_anchor_hash"] != prev:
            return None
        stored_anchor = row["anchor_hash"]
        if _sha256_hex(_canonical_bytes(_anchor_core(row))) != stored_anchor:
            return None
        prev = stored_anchor
        head_hash = stored_anchor
        max_seq = int(row["seq"])
    if max_seq != attested_seq:
        return None
    return head_hash


def verify_attestation_against_live(  # noqa: PLR0911, PLR0913 — guard clauses + offline parita s app.worm
    anchor_rows: Sequence[dict[str, Any]],
    attestation: Attestation | None,
    *,
    require_signatures: bool,
    offline_kms_verify: Callable[[str, bytes | None], bool],
    verify_tsa: bool,
    tsa_verify: Callable[[str, bytes | None], Any],
) -> AttestationCheck:
    """JADROVY INVARIANT (offline parita s app.worm.verify_attestation_against_live).

    `attestation` = NAJNOVSIA atestacia z WORM store (autorita). `anchor_rows` = ZIVY
    epoch_anchor export/RDS. Overi KMS podpis + TSA atestacie a prepocita zivu kotvu
    po atestovany seq; menej riadkov / iny root -> TAMPERED. None atestacia ->
    attested=False (nic nefinalizovane).
    """
    if attestation is None:
        return AttestationCheck(ok=True, attested=False, reason="no_attestation")

    if attestation_digest(attestation.core) != attestation.digest:
        return AttestationCheck(ok=False, attested=True, reason="attestation_digest_mismatch")
    if require_signatures and not offline_kms_verify(attestation.digest, attestation.signature):
        return AttestationCheck(ok=False, attested=True, reason="attestation_signature_invalid")
    if verify_tsa:
        tsa_res = tsa_verify(attestation.digest, attestation.tsa_token)
        if not tsa_res.ok:
            return AttestationCheck(
                ok=False, attested=True, reason=f"attestation_tsa:{tsa_res.reason}"
            )

    attested_seq = attestation.anchor_seq
    live_root = _live_root_at_seq(anchor_rows, attested_seq)
    if live_root is None:
        return AttestationCheck(ok=False, attested=True, reason="anchor_truncated_vs_attestation")
    if live_root != attestation.epoch_root_hash:
        return AttestationCheck(ok=False, attested=True, reason="anchor_attestation_mismatch")

    return AttestationCheck(ok=True, attested=True, covers_seq=attested_seq)


def anchor_entry_seq(
    anchor_rows: Sequence[dict[str, Any]], inspection_id: str, sealed_head_hash: str
) -> int | None:
    """Vrat seq epoch_anchor riadku, ktory viaze (inspection_id, sealed_head) (parita)."""
    best: int | None = None
    for row in anchor_rows:
        if (
            str(row["sealed_inspection_id"]) == inspection_id
            and row["sealed_head_hash"] == sealed_head_hash
        ):
            seq = int(row["seq"])
            if best is None or seq > best:
                best = seq
    return best


# --------------------------------------------------------------------------- #
# Loadery atestacie (export adresar/subor alebo S3)
# --------------------------------------------------------------------------- #
def load_latest_attestation_from_dir(path: Path) -> Attestation | None:
    """Nacitaj NAJNOVSIU atestaciu z lokalneho exportu WORM objektov.

    `path` moze byt:
      - adresar s objektmi <anchor_seq>.json (vyberie najvyssi seq),
      - jeden JSON subor (priamo jedna atestacia).
    """
    if path.is_dir():
        best_seq = -1
        best: Attestation | None = None
        for f in path.glob("*.json"):
            try:
                att = Attestation.from_json_bytes(f.read_bytes())
            except (ValueError, KeyError):
                continue
            if att.anchor_seq > best_seq:
                best_seq = att.anchor_seq
                best = att
        return best
    return Attestation.from_json_bytes(path.read_bytes())


def load_latest_attestation_from_s3(bucket: str, region: str) -> Attestation | None:
    """Nacitaj NAJNOVSIU atestaciu z S3 WORM bucketu (anchor-attestations/<seq>.json)."""
    import boto3

    s3: Any = boto3.client("s3", region_name=region)
    best_seq = -1
    best_key: str | None = None
    paginator = s3.get_paginator("list_objects_v2")
    for page in paginator.paginate(Bucket=bucket, Prefix=ATTESTATION_KEY_PREFIX):
        for obj in page.get("Contents", []):
            key = str(obj["Key"])
            name = key[len(ATTESTATION_KEY_PREFIX) :]
            if not name.endswith(".json"):
                continue
            try:
                seq = int(name[: -len(".json")])
            except ValueError:
                continue
            if seq > best_seq:
                best_seq = seq
                best_key = key
    if best_key is None:
        return None
    body: bytes = s3.get_object(Bucket=bucket, Key=best_key)["Body"].read()
    return Attestation.from_json_bytes(body)
