"""
Quick CFD Solver (Ray-Based, Navier-Stokes-Free)

Features:
- Real STL import
- Z-axis slicing
- Surface-normal & curvature encoding
- Streamline-style ray probes
- 4D ray carry-over between slices
- Pressure / velocity / drag heatmaps

Author intent:
Practical, high-resolution flow estimates with calibration support.
"""

import json
import os

import matplotlib.pyplot as plt
import numpy as np
import trimesh

try:
    from numba import njit, prange
    NUMBA_AVAILABLE = True
except ImportError:  # pragma: no cover - optional dependency
    NUMBA_AVAILABLE = False
    njit = None
    prange = range

EPS = 1e-8
DEFAULT_FLOW_DIRECTION = np.array([0.0, 0.0, 1.0])


# ==========================================================
# DATA STRUCTURES
# ==========================================================

class SliceBuffer:
    def __init__(self, index, z0, z1, r1, r2):
        self.index = index
        self.z0 = z0
        self.z1 = z1

        self.surface_normals = np.zeros((r1, r2, 3))
        self.curvature = np.zeros((r1, r2))

        self.pressure = np.zeros((r1, r2))
        self.velocity = np.zeros((r1, r2))
        self.drag = np.zeros((r1, r2))


# ==========================================================
# NUMBA ACCELERATION
# ==========================================================

if NUMBA_AVAILABLE:
    @njit(parallel=True, fastmath=True)
    def process_slice_numba(
        surface_normals,
        curvature,
        directions,
        energy,
        angular_velocity,
        angular_inertia,
        boundary_layer,
        surface_adherence,
        pressure,
        velocity,
        drag,
        flow_direction,
    ):
        r1, r2 = curvature.shape
        for i in prange(r1):
            for j in range(r2):
                normal = surface_normals[i, j]
                direction = directions[i, j]

                norm_a = np.sqrt(direction[0] ** 2 + direction[1] ** 2 + direction[2] ** 2) + EPS
                norm_b = np.sqrt(normal[0] ** 2 + normal[1] ** 2 + normal[2] ** 2) + EPS
                dot = (
                    (direction[0] / norm_a) * (normal[0] / norm_b)
                    + (direction[1] / norm_a) * (normal[1] / norm_b)
                    + (direction[2] / norm_a) * (normal[2] / norm_b)
                )
                if dot > 1.0:
                    dot = 1.0
                elif dot < -1.0:
                    dot = -1.0
                angle = np.arccos(dot)

                angular_velocity[i, j] += angle
                angular_inertia[i, j] *= np.exp(-angle)

                loss = angle ** 2
                energy[i, j] *= np.exp(-loss)

                boundary_layer[i, j] += curvature[i, j] * 0.01
                adherence = 1.0 - angle
                if adherence < 0.0:
                    adherence = 0.0
                surface_adherence[i, j] += adherence

                inertia = angular_inertia[i, j]
                direction = direction * inertia + normal * (1.0 - inertia)
                norm_dir = np.sqrt(direction[0] ** 2 + direction[1] ** 2 + direction[2] ** 2) + EPS
                directions[i, j] = direction / norm_dir

                pressure[i, j] = loss
                velocity[i, j] = energy[i, j]
                flow_alignment = -(
                    normal[0] * flow_direction[0]
                    + normal[1] * flow_direction[1]
                    + normal[2] * flow_direction[2]
                )
                if flow_alignment < 0.0:
                    flow_alignment = 0.0
                drag[i, j] = loss * flow_alignment


# ==========================================================
# GEOMETRY HELPERS
# ==========================================================

def angle_between(a, b):
    a = a / (np.linalg.norm(a) + EPS)
    b = b / (np.linalg.norm(b) + EPS)
    return np.arccos(np.clip(np.dot(a, b), -1.0, 1.0))


def sample_surface_normal(mesh, point):
    """
    Finds closest triangle and returns its normal.
    """
    result = mesh.nearest.on_surface([point])
    if len(result) == 3:
        _, _, face_index = result
    else:
        _, face_index = result
    return mesh.face_normals[face_index[0]]


def normalize_flow_direction(cfg):
    direction = np.array(cfg.get("flow_direction", DEFAULT_FLOW_DIRECTION), dtype=float)
    if direction.shape != (3,):
        raise ValueError("flow_direction must be a 3-element list or array.")
    norm = np.linalg.norm(direction)
    if norm <= EPS:
        raise ValueError("flow_direction must be a non-zero vector.")
    return direction / norm


def compute_sampling_frame(mesh):
    xmin, ymin = mesh.bounds[0][:2]
    xmax, ymax = mesh.bounds[1][:2]
    center = np.array([(xmin + xmax) * 0.5, (ymin + ymax) * 0.5])
    return center, np.array([xmin, ymin]), np.array([xmax, ymax])


def max_radius_for_angle(center, bounds_min, bounds_max, direction):
    dx, dy = direction
    if abs(dx) < EPS:
        tx = np.inf
    elif dx > 0:
        tx = (bounds_max[0] - center[0]) / dx
    else:
        tx = (bounds_min[0] - center[0]) / dx

    if abs(dy) < EPS:
        ty = np.inf
    elif dy > 0:
        ty = (bounds_max[1] - center[1]) / dy
    else:
        ty = (bounds_min[1] - center[1]) / dy

    return min(tx, ty)


def compute_sector_radii(mesh, r2):
    center, bounds_min, bounds_max = compute_sampling_frame(mesh)
    radii = np.zeros(r2, dtype=float)
    for j in range(r2):
        angle = (2.0 * np.pi * j) / r2
        direction = np.array([np.cos(angle), np.sin(angle)])
        radii[j] = max_radius_for_angle(center, bounds_min, bounds_max, direction)
    return center, radii


def compute_cell_area(mesh, r1, r2):
    if r1 < 2 or r2 < 2:
        raise ValueError("r1_resolution and r2_resolution must be at least 2.")
    _, radii = compute_sector_radii(mesh, r2)
    dtheta = (2.0 * np.pi) / r2
    cell_area = np.zeros((r1, r2), dtype=float)
    for j in range(r2):
        for i in range(r1):
            r_outer = radii[j] * (i / (r1 - 1))
            r_inner = radii[j] * ((i - 1) / (r1 - 1)) if i > 0 else 0.0
            cell_area[i, j] = 0.5 * (r_outer ** 2 - r_inner ** 2) * dtheta
    return cell_area


def unit_scale_to_meters(units):
    unit_map = {
        "m": 1.0,
        "meter": 1.0,
        "meters": 1.0,
        "cm": 0.01,
        "centimeter": 0.01,
        "centimeters": 0.01,
        "mm": 0.001,
        "millimeter": 0.001,
        "millimeters": 0.001,
        "in": 0.0254,
        "inch": 0.0254,
        "inches": 0.0254,
        "ft": 0.3048,
        "foot": 0.3048,
        "feet": 0.3048,
    }
    if units is None:
        raise ValueError("units must be provided to compute forces in SI.")
    scale = unit_map.get(units.lower())
    if scale is None:
        raise ValueError(f"Unsupported units: {units}. Use m, cm, mm, in, or ft.")
    return scale


def compute_reference_area(mesh, flow_direction, scale):
    extents = (mesh.bounds[1] - mesh.bounds[0]) * scale
    dx, dy, dz = np.abs(extents)
    flow = np.abs(flow_direction)
    return flow[0] * dy * dz + flow[1] * dx * dz + flow[2] * dx * dy


def safe_ratio(numerator, denominator):
    if abs(denominator) <= EPS:
        return None
    return numerator / denominator


def load_calibration(cfg):
    calibration_path = cfg.get("calibration_path")
    if calibration_path is None:
        calibration_path = os.path.join(cfg["output_path"], "calibration.json")
    if not os.path.exists(calibration_path):
        return {
            "path": calibration_path,
            "drag_scale": 1.0,
            "downforce_scale": 1.0,
        }
    with open(calibration_path, "r") as f:
        calibration = json.load(f)

    drag_scales = []
    downforce_scales = []
    for entry in calibration.get("entries", []):
        drag_ratio = safe_ratio(
            entry.get("expected_drag_coefficient"),
            entry.get("computed_drag_coefficient"),
        )
        if drag_ratio is not None:
            drag_scales.append(drag_ratio)
        downforce_ratio = safe_ratio(
            entry.get("expected_downforce"),
            entry.get("computed_downforce"),
        )
        if downforce_ratio is not None:
            downforce_scales.append(downforce_ratio)

    drag_scale = float(np.mean(drag_scales)) if drag_scales else 1.0
    downforce_scale = float(np.mean(downforce_scales)) if downforce_scales else 1.0
    return {
        "path": calibration_path,
        "drag_scale": drag_scale,
        "downforce_scale": downforce_scale,
    }


def append_calibration_entry(path, entry):
    payload = {"version": 1, "entries": []}
    if os.path.exists(path):
        with open(path, "r") as f:
            payload = json.load(f)
    payload.setdefault("version", 1)
    payload.setdefault("entries", [])
    payload["entries"].append(entry)
    with open(path, "w") as f:
        json.dump(payload, f, indent=2)
        f.write("\n")


# ==========================================================
# STAGE 1: STL IMPORT + SLICING
# ==========================================================

def load_mesh(path):
    mesh = trimesh.load(path, force="mesh")
    mesh.remove_unreferenced_vertices()
    return mesh


def generate_slices(mesh, cfg):
    zmin, zmax = mesh.bounds[:, 2]
    slices = []

    z = zmin
    index = 0

    while z < zmax:
        slices.append(
            SliceBuffer(
                index,
                z,
                z + cfg["slice_depth"],
                cfg["r1_resolution"],
                cfg["r2_resolution"],
            )
        )
        z += cfg["slice_depth"]
        index += 1

    return slices


def validate_config(cfg):
    required = [
        "stl_path",
        "slice_depth",
        "r1_resolution",
        "r2_resolution",
        "output_path",
        "airspeed",
        "air_density",
        "units",
    ]
    missing = [key for key in required if key not in cfg]
    if missing:
        raise KeyError(f"Missing required config keys: {', '.join(missing)}")
    if cfg["slice_depth"] <= 0:
        raise ValueError("slice_depth must be positive.")
    if cfg["r1_resolution"] < 2 or cfg["r2_resolution"] < 2:
        raise ValueError("r1_resolution and r2_resolution must be at least 2.")
    if cfg["airspeed"] <= 0:
        raise ValueError("airspeed must be positive.")
    if cfg["air_density"] <= 0:
        raise ValueError("air_density must be positive.")


def validate_mesh(mesh, stl_path):
    if mesh.is_empty:
        raise ValueError(f"Mesh loaded from {stl_path} is empty.")
    if mesh.faces.size == 0 or mesh.vertices.size == 0:
        raise ValueError(f"Mesh loaded from {stl_path} has no faces or vertices.")


def normalize_slice_geometry(mesh, slice_buf, flow_direction):
    """
    Converts STL geometry into:
    - surface normal field
    - curvature (angular deviation from flow)
    """

    r1, r2, _ = slice_buf.surface_normals.shape
    center, radii = compute_sector_radii(mesh, r2)

    for i in range(r1):
        for j in range(r2):
            radial = radii[j] * (i / (r1 - 1))
            angle = (2.0 * np.pi * j) / r2
            x = center[0] + radial * np.cos(angle)
            y = center[1] + radial * np.sin(angle)
            z = (slice_buf.z0 + slice_buf.z1) * 0.5

            point = np.array([x, y, z])
            normal = sample_surface_normal(mesh, point)

            slice_buf.surface_normals[i, j] = normal
            slice_buf.curvature[i, j] = angle_between(normal, flow_direction)


# ==========================================================
# STAGE 2: 4D RAY PROPAGATION
# ==========================================================

def initialize_ray_field(r1, r2, flow_direction):
    directions = np.empty((r1, r2, 3), dtype=float)
    directions[:, :, 0] = flow_direction[0]
    directions[:, :, 1] = flow_direction[1]
    directions[:, :, 2] = flow_direction[2]
    return {
        "directions": directions,
        "energy": np.ones((r1, r2), dtype=float),
        "angular_velocity": np.zeros((r1, r2), dtype=float),
        "angular_inertia": np.ones((r1, r2), dtype=float),
        "boundary_layer": np.zeros((r1, r2), dtype=float),
        "surface_adherence": np.zeros((r1, r2), dtype=float),
    }


def process_slice_python(slice_buf, rays, flow_direction):
    r1, r2 = slice_buf.curvature.shape

    directions = rays["directions"]
    energy = rays["energy"]
    angular_velocity = rays["angular_velocity"]
    angular_inertia = rays["angular_inertia"]
    boundary_layer = rays["boundary_layer"]
    surface_adherence = rays["surface_adherence"]

    for i in range(r1):
        for j in range(r2):
            normal = slice_buf.surface_normals[i, j]
            curvature = slice_buf.curvature[i, j]

            angle = angle_between(directions[i, j], normal)

            angular_velocity[i, j] += angle
            angular_inertia[i, j] *= np.exp(-angle)

            loss = angle**2
            energy[i, j] *= np.exp(-loss)

            boundary_layer[i, j] += curvature * 0.01
            surface_adherence[i, j] += max(0.0, 1.0 - angle)

            inertia = angular_inertia[i, j]
            directions[i, j] = directions[i, j] * inertia + normal * (1.0 - inertia)
            directions[i, j] /= np.linalg.norm(directions[i, j]) + EPS

            slice_buf.pressure[i, j] = loss
            slice_buf.velocity[i, j] = energy[i, j]
            flow_alignment = -float(np.dot(normal, flow_direction))
            if flow_alignment < 0.0:
                flow_alignment = 0.0
            slice_buf.drag[i, j] = loss * flow_alignment


def process_slice(slice_buf, rays, flow_direction, use_numba=True):
    if NUMBA_AVAILABLE and use_numba:
        process_slice_numba(
            slice_buf.surface_normals,
            slice_buf.curvature,
            rays["directions"],
            rays["energy"],
            rays["angular_velocity"],
            rays["angular_inertia"],
            rays["boundary_layer"],
            rays["surface_adherence"],
            slice_buf.pressure,
            slice_buf.velocity,
            slice_buf.drag,
            flow_direction,
        )
    else:
        process_slice_python(slice_buf, rays, flow_direction)


# ==========================================================
# STAGE 3: OUTPUT + VISUALIZATION
# ==========================================================

def save_matrix(mat, path):
    np.savetxt(path, mat, fmt="%.6e")


def plot_heatmap(mat, title, path):
    plt.figure(figsize=(6, 5))
    plt.imshow(mat, origin="lower", cmap="inferno")
    plt.colorbar()
    plt.title(title)
    plt.tight_layout()
    plt.savefig(path)
    plt.close()


# ==========================================================
# MAIN PIPELINE
# ==========================================================

def main():
    with open("config.json", "r") as f:
        cfg = json.load(f)

    validate_config(cfg)
    flow_direction = normalize_flow_direction(cfg)
    os.makedirs(cfg["output_path"], exist_ok=True)

    mesh = load_mesh(cfg["stl_path"])
    validate_mesh(mesh, cfg["stl_path"])
    slices = generate_slices(mesh, cfg)
    cell_area = compute_cell_area(mesh, cfg["r1_resolution"], cfg["r2_resolution"])
    units = cfg["units"]
    scale = unit_scale_to_meters(units)
    cell_area_m2 = cell_area * scale * scale
    cell_area_mean = float(np.mean(cell_area))
    cell_area_m2_mean = float(np.mean(cell_area_m2))
    density = cfg["air_density"]
    airspeed = cfg["airspeed"]
    ray_weight = cfg.get("ray_weight", 1.0)
    dynamic_pressure = 0.5 * density * airspeed * airspeed
    reference_area = compute_reference_area(mesh, flow_direction, scale)

    ray_field = initialize_ray_field(
        cfg["r1_resolution"],
        cfg["r2_resolution"],
        flow_direction,
    )
    use_numba = cfg.get("numba_parallel", True)

    total_force = np.zeros(3)
    for s in slices:
        print(f"Processing slice {s.index}")
        normalize_slice_geometry(mesh, s, flow_direction)
        process_slice(s, ray_field, flow_direction, use_numba=use_numba)

        base = f"{cfg['output_path']}/slice_{s.index}"

        save_matrix(s.pressure, base + "_pressure.txt")
        save_matrix(s.velocity, base + "_velocity.txt")
        save_matrix(s.drag, base + "_drag.txt")

        plot_heatmap(s.pressure, f"Pressure Slice {s.index}", base + "_pressure.png")
        plot_heatmap(s.velocity, f"Velocity Slice {s.index}", base + "_velocity.png")
        plot_heatmap(s.drag, f"Drag Slice {s.index}", base + "_drag.png")

    total_drag = sum(np.sum(s.drag) for s in slices)
    total_pressure_metric = sum(np.sum(s.pressure) for s in slices)
    for s in slices:
        pressure_force = s.pressure * dynamic_pressure * cell_area_m2 * ray_weight
        total_force += np.sum(pressure_force[..., None] * s.surface_normals, axis=(0, 1))

    drag_force = float(-np.dot(total_force, flow_direction))
    world_up = np.array([0.0, 1.0, 0.0])
    if abs(np.dot(world_up, flow_direction)) > 1.0 - 1e-3:
        world_up = np.array([1.0, 0.0, 0.0])
    lift_axis = world_up - np.dot(world_up, flow_direction) * flow_direction
    lift_axis_norm = np.linalg.norm(lift_axis)
    if lift_axis_norm > EPS:
        lift_axis /= lift_axis_norm
    else:
        lift_axis = np.array([0.0, 1.0, 0.0])
    lift_force = float(np.dot(total_force, lift_axis))
    downforce = float(-lift_force) if lift_force < 0 else 0.0

    coefficient_denominator = dynamic_pressure * reference_area
    calibration = load_calibration(cfg)
    calibrated_drag_force = drag_force * calibration["drag_scale"]
    calibration_scale = total_drag / (total_pressure_metric + EPS)
    calibrated_drag_force = drag_force * calibration_scale
    calibrated_drag_coefficient = (
        calibrated_drag_force / coefficient_denominator if coefficient_denominator > EPS else 0.0
    )
    if coefficient_denominator > EPS:
        drag_coefficient = drag_force / coefficient_denominator
        lift_coefficient = lift_force / coefficient_denominator
    else:
        drag_coefficient = 0.0
        lift_coefficient = 0.0
    calibrated_drag_coefficient = drag_coefficient * calibration["drag_scale"]
    calibrated_downforce = downforce * calibration["downforce_scale"]

    summary_payload = {
        "total_estimated_drag_metric": float(total_drag),
        "total_pressure_metric": float(total_pressure_metric),
        "calibration_drag_scale": float(calibration["drag_scale"]),
        "calibration_downforce_scale": float(calibration["downforce_scale"]),
        "calibration_scale": float(calibration_scale),
        "total_force_vector": total_force.tolist(),
        "drag_force": drag_force,
        "drag_force_proxy": drag_force,
        "calibrated_drag_force_estimate": float(calibrated_drag_force),
        "lift_force": lift_force,
        "lift_force_proxy": lift_force,
        "downforce": downforce,
        "downforce_proxy": downforce,
        "calibrated_downforce_estimate": float(calibrated_downforce),
        "cell_area": cell_area_mean,
        "cell_area_m2": cell_area_m2_mean,
        "density": float(density),
        "airspeed": float(airspeed),
        "dynamic_pressure": float(dynamic_pressure),
        "reference_area": float(reference_area),
        "drag_coefficient": float(drag_coefficient),
        "calibrated_drag_coefficient_estimate": float(calibrated_drag_coefficient),
        "lift_coefficient": float(lift_coefficient),
        "ray_weight": float(ray_weight),
        "slice_count": len(slices),
        "slice_depth": cfg["slice_depth"],
        "r1_resolution": cfg["r1_resolution"],
        "r2_resolution": cfg["r2_resolution"],
        "numba_parallel": bool(use_numba and NUMBA_AVAILABLE),
        "units": cfg.get("units"),
        "stl_path": cfg["stl_path"],
        "flow_direction": flow_direction.tolist(),
    }
    with open(cfg["output_path"] + "/summary.txt", "w") as f:
        f.write(json.dumps(summary_payload, indent=2))
        f.write("\n")

    append_calibration_entry(
        calibration["path"],
        {
            "label": cfg.get("model_label", os.path.basename(cfg["stl_path"])),
            "stl_path": cfg["stl_path"],
            "computed_drag_coefficient": float(drag_coefficient),
            "computed_downforce": float(downforce),
            "expected_drag_coefficient": None,
            "expected_downforce": None,
        },
    )

    print("Estimation complete.")
    print("Total drag metric:", total_drag)
    print("Drag force (N):", drag_force)
    print("Calibrated drag force (N):", calibrated_drag_force)
    print("Lift force (N):", lift_force)
    print("Downforce (N):", downforce)
    print("Calibrated downforce (N):", calibrated_downforce)
    print("Drag coefficient:", drag_coefficient)
    print("Calibrated drag coefficient:", calibrated_drag_coefficient)
    print("Lift coefficient:", lift_coefficient)


if __name__ == "__main__":
    main()
