"""RocketPy Study 1 + 2: execute with Python; only trajectory figures are saved.

Run: python simulate.py --output results-rerun
Inputs are vendored from RocketPy v1.13.0; no network is used during execution.
"""

import argparse
import copy
import csv
import hashlib
import json
import os
import platform
import tempfile
import warnings
from contextlib import redirect_stdout
from datetime import datetime, timezone
from importlib.metadata import version
from pathlib import Path

os.environ.setdefault("MPLCONFIGDIR", str(Path(tempfile.gettempdir()) / "rp12-matplotlib"))
import matplotlib

matplotlib.use("Agg")
import matplotlib.pyplot as plt
from matplotlib.ticker import MaxNLocator
import numpy as np
from rocketpy import Environment, Flight, Rocket, SolidMotor
from rocketpy.simulation import FlightDataExporter
from rocketpy.utilities import apogee_by_mass, liftoff_speed_by_mass

BASE = Path(__file__).resolve().parent
SEED = 20261004
STATE_FIELDS = "x y z vx vy vz e0 e1 e2 e3 w1 w2 w3".split()
# Every time-domain quantity behind the Study 2 plot families is evaluated.
FIELDS = STATE_FIELDS + "speed ax ay az acceleration path_angle attitude_angle lateral_attitude_angle psi theta phi alpha1 alpha2 alpha3 R1 R2 R3 M1 M2 M3 aerodynamic_lift aerodynamic_drag aerodynamic_bending_moment aerodynamic_spin_moment kinetic_energy rotational_energy translational_energy potential_energy total_energy thrust_power drag_power mach_number reynolds_number pressure dynamic_pressure total_pressure angle_of_attack partial_angle_of_attack angle_of_sideslip stability_margin".split()


def make_environment(comparison=False):
    env = Environment(latitude=35.5721, longitude=129.1822, elevation=53,
                      date=(2026, 10, 4, 12), timezone="Asia/Seoul", datum="WGS84")
    # These are assumptions, not an observed UNIST weather profile.
    env.set_atmospheric_model(type="custom_atmosphere", temperature=293.15,
                              wind_u=0 if comparison else 2.83,
                              wind_v=-5 if comparison else -2.83)
    return env


def drogue_trigger(p, h, y):
    return y[5] < 0


def main_trigger(p, h, y):
    return y[5] < 0 and h < 800


def make_rocket(fin_factor=None, callable_triggers=False):
    np.random.seed(SEED)
    motor = SolidMotor(
        thrust_source=str(BASE / "inputs/Cesaroni_M1670.eng"),
        dry_mass=1.815, dry_inertia=(0.125, 0.125, 0.002),
        nozzle_radius=0.033, grain_number=5, grain_density=1815,
        grain_outer_radius=0.033, grain_initial_inner_radius=0.015,
        grain_initial_height=0.120, grain_separation=0.005,
        grains_center_of_mass_position=0.397, center_of_dry_mass_position=0.317,
        nozzle_position=0, burn_time=3.9, throat_radius=0.011,
        coordinate_system_orientation="nozzle_to_combustion_chamber",
    )
    rocket = Rocket(
        radius=0.0635, mass=14.426, inertia=(6.321, 6.321, 0.034),
        power_off_drag=str(BASE / "inputs/powerOffDragCurve.csv"),
        power_on_drag=str(BASE / "inputs/powerOnDragCurve.csv"),
        center_of_mass_without_motor=0, coordinate_system_orientation="tail_to_nose",
    )
    rocket.add_motor(motor, position=-1.255)
    rocket.add_nose(length=0.55829, kind="von karman", position=1.278)
    if fin_factor is None:
        rocket.add_trapezoidal_fins(
            n=4, root_chord=0.120, tip_chord=0.060, span=0.110,
            position=-1.04956, cant_angle=0.5,
            airfoil=(str(BASE / "inputs/NACA0012-radians.txt"), "radians"),
        )
    else:
        # Rebuild each configuration so no tail is removed and no fins accumulate.
        rocket.add_trapezoidal_fins(n=4, root_chord=0.120, tip_chord=0.040,
                                    span=0.100, position=-1.04956 * fin_factor)
    rocket.add_tail(top_radius=0.0635, bottom_radius=0.0435,
                    length=0.060, position=-1.194656)
    rocket.set_rail_buttons(upper_button_position=0.0818,
                            lower_button_position=-0.618, angular_position=45)
    for name, cd_s, trigger in [
        ("Main", 10.0, main_trigger if callable_triggers else 800),
        ("Drogue", 1.0, drogue_trigger if callable_triggers else "apogee"),
    ]:
        rocket.add_parachute(name=name, cd_s=cd_s, trigger=trigger,
                            sampling_rate=105, lag=1.5, noise=(0, 8.3, 0.5),
                            radius=1.5, height=1.5, porosity=0.0432)
    assert len(rocket.aerodynamic_surfaces) == 3
    return rocket


def fly(rocket, env, **kwargs):
    return Flight(rocket=rocket, environment=env, rail_length=5.2,
                  heading=0, verbose=False, **kwargs)


def write_rows(path, rows):
    with path.open("w", encoding="utf-8", newline="") as f:
        writer = csv.DictWriter(f, fieldnames=list(rows[0]))
        writer.writeheader()
        writer.writerows(rows)


def write_json(path, value):
    path.write_text(json.dumps(value, ensure_ascii=False, indent=2, allow_nan=False) + "\n")


def event(flight, t):
    return {
        "time_s": float(t), "agl_m": float(flight.z(t) - flight.env.elevation),
        "x_east_m": float(flight.x(t)), "y_north_m": float(flight.y(t)),
        "speed_m_s": float(flight.speed(t)), "vz_m_s": float(flight.vz(t)),
        "mach": float(flight.mach_number(t)),
        "static_margin_c": float(flight.rocket.static_margin(t)),
        "stability_margin_c": float(flight.stability_margin(t)),
        "aoa_deg": float(flight.angle_of_attack(t)),
        "attitude_deg": float(flight.attitude_angle(t)),
        "path_angle_deg": float(flight.path_angle(t)),
    }


def run(out):
    manifest = json.loads((BASE / "input-manifest.json").read_text())
    for entry in manifest["files"]:
        assert hashlib.sha256((BASE / entry["file"]).read_bytes()).hexdigest() == entry["sha256"]
    env, rocket = make_environment(), make_rocket()
    print("Running baseline flight...", flush=True)
    flight = fly(rocket, env, inclination=85, max_time_step=0.05, max_time=600)
    assert 0 < flight.out_of_rail_time < 3.9 < flight.apogee_time < flight.t_final < 600
    assert abs(flight.z(flight.t_final) - env.elevation) < 0.01
    assert len(flight.parachute_events) == 2
    assert np.isfinite(flight.solution_array).all()
    flight_prints = ["initial_conditions", "surface_wind_conditions", "launch_rail_conditions",
                     "out_of_rail_conditions", "burn_out_conditions", "apogee_conditions",
                     "events_registered", "impact_conditions", "maximum_values"]
    with (out / "model_info.txt").open("w") as f, redirect_stdout(f):
        env.prints.all()
        rocket.motor.prints.all()
        rocket.prints.all()
    with (out / "flight_prints.txt").open("w") as f, redirect_stdout(f):
        for method in flight_prints:
            getattr(flight.prints, method)()

    write_rows(out / "environment_profile.csv", [
        {"altitude_asl_m": float(z), "altitude_agl_m": float(z - env.elevation),
         **{name: float(getattr(env, name)(z)) for name in
            ["pressure", "temperature", "density", "speed_of_sound", "wind_velocity_x", "wind_velocity_y"]}}
        for z in np.linspace(env.elevation, env.elevation + 4000, 81)
    ])
    motor_fields = "thrust total_mass propellant_mass mass_flow_rate exhaust_velocity center_of_mass I_11 I_22 I_33 grain_inner_radius grain_height burn_area burn_rate".split()
    motor_rows = [{"time_s": float(t), **{name: float(getattr(rocket.motor, name)(t))
                                         for name in motor_fields}}
                  for t in np.linspace(0, 3.9, 391)]
    assert np.isfinite([[row[k] for k in row] for row in motor_rows]).all()
    write_rows(out / "motor_profile.csv", motor_rows)
    np.savetxt(out / "motor_Kn.csv", rocket.motor.Kn.source, delimiter=",",
               header=",".join(rocket.motor.Kn.__inputs__ + rocket.motor.Kn.__outputs__), comments="")
    write_rows(out / "drag_curves.csv", [
        {"mach": float(m), "power_on_Cd": float(rocket.power_on_drag_by_mach(m)),
         "power_off_Cd": float(rocket.power_off_drag_by_mach(m))} for m in np.linspace(0, 2, 201)
    ])

    events = [{"event": name, **event(flight, t)} for name, t in [
        ("ignition", 0), ("rail_exit", flight.out_of_rail_time),
        ("burnout", 3.9), ("apogee", flight.apogee_time), ("impact", flight.t_final),
    ]]
    parachutes = []
    for t, parachute in flight.parachute_events:
        parachutes.append({"name": parachute.name, "trigger_s": float(t),
                           "inflation_s": float(t + parachute.lag),
                           "trigger_agl_m": float(flight.z(t) - env.elevation),
                           "inflation_agl_m": float(flight.z(t + parachute.lag) - env.elevation)})
    write_rows(out / "events.csv", events)
    write_rows(out / "parachutes.csv", parachutes)

    # Evaluate each physical quantity on the solver's actual time grid; retain units.
    series = {}
    extrema = []
    for name in FIELDS:
        func = getattr(flight, name)
        values = np.asarray(func(flight.time), dtype=float)
        assert np.isfinite(values).all(), name
        series[name] = values
        for period, mask in [
            ("full_flight", np.ones(len(flight.time), dtype=bool)),
            ("rail_exit_to_burnout", (flight.time >= flight.out_of_rail_time) & (flight.time <= 3.9)),
            ("burnout_to_apogee", (flight.time >= 3.9) & (flight.time <= flight.apogee_time)),
        ]:
            times, data = flight.time[mask], values[mask]
            low, high = int(np.argmin(data)), int(np.argmax(data))
            extrema.append({"quantity": name, "label": func.__outputs__[0], "period": period,
                            "min": float(data[low]), "min_time_s": float(times[low]),
                            "max": float(data[high]), "max_time_s": float(times[high])})
    write_rows(out / "extrema.csv", extrema)
    exporter = FlightDataExporter(flight)
    # Explicit fields avoid the missing velocity labels in 1.13.0's no-argument CSV header.
    exporter.export_data(str(out / "full_state.csv"), *STATE_FIELDS)
    exporter.export_data(str(out / "selected.csv"), "angle_of_attack", "mach_number")
    exporter.export_data(str(out / "selected_1s.csv"), "angle_of_attack", "mach_number", time_step=1.0)
    exporter.export_data(str(out / "diagnostics_0.1s.csv"), *FIELDS, time_step=0.1)
    exporter.export_kml(file_name=str(out / "trajectory.kml"), extrude=True,
                        altitude_mode="relativetoground")
    np.savetxt(out / "speed_source.csv", flight.speed.source, delimiter=",",
               header="time_s,speed_m_s", comments="")
    with (out / "full_state.csv").open() as f:
        assert len(next(csv.reader(f))) == 14
    assert np.allclose(np.loadtxt(out / "full_state.csv", delimiter=",", skiprows=1),
                       flight.solution_array, atol=5.1e-7, rtol=0)
    times = np.linspace(0, 3.9, 391)
    static_rows = [{"time_s": float(t), "mass_kg": float(rocket.total_mass(t)),
                    "cg_m": float(rocket.center_of_mass(t)),
                    "cp_m_at_mach0": float(rocket.cp_position(0)),
                    "static_margin_c": float(rocket.static_margin(t))} for t in times]
    write_rows(out / "static_margin.csv", static_rows)
    inertia = {str(t): {"tensor_kg_m2": [list(row) for row in rocket.get_inertia_tensor_at_time(t)],
                        "derivative_kg_m2_s": [list(row) for row in rocket.get_inertia_tensor_derivative_at_time(t)]}
               for t in [0, 0.5, 3.9]}
    write_json(out / "inertia.json", inertia)
    rail = {}
    for name in ["rail_button1_normal_force", "rail_button1_shear_force",
                 "rail_button2_normal_force", "rail_button2_shear_force"]:
        source = np.asarray(getattr(flight, name).source)
        np.savetxt(out / f"{name}.csv", source, delimiter=",", header="time_s,force_N", comments="")
        index = int(np.argmax(abs(source[:, 1])))
        rail[name] = {"max_abs_N": float(abs(source[index, 1])), "time_s": float(source[index, 0])}

    frequency_rows, frequency_peaks = [], {}
    for name in ["attitude_frequency_response", "omega1_frequency_response",
                 "omega2_frequency_response", "omega3_frequency_response"]:
        source = np.asarray(getattr(flight, name).source)
        positive = source[source[:, 0] > 0]
        peak = positive[np.argmax(positive[:, 1])]
        frequency_peaks[name] = {"frequency_Hz": float(peak[0]), "amplitude": float(peak[1])}
        frequency_rows.extend({"quantity": name, "frequency_Hz": float(x), "amplitude": float(y)}
                              for x, y in positive)
    write_rows(out / "frequency_response.csv", frequency_rows)

    print("Running mass sweeps (10 + 10 flights)...", flush=True)
    mass_before = rocket.mass
    analysis_flight = copy.deepcopy(flight)
    # Both helpers stop at apogee. Recovery is inactive in that interval; removing
    # its callbacks avoids reusing pressure-sensor histories across helper runs
    # in RocketPy 1.13.0. The rocket mass and inertias are unchanged by clear().
    analysis_flight.rocket.parachutes.clear()
    apogee = apogee_by_mass(copy.deepcopy(analysis_flight), 5, 20, points=10, plot=False)
    liftoff = liftoff_speed_by_mass(copy.deepcopy(analysis_flight), 5, 20, points=10, plot=False)
    assert rocket.mass == mass_before
    assert len(rocket.parachutes) == 2
    assert np.allclose(apogee.source[:, 0], liftoff.source[:, 0])
    mass_rows = [{"mass_without_motor_kg": float(m), "apogee_agl_m": float(a),
                  "rail_exit_speed_m_s": float(v)}
                 for (m, a), (_, v) in zip(apogee.source, liftoff.source)]
    write_rows(out / "mass_sweep.csv", mass_rows)

    print("Running five fin-position cases...", flush=True)
    fin_rows, fin_samples = [], []
    for factor in [-0.5, -0.2, 0.1, 0.4, 0.7]:
        comparison_rocket = make_rocket(fin_factor=factor)
        with warnings.catch_warnings(record=True) as caught:
            warnings.simplefilter("always")
            trial = fly(comparison_rocket, make_environment(comparison=True), inclination=90,
                        max_time_step=0.01, max_time=5, terminate_on_apogee=True)
        grid = np.linspace(trial.out_of_rail_time, trial.t_final, 1001)
        row = {"factor": factor, "fin_position_m": -1.04956 * factor,
               "static_margin_ignition_c": float(comparison_rocket.static_margin(0)),
               "static_margin_rail_exit_c": float(comparison_rocket.static_margin(trial.out_of_rail_time)),
               "static_margin_end_c": float(comparison_rocket.static_margin(trial.t_final)),
               "end_time_s": float(trial.t_final),
               "attitude_at_1.5s_deg": float(trial.attitude_angle(min(1.5, trial.t_final))),
               "attitude_end_deg": float(trial.attitude_angle(trial.t_final)),
               "max_aoa_deg": float(np.max(trial.angle_of_attack(grid))),
               "max_transverse_rate_rad_s": float(np.max(np.hypot(trial.w1(grid), trial.w2(grid)))),
               "warnings": " | ".join(sorted({str(w.message) for w in caught}))}
        assert np.isfinite(trial.solution_array).all()
        fin_rows.append(row)
        for t in np.linspace(0, trial.t_final, 501):
            fin_samples.append({"factor": factor, **event(trial, t),
                                "w1_rad_s": float(trial.w1(t)), "w2_rad_s": float(trial.w2(t)),
                                "w3_rad_s": float(trial.w3(t))})
    write_rows(out / "fin_sweep.csv", fin_rows)
    write_rows(out / "fin_responses.csv", fin_samples)

    print("Checking callable triggers and integration step sensitivity...", flush=True)
    callable_flight = fly(make_rocket(callable_triggers=True), make_environment(),
                          inclination=85, max_time_step=0.05)
    builtin_events = [(float(t), p.name) for t, p in flight.parachute_events]
    callable_events = [(float(t), p.name) for t, p in callable_flight.parachute_events]
    assert builtin_events == callable_events
    assert abs(callable_flight.apogee - flight.apogee) < 1e-6
    refined = fly(make_rocket(), make_environment(), inclination=85,
                  max_time_step=0.025, terminate_on_apogee=True)
    sensitivity = {"max_step_baseline_s": 0.05, "max_step_refined_s": 0.025,
                   "apogee_difference_m": float(refined.apogee - flight.apogee),
                   "rail_exit_speed_difference_m_s": float(refined.out_of_rail_velocity - flight.out_of_rail_velocity)}
    assert abs(sensitivity["apogee_difference_m"]) < 1.0
    assert abs(sensitivity["rail_exit_speed_difference_m_s"]) < 0.05
    quaternion_error = float(np.max(abs(np.sum(flight.solution_array[:, 7:11] ** 2, axis=1) - 1)))
    assert quaternion_error < 1e-3
    assert np.linalg.eigvalsh(inertia["0.5"]["tensor_kg_m2"]).min() > 0

    # Only the trajectory is plotted; all other results are tables or raw data.
    plt.rcParams.update({"font.size": 12})
    fig = plt.figure(figsize=(7.5, 7.2), layout="constrained")
    ax = fig.add_subplot(111, projection="3d")
    ascent = flight.time <= flight.apogee_time
    for mask, label, color in [(ascent, "Ascent", "#146d9a"), (~ascent, "Descent", "#d26822")]:
        ax.plot(series["x"][mask], series["y"][mask], (series["z"][mask] - 53), color=color, lw=2, label=label)
    for name, t, marker in [("Launch", 0, "o"), ("Apogee", flight.apogee_time, "^"), ("Landing", flight.t_final, "s")]:
        x, y, z = float(flight.x(t)), float(flight.y(t)), float(flight.z(t) - 53)
        ax.scatter(x, y, z, marker=marker, s=45, color="#263947")
        ax.text(x + 40, y, z, name, fontsize=10)
    ax.set(xlabel="East (m)", ylabel="North (m)", zlabel="Altitude AGL (m)")
    ax.zaxis.labelpad = 20
    ax.view_init(elev=22, azim=-60)
    ax.set_box_aspect((1, 1, 1.35))
    ax.xaxis.set_major_locator(MaxNLocator(5))
    ax.yaxis.set_major_locator(MaxNLocator(5))
    ax.zaxis.set_major_locator(MaxNLocator(6))
    ax.legend(loc="upper left", frameon=False)
    fig.suptitle("Calisto + Cesaroni M1670\nUNIST coordinate scenario", fontsize=15)
    fig.supxlabel("Assumed NW wind 4.00 m/s; 20 C; ISA pressure\nSimulation only | Horizontal and vertical axis scales differ", fontsize=10)
    fig.savefig(out / "trajectory.png", dpi=180)
    fig.savefig(out / "trajectory.svg")
    plt.close(fig)

    summary = {
        "executed_at_utc": datetime.now(timezone.utc).isoformat(),
        "versions": {"python": platform.python_version(), **{p: version(p) for p in ["rocketpy", "numpy", "scipy", "matplotlib"]}},
        "seed": SEED, "source_commit": manifest["source_commit"],
        "environment": {"latitude": env.latitude, "longitude": env.longitude, "elevation_m": env.elevation,
                        "date_KST": "2026-10-04T12:00:00+09:00", "temperature_K": 293.15,
                        "wind_u_m_s": 2.83, "wind_v_m_s": -2.83, "pressure": "ISA",
                        "surface_pressure_Pa": float(env.pressure(53)), "surface_density_kg_m3": float(env.density(53))},
        "motor": {"total_impulse_N_s": float(rocket.motor.total_impulse),
                  "initial_propellant_mass_kg": float(rocket.motor.propellant_initial_mass),
                  "burnout_time_s": float(rocket.motor.burn_out_time)},
        "initial_mass_kg": float(rocket.total_mass(0)), "burnout_mass_kg": float(rocket.total_mass(3.9)),
        "events": events, "parachutes": parachutes, "rail_forces": rail,
        "frequency_peaks": frequency_peaks, "mass_sweep": mass_rows, "fin_sweep": fin_rows,
        "quaternion_norm_squared_max_error": quaternion_error,
        "callable_trigger_events_match": True, "solver_sensitivity": sensitivity,
        "solver_rows": len(flight.time), "time_domain_quantities": len(FIELDS),
        "print_methods_executed": flight_prints,
        "baseline_mass_preserved": rocket.mass == 14.426,
    }
    write_json(out / "summary.json", summary)
    print(json.dumps({"apogee_AGL_m": flight.apogee - env.elevation,
                      "flight_time_s": flight.t_final, "output": str(out)}, indent=2), flush=True)


if __name__ == "__main__":
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--output", type=Path, default=BASE / "results")
    args = parser.parse_args()
    args.output.mkdir(parents=True, exist_ok=False)
    with warnings.catch_warnings(record=True) as recorded:
        warnings.simplefilter("always")
        run(args.output)
    write_json(args.output / "warnings.json", [
        {"category": w.category.__name__, "message": str(w.message),
         "source": Path(w.filename).name, "line": w.lineno}
        for w in recorded
    ])
