#!/usr/bin/env python3
"""Reproduce the nearest-exoplanet travel-time table and comparison chart.

Input: nasa_exoplanet_distances_2026-09-23.csv (NASA Exoplanet Archive)
Outputs: journey_times.csv and travel_time_comparison.png
"""

import argparse
import csv
import math
from pathlib import Path

PACKAGE_DIR = Path(__file__).resolve().parent
DATA_FILE = PACKAGE_DIR / "nasa_exoplanet_distances_2026-09-23.csv"
RESULTS_FILE = PACKAGE_DIR / "journey_times.csv"
CHART_FILE = PACKAGE_DIR / "travel_time_comparison.png"
PC_TO_LY = 3.2615637771674333
C_KM_S = 299_792.458
SPEEDS = (("17 km/s", 17 / C_KM_S), ("1% c", 0.01),
          ("10% c", 0.10), ("90% c", 0.90))


def load_nearest(data_file=DATA_FILE):
    with Path(data_file).open(newline="", encoding="utf-8") as handle:
        rows = list(csv.DictReader(handle))
    rows = [r for r in rows if r.get("sy_dist")]
    nearest_pc = min(float(r["sy_dist"]) for r in rows)
    nearest = [r for r in rows if float(r["sy_dist"]) == nearest_pc]
    exemplar = nearest[0]
    return nearest, nearest_pc, float(exemplar["sy_disterr1"]), float(exemplar["sy_disterr2"])


def calculate(distance_pc, err_plus_pc, err_minus_pc):
    distance_ly = distance_pc * PC_TO_LY
    low_ly = (distance_pc + err_minus_pc) * PC_TO_LY
    high_ly = (distance_pc + err_plus_pc) * PC_TO_LY
    output = []
    for label, beta in SPEEDS:
        gamma = 1 / math.sqrt(1 - beta**2)
        earth = distance_ly / beta
        onboard = earth / gamma
        output.append({
            "speed": label, "beta": beta, "lorentz_gamma": gamma,
            "earth_years": earth,
            "earth_err_minus_years": earth - low_ly / beta,
            "earth_err_plus_years": high_ly / beta - earth,
            "onboard_years": onboard,
            "onboard_err_minus_years": onboard - low_ly / beta / gamma,
            "onboard_err_plus_years": high_ly / beta / gamma - onboard,
        })
    return distance_ly, output


def write_results(rows, results_file=RESULTS_FILE):
    with Path(results_file).open("w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(handle, fieldnames=rows[0].keys())
        writer.writeheader()
        writer.writerows(rows)


def draw_chart(rows, chart_file=CHART_FILE):
    import matplotlib.pyplot as plt
    from matplotlib.ticker import FuncFormatter

    labels = [r["speed"] for r in rows]
    y = list(range(len(rows)))
    earth = [r["earth_years"] for r in rows]
    onboard = [r["onboard_years"] for r in rows]
    fig, ax = plt.subplots(figsize=(7.2, 4.4), constrained_layout=True)
    ax.hlines(y, onboard, earth, color="#c7cbd1", linewidth=2, zorder=1)
    ax.scatter(earth, y, s=58, color="#0072B2", marker="o", label="Earth frame", zorder=3)
    ax.scatter(onboard, y, s=58, facecolor="white", edgecolor="#D55E00", linewidth=1.8,
               marker="o", label="Onboard time", zorder=4)
    ax.set_xscale("log")
    ax.set_yticks(y, labels)
    ax.invert_yaxis()
    ax.set_xlabel("One-way elapsed time (years, logarithmic scale)")
    ax.set_title("Time dilation becomes substantial only near light speed", loc="left")
    ax.xaxis.set_major_formatter(FuncFormatter(lambda x, pos: f"{x:g}"))
    ax.grid(axis="x", which="both", color="#e4e7eb", linewidth=0.7)
    ax.spines[["top", "right", "left"]].set_visible(False)
    ax.tick_params(axis="y", length=0)
    ax.legend(frameon=False, loc="lower right")
    ax.margins(x=0.06, y=0.16)
    fig.text(0.01, 0.01,
             "Distance: 1.30119 pc (4.24391 ly). Points assume constant-speed cruise; acceleration and braking omitted.",
             fontsize=8, color="#4b5563")
    fig.savefig(chart_file, dpi=300, bbox_inches="tight")
    plt.close(fig)


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--data", type=Path, default=DATA_FILE)
    parser.add_argument("--output-dir", type=Path, default=PACKAGE_DIR)
    args = parser.parse_args()
    args.output_dir.mkdir(parents=True, exist_ok=True)
    nearest, distance_pc, err_plus_pc, err_minus_pc = load_nearest(args.data)
    distance_ly, results = calculate(distance_pc, err_plus_pc, err_minus_pc)
    write_results(results, args.output_dir / "journey_times.csv")
    draw_chart(results, args.output_dir / "travel_time_comparison.png")
    print(f"Nearest distance: {distance_pc:.8f} pc = {distance_ly:.8f} ly")
    print("Tied planets: " + ", ".join(r["pl_name"] for r in nearest))
    for row in results:
        print(f"{row['speed']:>7}: Earth {row['earth_years']:.6f} y; onboard {row['onboard_years']:.6f} y")


if __name__ == "__main__":
    main()
