"""
build sparse ephemeris lookup table builder (GCRS, Earth-centered!).
"""

from datetime import datetime, timedelta, timezone
import numpy as np
from skyfield.api import load
from tqdm import tqdm

BODIES = [
    "sun",
    "moon",
    "mercury",
    "venus",
    "mars",
    "jupiter BARYCENTER",
    "saturn BARYCENTER",
    "uranus BARYCENTER",
    "neptune BARYCENTER",
    "pluto BARYCENTER",
]

START_YEAR = 2026
START_MONTH = 1
START_DAY = 1
START_HOUR = 0
START_MINUTE = 0
START_SECOND = 0

YEARS = 1

# sample spacing (seconds)
RESOLUTION_SECONDS = 60 * 60 * 24 * 7 * 6  # 6 weeks

OUTPUT_DIR = "data"
C_HEADER_PATH = f"{OUTPUT_DIR}/ephemeris_tables.h"


def c_identifier(name: str) -> str:
    return name.lower().replace(" ", "_")


def write_c_array_1d(f, ctype, name, array):
    f.write(f"const {ctype} {name}[{len(array)}] = {{\n")
    for i, v in enumerate(array):
        if ctype == "int64_t":
            f.write(f"  {int(v)}")
        else:
            f.write(f"  {float(v):.9e}f")
        if i < len(array) - 1:
            f.write(",")
        f.write("\n")
    f.write("};\n\n")


def write_c_array_2d(f, ctype, name, array):
    rows, cols = array.shape
    f.write(f"const {ctype} {name}[{rows}][{cols}] = {{\n")
    for r in range(rows):
        f.write("  { ")
        for c in range(cols):
            v = float(array[r, c])
            f.write(f"{v:.9e}f")
            if c < cols - 1:
                f.write(", ")
        f.write(" }")
        if r < rows - 1:
            f.write(",")
        f.write("\n")
    f.write("};\n\n")


ts = load.timescale()
eph = load("de421.bsp")
earth = eph["earth"]

start_dt = datetime(
    START_YEAR,
    START_MONTH,
    START_DAY,
    START_HOUR,
    START_MINUTE,
    START_SECOND,
    tzinfo=timezone.utc,
)

num_samples = int(YEARS * 365 * 24 * 3600 // RESOLUTION_SECONDS)

datetimes = [
    start_dt + timedelta(seconds=i * RESOLUTION_SECONDS)
    for i in range(num_samples)
]

times = ts.from_datetimes(datetimes)

unix_time = np.array(
    [int(dt.timestamp()) for dt in datetimes],
    dtype=np.int64
)

np.save(f"{OUTPUT_DIR}/time_table.npy", unix_time)

tables = {}

for body_name in tqdm(BODIES):
    print(f"Building table for {body_name} ...")

    body = eph[body_name]
    obs = earth.at(times).observe(body)

    table = np.empty((num_samples, 6), dtype=np.float32)
    table[:, 0:3] = obs.position.au.T
    table[:, 3:6] = obs.velocity.au_per_d.T

    key = c_identifier(body_name)
    tables[key] = table

    np.save(f"{OUTPUT_DIR}/{key}_sv_table.npy", table)

print("ephemeris complete :)")

print("Writing C header...")

with open(C_HEADER_PATH, "w") as f:
    f.write("// AUTO-GENERATED FILE — DO NOT EDIT\n")
    f.write("// Ephemeris lookup tables (GCRS, Earth-centered)\n\n")

    f.write("#pragma once\n\n")
    f.write("#include <stdint.h>\n\n")

    f.write(f"#define EPHEMERIS_SAMPLES {num_samples}\n")
    f.write(f"#define EPHEMERIS_COLUMNS 6\n\n")

    write_c_array_1d(
        f,
        ctype="int64_t",
        name="time_table",
        array=unix_time,
    )

    for name, table in tables.items():
        write_c_array_2d(
            f,
            ctype="float",
            name=f"{name}_sv",
            array=table,
        )

print(f"C .h written done")
