#!/usr/bin/env python3
"""Compute the Voigt and Reuss equivalent core moduli for the homogeneous control.

The equivalent modulus must come from the SAME element-wise group assignment the
heterogeneous solve uses, otherwise the control is not comparable. This script
therefore calls the upstream loader on the real mesh rather than re-deriving the
grayscale mapping:

    load_reynolds_cellwise_heterogeneity(volume, points, cells, cell_tags, ...)

It then reports, weighted by reference tetrahedron volume over the core region:

    Voigt (upper bound)  G_V = sum(f_i * G_i)
    Reuss (lower bound)  G_R = 1 / sum(f_i / G_i)

and converts each to a Young modulus using the same K/G ratio the heterogeneous
run uses (default 30 -> nu = 0.4835), so the homogeneous control keeps the same
compressibility as the material it replaces.

Any physically admissible homogenisation lies between Reuss and Voigt, so
running the control at both brackets the answer without having to argue for one
particular averaging scheme.
"""

from __future__ import annotations

import argparse
import json
import sys
from pathlib import Path

import numpy as np


def _find_mesh(project_root: Path, mesh_name: str) -> Path:
    """Locate the gmsh file for a named mesh, tolerating layout differences."""
    candidates = sorted(project_root.rglob("*.msh"))
    if not candidates:
        raise SystemExit(
            f"No .msh found under {project_root}. The submission zip does not "
            "ship the mesh generator; run this against the real working tree."
        )
    exact = [path for path in candidates if mesh_name in str(path)]
    if exact:
        return exact[0]
    if len(candidates) == 1:
        return candidates[0]
    raise SystemExit(
        f"Could not disambiguate a mesh for {mesh_name!r}. Candidates:\n  "
        + "\n  ".join(str(path) for path in candidates)
        + "\nPass --mesh-path explicitly."
    )


def main() -> int:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--project-root", type=Path, required=True)
    parser.add_argument("--mesh-name", default="intermediate_h0p75")
    parser.add_argument("--mesh-path", type=Path, default=None)
    parser.add_argument(
        "--template-volume",
        type=Path,
        default=None,
        help="Chromatin .npz (default: search the project tree).",
    )
    parser.add_argument("--outer-axes-um", type=float, nargs=3, default=(7.5, 5.5, 3.5))
    parser.add_argument("--shell-thickness-um", type=float, default=0.12)
    parser.add_argument("--group-count", type=int, default=10)
    parser.add_argument("--minimum-shear-pa", type=float, default=170.0)
    parser.add_argument("--maximum-shear-pa", type=float, default=4000.0)
    parser.add_argument("--bulk-to-shear-ratio", type=float, default=30.0)
    parser.add_argument("--json-out", type=Path, default=None)
    args = parser.parse_args()

    source_dir = args.project_root / "05_Source_Data"
    if not (source_dir / "nonlinear_solver.py").exists():
        source_dir = args.project_root
    for path in (str(source_dir), str(args.project_root / "src")):
        if path not in sys.path:
            sys.path.insert(0, path)

    import meshio  # noqa: E402
    from nuclear_envelope_fem.chromatin_heterogeneity import (  # noqa: E402
        load_reynolds_cellwise_heterogeneity,
    )
    from nuclear_envelope_fem import tags  # noqa: E402

    mesh_path = args.mesh_path or _find_mesh(args.project_root, args.mesh_name)
    volume_path = args.template_volume
    if volume_path is None:
        matches = sorted(args.project_root.rglob("*chromatin_50cube.npz"))
        if not matches:
            raise SystemExit("Chromatin .npz not found; pass --template-volume.")
        volume_path = matches[0]

    print(f"mesh   : {mesh_path}")
    print(f"volume : {volume_path}")

    mesh = meshio.read(mesh_path)
    points = np.asarray(mesh.points, dtype=float)
    cells = None
    cell_tags = None
    for block, data in zip(mesh.cells, mesh.cell_data.get("gmsh:physical", []), strict=False):
        if block.type == "tetra":
            cells = np.asarray(block.data, dtype=np.int64)
            cell_tags = np.asarray(data, dtype=np.int32)
    if cells is None:
        raise SystemExit(f"No tetrahedral block in {mesh_path}")

    inner_axes = np.asarray(args.outer_axes_um, dtype=float) - float(args.shell_thickness_um)
    heterogeneity = load_reynolds_cellwise_heterogeneity(
        str(volume_path),
        points,
        cells,
        cell_tags,
        inner_axes,
        group_count=args.group_count,
        minimum_shear_pa=args.minimum_shear_pa,
        maximum_shear_pa=args.maximum_shear_pa,
        bulk_to_shear_ratio=args.bulk_to_shear_ratio,
    )

    core = cell_tags == tags.CORE_VOLUME
    tetra = points[cells[core]]
    volumes = np.abs(
        np.einsum(
            "ij,ij->i",
            tetra[:, 1] - tetra[:, 0],
            np.cross(tetra[:, 2] - tetra[:, 0], tetra[:, 3] - tetra[:, 0]),
        )
    ) / 6.0
    groups = np.asarray(heterogeneity.group_ids, dtype=np.int64)[core]

    shear_by_group = []
    for material in heterogeneity.materials_by_group:
        modulus = getattr(material, "shear_modulus_pa", None)
        if modulus is None:
            young = getattr(material, "young_modulus_pa", None)
            poisson = getattr(material, "poisson_ratio", None)
            if young is None or poisson is None:
                raise SystemExit(
                    "Cannot read a shear modulus off the material objects; "
                    "inspect materials_by_group and adjust this script."
                )
            modulus = young / (2.0 * (1.0 + poisson))
        shear_by_group.append(float(modulus))
    shear_by_group = np.asarray(shear_by_group, dtype=float)

    total = volumes.sum()
    fractions = np.zeros(len(shear_by_group), dtype=float)
    for index in range(len(shear_by_group)):
        fractions[index] = volumes[groups == index].sum() / total

    present = fractions > 0.0
    voigt = float(np.sum(fractions[present] * shear_by_group[present]))
    reuss = float(1.0 / np.sum(fractions[present] / shear_by_group[present]))
    poisson = (3.0 * args.bulk_to_shear_ratio - 2.0) / (6.0 * args.bulk_to_shear_ratio + 2.0)

    report = {
        "mesh": str(mesh_path),
        "template_volume": str(volume_path),
        "core_tetrahedra": int(core.sum()),
        "poisson_ratio_from_bulk_shear_ratio": poisson,
        "group_volume_fractions": fractions.tolist(),
        "group_shear_moduli_pa": shear_by_group.tolist(),
        "voigt_shear_pa": voigt,
        "reuss_shear_pa": reuss,
        "voigt_young_pa": 2.0 * voigt * (1.0 + poisson),
        "reuss_young_pa": 2.0 * reuss * (1.0 + poisson),
    }

    print("\nVolume-weighted equivalents over the core region")
    print(f"  Poisson ratio (K/G = {args.bulk_to_shear_ratio:g}) : {poisson:.4f}")
    print(f"  Voigt  G = {voigt:9.1f} Pa   ->  E = {report['voigt_young_pa']:9.1f} Pa")
    print(f"  Reuss  G = {reuss:9.1f} Pa   ->  E = {report['reuss_young_pa']:9.1f} Pa")
    print("\nFeed both into 03_run_homogeneous_control.py:")
    print(
        f"  --core-young-pa {report['reuss_young_pa']:.1f} {report['voigt_young_pa']:.1f} "
        "--label reuss voigt"
    )

    if args.json_out:
        args.json_out.write_text(json.dumps(report, indent=2), encoding="utf-8")
        print(f"\nWrote {args.json_out}")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
