
import pandas as pd
from astroquery.jplhorizons import Horizons
from astropy import units as u
import numpy as np
from astropy.time import Time
from astropy.utils import iers

# 自动下载并打开 IERS C04 表
iers.conf.auto_download = True
iers_a = iers.IERS_Auto.open()

def fetch_mercury(start="2010-01-01", stop="2023-12-31", step="1d"):
    obj = Horizons(id=199, location='500@0',
                   epochs={'start': start, 'stop': stop, 'step': step})
    eph = obj.ephemerides()
    df = (
        eph.to_pandas()[["datetime_str", "RA", "DEC"]]
           .rename(columns={
               "datetime_str": "datetime",
               "RA":           "RA_deg",
               "DEC":          "DEC_deg",
           })
    )
    df["datetime"] = pd.to_datetime(df["datetime"])
    df.set_index("datetime", inplace=True)
    df["e"] = obj.elements()[0]["e"]
    return df

def merge_with_iers(df):
    # 索引 → MJD
    times   = Time(df.index.to_pydatetime(), scale='utc')
    mjd     = times.mjd

    tbl     = iers_a
    # 把带单位的列剥成纯 float
    mjd_tab = np.array(tbl['MJD'],       dtype=float)
    ut1_tab = np.array(tbl['UT1_UTC_A'], dtype=float)
    lod_tab = np.array(tbl['LOD_A'],     dtype=float)

    # 插值
    ut1_vals = np.interp(mjd, mjd_tab, ut1_tab)
    lod_vals = np.interp(mjd, mjd_tab, lod_tab)

    df2         = df.copy()
    df2['dUT1'] = ut1_vals
    df2['LOD']  = lod_vals
    return df2

def load_drift_data():
    df = fetch_mercury()
    return merge_with_iers(df)
