#!/usr/bin/env python3
"""mpmath O3 for mechanical.statics.planar_equilibrium. ΣF and ΣM about the origin."""
from __future__ import annotations
import json, math, shutil
from pathlib import Path
import mpmath as mp


def split_rows(raw: str):
    return [row.strip() for row in str(raw).replace("\n", ";").split(";") if row.strip()]


def parse_forces(raw: str):
    out = []
    for row in split_rows(raw):
        parts = [p.strip() for p in row.split(",")]
        unknown = ""
        if len(parts) == 5:
            unknown = parts[4].lower()
        elif len(parts) == 3 and any(c.isalpha() for c in parts[2]):
            unknown = parts[2].lower()
        has_pos = len(parts) >= 4 or (len(parts) == 3 and not unknown)
        fx = mp.mpf(parts[0])
        fy = mp.mpf(parts[1])
        x = mp.mpf(parts[2]) if has_pos and parts[2] and not any(c.isalpha() for c in parts[2]) else mp.mpf(0)
        y = mp.mpf(parts[3]) if has_pos and len(parts) > 3 and parts[3] != "" else mp.mpf(0)
        out.append({"fx": fx, "fy": fy, "x": x, "y": y, "unknown": unknown})
    return out


def parse_moments(raw: str | None):
    if not raw:
        return []
    out = []
    for row in split_rows(raw):
        parts = [p.strip() for p in row.split(",")]
        unknown = len(parts) > 1 and parts[1].lower() == "m"
        out.append({"m": mp.mpf(parts[0]), "unknown": unknown})
    return out


def round_smart(n):
    x = float(n)
    if not math.isfinite(x):
        return x
    if abs(x) < 1e-15:
        return 0.0
    # Match IUT Number(n.toPrecision(12))
    return float(format(x, ".12g"))


def moment_of(x, y, fx, fy):
    return mp.mpf(str(round_smart(float(x * fy - y * fx))))


def assemble(forces, moments):
    known_x = mp.mpf(0)
    known_y = mp.mpf(0)
    known_m = mp.mpf(0)
    columns = []
    for i, force in enumerate(forces):
        fx, fy = force["fx"], force["fy"]
        unk = force["unknown"]
        if unk in ("fx", "fx+fy", "both", "f"):
            columns.append({"id": f"F{i+1}.fx", "coeff": [mp.mpf(1), mp.mpf(0), -force["y"]]})
            fx = mp.mpf(0)
        if unk in ("fy", "fx+fy", "both", "f"):
            columns.append({"id": f"F{i+1}.fy", "coeff": [mp.mpf(0), mp.mpf(1), force["x"]]})
            fy = mp.mpf(0)
        if fx != 0 or fy != 0:
            known_x = mp.mpf(str(round_smart(float(known_x + fx))))
            known_y = mp.mpf(str(round_smart(float(known_y + fy))))
            known_m += moment_of(force["x"], force["y"], fx, fy)
    for i, moment in enumerate(moments):
        if moment["unknown"]:
            columns.append({"id": f"M{i+1}", "coeff": [mp.mpf(0), mp.mpf(0), mp.mpf(1)]})
        else:
            known_m += moment["m"]
    return known_x, known_y, known_m, columns


def solve_n(At, b):
    n = len(b)
    M = [[At[r][c] for c in range(n)] + [b[r]] for r in range(n)]
    for col in range(n):
        pivot = max(range(col, n), key=lambda r: abs(M[r][col]))
        if M[pivot][col] == 0:
            raise ValueError("singular")
        M[col], M[pivot] = M[pivot], M[col]
        div = M[col][col]
        for c in range(col, n + 1):
            M[col][c] /= div
        for r in range(n):
            if r == col:
                continue
            factor = M[r][col]
            for c in range(col, n + 1):
                M[r][c] -= factor * M[col][c]
    return [M[r][n] for r in range(n)]


def solve_unknowns(columns, known_x, known_y, known_m):
    n = len(columns)
    b3 = [-known_x, -known_y, -known_m]
    if n == 3:
        At = [[columns[c]["coeff"][r] for c in range(3)] for r in range(3)]
        return solve_n(At, b3)
    tol = mp.mpf("1e-8")
    last = None
    for index in ([0, 1], [0, 2], [1, 2]):
        At = [[columns[c]["coeff"][index[r]] for c in range(2)] for r in range(2)]
        try:
            x = solve_n(At, [b3[i] for i in index])
        except ValueError:
            continue
        residual = []
        for row in range(3):
            acc = sum(columns[k]["coeff"][row] * x[k] for k in range(2))
            residual.append(acc - b3[row])
        worst = max(abs(v) for v in residual)
        scale = 1 + max(abs(v) for v in b3)
        if worst <= tol * scale:
            return x
        last = worst
    raise ValueError(f"unsatisfiable 2-unknown residual {last}")


def pack_num(v):
    return {"f64": float(v), "decimal": mp.nstr(v, 40, strip_zeros=False)}


def row(vid, inputs, kind):
    forces = parse_forces(inputs.get("forces") or inputs.get("F") or "")
    moments = parse_moments(inputs.get("moments") or inputs.get("M"))
    known_x, known_y, known_m, columns = assemble(forces, moments)
    resultant = mp.mpf(str(round_smart(float(mp.sqrt(known_x * known_x + known_y * known_y)))))
    unknowns = []
    if inputs.get("mode") == "equilibrium" and columns:
        values = solve_unknowns(columns, known_x, known_y, known_m)
        unknowns = [{"id": col["id"], "value_f64": float(values[i]), "value_decimal": mp.nstr(values[i], 40, strip_zeros=False)} for i, col in enumerate(columns)]
    return {
        "id": vid,
        "kind": kind,
        "inputs": inputs,
        "resultant_n_f64": float(resultant),
        "resultant_n_decimal": mp.nstr(resultant, 40, strip_zeros=False),
        "sum_fx_f64": float(known_x),
        "sum_fx_decimal": mp.nstr(known_x, 40, strip_zeros=False),
        "sum_fy_f64": float(known_y),
        "sum_fy_decimal": mp.nstr(known_y, 40, strip_zeros=False),
        "sum_m_f64": float(known_m),
        "sum_m_decimal": mp.nstr(known_m, 40, strip_zeros=False),
        "unknowns": unknowns,
    }


def main():
    mp.mp.dps = 80
    here = Path(__file__).resolve().parent
    vectors = [
        row("o3-345", {"mode": "resultant", "forces": "3,0; 0,4"}, "resultant"),
        row("o3-moment", {"mode": "resultant", "forces": "0,-10,2,0"}, "moment"),
        row("o3-bracket", {"mode": "equilibrium", "forces": "0,-100,2,0; 0,0,0,0,fx+fy", "moments": "0,m"}, "equilibrium3"),
        row("o3-two", {"mode": "equilibrium", "forces": "0,-10; 0,0,fx+fy"}, "equilibrium2"),
        row("o3-awkward", {"mode": "resultant", "forces": "1.37,-2.5,0.83,0.2"}, "awkward"),
        row("o3-alias", {"mode": "resultant", "F": "5,0; -3,4"}, "alias"),
        row("o3-moments-only", {"mode": "resultant", "moments": "12; -4"}, "moments"),
        row("o3-cancel", {"mode": "resultant", "forces": "2,3; -2,-3"}, "cancel"),
    ]
    table = {
        "family": "planar_equilibrium",
        "generator_id": "planar-equilibrium-mpmath-o3",
        "generator_version": "1.0.0",
        "seed": "20260927.planar-equilibrium-o3",
        "mpmath_dps": 80,
        "precision_bits": int(80 * math.log2(10)),
        "library": f"mpmath {mp.__version__}",
        "notes": "ΣFx, ΣFy, ΣM about origin; M = x Fy − y Fx. Not beam deflection.",
        "vectors": vectors,
    }
    dest = here / "planar-equilibrium-o3-tables.json"
    dest.write_text(json.dumps(table, indent=2) + "\n", encoding="utf-8")
    public = here.parents[3] / "public/developers/cvp/reproduce"
    public.mkdir(parents=True, exist_ok=True)
    shutil.copy2(dest, public / "planar-equilibrium-o3-tables.json")
    shutil.copy2(Path(__file__), public / "generate-planar-equilibrium-o3.py")
    print(f"wrote {dest} ({len(vectors)} vectors)")


if __name__ == "__main__":
    main()
