"""Verify a saved Montora replay result using ONLY a public key. No Montora login, no server code.

Standalone on purpose: it depends on the Python standard library and ``cryptography`` only, so an
auditor can read all of it and run it anywhere.

    # with the published keyring (a file, or the URL of the well-known endpoint)
    python verify_replay_result.py result.json --keyring https://api.example.com/.well-known/montora-replay-signing-key
    python verify_replay_result.py report.json --keyring keyring.json

    # with a single public key you already trust (hex of the 32-byte Ed25519 key)
    python verify_replay_result.py result.json --public-key 3b6a27bc...

The input may be one FidelityResult, a JSON list of them, or a fidelity report
(``{"chunks": {"<chunkId>": <result>, ...}}``).

What is checked, per result:
  1. the reproduction key is recomputed from the result body and must match (tamper check);
  2. the signature must be ``ed25519:<hex>`` and verify over the canonical body under the public
     key whose id is ``provenance.signing.keyId`` (inside the signed body);
  3. if the keyring carries ``engineVersion`` and ``revocation`` (the well-known document does),
     the result is also checked for revocation (revoked engine version / reproduction key, or a
     superseded engine version).

Canonical body: the result JSON without ``signature`` and ``signatureMeta``, serialised with
sorted keys, ``(",", ":")`` separators and ASCII escaping, UTF-8 encoded.

A key the keyring marks ``"status": "compromised"`` fails verification (COMPROMISED_KEY): whoever
holds a leaked key can forge evidence, so nothing it signed can be relied on.

Exit status: 0 every result valid and current; 1 any result invalid, unsigned or signed by a
compromised key; 3 all valid but
at least one revoked; 2 usage / input error.
"""

from __future__ import annotations

import argparse
import hashlib
import json
import sys
import urllib.request

from cryptography.exceptions import InvalidSignature
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PublicKey

UNSIGNED_FIELDS = ("signature", "signatureMeta")


def canonical(result: dict) -> bytes:
    body = {k: v for k, v in result.items() if k not in UNSIGNED_FIELDS}
    return json.dumps(body, sort_keys=True, separators=(",", ":"), default=str).encode("utf-8")


def reproduction_key(result: dict) -> str:
    prov = result.get("provenance") or {}
    material = {
        "engineVersion": result.get("engineVersion"),
        "referenceTier": result.get("referenceTier"),
        "recordingDigest": prov.get("recordingDigest"),
        "migratedHash": prov.get("migratedHash") or prov.get("migrationHash"),
        "score": result.get("score"),
        "gates": result.get("gates"),
        "total": result.get("total"),
        "agreements": result.get("agreements"),
        "errors": result.get("errors"),
        "disagreements": result.get("disagreements"),
    }
    text = json.dumps(material, sort_keys=True, separators=(",", ":"), default=str)
    return "sha256:" + hashlib.sha256(text.encode("utf-8")).hexdigest()


def verify_signature(result: dict, public_key_hex: str) -> bool:
    alg, _, hexsig = str(result.get("signature", "")).partition(":")
    if alg != "ed25519":
        return False
    try:
        key = Ed25519PublicKey.from_public_bytes(bytes.fromhex(public_key_hex))
        key.verify(bytes.fromhex(hexsig), canonical(result))
        return True
    except (InvalidSignature, ValueError):
        return False


def revocation(result: dict, keyring: dict) -> str | None:
    """A reason string if the keyring's revocation data revokes this result, else None."""
    current = keyring.get("engineVersion")
    rev = keyring.get("revocation") or {}
    if not current:
        return None
    rkey = result.get("reproductionKey")
    if rkey and rkey in (rev.get("revokedReproductionKeys") or {}):
        return f"reproduction key revoked: {rev['revokedReproductionKeys'][rkey]}"
    version = result.get("engineVersion")
    if not version:
        return "no engine version recorded"
    if version in (rev.get("revokedEngineVersions") or {}):
        return f"engine {version} revoked: {rev['revokedEngineVersions'][version]}"
    if version != current and version not in (rev.get("compatiblePriorEngineVersions") or []):
        return f"engine {version} superseded by {current}"
    return None


def verify(result: dict, keys: dict, keyring: dict, trust_any_key_id: bool) -> tuple[str, str]:
    """(status, detail). status: VALID | REVOKED | INVALID | UNSIGNED | UNKNOWN_KEY | COMPROMISED_KEY."""
    if reproduction_key(result) != result.get("reproductionKey"):
        return "INVALID", "reproduction key does not match the result body (tampered)"
    sig = str(result.get("signature", ""))
    if not sig.startswith("ed25519:"):
        return "UNSIGNED", f"no Ed25519 signature ({sig.split(':', 1)[0] or 'none'}); nothing to verify"
    key_id = ((result.get("provenance") or {}).get("signing") or {}).get("keyId") \
        or (result.get("signatureMeta") or {}).get("keyId") or ""
    pub = keys.get(key_id)
    if pub is None and trust_any_key_id and len(keys) == 1:
        pub = next(iter(keys.values()))
    if pub is None:
        return "UNKNOWN_KEY", f"no public key for key id {key_id!r}"
    if not verify_signature(result, pub):
        return "INVALID", f"signature does not verify under key {key_id!r}"
    for k in (keyring or {}).get("keys", []):
        if k.get("keyId") == key_id and k.get("status") == "compromised":
            return "COMPROMISED_KEY", (f"signature valid, but key {key_id} was reported compromised on "
                                       f"{k.get('compromisedAt', 'unknown')}; do not rely on this result")
    reason = revocation(result, keyring)
    if reason:
        return "REVOKED", f"signature valid (key {key_id}), but {reason}"
    return "VALID", f"key {key_id}, engine {result.get('engineVersion')}"


def _load_json(source: str):
    if source.startswith(("http://", "https://")):
        with urllib.request.urlopen(source, timeout=20) as resp:  # noqa: S310 — user-supplied URL
            return json.loads(resp.read().decode("utf-8"))
    with open(source, encoding="utf-8") as fh:
        return json.load(fh)


def _results(doc) -> list[tuple[str, dict]]:
    if isinstance(doc, list):
        return [(str(r.get("chunkId", i)), r) for i, r in enumerate(doc)]
    if isinstance(doc, dict) and isinstance(doc.get("chunks"), dict):
        return [(str(k), v) for k, v in doc["chunks"].items()]
    if isinstance(doc, dict):
        return [(str(doc.get("chunkId", "result")), doc)]
    raise ValueError("input is not a result, a list of results, or a fidelity report")


def main(argv: list[str] | None = None) -> int:
    parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    parser.add_argument("result", help="path to the saved result / report JSON")
    src = parser.add_mutually_exclusive_group(required=True)
    src.add_argument("--keyring", help="keyring JSON file or URL (GET /.well-known/montora-replay-signing-key)")
    src.add_argument("--public-key", help="a trusted Ed25519 public key, hex")
    parser.add_argument("--key-id", default="", help="key id for --public-key (default: accept any)")
    args = parser.parse_args(argv)

    try:
        doc = _load_json(args.result)
        items = _results(doc)
        if args.keyring:
            keyring = _load_json(args.keyring)
            keys = {k["keyId"]: k["publicKeyHex"] for k in keyring.get("keys", [])}
            trust_any = False
        else:
            keyring = {}
            keys = {args.key_id or "*": args.public_key.strip()}
            trust_any = not args.key_id
    except (OSError, ValueError, KeyError) as exc:
        print(f"error: {exc}", file=sys.stderr)
        return 2

    statuses = []
    for chunk_id, result in items:
        status, detail = verify(result, keys, keyring, trust_any)
        statuses.append(status)
        print(f"{chunk_id}: {status} - {detail}")
    if any(s in ("INVALID", "UNSIGNED", "UNKNOWN_KEY", "COMPROMISED_KEY") for s in statuses):
        return 1
    if any(s == "REVOKED" for s in statuses):
        return 3
    return 0


if __name__ == "__main__":
    sys.exit(main())
