#!/usr/bin/env python3
"""Reproduce the frozen nine-cell, 2022 upstream-value-added study offline.

Use the official regular ICIO CSV, never the extended/firm-split table. All
country-industry nodes are retained. The headline is basic-price value added;
intermediate-product net taxes are a separate account, not relabelled as VA.
"""
import os

# Modest CPU use, including on machines with a large default BLAS thread count.
for _name in ("OPENBLAS_NUM_THREADS", "MKL_NUM_THREADS", "OMP_NUM_THREADS", "BLIS_NUM_THREADS"):
    os.environ[_name] = "2"

import argparse
import csv
import hashlib
import importlib.metadata
import json
from pathlib import Path
import platform
import sys
import warnings

import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
from openpyxl import load_workbook
import pandas as pd
from scipy.linalg import lu_factor, lu_solve
from scipy.linalg.lapack import get_lapack_funcs

HERE = Path(__file__).resolve().parent
FD = ["HFCE", "NPISH", "GGFC", "GFCF", "INVNT", "DPABR"]
EXPORTERS = ["JPN", "KOR", "DEU"]
INDUSTRIES = ["C26", "C27", "C28"]
COUNTRY_NAMES = {"JPN": "Japan", "KOR": "Korea", "DEU": "Germany"}
SECTOR_SHORT = {"C26": "Electronics & optical", "C27": "Electrical equipment", "C28": "Machinery"}


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


def save_json(path, data):
    path.write_text(json.dumps(data, indent=2, allow_nan=False) + "\n")


def sample_negative(array, row_labels, col_labels=None, count=12):
    where = np.argwhere(array < 0)
    samples = []
    for loc in where[:count]:
        i = int(loc[0])
        item = {"row": row_labels[i], "value": float(array[tuple(loc)])}
        if col_labels is not None:
            item["column"] = col_labels[int(loc[1])]
        samples.append(item)
    return {"count": int(len(where)), "minimum": float(np.min(array)), "sum_negative": float(array[array < 0].sum()), "examples": samples}


def read_documentation(raw_dir):
    with warnings.catch_warnings():
        warnings.simplefilter("ignore", UserWarning)
        workbook = load_workbook(raw_dir / "ReadMe_ICIO_small.xlsx", read_only=True, data_only=True)
    rows = [row[2] for row in workbook["RowItems"].values if len(row) > 2 and isinstance(row[1], (int, float))]
    columns = [row[2] for row in workbook["ColItems"].values if len(row) > 2 and isinstance(row[1], (int, float))]
    countries, sectors = {}, {}
    for row in workbook["Area_Activities"].values:
        for code_pos, name_pos in ((2, 3), (5, 6)):
            if isinstance(row[code_pos], str) and len(row[code_pos]) == 3 and row[code_pos].isupper():
                countries[row[code_pos]] = row[name_pos]
        if isinstance(row[8], (int, float)) and isinstance(row[9], str):
            sectors[row[9]] = {"name": row[10], "isic_rev4": str(row[11])}
    units = str(workbook["ReadMe"]["C11"].value)
    if "current USD, million" not in units:
        raise ValueError("Unexpected or unverified data units: " + units)
    workbook.close()
    return rows, columns, countries, sectors, units


def balance_description(gaps, tolerance, labels):
    index = int(np.argmax(np.abs(gaps)))
    return {"max_abs_usd_mn": float(np.abs(gaps).max()), "max_abs_node": labels[index], "signed_sum_usd_mn": float(gaps.sum()), "tolerance_usd_mn": tolerance, "exceedances": int((np.abs(gaps) > tolerance).sum())}


def make_charts(results, output):
    plt.rcParams.update({"font.family": "DejaVu Sans", "font.size": 10, "axes.spines.top": False, "axes.spines.right": False, "svg.fonttype": "none", "svg.hashsalt": "statecrafts-2022-icio-study"})
    china, usa = "#bd452f", "#27698c"
    fig, axes = plt.subplots(1, 3, figsize=(12.8, 5.3), sharey=True)
    ymax = max(results.chn_va_share_pct.max(), results.usa_va_share_pct.max()) * 1.30
    for ax, exporter in zip(axes, EXPORTERS):
        r = results[results.exporter == exporter]
        idx = np.arange(3)
        a = ax.bar(idx - .19, r.chn_va_share_pct, .36, color=china, label="China-origin VA")
        b = ax.bar(idx + .19, r.usa_va_share_pct, .36, color=usa, label="U.S.-origin VA")
        ax.bar_label(a, fmt="%.2f", padding=4, fontsize=9)
        ax.bar_label(b, fmt="%.2f", padding=4, fontsize=9)
        ax.set_xticks(idx, ["C26\nElectronics", "C27\nElectrical", "C28\nMachinery"])
        ax.set_title(COUNTRY_NAMES[exporter], fontweight="bold")
        ax.set_ylim(0, ymax)
        ax.grid(axis="y", alpha=.2)
        ax.set_axisbelow(True)
    axes[0].set_ylabel("Origin value added / sector gross exports (%)")
    handles, labels = axes[0].get_legend_handles_labels()
    fig.legend(handles, labels, loc="upper right", bbox_to_anchor=(.97, .92), ncol=2, frameon=False)
    fig.suptitle("China and U.S. upstream value added in manufacturing exports — 2022", fontsize=15, x=.04, ha="left", fontweight="bold")
    fig.text(.04, .055, "Source: Statecrafts calculations, OECD regular ICIO 2025 edition, January 2026 revision. Basic-price VA; taxes separate.", fontsize=8.5)
    fig.text(.04, .023, "Historical accounting benchmark, not a measurement of current dependence, substitutability, or wartime access.", fontsize=8.5)
    fig.subplots_adjust(left=.075, right=.98, top=.76, bottom=.22, wspace=.12)
    for ext in ("png", "svg"):
        meta = {"Software": "Statecrafts research"} if ext == "png" else {"Date": None, "Creator": "Statecrafts research"}
        fig.savefig(output / ("china-us-shares-2022." + ext), dpi=180, metadata=meta)
    plt.close(fig)

    fig, ax = plt.subplots(figsize=(12.8, 6.5))
    y = np.arange(9)
    labels = [f"{COUNTRY_NAMES[row.exporter]} · {row.industry}" for row in results.itertuples()]
    parts = [("Domestic VA", "domestic_va_usd_mn", "#c4cbd1"), ("China-origin VA", "chn_va_usd_mn", china), ("U.S.-origin VA", "usa_va_usd_mn", usa), ("Other foreign VA", "other_foreign_va_usd_mn", "#5e8f75"), ("Net intermediate taxes", "total_tls_usd_mn", "#d4ab47"), ("Unallocated source balance", "unallocated_source_balance_usd_mn", "#383e45")]
    left = np.zeros(9)
    for label, field, color in parts:
        values = results[field].to_numpy() / results.gross_exports_usd_mn.to_numpy() * 100
        ax.barh(y, values, left=left, height=.62, label=label, color=color)
        left += values
    ax.set_yticks(y, labels)
    ax.invert_yaxis()
    ax.set_xlim(0, 100.05)
    ax.set_xlabel("Share of each gross-export bundle (%)")
    ax.set_title("Where the value in each export bundle originated — 2022", fontsize=15, fontweight="bold", loc="left", pad=18)
    ax.legend(loc="upper center", bbox_to_anchor=(.45, -.14), ncol=3, frameon=False, fontsize=9)
    ax.spines["left"].set_visible(False)
    ax.spines["bottom"].set_visible(False)
    ax.tick_params(axis="y", length=0)
    fig.text(.04, .035, "Basic-price VA and net taxes are separate. Source-balance residuals (too small to see here) are retained, not assigned to an origin.", fontsize=8.5)
    fig.text(.04, .012, "OECD regular ICIO, 2025 edition / January 2026 revision. Each bar is separate; do not add the nine export bundles into a coalition total.", fontsize=8.5)
    fig.subplots_adjust(left=.16, right=.97, top=.88, bottom=.25)
    for ext in ("png", "svg"):
        meta = {"Software": "Statecrafts research"} if ext == "png" else {"Date": None, "Creator": "Statecrafts research"}
        fig.savefig(output / ("origin-accounting-2022." + ext), dpi=180, metadata=meta)
    plt.close(fig)


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--raw-dir", type=Path, required=True)
    parser.add_argument("--output-dir", type=Path, default=HERE)
    parser.add_argument("--manifest", type=Path, default=HERE / "source-manifest.json")
    parser.add_argument("--skip-charts", action="store_true")
    args = parser.parse_args()
    args.output_dir.mkdir(parents=True, exist_ok=True)
    manifest = json.loads(args.manifest.read_text())
    scope = json.loads((HERE / "frozen-scope.json").read_text())
    if scope["exporters"] != EXPORTERS or scope["industries"] != INDUSTRIES or scope["observation_year"] != 2022:
        raise ValueError("Frozen scope changed")
    table = args.raw_dir / manifest["extracted_table"]["filename"]
    table_hash = digest(table)
    if table_hash != manifest["extracted_table"]["sha256"]:
        raise ValueError("Raw table checksum differs from the manifest")
    readme_entry = next(x for x in manifest["sources"] if x["filename"] == "ReadMe_ICIO_small.xlsx")
    if digest(args.raw_dir / "ReadMe_ICIO_small.xlsx") != readme_entry["sha256"]:
        raise ValueError("ReadMe checksum differs from the manifest")
    expected_rows, expected_cols, countries, sectors, units = read_documentation(args.raw_dir)
    print("Reading " + table.name, flush=True)
    frame = pd.read_csv(table, index_col=0)
    labels = list(frame.index)
    cols = list(frame.columns)
    if labels != expected_rows or cols != expected_cols:
        raise ValueError("CSV labels differ from official ReadMe sector-code labels")
    industries = labels[:-3]
    if labels[-3:] != ["TLS", "VA", "OUT"] or cols[-1] != "OUT" or cols[:len(industries)] != industries:
        raise ValueError("Unexpected account structure")
    n = len(industries)
    origin_codes = np.array([label.split("_", 1)[0] for label in industries])
    industry_codes = [label.split("_", 1)[1] for label in industries]
    origins = list(dict.fromkeys(origin_codes.tolist()))
    if n != 4050 or len(origins) != 81 or set(industry_codes) != set(sectors):
        raise ValueError("This is not the expected full regular ICIO (81 origins × 50 sectors)")
    if "CN1" in origins or "CN2" in origins or "MX1" in origins or "MX2" in origins:
        raise ValueError("Extended ICIO must not replace the frozen regular version")
    expected_fd = [country + "_" + use for country in origins for use in FD]
    fd_cols = cols[n:-1]
    if fd_cols != expected_fd:
        raise ValueError("Unexpected final demand structure")
    numeric = frame.to_numpy(dtype=np.float64)
    z = numeric[:n, :n]
    final = numeric[:n, n:-1]
    x = numeric[-1, :n]
    x_col = numeric[:n, -1]
    v = numeric[-2, :n]
    tax = numeric[-3, :n]
    for array in (z, final, x, x_col, v, tax):
        if not np.isfinite(array).all():
            raise ValueError("Missing or non-finite observation in a required matrix block")
    if (x < 0).any():
        raise ValueError("Negative gross output cannot be normalized")
    zero = x == 0
    if np.any(z[:, zero]) or np.any(z[zero, :]) or np.any(final[zero, :]) or np.any(v[zero]) or np.any(tax[zero]):
        raise ValueError("A zero-output node has nonzero transactions/accounts")

    # Monetary CSV rounding is documented in source-format.json, frozen after
    # inspecting numeric formatting but before computing any cell estimates.
    source_format = json.loads((HERE / "source-format.json").read_text())
    rounding_unit = float(source_format["numeric_rounding_unit_usd_mn"])
    row_tolerance = (z.shape[1] + final.shape[1] + 1) * rounding_unit / 2 + 1e-7
    col_tolerance = (z.shape[0] + 3) * rounding_unit / 2 + 1e-7
    row_gap = x_col - z.sum(axis=1) - final.sum(axis=1)
    col_gap = x - z.sum(axis=0) - v - tax
    out_gap = x - x_col
    validation = {
        "status": "pending",
        "source_table_sha256": table_hash,
        "shape": {"production_nodes": n, "origins": len(origins), "industries_per_origin": len(sectors), "final_demand_columns": len(fd_cols)},
        "labels_match_official_readme": True,
        "numeric_sheet_ids_used": False,
        "finite_required_data": True,
        "units": "current USD million",
        "source_numeric_rounding_unit_usd_mn": rounding_unit,
        "balance_tolerance_rule": "number of independently rounded terms × half the source rounding unit + 1e-7 USD million floating-point allowance",
        "output_identity": balance_description(row_gap, row_tolerance, industries),
        "input_identity_including_explicit_tls": balance_description(col_gap, col_tolerance, industries),
        "output_row_vs_column": balance_description(out_gap, rounding_unit + 1e-7, industries),
        "zero_output_nodes": [industries[i] for i in np.flatnonzero(zero)],
        "zero_output_handling": "Verified no flows, VA or TLS; coefficient column set to zero without dropping any origin",
        "negative_values": {
            "intermediate_use": sample_negative(z, industries, industries),
            "final_demand": sample_negative(final, industries, fd_cols),
            "value_added_basic_prices": sample_negative(v, industries),
            "intermediate_net_taxes": sample_negative(tax, industries),
            "final_demand_by_category": {use: sample_negative(final[:, j::6], industries, fd_cols[j::6]) for j, use in enumerate(FD)},
        },
        "negative_value_policy": "No clipping or netting away; explicit source values retained and negative inventory/net-tax entries reported separately",
    }
    if validation["output_row_vs_column"]["exceedances"]:
        save_json(args.output_dir / "validation.json", validation)
        raise ValueError("Output row differs from output column")
    source_balance_warnings = any(validation[key]["exceedances"] for key in ("output_identity", "input_identity_including_explicit_tls"))
    validation["source_balance_status"] = "does_not_close_within_rounding" if source_balance_warnings else "within_rounding_envelope"
    validation["source_balance_handling"] = "Keep all published matrix entries unchanged. Carry x - column_sum(Z) - VA - TLS as an unallocated source-balance account; never relabel it as rounding or allocate it to any origin. Report both net propagated residual and propagated absolute source-gap mass."
    pd.DataFrame({"node": industries, "output_usd_mn": x, "output_less_row_uses_usd_mn": row_gap, "output_less_input_va_tls_usd_mn": col_gap, "output_row_minus_column_usd_mn": out_gap, "row_exceeds_rounding_envelope": np.abs(row_gap) > row_tolerance, "column_exceeds_rounding_envelope": np.abs(col_gap) > col_tolerance}).to_csv(args.output_dir / "source-balance-diagnostics.csv", index=False, float_format="%.17g")
    inverse_x = np.divide(1.0, x, out=np.zeros_like(x), where=x != 0)
    a = z * inverse_x[None, :]
    system = np.eye(n) - a
    target_labels = [exporter + "_" + industry for exporter in EXPORTERS for industry in INDUSTRIES]
    e = np.zeros((n, len(target_labels)))
    target_info = []
    for k, label in enumerate(target_labels):
        exporter, industry = label.split("_", 1)
        index = industries.index(label)
        foreign_intermediate = origin_codes != exporter
        foreign_final = np.array([c.split("_", 1)[0] != exporter for c in fd_cols])
        intermediate_exports = float(z[index, foreign_intermediate].sum())
        final_exports = float(final[index, foreign_final].sum())
        exports = intermediate_exports + final_exports
        if exports <= 0:
            raise ValueError("Selected cell has a nonpositive gross-export denominator")
        e[index, k] = exports
        target_info.append((exporter, industry, intermediate_exports, final_exports, exports))
    print("Solving full 4,050-node system for nine right-hand sides (2 BLAS threads)", flush=True)
    factor = lu_factor(system, check_finite=False)
    gecon = get_lapack_funcs("gecon", (system,))
    reciprocal_condition, condition_info = gecon(factor[0], np.linalg.norm(system, 1), norm="1")
    if condition_info != 0 or reciprocal_condition <= 0:
        raise ValueError("Unable to estimate Leontief-system conditioning")
    q = lu_solve(factor, e, check_finite=False)
    residual = system @ q - e
    solver_relative = np.max(np.abs(residual), axis=0) / np.max(np.abs(e), axis=0)
    if float(solver_relative.max()) > scope["tolerances"]["solver_max_abs_relative_to_rhs"]:
        raise ValueError("Leontief solver residual exceeds frozen tolerance")
    va_contributions = (v * inverse_x)[:, None] * q
    tax_contributions = (tax * inverse_x)[:, None] * q
    # This unallocated account includes source imbalance beyond rounding. It is
    # not an attributed VA origin and does not modify the published VA vector.
    source_balance_contributions = (col_gap * inverse_x)[:, None] * q
    absolute_source_gap_propagated = (np.abs(col_gap) * inverse_x)[:, None] * np.abs(q)
    results, origin_rows, reconciliation = [], [], []
    for k, (exporter, industry, intermediate_exports, final_exports, exports) in enumerate(target_info):
        va_origin = {origin: float(va_contributions[origin_codes == origin, k].sum()) for origin in origins}
        tax_origin = {origin: float(tax_contributions[origin_codes == origin, k].sum()) for origin in origins}
        total_va, total_tax = float(va_contributions[:, k].sum()), float(tax_contributions[:, k].sum())
        source_balance = float(source_balance_contributions[:, k].sum())
        gap = exports - total_va - total_tax - source_balance
        results.append({
            "exporter": exporter, "industry": industry, "year": 2022,
            "export_intermediate_usd_mn": intermediate_exports, "export_final_usd_mn": final_exports, "gross_exports_usd_mn": exports,
            "chn_va_usd_mn": va_origin["CHN"], "chn_va_share_pct": va_origin["CHN"] / exports * 100,
            "usa_va_usd_mn": va_origin["USA"], "usa_va_share_pct": va_origin["USA"] / exports * 100,
            "chn_minus_usa_pp": (va_origin["CHN"] - va_origin["USA"]) / exports * 100,
            "domestic_va_usd_mn": va_origin[exporter],
            "other_foreign_va_usd_mn": sum(value for origin, value in va_origin.items() if origin not in (exporter, "CHN", "USA")),
            "total_va_usd_mn": total_va, "total_tls_usd_mn": total_tax,
            "unallocated_source_balance_usd_mn": source_balance, "reconciliation_residual_usd_mn": gap,
            "chn_va_plus_tls_share_pct": (va_origin["CHN"] + tax_origin["CHN"]) / exports * 100,
            "usa_va_plus_tls_share_pct": (va_origin["USA"] + tax_origin["USA"]) / exports * 100,
            "source_table_sha256": table_hash,
        })
        for origin in origins:
            origin_rows.append({"exporter": exporter, "industry": industry, "year": 2022, "origin": origin, "is_domestic": origin == exporter, "origin_va_usd_mn": va_origin[origin], "origin_va_share_pct": va_origin[origin] / exports * 100, "origin_tls_usd_mn": tax_origin[origin], "origin_tls_share_pct": tax_origin[origin] / exports * 100, "origin_va_plus_tls_share_pct": (va_origin[origin] + tax_origin[origin]) / exports * 100})
        reconciliation.append({"cell": exporter + "_" + industry, "export_denominator_usd_mn": exports, "all_origin_basic_price_va_usd_mn": total_va, "explicit_intermediate_tls_usd_mn": total_tax, "unallocated_source_balance_usd_mn": source_balance, "unallocated_source_balance_share_pct": source_balance / exports * 100, "propagated_absolute_source_gap_usd_mn": float(absolute_source_gap_propagated[:, k].sum()), "propagated_absolute_source_gap_share_pct": float(absolute_source_gap_propagated[:, k].sum()) / exports * 100, "numerical_residual_usd_mn": gap, "numerical_residual_relative": gap / exports, "solver_max_abs_relative_to_rhs": float(solver_relative[k])})
    validation.update({
        "solver": "SciPy LU factorization plus nine forward solves; no explicit matrix inverse",
        "blas_threads": 2,
        "system_condition_number_1norm_estimate": float(1 / reciprocal_condition),
        "solver_max_abs_relative_to_rhs": float(solver_relative.max()),
        "nonnegative_output_responses_with_tolerance": bool(q.min() >= -1e-9),
        "min_output_response_usd_mn": float(q.min()),
        "reconciliation": reconciliation,
        "max_reconciliation_residual_relative": max(abs(item["numerical_residual_relative"]) for item in reconciliation),
        "tax_convention": "Headline basic-price VA excludes TLS. TLS is an explicit separate account; supplementary VA+TLS percentages are labelled and never silently substituted.",
        "no_residual_assigned_to_origin": True,
        "no_coalition_sum_or_confidence_interval": True,
    })
    if validation["max_reconciliation_residual_relative"] > scope["tolerances"]["all_origin_reconciliation_with_explicit_tax_and_rounding_relative"]:
        raise ValueError("All-origin VA + TLS + unallocated source balance does not reconcile")
    if validation["negative_values"]["intermediate_use"]["count"]:
        raise ValueError("Unexplained negative intermediate-use entries require review before publishing")
    validation["numerical_status"] = "passed"
    validation["status"] = "passed_with_source_balance_warnings" if source_balance_warnings else "passed"
    data = pd.DataFrame(results)
    data.to_csv(args.output_dir / "results.csv", index=False, float_format="%.17g")
    pd.DataFrame(origin_rows).to_csv(args.output_dir / "all-origin-reconciliation.csv", index=False, float_format="%.17g")
    save_json(args.output_dir / "validation.json", validation)
    save_json(args.output_dir / "mappings.json", {"countries": countries, "industries": sectors, "final_demand_categories": FD, "readme_units_text": units, "production_labels": industries, "numeric_readme_ids_used": False})
    high = data.loc[data.chn_va_share_pct.idxmax()]
    low = data.loc[data.chn_va_share_pct.idxmin()]
    summary = {"study_id": scope["study_id"], "observation_year": 2022, "dataset_edition": scope["edition"], "units": "current USD million", "measure": "Origin basic-price value added as a percentage of sector gross exports; intermediate-product net taxes excluded from headline VA and separately reconciled", "source_table_sha256": table_hash, "scope_frozen_at_utc": scope["frozen_at_utc"], "cell_count": len(data), "china_above_us_cell_count": int((data.chn_minus_usa_pp > 0).sum()), "china_share_range_pct": [float(low.chn_va_share_pct), float(high.chn_va_share_pct)], "china_share_highest_cell": {"exporter": high.exporter, "industry": high.industry, "share_pct": float(high.chn_va_share_pct)}, "china_share_lowest_cell": {"exporter": low.exporter, "industry": low.industry, "share_pct": float(low.chn_va_share_pct)}, "results": results, "validation_status": validation["status"], "claim_limits": scope["limitations"], "aggregation_warning": "Do not sum gross-export bundles into unique coalition value added; upstream flows can cross borders repeatedly", "confidence_intervals": None}
    save_json(args.output_dir / "summary.json", summary)
    environment = {"python": sys.version, "platform": platform.platform(), "implementation": platform.python_implementation(), "packages": {pkg: importlib.metadata.version(pkg) for pkg in ("numpy", "scipy", "pandas", "matplotlib", "openpyxl")}, "blas_thread_limit": 2, "numpy_blas": np.__config__.CONFIG.get("Build Dependencies", {}).get("blas", {}), "entrypoint": "analyze.py", "analysis_script_sha256": digest(Path(__file__)), "scope_sha256": digest(HERE / "frozen-scope.json")}
    save_json(args.output_dir / "environment.json", environment)
    if not args.skip_charts:
        make_charts(data, args.output_dir)
    print(data[["exporter", "industry", "gross_exports_usd_mn", "chn_va_share_pct", "usa_va_share_pct", "chn_minus_usa_pp"]].to_string(index=False), flush=True)
    print("Validation: " + validation["status"], flush=True)


if __name__ == "__main__":
    main()
