#!/usr/bin/env python3
"""Check a TickerTrac record entry yourself, without trusting TickerTrac's answer.

    python3 verify.py https://tickertrac.com/record/entry/<entry id>
    python3 verify.py <entry id>
    python3 verify.py --entry entry.json --keys keys.json [--record record.json]

The first two forms fetch the entry, the published keys and, when the record is listed,
the whole record from the public API. The third runs on saved JSON and touches the
network not at all.

What it recomputes, from the public JSON alone:

  payload_hash        SHA-256 of the canonical payload equals the recorded payload hash
                      (skipped while an entry is sealed: the text is not public yet)
  entry_hash          SHA-256 of the entry header equals the recorded entry hash
  chain_link          the entry names the previous entry's hash (needs the listed record)
  signing_key_known   the signing key id is one TickerTrac publishes
  signature           the Ed25519 signature over the entry hash verifies with that key
  anchor_covers       a daily anchor lists this record's head at or past this entry
  anchor_digest       the anchor's digest recomputes from the heads it lists
  anchor_token        the authority's token covers that digest; its time is when the entry
                      provably existed (pass --tsa-ca to also verify the signature via openssl)

Canonical form: JSON with keys sorted, separators "," and ":", UTF-8, non-ASCII kept as is.
Header fields: schema_version, record_id, sequence_number, kind, payload_sha256,
previous_entry_sha256, recorded_at (UTC, six fractional digits, "Z").

Exit status 0: every check that could run passed. 1: a check failed. 2: could not run.
Needs Python 3.9+ and the cryptography package:  python3 -m pip install cryptography
"""

from __future__ import annotations

import argparse
import base64
import datetime as dt
import hashlib
import json
import os
import re
import sys
import urllib.error
import urllib.request

SCHEMA_VERSION = "tickertrac-record/v1"
DEFAULT_API = "https://api.tickertrac.com"
ENTRY_ID = re.compile(r"[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}", re.IGNORECASE)
TIMESTAMP = re.compile(
    r"^(\d{4})-(\d{2})-(\d{2})[T ](\d{2}):(\d{2}):(\d{2})(?:\.(\d{1,6}))?(Z|z|[+-]\d{2}:?\d{2})$"
)


def canonical_bytes(value) -> bytes:
    return json.dumps(value, sort_keys=True, separators=(",", ":"), ensure_ascii=False).encode("utf-8")


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


def iso_utc(value: str) -> str:
    """Normalise any ISO-8601 rendering of the recorded time to the form the server hashed."""
    match = TIMESTAMP.match(value.strip())
    if not match:
        raise ValueError(f"unrecognised timestamp: {value!r}")
    year, month, day, hour, minute, second, fraction, zone = match.groups()
    micros = int((fraction or "0").ljust(6, "0"))
    if zone in ("Z", "z"):
        offset = dt.timedelta(0)
    else:
        sign = 1 if zone[0] == "+" else -1
        digits = zone[1:].replace(":", "")
        offset = sign * dt.timedelta(hours=int(digits[:2]), minutes=int(digits[2:]))
    moment = dt.datetime(
        int(year), int(month), int(day), int(hour), int(minute), int(second), micros,
        tzinfo=dt.timezone(offset),
    ).astimezone(dt.timezone.utc)
    return moment.strftime("%Y-%m-%dT%H:%M:%S.%f") + "Z"


def fetch_json(url: str):
    request = urllib.request.Request(url, headers={"User-Agent": "tickertrac-record-verify/1", "Accept": "application/json"})
    with urllib.request.urlopen(request, timeout=30) as response:  # noqa: S310 - https URL built from a fixed API base
        return json.load(response)


def load(path: str):
    with open(path, encoding="utf-8") as handle:
        return json.load(handle)


class Report:
    def __init__(self) -> None:
        self.rows: list[tuple[str, str, str]] = []

    def add(self, name: str, status: str, detail: str) -> None:
        self.rows.append((name, status, detail))

    @property
    def failed(self) -> bool:
        return any(status == "fail" for _, status, _ in self.rows)

    def print(self) -> None:
        width = max(len(name) for name, _, _ in self.rows)
        for name, status, detail in self.rows:
            print(f"{name.ljust(width)}  {status.upper():4}  {detail}")


def verify(public_entry: dict, keys: dict, record: dict | None, *, report: Report) -> None:
    entry = public_entry["entry"]
    record_id = public_entry["profile"]["record_id"]

    payload = entry.get("canonical_payload")
    if payload is None:
        report.add("payload_hash", "skip", "Text not public: only the hashes can be checked")
    elif sha256_hex(canonical_bytes(payload)) == entry["payload_sha256"]:
        report.add("payload_hash", "pass", "The published text hashes to the recorded payload hash")
    else:
        report.add("payload_hash", "fail", "The published text does NOT hash to the recorded payload hash")

    header = {
        "schema_version": SCHEMA_VERSION,
        "record_id": record_id,
        "sequence_number": entry["sequence_number"],
        "kind": entry["kind"],
        "payload_sha256": entry["payload_sha256"],
        "previous_entry_sha256": entry.get("previous_entry_sha256"),
        "recorded_at": iso_utc(entry["recorded_at"]),
    }
    if sha256_hex(canonical_bytes(header)) == entry["entry_sha256"]:
        report.add("entry_hash", "pass", f"The entry header hashes to {entry['entry_sha256'][:12]}…")
    else:
        report.add("entry_hash", "fail", "The entry header does NOT hash to the recorded entry hash")

    sequence = entry["sequence_number"]
    previous = entry.get("previous_entry_sha256")
    if sequence == 1:
        if previous is None:
            report.add("chain_link", "pass", "First entry of this record; it names no predecessor")
        else:
            report.add("chain_link", "fail", "A first entry must not name a predecessor")
    elif record is None:
        report.add("chain_link", "skip", "The record is not listed (or was not supplied), so the previous entry's hash cannot be compared")
    else:
        by_sequence = {item["sequence_number"]: item for item in record.get("entries", [])}
        before = by_sequence.get(sequence - 1)
        if before is None:
            report.add("chain_link", "fail", f"Entry {sequence - 1} is missing from the record")
        elif before["entry_sha256"] == previous:
            report.add("chain_link", "pass", f"Links to entry {sequence - 1}'s hash {str(previous)[:12]}…")
        else:
            report.add("chain_link", "fail", f"Entry {sequence - 1}'s hash does not match what this entry names")

    published = {key["key_id"]: key["public_key_base64"] for key in keys.get("keys", [])}
    public_key_b64 = published.get(entry["signing_key_id"])
    if public_key_b64 is None:
        report.add("signing_key_known", "fail", f"Key {entry['signing_key_id']} is not among the published keys")
        report.add("signature", "skip", "No published key to check the signature against")
        return
    report.add("signing_key_known", "pass", f"Signed with published key {entry['signing_key_id']}")

    try:
        from cryptography.exceptions import InvalidSignature
        from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PublicKey
    except ImportError:
        report.add("signature", "skip", "Install the cryptography package to check the signature: python3 -m pip install cryptography")
        return
    try:
        Ed25519PublicKey.from_public_bytes(base64.b64decode(public_key_b64, validate=True)).verify(
            base64.b64decode(entry["signature"], validate=True), entry["entry_sha256"].encode("ascii")
        )
    except (InvalidSignature, ValueError, TypeError):
        report.add("signature", "fail", "The signature does NOT verify over the entry hash with that key")
    else:
        report.add("signature", "pass", "Ed25519 signature over the entry hash verifies")


# ---------------------------------------------------------------------------
# Proof of time: the daily anchor and the authority's token (structure only, plus
# openssl for the signature when a root chain is given)
# ---------------------------------------------------------------------------

OID_SIGNED_DATA = "1.2.840.113549.1.7.2"
OID_TST_INFO = "1.2.840.113549.1.9.16.1.4"


def _tlv(data: bytes, offset: int):
    if offset + 2 > len(data):
        raise ValueError("truncated DER")
    tag = data[offset]
    length = data[offset + 1]
    body_start = offset + 2
    if length & 0x80:
        count = length & 0x7F
        if count == 0 or count > 4 or body_start + count > len(data):
            raise ValueError("bad DER length")
        length = int.from_bytes(data[body_start : body_start + count], "big")
        body_start += count
    body_end = body_start + length
    if body_end > len(data):
        raise ValueError("DER element runs past the end")
    return tag, body_start, body_end


def _children(data: bytes, start: int, end: int):
    out = []
    offset = start
    while offset < end:
        tag, body_start, body_end = _tlv(data, offset)
        out.append((tag, body_start, body_end))
        offset = body_end
    return out


def _oid(data: bytes, start: int, end: int) -> str:
    body = data[start:end]
    parts = [str(body[0] // 40), str(body[0] % 40)]
    value = 0
    for byte in body[1:]:
        value = (value << 7) | (byte & 0x7F)
        if not byte & 0x80:
            parts.append(str(value))
            value = 0
    return ".".join(parts)


def _pick(children, index, tag, what):
    if index >= len(children) or children[index][0] != tag:
        raise ValueError(f"unexpected token structure at {what}")
    return children[index][1], children[index][2]


def parse_token(token: bytes) -> dict:
    """Walk an RFC 3161 token (a CMS ContentInfo) to its TSTInfo: the digest it covers and the time."""
    tag, start, end = _tlv(token, 0)
    if tag != 0x30:
        raise ValueError("token is not a SEQUENCE")
    content = _children(token, start, end)
    oid_start, oid_end = _pick(content, 0, 0x06, "contentType")
    if _oid(token, oid_start, oid_end) != OID_SIGNED_DATA:
        raise ValueError("token is not SignedData")
    explicit = _children(token, *_pick(content, 1, 0xA0, "content"))
    signed = _children(token, *_pick(explicit, 0, 0x30, "SignedData"))
    eci = _children(token, *_pick(signed, 2, 0x30, "EncapsulatedContentInfo"))
    type_start, type_end = _pick(eci, 0, 0x06, "eContentType")
    if _oid(token, type_start, type_end) != OID_TST_INFO:
        raise ValueError("token content is not a TSTInfo")
    octets = _children(token, *_pick(eci, 1, 0xA0, "eContent"))
    octet_start, octet_end = _pick(octets, 0, 0x04, "eContent OCTET STRING")
    tst_tag, tst_start, tst_end = _tlv(token, octet_start)
    if tst_tag != 0x30:
        raise ValueError("TSTInfo is not a SEQUENCE")
    tst = _children(token, tst_start, tst_end)
    imprint = _children(token, *_pick(tst, 2, 0x30, "messageImprint"))
    hashed_start, hashed_end = _pick(imprint, 1, 0x04, "hashedMessage")
    time_start, time_end = _pick(tst, 4, 0x18, "genTime")
    text = token[time_start:time_end].decode("ascii")
    if not text.endswith("Z") or len(text) < 15:
        raise ValueError(f"genTime is not UTC: {text!r}")
    when = f"{text[0:4]}-{text[4:6]}-{text[6:8]}T{text[8:10]}:{text[10:12]}:{text[12:14]}Z"
    return {"digest": token[hashed_start:hashed_end].hex(), "time": when}


def verify_anchor(public_entry: dict, anchor: dict | None, *, report: Report, tsa_ca: str | None) -> None:
    entry = public_entry["entry"]
    record_id = public_entry["profile"]["record_id"]
    if anchor is None:
        report.add("anchor_covers", "skip", "No anchor supplied or the entry is not yet covered by a public timestamp (anchors are taken daily)")
        return
    heads = anchor.get("heads") or []
    mine = next((head for head in heads if head.get("record_id") == record_id), None)
    if mine is None:
        report.add("anchor_covers", "fail", "The anchor lists no head for this record")
    elif int(mine["sequence_number"]) >= int(entry["sequence_number"]):
        report.add("anchor_covers", "pass", f"Anchor {anchor.get('id')} lists this record's head at entry {mine['sequence_number']} (this is entry {entry['sequence_number']})")
    else:
        report.add("anchor_covers", "fail", f"The anchor's head for this record is entry {mine['sequence_number']}, before this entry")

    document = {
        "schema_version": anchor.get("schema_version"),
        "heads": sorted(
            ({"record_id": h["record_id"], "entry_sha256": h["entry_sha256"], "sequence_number": int(h["sequence_number"])} for h in heads),
            key=lambda h: h["record_id"],
        ),
    }
    digest_ok = sha256_hex(canonical_bytes(document)) == anchor.get("digest")
    report.add(
        "anchor_digest",
        "pass" if digest_ok else "fail",
        "The listed heads hash to the anchor's digest" if digest_ok else "The listed heads do NOT hash to the anchor's digest",
    )

    token_b64 = anchor.get("token_base64")
    if not token_b64:
        report.add("anchor_token", "skip", "The anchor has no token yet (the authority was unreachable that day; it is retried daily)")
        return
    try:
        token = base64.b64decode(token_b64, validate=True)
        parsed = parse_token(token)
    except (ValueError, TypeError) as exc:
        report.add("anchor_token", "fail", f"The token could not be read: {exc}")
        return
    if parsed["digest"] != anchor.get("digest"):
        report.add("anchor_token", "fail", "The authority's token covers a different digest than the anchor claims")
        return
    report.add("anchor_token", "pass", f"Existed by {parsed['time']} according to {anchor.get('tsa_url')} (token imprint matches the digest)")

    if not tsa_ca:
        report.add("anchor_signature", "skip", "Pass --tsa-ca <authority root chain PEM> to verify the token's signature with openssl")
        return
    import shutil
    import subprocess
    import tempfile

    openssl = shutil.which("openssl")
    if openssl is None:
        report.add("anchor_signature", "skip", "openssl is not installed; the token's signature was not verified")
        return
    with tempfile.NamedTemporaryFile(suffix=".tsr", delete=False) as handle:
        handle.write(token)
        path = handle.name
    try:
        result = subprocess.run(
            [openssl, "ts", "-verify", "-digest", anchor["digest"], "-in", path, "-token_in", "-CAfile", tsa_ca],
            capture_output=True, text=True, timeout=60, check=False,
        )
    finally:
        os.unlink(path)
    if result.returncode == 0:
        report.add("anchor_signature", "pass", "openssl verified the authority's signature over the digest")
    else:
        report.add("anchor_signature", "fail", (result.stderr or result.stdout).strip().splitlines()[-1] if (result.stderr or result.stdout).strip() else "openssl rejected the token")


def main(argv: list[str]) -> int:
    parser = argparse.ArgumentParser(description=__doc__.splitlines()[0])
    parser.add_argument("entry", nargs="?", help="entry id or a tickertrac.com/record/entry/<id> link")
    parser.add_argument("--api", default=DEFAULT_API, help=f"public API base (default {DEFAULT_API})")
    parser.add_argument("--entry", dest="entry_file", help="saved JSON of GET /record/entries/<id>")
    parser.add_argument("--keys", dest="keys_file", help="saved JSON of GET /record/keys")
    parser.add_argument("--record", dest="record_file", help="saved JSON of GET /record/<handle> (for the chain check)")
    parser.add_argument("--anchor", dest="anchor_file", help="saved JSON of GET /record/anchors/<id> (for the proof of time)")
    parser.add_argument("--no-anchor", action="store_true", help="online mode: do not fetch the covering anchor")
    parser.add_argument("--tsa-ca", dest="tsa_ca", help="PEM chain of the timestamp authority; verifies the token's signature with openssl")
    args = parser.parse_args(argv)

    report = Report()
    try:
        if args.entry_file or args.keys_file:
            if not (args.entry_file and args.keys_file):
                parser.error("--entry and --keys go together")
            public_entry = load(args.entry_file)
            keys = load(args.keys_file)
            record = load(args.record_file) if args.record_file else None
            anchor = load(args.anchor_file) if args.anchor_file else None
        else:
            if not args.entry:
                parser.error("give an entry id or link, or --entry and --keys files")
            match = ENTRY_ID.search(args.entry)
            if not match:
                parser.error("no entry id found in that argument")
            api = args.api.rstrip("/")
            public_entry = fetch_json(f"{api}/record/entries/{match.group(0).lower()}")
            keys = fetch_json(f"{api}/record/keys")
            record = None
            handle = public_entry.get("profile", {}).get("handle")
            if handle:
                try:
                    record = fetch_json(f"{api}/record/{handle}")
                except urllib.error.HTTPError as exc:
                    if exc.code != 404:
                        raise
            anchor = None
            anchored = public_entry.get("anchored")
            if anchored and not args.no_anchor:
                anchor = fetch_json(anchored["anchor_url"])
        verify(public_entry, keys, record, report=report)
        verify_anchor(public_entry, anchor, report=report, tsa_ca=args.tsa_ca)
    except (OSError, ValueError, KeyError, TypeError) as exc:
        print(f"could not verify: {exc}", file=sys.stderr)
        return 2

    entry = public_entry["entry"]
    print(f"entry {entry['id']}  #{entry['sequence_number']}  {entry['kind']}  recorded {entry['recorded_at']}")
    report.print()
    if report.failed:
        print("RESULT: FAILED. This entry does not check out.")
        return 1
    skipped = [name for name, status, _ in report.rows if status == "skip"]
    if skipped:
        print(f"RESULT: passed every check that could run (skipped: {', '.join(skipped)}).")
    else:
        print("RESULT: passed every check.")
    return 0


if __name__ == "__main__":
    sys.exit(main(sys.argv[1:]))
