"""Verify a ParetoAlpha public-filing proof without trusting ParetoAlpha. Standard library only.

    python3 verify_proof.py 1e3df8bf8e6c                        # a short link code or a 64-hex proof hash
    python3 verify_proof.py 1e3df8bf8e6c --fetch                # and, for a 990-PF, fetch the return from irs.gov yourself
    python3 verify_proof.py 1e3df8bf8e6c --return 2024.xml      # or pass a copy of the IRS return you already have

What it checks, and against what:
  1. The body.     Fetches the sealed body (/api/proof/body) and recomputes its hash with the canonical form in
                   verify.py (beside this file). A body that was changed after sealing no longer matches its hash.
  2. The ledger.   Asks /api/proof/verify whether that hash is on the ledger and when it was first recorded.
  3. The source.   The return the IRS published: --fetch reads that one file from irs.gov (the zip and offset are
                   on the proof page; Deflate64 is inflated here, since Python's zipfile cannot), or --return takes
                   a copy you have. Its SHA-256 must equal the body's sourceSha256, and
  4. The math.     re-derives the return's own arithmetic from the XML, independently of ParetoAlpha's code: the 5%
                   minimum investment return (prorated by days/365 for a short year, or by days/366 in a leap year
                   when only that reconciles, the reading of the IRS instructions for Part IX line 6; a 52-53-week
                   year reported at the month end, or of exactly 364 days, is read as a full year when the filed line
                   reconciles to the unprorated 5%), the
                   excise tax at the rate the form prescribes for the filer (1.39% of net investment income, or 4% of
                   Part I line 12 column (b) for a foreign organization, IRC 4948(a)), and the Part X distributable
                   amount; and compares each with what the proof says was filed and recomputed.

Steps 1, 3 and 4 need nothing from ParetoAlpha except the body, which you can also paste in with --body FILE.json.
The seal format is published at https://paretoalphasystems.com/standard.

A body sealed from version 2 names the code that computed it (body.build.digest). When the arithmetic re-derived here
disagrees with the seal, the verdict line says "sealed under rules <digest>, re-derived under <these rules>": a proof
sealed under earlier rules is superseded, not false; /api/proof/verify names the proof that supersedes it.
"""
import argparse, hashlib, json, re, sys, urllib.request, zlib
from datetime import date
from decimal import Decimal, getcontext
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parent))
from verify import content_sha256  # the canonical form, one definition shared by every verifier

getcontext().prec = 50
SITE = "https://paretoalphasystems.com"
# The rules this file re-derives with; named beside a seal's own build digest when the two disagree.
RULES = "verify_proof.py rules of 2026-09-27"


# ─── Fetching the IRS return yourself (--fetch) ─────────────────────────────
# The IRS publishes 990-PF returns inside large zips, some members compressed with Deflate64 (zip method 9), which
# Python's zipfile refuses. So this reads only the one member, with an HTTP Range request straight to irs.gov, and
# inflates it here. The locator (which zip, where in it) comes from the proof page and is not trusted: the bytes must
# match the member's CRC-32 and then the SHA-256 sealed inside the proof, or the check fails.

LOCAL_SLACK = 30 + 256 + 1024

_LEN_BASE = [3, 4, 5, 6, 7, 8, 9, 10, 11, 13, 15, 17, 19, 23, 27, 31, 35, 43, 51, 59, 67, 83, 99, 115, 131, 163, 195, 227, 258]
_LEN_EXTRA = [0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 2, 2, 2, 2, 3, 3, 3, 3, 4, 4, 4, 4, 5, 5, 5, 5, 0]
_DIST_BASE = [1, 2, 3, 4, 5, 7, 9, 13, 17, 25, 33, 49, 65, 97, 129, 193, 257, 385, 513, 769, 1025, 1537, 2049, 3073,
              4097, 6145, 8193, 12289, 16385, 24577, 32769, 49153]
_DIST_EXTRA = [0, 0, 0, 0, 1, 1, 2, 2, 3, 3, 4, 4, 5, 5, 6, 6, 7, 7, 8, 8, 9, 9, 10, 10, 11, 11, 12, 12, 13, 13, 14, 14]


def inflate_raw(data: bytes, deflate64: bool = False) -> bytes:
    """Raw DEFLATE (RFC 1951), and Deflate64 when asked: a 64 KiB window, length code 285 = 3 + 16 extra bits, and
    distance codes 30-31. A small, plain decoder: correctness over speed; a return is tens of kilobytes."""
    pos, bit, out = 0, 0, bytearray()

    def bits(n):
        nonlocal pos, bit
        v = 0
        for i in range(n):
            if pos >= len(data):
                raise ValueError("compressed data ended early")
            v |= ((data[pos] >> bit) & 1) << i
            bit += 1
            if bit == 8:
                bit, pos = 0, pos + 1
        return v

    def table(lengths):
        counts, codes = [0] * 16, {}
        for L in lengths:
            counts[L] += 1
        counts[0], code, nxt = 0, 0, [0] * 16
        for L in range(1, 16):
            code = (code + counts[L - 1]) << 1
            nxt[L] = code
        for sym, L in enumerate(lengths):
            if L:
                codes[(L, nxt[L])] = sym
                nxt[L] += 1
        return codes

    def decode(t):
        code = 0
        for L in range(1, 16):
            code = (code << 1) | bits(1)
            if (L, code) in t:
                return t[(L, code)]
        raise ValueError("invalid Huffman code")

    len_base, len_extra = list(_LEN_BASE), list(_LEN_EXTRA)
    if deflate64:
        len_base[28], len_extra[28] = 3, 16
    fixed_lit = table([8] * 144 + [9] * 112 + [7] * 24 + [8] * 8)
    fixed_dist = table([5] * 32)
    final = 0
    while not final:
        final, kind = bits(1), bits(2)
        if kind == 0:
            if bit:
                bit, pos = 0, pos + 1
            n = data[pos] | (data[pos + 1] << 8)
            pos += 4
            out += data[pos:pos + n]
            pos += n
            continue
        if kind == 1:
            lit, dist = fixed_lit, fixed_dist
        elif kind == 2:
            hlit, hdist, hclen = bits(5) + 257, bits(5) + 1, bits(4) + 4
            order = [16, 17, 18, 0, 8, 7, 9, 6, 10, 5, 11, 4, 12, 3, 13, 2, 14, 1, 15]
            cl = [0] * 19
            for i in range(hclen):
                cl[order[i]] = bits(3)
            clt, lengths = table(cl), []
            while len(lengths) < hlit + hdist:
                sym = decode(clt)
                if sym < 16:
                    lengths.append(sym)
                elif sym == 16:
                    lengths += [lengths[-1]] * (3 + bits(2))
                elif sym == 17:
                    lengths += [0] * (3 + bits(3))
                else:
                    lengths += [0] * (11 + bits(7))
            lit, dist = table(lengths[:hlit]), table(lengths[hlit:])
        else:
            raise ValueError("invalid block type")
        while True:
            sym = decode(lit)
            if sym < 256:
                out.append(sym)
            elif sym == 256:
                break
            else:
                i = sym - 257
                n = len_base[i] + bits(len_extra[i])
                d = decode(dist)
                if d >= (32 if deflate64 else 30):
                    raise ValueError("invalid distance code")
                back = _DIST_BASE[d] + bits(_DIST_EXTRA[d])
                if back > len(out):
                    raise ValueError("distance reaches before the start")
                for _ in range(n):
                    out.append(out[-back])
    return bytes(out)


def member_from_local(buf: bytes, m: dict) -> bytes:
    """A read that starts at a zip member's local header -> its bytes, size- and CRC-checked."""
    if len(buf) < 30 or buf[:4] != b"PK\x03\x04":
        raise ValueError("no local file header at the recorded offset")
    start = 30 + int.from_bytes(buf[26:28], "little") + int.from_bytes(buf[28:30], "little")
    data = buf[start:start + m["compressedSize"]]
    if len(data) < m["compressedSize"]:
        raise ValueError("the read ended before the member did")
    if m["method"] == 0:
        out = data
    elif m["method"] == 8:
        out = zlib.decompress(data, -15)
    elif m["method"] == 9:
        out = inflate_raw(data, deflate64=True)
    else:
        raise ValueError(f"compression method {m['method']}")
    if len(out) != m["size"]:
        raise ValueError(f"size {len(out)} is not the recorded {m['size']}")
    if (zlib.crc32(out) & 0xFFFFFFFF) != (m["crc"] & 0xFFFFFFFF):
        raise ValueError("CRC-32 does not match: the bytes changed in transit")
    return out


def fetch_return(loc: dict) -> bytes:
    """The one IRS member, read with a single Range request to irs.gov and inflated here."""
    if not str(loc.get("zipUrl", "")).startswith("https://apps.irs.gov/"):
        raise ValueError("the locator does not point at apps.irs.gov")
    m = loc["member"]
    a, b = m["headerOffset"], m["headerOffset"] + LOCAL_SLACK + m["compressedSize"] - 1
    req = urllib.request.Request(loc["zipUrl"], headers={"Range": f"bytes={a}-{b}", "User-Agent": "verify_proof.py"})
    with urllib.request.urlopen(req, timeout=60) as r:
        if r.status != 206:
            raise ValueError(f"irs.gov answered {r.status} to a range request")
        return member_from_local(r.read(), m)


def get_json(url: str):
    with urllib.request.urlopen(urllib.request.Request(url, headers={"User-Agent": "verify_proof.py"}), timeout=20) as r:
        return json.loads(r.read().decode())


def resolve(code: str) -> str:
    code = code.strip().lower().rsplit("/", 1)[-1]
    if re.fullmatch(r"[0-9a-f]{64}", code):
        return code
    if re.fullmatch(r"[0-9a-f]{12}", code):
        req = urllib.request.Request(f"{SITE}/p/{code}", method="HEAD", headers={"User-Agent": "verify_proof.py"})
        opener = urllib.request.build_opener(type("NoRedirect", (urllib.request.HTTPRedirectHandler,), {"redirect_request": lambda *a, **k: None}))
        try:
            opener.open(req, timeout=20)
        except urllib.error.HTTPError as e:
            loc = e.headers.get("Location", "")
            m = re.search(r"[0-9a-f]{64}", loc)
            if m:
                return m.group(0)
    sys.exit(f"not a proof hash or short-link code: {code}")


def group(xml: str, name: str):
    m = re.search(r"<%s(?:\s[^>]*)?>(.*?)</%s>" % (name, name), xml, re.S)
    return m.group(1) if m else None


def val(block, tag):
    if block is None:
        return None
    m = re.search(r"<%s>([^<]*)</%s>" % (tag, tag), block)
    return Decimal(m.group(1)) if m else None


def checked(block, tag) -> bool:
    """A check box on the form: X, 1 or true."""
    return bool(re.search(r"<%s>(X|1|true)<" % tag, block or "", re.I))


def month_end(iso: str) -> bool:
    d = date.fromisoformat(iso)
    nxt = date(d.year + (d.month == 12), d.month % 12 + 1, 1)
    return (nxt - d).days == 1


def twelve_month_days(begin: str) -> int:
    """Days in the twelve months that start on `begin`: 366 when they cross a 29 February, 365 otherwise."""
    b = date.fromisoformat(begin)
    try:
        nxt = b.replace(year=b.year + 1)
    except ValueError:  # 29 February: the next year has none
        nxt = date(b.year + 1, 3, 1)
    return (nxt - b).days


def touches_leap_year(begin: str, end: str) -> bool:
    """Whether any calendar year the period touches is a leap year (the instructions' 'in a leap year')."""
    leap = lambda y: (y % 4 == 0 and y % 100 != 0) or y % 400 == 0
    return any(leap(y) for y in range(int(begin[:4]), int(end[:4]) + 1))


def recompute_990pf(xml: str):
    """The return's own arithmetic, from the XML alone. Keys match the proof's check ids."""
    hdr = group(xml, "ReturnHeader")
    begin = re.search(r"<TaxPeriodBeginDt>([^<]*)<", hdr or "")
    end = re.search(r"<TaxPeriodEndDt>([^<]*)<", hdr or "").group(1)
    days = (date.fromisoformat(end) - date.fromisoformat(begin.group(1))).days + 1 if begin else None
    pf = group(xml, "IRS990PF")
    mir, dist = group(pf, "MinimumInvestmentReturnGrp"), group(pf, "DistributableAmountGrp")
    ex, rev = group(pf, "ExciseTaxBasedOnInvstIncmGrp"), group(pf, "AnalysisOfRevenueAndExpenses")
    out = {}
    n, m = val(mir, "NetVlNoncharitableAssetsAmt"), val(mir, "MinimumInvestmentReturnAmt")
    # Fewer than 364 days is a short year, prorated, except in the 52-53-week band (IRC 441(f)): a year of whole weeks
    # that does not begin on the 1st is reported at the calendar month end, so its e-file spans 357-363 days. There the
    # filed line is read as a full year when it reconciles (within $1) to the unprorated 5%, and prorated otherwise.
    band = bool(days and 357 <= days <= 364 and begin and ((int(begin.group(1)[8:10]) != 1 and month_end(end)) or days == 364))
    week_year = band and n is not None and m is not None and abs(m - n * Decimal("0.05")) <= 1
    # Short = fewer days than the twelve months that start on the first day (365, or 366 across a 29 February).
    year_len = twelve_month_days(begin.group(1)) if begin else 365
    short = bool(days and days < year_len) and not week_year
    if n is not None and m is not None:
        tol = abs(n) * Decimal("0.0000025") + 1
        denom = 365
        # The leap-year reading of the IRS instructions (Part IX line 6: "365, or 366 in a leap year"): only where the
        # period touches a leap year and the filed line reconciles to days/366 while days/365 does not.
        if short and touches_leap_year(begin.group(1), end):
            at365, at366 = n * Decimal("0.05") * Decimal(days) / 365, n * Decimal("0.05") * Decimal(days) / 366
            if abs(m - at365) > tol and abs(m - at366) <= tol:
                denom = 366
        out["minimum_investment_return"] = (m, n * Decimal("0.05") * (Decimal(days) / denom if short else 1))
        # The largest residual a day fraction rounded to four decimals can explain, plus whole-dollar rounding.
        out["_mir_tolerance"] = tol if short else None
    # Part V line 1: 1.39% of net investment income for a domestic foundation; for a foreign organization (header item
    # D1) the form prescribes 4% of Part I line 12, column (b) (IRC 4948(a)). No key when the base line is absent.
    t = val(ex, "InvestmentIncomeExciseTaxAmt")
    if checked(pf, "ForeignOrganizationInd"):
        gross = val(rev, "TotalNetInvstIncmAmt")
        if gross is not None and t is not None:
            out["excise_tax"] = (t, gross * Decimal("0.04"))
    else:
        nii = val(rev, "NetInvestmentIncomeAmt")
        if nii is not None and t is not None:
            out["excise_tax"] = (t, nii * Decimal("0.0139"))
    # Part X does not apply to a private operating foundation (its own statement) or to a filer that checks the Part X
    # box (4942(j)(3)/(j)(5) foundations and certain foreign organizations): two boxes, read apart.
    operating = checked(pf, "PrivateOperatingFoundationInd")
    part_x_box = checked(dist, "Sect4942j3j5FndtnAndFrgnOrgInd")
    if not operating and not part_x_box and val(dist, "DistributableAsAdjustedAmt") is not None and val(dist, "MinimumInvestmentReturnAmt") is not None:
        z = lambda v: v if v is not None else Decimal(0)
        out["distributable_amount"] = (val(dist, "DistributableAsAdjustedAmt"),
                                       val(dist, "MinimumInvestmentReturnAmt") - z(val(dist, "TotalTaxAmt")) + z(val(dist, "RecoveriesQualfiedDistriAmt")) - z(val(dist, "DeductionFromDistributableAmt")))
    return out


def expected_verdict(check_id, filed, recomputed, mine):
    """The verdict the proof should carry, re-derived here rather than trusted."""
    d = abs(filed - recomputed)
    if check_id == "distributable_amount" and filed == 0 and recomputed < 0:
        return "floored"
    if d == 0:
        return "exact"
    if d <= 1:
        return "rounding"
    tol = mine.get("_mir_tolerance") if check_id == "minimum_investment_return" else None
    return "prorated" if tol is not None and d <= tol else "differs"


def main():
    ap = argparse.ArgumentParser(description=__doc__.split("\n")[0])
    ap.add_argument("proof", help="64-hex proof hash, 12-hex short-link code, or a /p/ or /proof/ link")
    ap.add_argument("--return", dest="ret", help="the IRS 990-PF XML the proof was read from")
    ap.add_argument("--body", help="a saved body JSON instead of fetching it")
    ap.add_argument("--fetch", action="store_true", help="for a 990-PF: fetch the return yourself from irs.gov (one Range request) instead of --return")
    ap.add_argument("--offline", action="store_true", help="with --body and --return: skip the ledger; check body, source and math only")
    ap.add_argument("--current-build", help="the 990-PF engine digest of the code now deployed, if you know it, to name beside an older seal's")
    a = ap.parse_args()
    ok = True
    # VERIFIED means what was checked; a check skipped is said, never implied. Each skipped step is named here so the
    # closing line can say which: a reader who passed --return has checked the return, whatever the ledger did.
    skipped = {}
    h = resolve(a.proof)
    served = json.loads(Path(a.body).read_text()) if a.body else get_json(f"{SITE}/api/proof/body?hash={h}")
    body = served["body"]

    got = content_sha256(body)
    print(f"1 body    {'OK ' if got == h else 'FAIL'}  sha256(canonical(body)) = {got}")
    ok &= got == h
    if a.offline:
        print("2 ledger  --   skipped (--offline)")
        skipped["ledger"] = "--offline"
    else:
        try:
            led = get_json(f"{SITE}/api/proof/verify?hash={h}")
            print(f"2 ledger  {'OK ' if led.get('known') else 'FAIL'}  on the ledger: {led.get('known')}, kind: {led.get('kind')}, first recorded {led.get('firstSeen')}")
            ok &= bool(led.get("known"))
        except Exception as e:  # the ledger check is the only one that needs ParetoAlpha to answer
            print(f"2 ledger  --   could not reach the ledger ({e}); steps 1, 3 and 4 do not need it")
            skipped["ledger"] = "the ledger was unreachable"

    if body.get("kind") == "public-filing-990pf":
        src = body["source"]
        print(f"          990-PF of {body['filer']['name']} (EIN {body['filer']['ein']}), IRS object {src['objectId']}, period to {src['period']['end']}")
        # Body v2 names the code that computed it. A seal made under earlier rules is superseded, not false: when the
        # re-derivation disagrees, the line below says under which rules each side was computed.
        sealed_build = (body.get("build") or {}).get("digest") if isinstance(body.get("build"), dict) else None
        current = a.current_build or (served.get("currentBuild") or {}).get("pfDigest") if isinstance(served.get("currentBuild"), dict) else a.current_build
        if sealed_build:
            commit = (body.get("build") or {}).get("commit")
            print(f"          sealed under rules {sealed_build}{f' (commit {str(commit)[:12]})' if commit else ''} (body v{body.get('version')})")
        rules_note = ""
        if sealed_build:
            rules_note = f" sealed under rules {sealed_build[:12]}..., re-derived under {current[:12] + '...' if current and current != sealed_build else (RULES if not current else 'the same rules')}"
        elif body.get("version") == 1:
            # A v1 body predates the build stamp: it can only be said to have been sealed under the rules of its day.
            rules_note = f" sealed as body v1 (rules before 2026-09-27, no build stamp), re-derived under {RULES}"
        raw = None
        if a.ret:
            raw = Path(a.ret).read_bytes()
        elif a.fetch:
            loc = served.get("locator")
            if not loc:
                print("3 source  FAIL  no locator for this proof; download the return and pass --return")
                ok = False
            else:
                print(f"          fetching {loc['member']['name']} from {loc['zipUrl']} (one Range request, inflated here)")
                raw = fetch_return(loc)
        if raw is None:
            if not a.fetch:
                print("3 source  --   pass --fetch (or --return <the IRS XML>) to check the source and re-derive the arithmetic")
            skipped["source"] = "no --return or --fetch"
        else:
            sha = hashlib.sha256(raw).hexdigest()
            print(f"3 source  {'OK ' if sha == src['sourceSha256'] else 'FAIL'}  sha256(return) = {sha}")
            ok &= sha == src["sourceSha256"]
            mine = recompute_990pf(raw.decode("utf-8", "replace"))
            for c in body.get("checks", []):
                if c["id"].startswith("_") or c["id"] not in mine:
                    print(f"4 math    FAIL  {c['id']}: the proof has it, the return does not")
                    ok = False
                    continue
                filed, recomputed = mine[c["id"]]
                same = Decimal(c["filed"]) == filed and abs(Decimal(c["recomputed"]) - recomputed) < Decimal("1e-20")
                want = expected_verdict(c["id"], filed, recomputed, mine)
                print(f"4 math    {'OK ' if same else 'FAIL'}  {c['id']}: filed {filed}, recomputed {recomputed.normalize():f}{'' if same else ';' + rules_note}")
                print(f"  verdict {'OK ' if want == c['verdict'] else 'FAIL'}  the proof says {c['verdict']}; re-derived here: {want}{'' if want == c['verdict'] else ';' + rules_note}")
                ok &= same and want == c["verdict"]
            if not ok and rules_note:
                print(f"          a seal and a re-derivation that disagree:{rules_note}. An older seal reads as superseded, not false;"
                      f" {SITE}/api/proof/verify?hash={h} names the proof that supersedes it, if one was sealed.")
    if not ok:
        print("\nNOT VERIFIED")
    elif "source" in skipped:
        print("\nVERIFIED AS FAR AS CHECKED: the body is proven unchanged since sealing but not yet proven to match the IRS\n"
              "return. Pass --return or --fetch to check the source and re-derive the arithmetic yourself."
              + (f" The ledger was also skipped ({skipped['ledger']})." if "ledger" in skipped else ""))
    elif "ledger" in skipped:
        print(f"\nVERIFIED (ledger not checked): body, source and arithmetic were re-derived here and match the seal; whether\n"
              f"ParetoAlpha recorded this hash was not checked ({skipped['ledger']}). Steps 1, 3 and 4 need nothing from ParetoAlpha.")
    else:
        print("\nVERIFIED")
    sys.exit(0 if ok else 1)


if __name__ == "__main__":
    main()
