"""
Quick CFD Estimator (Estimator-Grade, Not a Solver)

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:
Fast iteration, intuition building, comparative design work.
"""

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

EPS = 1e-8
DEFAULT_FLOW_DIRECTION = np.array([0.0, 0.0, 1.0])
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))


class RayState:
    """
    Persistent ray (4D: x, y, z, slice index)
    """

    def __init__(self, flow_direction):
        self.direction = flow_direction.copy()
    def __init__(self, flow_direction):
        self.direction = flow_direction.copy()
        self.energy = 1.0

        self.angular_velocity = 0.0
        self.angular_inertia = 1.0
        self.boundary_layer = 0.0
        self.surface_adherence = 0.0


# ==========================================================
# 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_cell_area(mesh, r1, r2):
    xmin, ymin = mesh.bounds[0][:2]
    xmax, ymax = mesh.bounds[1][:2]
    if r1 < 2 or r2 < 2:
        raise ValueError("r1_resolution and r2_resolution must be at least 2.")
    dx = (xmax - xmin) / (r1 - 1)
    dy = (ymax - ymin) / (r2 - 1)
    return dx * dy


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 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_cell_area(mesh, r1, r2):
    xmin, ymin = mesh.bounds[0][:2]
    xmax, ymax = mesh.bounds[1][:2]
    if r1 < 2 or r2 < 2:
        raise ValueError("r1_resolution and r2_resolution must be at least 2.")
    dx = (xmax - xmin) / (r1 - 1)
    dy = (ymax - ymin) / (r2 - 1)
    return dx * dy


# ==========================================================
# 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):
        z += cfg["slice_depth"]
        index += 1
        return slices


def validate_config(cfg):
    required = ["stl_path", "slice_depth", "r1_resolution", "r2_resolution", "output_path"]
    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.")


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
    xmin, ymin = mesh.bounds[0][:2]
    xmax, ymax = mesh.bounds[1][:2]

    for i in range(r1):
        for j in range(r2):
            x = xmin + (xmax - xmin) * (i / (r1 - 1))
            y = ymin + (ymax - ymin) * (j / (r2 - 1))
            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)
            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):
    return [[RayState(flow_direction) for _ in range(r2)] for _ in range(r1)]
def initialize_ray_field(r1, r2, flow_direction):
    return [[RayState(flow_direction) for _ in range(r2)] for _ in range(r1)]


def process_ray(ray, normal, curvature, cfg):
    angle = angle_between(ray.direction, normal)

    ray.angular_velocity += angle
    ray.angular_inertia *= np.exp(-angle)

    loss = angle**2
    ray.energy *= np.exp(-loss)

    ray.boundary_layer += curvature * 0.01
    ray.surface_adherence += max(0.0, 1.0 - angle)

    # Direction relaxation toward surface
    ray.direction = (
        ray.direction * ray.angular_inertia +
        normal * (1.0 - ray.angular_inertia)
    )
    ray.direction /= np.linalg.norm(ray.direction) + EPS

    return loss


def process_slice(slice_buf, rays, cfg):
    r1, r2 = slice_buf.curvature.shape

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

            loss = process_ray(ray, normal, curvature, cfg)

            slice_buf.pressure[i, j] = loss
            slice_buf.velocity[i, j] = ray.energy
            slice_buf.drag[i, j] = loss * curvature


# ==========================================================
# 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
    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
    )

    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, cfg)

        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)
    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
    drag_coefficient = drag_force / (dynamic_pressure * reference_area)
    lift_coefficient = lift_force / (dynamic_pressure * reference_area)

    summary_payload = {
        "total_estimated_drag_metric": float(total_drag),
        "total_force_vector": total_force.tolist(),
        "drag_force_proxy": drag_force,
        "lift_force_proxy": lift_force,
        "downforce_proxy": downforce,
        "cell_area": float(cell_area),
        "cell_area_m2": float(cell_area_m2),
        "density": float(density),
        "airspeed": float(airspeed),
        "dynamic_pressure": float(dynamic_pressure),
        "reference_area": float(reference_area),
        "drag_coefficient": float(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"],
        "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")

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


if __name__ == "__main__":
    main()
