#!/usr/bin/env python3
"""Check published CSV dimensions, validation diagnostics and actual SVG geometry.

No analysis functions are imported. SVG bars are measured against axis tick
coordinates, not checked merely against their printed numeric labels.
"""
import argparse
import csv
import hashlib
import json
from pathlib import Path
import re
import xml.etree.ElementTree as ET

NS = {"s": "http://www.w3.org/2000/svg"}
EXPORTERS = ("JPN", "KOR", "DEU")
NAMES = ("Japan", "Korea", "Germany")
INDUSTRIES = ("C26", "C27", "C28")
NUMBER = r"[-+]?(?:\d+\.?\d*|\.\d+)(?:[eE][-+]?\d+)?"


def records(path):
    with open(path, newline="") as f:
        return list(csv.DictReader(f))


def texts(element):
    return [t.text or "" for t in element.findall(".//s:text", NS)]


def groups(root, prefix):
    return [g for g in root.findall(".//s:g", NS) if g.get("id", "").startswith(prefix)]


def rects(axes, color):
    found = []
    for group in axes.findall("s:g", NS):
        if not group.get("id", "").startswith("patch_"):
            continue
        p = group.find("s:path", NS)
        if p is not None and p.get("style") == "fill: " + color:
            values = list(map(float, re.findall(NUMBER, p.get("d", ""))))
            if len(values) == 8:
                xx, yy = values[::2], values[1::2]
                found.append({"left": min(xx), "right": max(xx), "top": min(yy), "bottom": max(yy)})
    return found


def tick_scale(axes, axis):
    ticks = []
    for group in groups(axes, axis + "tick_"):
        words = texts(group)
        if not words or not re.fullmatch(NUMBER, words[0]):
            continue
        uses = group.findall(".//s:use", NS)
        if uses:
            ticks.append((float(words[0]), float(uses[0].attrib[axis])))
    ticks.sort()
    if len(ticks) < 2:
        raise ValueError("Insufficient tick coordinates")
    p0, p1 = ticks[0], ticks[-1]
    scale = (p1[1] - p0[1]) / (p1[0] - p0[0])
    return p0[1] - p0[0] * scale, scale


def main(args):
    checks = []
    def test(name, ok, detail):
        checks.append({"name": name, "status": "PASS" if ok else "FAIL", "detail": detail})
    def compare(name, actual, expected, tolerance):
        delta = abs(float(actual) - float(expected))
        test(name, delta <= tolerance, {"measured": float(actual), "expected": float(expected),
                                       "absolute_difference": delta, "tolerance": tolerance})

    independent = records(args.out / "independent-results.csv")
    reported = records(args.study / "results.csv")
    origin_rows = records(args.study / "all-origin-reconciliation.csv")
    by_cell = {(r["exporter"], r["industry"]): r for r in independent}
    test("results_and_origins_year", all(r["year"] == "2022" for r in reported + origin_rows), "Every row is 2022")
    test("all_origin_domestic_labels", all(r["is_domestic"].lower() == str(r["origin"] == r["exporter"]).lower() for r in origin_rows), "All 729 origin rows")
    test("headline_positive_ordering", all(float(r["chn_va_share_pct"]) > float(r["usa_va_share_pct"]) > 0 for r in independent), "All nine, with no sample dropping")
    validation = json.loads((args.study / "validation.json").read_text())
    for row in validation["reconciliation"]:
        key = tuple(row["cell"].split("_", 1))
        r = by_cell[key]
        compare(row["cell"] + ":absolute_source_gap", r["diagnostic_abs_source_gap_usd_mn"], row["propagated_absolute_source_gap_usd_mn"], 1e-6)
    test("source_balance_failure_disclosed", validation["source_balance_status"] == "does_not_close_within_rounding", validation["source_balance_status"])
    test("qualification_status", validation["status"] == "passed_with_source_balance_warnings", validation["status"])

    share_path = args.study / "china-us-shares-2022.svg"
    share_root = ET.parse(share_path).getroot()
    panels = groups(share_root, "axes_")
    test("share_chart_three_panels", len(panels) == 3, len(panels))
    zero, scale = tick_scale(panels[0], "y")
    for ax, exporter, name in zip(panels, EXPORTERS, NAMES):
        words = texts(ax)
        test(exporter + ":share_chart_panel_label", name in words and all(i in words for i in INDUSTRIES), words)
        for color, field in (("#bd452f", "chn_va_share_pct"), ("#27698c", "usa_va_share_pct")):
            bars = sorted(rects(ax, color), key=lambda r: r["left"])
            test(exporter + ":share_bar_count:" + field, len(bars) == 3, len(bars))
            for industry, bar in zip(INDUSTRIES, bars):
                expected = float(by_cell[(exporter, industry)][field])
                measured = (bar["top"] - zero) / scale
                compare(exporter + "_" + industry + ":share_geometry:" + field, measured, expected, 2e-6)
                test(exporter + "_" + industry + ":share_text:" + field, f"{expected:.2f}" in words, f"{expected:.2f}")
                compare(exporter + "_" + industry + ":zero_baseline:" + field, bar["bottom"], zero, 2e-6)
    all_share_words = " ".join(texts(share_root))
    test("share_chart_year_and_basis", "2022" in all_share_words and "Basic-price VA" in all_share_words and "taxes separate" in all_share_words,
         "Year and primary VA-only basis appear on exported figure")

    account_path = args.study / "origin-accounting-2022.svg"
    account_root = ET.parse(account_path).getroot()
    ax = groups(account_root, "axes_")[0]
    zero, scale = tick_scale(ax, "x")
    expected_order = [(e, i) for e in EXPORTERS for i in INDUSTRIES]
    labels = [texts(g)[0] for g in groups(ax, "ytick_")]
    expected_labels = [name + " · " + i for name in NAMES for i in INDUSTRIES]
    test("account_chart_row_order", labels == expected_labels, labels)
    left = {k: 0.0 for k in expected_order}
    for color, field in (("#c4cbd1", "domestic_va_usd_mn"), ("#bd452f", "chn_va_usd_mn"),
                         ("#27698c", "usa_va_usd_mn"), ("#5e8f75", "other_foreign_va_usd_mn"),
                         ("#d4ab47", "total_tls_usd_mn"), ("#383e45", "unallocated_source_balance_usd_mn")):
        bars = sorted(rects(ax, color), key=lambda r: r["top"])
        test("account_component_count:" + field, len(bars) == 9, len(bars))
        for key, bar in zip(expected_order, bars):
            r = by_cell[key]
            expected = float(r[field]) / float(r["gross_exports_usd_mn"]) * 100
            compare("_".join(key) + ":account_geometry:" + field, (bar["right"] - bar["left"]) / scale, expected, 2e-6)
            compare("_".join(key) + ":account_left:" + field, (bar["left"] - zero) / scale, left[key], 2e-6)
            left[key] += expected
    for key, total in left.items():
        compare("_".join(key) + ":account_chart_total", total, 100, 1e-9)
    all_account_words = " ".join(texts(account_root))
    test("account_chart_source_balance_caveat", "Source-balance residuals" in all_account_words and "not assigned to an origin" in all_account_words,
         "Residual transparency retained")
    test("account_chart_no_coalition_sum", "do not add the nine export bundles" in all_account_words, "Separate bundles")
    hashes = {p.name: hashlib.sha256(p.read_bytes()).hexdigest() for p in
              [share_path, account_path, args.study / "china-us-shares-2022.png", args.study / "origin-accounting-2022.png"]}
    failed = sum(c["status"] == "FAIL" for c in checks)
    report = {"status": "PASS" if not failed else "FAIL", "failed_count": failed, "check_count": len(checks),
              "figure_sha256": hashes, "checks": checks,
              "visual_inspection": "Both PNGs inspected by verifier agent: all panels/labels/legends/footnotes readable; no clipping observed. Not a human review."}
    (args.out / "output-and-figure-checks.json").write_text(json.dumps(report, indent=2) + "\n")
    print(json.dumps({k: v for k, v in report.items() if k not in ("checks", "figure_sha256")}))
    if failed:
        print(json.dumps([c for c in checks if c["status"] == "FAIL"], indent=2))
    return bool(failed)


if __name__ == "__main__":
    p = argparse.ArgumentParser(description=__doc__)
    p.add_argument("--study", type=Path, required=True)
    p.add_argument("--out", type=Path, default=Path(__file__).resolve().parent)
    raise SystemExit(main(p.parse_args()))
