#!/usr/bin/env python3
"""Independent ICIO verifier: stdlib CSV parsing and transposed origin solves.

Never imports the study analysis code or writes into its directory.
All monetary quantities are current USD million; shares are percentages.
"""
from __future__ import annotations

import argparse
import csv
import hashlib
import json
import math
import platform
import sys
from pathlib import Path

import numpy as np
import scipy
from scipy.linalg import solve

EXPORTERS = ("JPN", "KOR", "DEU")
INDUSTRIES = ("C26", "C27", "C28")
FD_ITEMS = {"HFCE", "NPISH", "GGFC", "GFCF", "INVNT", "DPABR"}
CHECKS = []


def digest(path):
    h = hashlib.sha256()
    with open(path, "rb") as f:
        for chunk in iter(lambda: f.read(1024 * 1024), b""):
            h.update(chunk)
    return h.hexdigest()


def check(name, passed, detail, fatal=False):
    CHECKS.append({"name": name, "status": "PASS" if bool(passed) else "FAIL", "detail": detail})
    if fatal and not passed:
        raise ValueError(f"{name}: {detail}")


def numeric(token):
    return float(token) if token.strip().lower() not in ("", "na", "nan") else float("nan")


def load_matrix(path):
    labels, arrays = [], []
    with open(path, newline="", encoding="utf-8-sig") as f:
        reader = csv.reader(f)
        header = next(reader)[1:]
        for row in reader:
            if not row:
                continue
            if len(row) != len(header) + 1:
                raise ValueError("Ragged matrix row: " + row[0])
            labels.append(row[0])
            arrays.append(np.fromiter((numeric(t) for t in row[1:]), dtype=np.float64, count=len(header)))
    return labels, header, np.vstack(arrays)


def maximum_abs(a):
    return float(np.max(np.abs(a), initial=0))


def compare(name, actual, expected, atol, rtol=0.0):
    delta = abs(float(actual) - float(expected))
    allowed = atol + rtol * abs(float(expected))
    check(name, delta <= allowed, {"independent": float(actual), "reported": float(expected),
                                  "absolute_difference": delta, "allowed": allowed})


def run(args):
    args.out.mkdir(parents=True, exist_ok=True)
    row_labels, column_labels, full = load_matrix(args.raw)
    check("unique_row_labels", len(set(row_labels)) == len(row_labels), len(row_labels), True)
    check("unique_column_labels", len(set(column_labels)) == len(column_labels), len(column_labels), True)
    industry_labels = [x for x in row_labels if "_" in x]
    n = len(industry_labels)
    check("regular_4050_nodes", n == 4050, n, True)
    check("special_rows", set(row_labels) - set(industry_labels) == {"TLS", "VA", "OUT"},
          sorted(set(row_labels) - set(industry_labels)), True)
    row_index = {x: i for i, x in enumerate(row_labels)}
    col_index = {x: i for i, x in enumerate(column_labels)}
    check("industry_column_identity", set(industry_labels) <= set(column_labels), "all production nodes in columns", True)
    ri = [row_index[x] for x in industry_labels]
    ci = [col_index[x] for x in industry_labels]
    origins = sorted(set(x.split("_", 1)[0] for x in industry_labels))
    industries = sorted(set(x.split("_", 1)[1] for x in industry_labels))
    check("regular_universe", len(origins) == 81 and len(industries) == 50 and
          set(industry_labels) == {f"{o}_{i}" for o in origins for i in industries},
          {"economies_including_ROW": len(origins), "industries": len(industries)}, True)
    fd_labels = [x for x in column_labels if x not in industry_labels and x != "OUT"]
    check("final_demand_universe", set(fd_labels) == {f"{o}_{f}" for o in origins for f in FD_ITEMS},
          {"columns": len(fd_labels), "categories": sorted(FD_ITEMS)}, True)
    z = full[np.ix_(ri, ci)]
    fd = full[np.ix_(ri, [col_index[x] for x in fd_labels])]
    output = full[ri, col_index["OUT"]]
    va = full[row_index["VA"], ci]
    tls = full[row_index["TLS"], ci]
    output_row = full[row_index["OUT"], ci]
    check("finite_required_values", all(np.isfinite(a).all() for a in (z, fd, output, va, tls, output_row)),
          {"nonfinite": {k: int((~np.isfinite(v)).sum()) for k, v in
                         {"Z": z, "FD": fd, "x": output, "VA": va, "TLS": tls, "OUT_row": output_row}.items()}}, True)
    check("nonnegative_output", bool((output >= 0).all()), float(output.min()), True)
    zero = output == 0
    check("zero_output_is_inert", maximum_abs(z[:, zero]) == 0 and maximum_abs(z[zero, :]) == 0
          and maximum_abs(fd[zero]) == 0 and maximum_abs(va[zero]) == 0 and maximum_abs(tls[zero]) == 0,
          {"count": int(zero.sum()), "labels": [industry_labels[i] for i in np.flatnonzero(zero)]}, True)
    sign_inventory = {k: {"negative_count": int((v < 0).sum()), "minimum": float(v.min())}
                      for k, v in {"Z": z, "FD": fd, "VA": va, "TLS": tls}.items()}
    check("no_negative_intermediate_values", sign_inventory["Z"]["negative_count"] == 0, sign_inventory["Z"])
    # Source entries are checked for published precision, not rebalanced. A failed
    # pure-rounding check remains visible even when residual accounting succeeds.
    tol_row = (n + len(fd_labels) + 1) * args.source_unit / 2
    tol_col = (n + 3) * args.source_unit / 2
    row_residual = output - z.sum(axis=1) - fd.sum(axis=1)
    col_residual = output - z.sum(axis=0) - va - tls
    compare("output_row_vs_column", maximum_abs(output_row - output), 0, args.source_unit)
    compare("output_row_balance", maximum_abs(row_residual), 0, tol_row)
    compare("output_column_balance", maximum_abs(col_residual), 0, tol_col)
    coef = np.zeros_like(output)
    coef[~zero] = 1 / output[~zero]
    a = z * coef[None, :]
    m = np.eye(n) - a
    masks = np.array([[x.startswith(o + "_") for o in origins] for x in industry_labels], dtype=float)
    # Independent direction: adjoint, 81 VA + 81 TLS origin RHS, signed source
    # residual and L1 absolute source-gap contribution (cancellation diagnostic).
    rhs = np.column_stack((masks * (va * coef)[:, None], masks * (tls * coef)[:, None],
                           col_residual * coef, np.abs(col_residual) * coef))
    h = solve(m.T, rhs, assume_a="gen", check_finite=True)
    solver_residual = maximum_abs(m.T @ h - rhs)
    compare("adjoint_solver_absolute_residual", solver_residual, 0, 1e-11)
    # World accounting identity over every active node, independently of nine export bundles.
    compare("global_origin_tax_residual_identity", maximum_abs(h[~zero, :-1].sum(axis=1) - 1), 0, 1e-10)
    nodes = {x: i for i, x in enumerate(industry_labels)}
    positions = {o: i for i, o in enumerate(origins)}
    independent_rows, origin_rows = [], []
    for exporter in EXPORTERS:
        for industry in INDUSTRIES:
            label = f"{exporter}_{industry}"
            idx = nodes[label]
            foreign_intermediate = [j for j, s in enumerate(industry_labels) if not s.startswith(exporter + "_")]
            foreign_final = [j for j, s in enumerate(fd_labels) if not s.startswith(exporter + "_")]
            domestic_i = [j for j, s in enumerate(industry_labels) if s.startswith(exporter + "_")]
            domestic_f = [j for j, s in enumerate(fd_labels) if s.startswith(exporter + "_")]
            ex_i = math.fsum(float(z[idx, j]) for j in foreign_intermediate)
            ex_f = math.fsum(float(fd[idx, j]) for j in foreign_final)
            ex = ex_i + ex_f
            output_minus_domestic = output[idx] - math.fsum(float(z[idx, j]) for j in domestic_i) - math.fsum(float(fd[idx, j]) for j in domestic_f)
            compare(label + ":export_complement", ex, output_minus_domestic, tol_row)
            check(label + ":positive_exports", ex > 0, ex)
            v = h[idx, :len(origins)] * ex
            t = h[idx, len(origins):2 * len(origins)] * ex
            source_residual = float(h[idx, -2] * ex)
            source_l1_gap = float(h[idx, -1] * ex)
            def amount(o): return float(v[positions[o]])
            def share(o): return float(v[positions[o]] / ex * 100)
            domestic = amount(exporter)
            total_v = math.fsum(map(float, v))
            total_t = math.fsum(map(float, t))
            residual = ex - total_v - total_t - source_residual
            compare(label + ":VA_TLS_source_residual_reconcile", residual, 0, max(1e-7, ex * 1e-10))
            dpabr = math.fsum(float(fd[idx, j]) for j in foreign_final if fd_labels[j].endswith("_DPABR"))
            result = {"exporter": exporter, "industry": industry, "year": 2022,
                      "export_intermediate_usd_mn": ex_i, "export_final_usd_mn": ex_f,
                      "gross_exports_usd_mn": ex, "chn_va_usd_mn": amount("CHN"), "chn_va_share_pct": share("CHN"),
                      "usa_va_usd_mn": amount("USA"), "usa_va_share_pct": share("USA"),
                      "chn_minus_usa_pp": share("CHN") - share("USA"), "domestic_va_usd_mn": domestic,
                      "other_foreign_va_usd_mn": total_v - domestic - amount("CHN") - amount("USA"),
                      "total_va_usd_mn": total_v, "total_tls_usd_mn": total_t,
                      "unallocated_source_balance_usd_mn": source_residual, "reconciliation_residual_usd_mn": residual,
                      "chn_va_plus_tls_share_pct": float((v[positions["CHN"]] + t[positions["CHN"]]) / ex * 100),
                      "usa_va_plus_tls_share_pct": float((v[positions["USA"]] + t[positions["USA"]]) / ex * 100),
                      "foreign_DPABR_usd_mn": dpabr,
                      "foreign_INVNT_usd_mn": math.fsum(float(fd[idx, j]) for j in foreign_final if fd_labels[j].endswith("_INVNT")),
                      "diagnostic_abs_source_gap_usd_mn": source_l1_gap}
            independent_rows.append(result)
            for j, origin in enumerate(origins):
                origin_rows.append({"exporter": exporter, "industry": industry, "origin": origin,
                                    "origin_va_usd_mn": float(v[j]), "origin_va_share_pct": float(v[j] / ex * 100),
                                    "origin_tls_usd_mn": float(t[j]), "origin_tls_share_pct": float(t[j] / ex * 100),
                                    "origin_va_plus_tls_share_pct": float((v[j] + t[j]) / ex * 100)})
    write_csv(args.out / "independent-results.csv", independent_rows)
    write_csv(args.out / "independent-all-origins.csv", origin_rows)
    metadata = {"raw_sha256": digest(args.raw), "raw_file": str(args.raw), "sign_inventory": sign_inventory,
                "source_rounding_unit_usd_mn": args.source_unit,
                "row_balance_max_abs_usd_mn": maximum_abs(row_residual),
                "column_balance_max_abs_usd_mn": maximum_abs(col_residual),
                "row_balance_allowed_usd_mn": tol_row, "column_balance_allowed_usd_mn": tol_col,
                "solver_residual_max_abs": solver_residual,
                "python": platform.python_version(), "numpy": np.__version__, "scipy": scipy.__version__}
    (args.out / "independent-metadata.json").write_text(json.dumps(metadata, indent=2) + "\n")
    if args.study:
        compare_report(args.study / "results.csv", independent_rows, ["exporter", "industry"], "results", args.out)
        compare_report(args.study / "all-origin-reconciliation.csv", origin_rows, ["exporter", "industry", "origin"], "origins", args.out)


def write_csv(path, records):
    with open(path, "w", newline="") as f:
        writer = csv.DictWriter(f, fieldnames=list(records[0]))
        writer.writeheader()
        writer.writerows(records)


def compare_report(path, independent, keys, prefix, out):
    check(prefix + ":reported_csv_exists", path.exists(), str(path))
    if not path.exists():
        return
    with open(path, newline="") as f:
        reported = list(csv.DictReader(f))
    make_key = lambda row: tuple(str(row[k]) for k in keys)
    known = {make_key(r): r for r in independent}
    report_map = {make_key(r): r for r in reported}
    check(prefix + ":complete_unique_keys", len(reported) == len(report_map) and set(report_map) == set(known),
          {"expected": len(known), "reported": len(reported)}, True)
    for key, truth in known.items():
        row = report_map[key]
        for field, expected in truth.items():
            if field in keys or field.startswith(("foreign_", "diagnostic_")):
                continue
            check(prefix + ":field_present:" + field, field in row, "/".join(key))
            if field in row:
                # Output rounding allowed: six decimals for USD million, nine for percentages.
                allowed = 1e-6 if "usd_mn" in field else 1e-8
                compare(prefix + ":" + "/".join(key) + ":" + field, expected, row[field], allowed, 1e-11)


if __name__ == "__main__":
    ap = argparse.ArgumentParser(description=__doc__)
    ap.add_argument("--raw", type=Path, required=True)
    ap.add_argument("--study", type=Path)
    ap.add_argument("--out", type=Path, default=Path(__file__).resolve().parent)
    ap.add_argument("--source-unit", type=float, default=0.0001,
                    help="One least-significant published monetary unit, declared before computation")
    args = ap.parse_args()
    exit_code = 0
    try:
        run(args)
    except Exception as exc:
        check("execution_completed", False, str(exc))
        exit_code = 2
    finally:
        args.out.mkdir(parents=True, exist_ok=True)
        failed = sum(c["status"] == "FAIL" for c in CHECKS)
        summary = {"status": "PASS" if not failed else "FAIL", "check_count": len(CHECKS),
                   "failed_count": failed, "checks": CHECKS}
        (args.out / "independent-checks.json").write_text(json.dumps(summary, indent=2) + "\n")
        print(json.dumps({k: v for k, v in summary.items() if k != "checks"}))
        if failed:
            print(json.dumps([c for c in CHECKS if c["status"] == "FAIL"][:12], indent=2))
    sys.exit(exit_code or (1 if failed else 0))
