"""Point-in-time Treasury-curve event study against BTC daily forward returns.

This is research only: it neither changes strategy inputs nor optimizes a
trading rule.  Treasury observations are shifted one BTC day before signals
are evaluated, so a same-day Treasury close can never affect that day's BTC
return.

Examples:
    python tools/run_treasury_curve_event_study.py --download-fred
    python tools/run_treasury_curve_event_study.py --yields /path/to/yields.csv
"""
from __future__ import annotations

import argparse
import json
from pathlib import Path
from typing import Iterable

import numpy as np
import pandas as pd


REPO = Path(__file__).resolve().parent.parent
DEFAULT_BTC = REPO / "data" / "mlp" / "COINBASE_BTCUSD, 1D.csv"
FRED_URL = "https://fred.stlouisfed.org/graph/fredgraph.csv?id=DGS2,DGS10,DGS30"
HORIZONS = (1, 5, 10, 20, 60)
IS_END = pd.Timestamp("2026-02-28")
OOS_START = pd.Timestamp("2026-03-01")


def _date_column(frame: pd.DataFrame) -> str:
    for column in frame.columns:
        if column.lower() in {"date", "time", "datetime", "observation_date"}:
            return column
    raise ValueError("Expected a date column named DATE, observation_date, time, or datetime.")


def load_yields(path_or_url: str | Path) -> pd.DataFrame:
    """Load FRED-format yields, dropping non-numeric missing observations."""
    frame = pd.read_csv(path_or_url)
    date_column = _date_column(frame)
    required = ("DGS2", "DGS10", "DGS30")
    missing = [column for column in required if column not in frame.columns]
    if missing:
        raise ValueError(f"Yield data is missing required columns: {', '.join(missing)}")
    frame = frame[[date_column, *required]].copy()
    frame[date_column] = pd.to_datetime(frame[date_column], errors="coerce")
    for column in required:
        frame[column] = pd.to_numeric(frame[column], errors="coerce")
    return frame.dropna().rename(columns={date_column: "time"}).sort_values("time").drop_duplicates("time")


def load_btc(path: Path) -> pd.DataFrame:
    frame = pd.read_csv(path)
    date_column = _date_column(frame)
    if "close" not in {column.lower() for column in frame.columns}:
        raise ValueError("BTC data needs a close column.")
    close_column = next(column for column in frame.columns if column.lower() == "close")
    frame = frame[[date_column, close_column]].copy()
    frame.columns = ["time", "close"]
    frame["time"] = pd.to_datetime(frame["time"], errors="coerce")
    frame["close"] = pd.to_numeric(frame["close"], errors="coerce")
    return frame.dropna().sort_values("time").drop_duplicates("time")


def align_point_in_time(btc: pd.DataFrame, yields: pd.DataFrame, availability_lag_days: int = 1) -> pd.DataFrame:
    """Join yields to BTC while making daily Treasury data usable next BTC day."""
    if availability_lag_days < 1:
        raise ValueError("availability_lag_days must be at least one day.")
    rates = yields.set_index("time")[["DGS2", "DGS10", "DGS30"]].reindex(pd.DatetimeIndex(btc["time"])).ffill()
    rates = rates.shift(availability_lag_days)
    return btc.set_index("time").join(rates).dropna().reset_index()


def build_signals(frame: pd.DataFrame, change_days: int, sharp_drop_bps: float) -> pd.DataFrame:
    """Create the two hypotheses fixed before examining forward returns."""
    output = frame.copy()
    # Positive 10s30s means the 10Y yield is above the 30Y yield.
    output["ten_above_thirty"] = output["DGS10"] > output["DGS30"]
    changes_bps = output[["DGS2", "DGS10", "DGS30"]].diff(change_days) * 100
    output["two_year_sharply_falling"] = changes_bps["DGS2"] <= -sharp_drop_bps
    output["all_yields_sharply_falling"] = (changes_bps <= -sharp_drop_bps).all(axis=1)
    return output


def _summary(frame: pd.DataFrame, mask: pd.Series, horizons: Iterable[int]) -> dict[str, object]:
    result: dict[str, object] = {"signal_days": int(mask.sum()), "horizons": {}}
    for horizon in horizons:
        returns = frame.loc[mask, f"forward_{horizon}d_return_pct"].dropna()
        result["horizons"][str(horizon)] = {
            "n": int(len(returns)),
            "mean_return_pct": round(float(returns.mean()), 3) if len(returns) else None,
            "median_return_pct": round(float(returns.median()), 3) if len(returns) else None,
            "positive_return_pct": round(float((returns > 0).mean() * 100), 1) if len(returns) else None,
        }
    return result


def study(btc: pd.DataFrame, yields: pd.DataFrame, *, change_days: int = 10, sharp_drop_bps: float = 25.0) -> dict[str, object]:
    frame = build_signals(align_point_in_time(btc, yields), change_days, sharp_drop_bps)
    for horizon in HORIZONS:
        frame[f"forward_{horizon}d_return_pct"] = (frame["close"].shift(-horizon) / frame["close"] - 1) * 100

    windows = {
        "IS": frame[frame["time"] <= IS_END],
        "OOS": frame[frame["time"] >= OOS_START],
        "full": frame,
    }
    signals = ("ten_above_thirty", "two_year_sharply_falling", "all_yields_sharply_falling")
    return {
        "method": {
            "yield_availability_lag_btc_days": 1,
            "ten_above_thirty": "DGS10 > DGS30",
            "two_year_sharply_falling": f"DGS2 declines at least {sharp_drop_bps:g} bp over {change_days} BTC days",
            "all_yields_sharply_falling": f"DGS2, DGS10, and DGS30 each decline at least {sharp_drop_bps:g} bp over {change_days} BTC days",
            "forward_return_horizons_btc_days": list(HORIZONS),
        },
        "coverage": {"start": str(frame["time"].min().date()), "end": str(frame["time"].max().date()), "bars": len(frame)},
        "windows": {
            name: {"baseline": _summary(window, pd.Series(True, index=window.index), HORIZONS),
                   **{signal: _summary(window, window[signal], HORIZONS) for signal in signals}}
            for name, window in windows.items()
        },
    }


def main() -> int:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--btc", type=Path, default=DEFAULT_BTC)
    parser.add_argument("--yields", help="Local FRED CSV with DATE,DGS2,DGS10,DGS30 columns.")
    parser.add_argument("--download-fred", action="store_true", help="Read the public FRED CSV directly; does not save it.")
    parser.add_argument("--change-days", type=int, default=10)
    parser.add_argument("--sharp-drop-bps", type=float, default=25.0)
    parser.add_argument("--output", type=Path, help="Optional JSON output path.")
    args = parser.parse_args()
    if bool(args.yields) == bool(args.download_fred):
        parser.error("provide exactly one of --yields or --download-fred")
    if args.change_days < 1 or args.sharp_drop_bps <= 0:
        parser.error("--change-days and --sharp-drop-bps must be positive")
    report = study(load_btc(args.btc), load_yields(FRED_URL if args.download_fred else args.yields),
                   change_days=args.change_days, sharp_drop_bps=args.sharp_drop_bps)
    text = json.dumps(report, indent=2) + "\n"
    print(text, end="")
    if args.output:
        args.output.parent.mkdir(parents=True, exist_ok=True)
        args.output.write_text(text)
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
