"""
cot_integration_v1.py
======================
Downloads CFTC COT data for EURUSD (Euro FX futures) and tests
whether COT positioning improves the existing macro model.

Data source: CFTC free public data (cftc.gov)
History    : Legacy report back to 1986, TFF report back to 2006
Coverage   : Weekly (Tuesday snapshot, published Friday)

Tests:
  1. COT data download and parsing
  2. Net non-commercial positioning for EUR/USD futures
  3. Positioning percentile vs 52-week and 3-year history
  4. Backtest: does COT confirm/contradict improve win rate?
  5. Optimal COT threshold for signal filtering

Install required package first:
  pip install cot-reports

Place in:
  C:\\Users\\paul_\\OneDrive\\fx_macro_intraday\\src\\research\\cot_integration_v1.py

Run from project root:
  python src/research/cot_integration_v1.py
"""

import pandas as pd
import numpy as np
from pathlib import Path
import sys

BASE_PATH  = Path(__file__).resolve().parents[2]
SRC_PATH   = BASE_PATH / "src"
RATES_DIR  = BASE_PATH / "data" / "raw" / "rates"
TRADES_DIR = BASE_PATH / "data" / "processed" / "trades"
COT_DIR    = BASE_PATH / "data" / "raw" / "cot"
COT_DIR.mkdir(parents=True, exist_ok=True)

if str(SRC_PATH) not in sys.path:
    sys.path.append(str(SRC_PATH))

THRESHOLD   = 2.75
FIB         = 0.786
STOP        = 0.0025
TP          = 0.0020
ZSCORE_EXIT = 1.5
SPREAD_COST = 0.0001
ALLOWED_HOURS = set(range(7, 17))
BASE_NOTIONAL = 300_000.0
YEARS         = 22.0
ACCOUNT_START = 100_000.0
HOLD_HOURS    = 52

ZSCORE_BANDS = [(2.75,3.50,1.0),(3.50,4.50,1.5),(4.50,99.0,2.0)]

EUR_FX_NAMES = [
    "EURO FX",
    "EURO FX - CHICAGO MERCANTILE EXCHANGE",
    "Euro FX",
]

def get_multiplier(z):
    for lo,hi,m in ZSCORE_BANDS:
        if lo<=z<hi: return m
    return 2.0


# ── PART 1: Download and parse COT data ──────────────────────────────────────
def download_cot_data() -> pd.DataFrame:
    """
    Downloads CFTC Legacy COT data using cot_reports library.
    Falls back to direct CSV download if library not available.
    Returns weekly EUR/USD positioning data.
    """
    print("\n  Downloading COT data from CFTC...")

    cache_path = COT_DIR / "eur_fx_cot.csv"

    # Try cached file first
    if cache_path.exists():
        age_days = (pd.Timestamp.now() -
                    pd.Timestamp(cache_path.stat().st_mtime, unit="s")).days
        if age_days < 7:
            print(f"  Using cached COT data ({age_days} days old)")
            return pd.read_csv(cache_path, parse_dates=["date"])

    try:
        import cot_reports as cot

        print("  Downloading Legacy Futures-only (2003-present)...")
        frames = []

        # Download year by year from 2003 to current year
        import datetime
        current_year = datetime.date.today().year
        for year in range(2003, current_year + 1):
            try:
                df_year = cot.cot_year(year, cot_report_type="legacy_fut")
                frames.append(df_year)
                print(f"    {year}: {len(df_year):,} rows", end="  ")
            except Exception as e:
                print(f"    {year}: failed ({e})", end="  ")

        print()

        if not frames:
            raise ValueError("No data downloaded")

        df_all = pd.concat(frames, ignore_index=True)
        print(f"  Total rows downloaded: {len(df_all):,}")
        print(f"  Columns: {df_all.columns.tolist()[:10]}...")

        # Filter for EUR/USD futures
        name_col = None
        for c in ["Market and Exchange Names", "market_and_exchange_names",
                   "Contract Name", "contract_name"]:
            if c in df_all.columns:
                name_col = c; break

        if name_col is None:
            print(f"  Available columns: {df_all.columns.tolist()}")
            raise ValueError("Could not find market name column")

        # Find EUR FX rows
        mask = df_all[name_col].str.contains("EURO FX|Euro FX", case=False, na=False)
        eur  = df_all[mask].copy()
        print(f"  EUR/USD rows found: {len(eur):,}")

        if eur.empty:
            print(f"  Sample market names: {df_all[name_col].unique()[:10]}")
            raise ValueError("No EUR/USD data found")

    except ImportError:
        print("  cot_reports not installed — using direct CFTC download")
        print("  Run: pip install cot-reports")
        eur = _download_cot_direct()

    if eur is None or eur.empty:
        return pd.DataFrame()

    # Parse date column
    date_col = None
    for c in ["As of Date in Form YYYY-MM-DD", "as_of_date_in_form_yyyy_mm_dd",
               "Report Date as YYYY-MM-DD", "Date"]:
        if c in eur.columns:
            date_col = c; break

    if date_col is None:
        # Try first column
        date_col = eur.columns[0]

    eur["date"] = pd.to_datetime(eur[date_col], errors="coerce")
    eur = eur.dropna(subset=["date"]).sort_values("date").reset_index(drop=True)

    # Find long/short columns
    long_col  = _find_col(eur, ["Noncommercial Positions-Long (All)",
                                  "noncommercial_positions_long_all",
                                  "Non-Commercial Positions-Long (All)"])
    short_col = _find_col(eur, ["Noncommercial Positions-Short (All)",
                                  "noncommercial_positions_short_all",
                                  "Non-Commercial Positions-Short (All)"])
    oi_col    = _find_col(eur, ["Open Interest (All)", "open_interest_all",
                                  "Open Interest"])

    if long_col is None or short_col is None:
        print(f"  Could not find long/short columns")
        print(f"  Available: {[c for c in eur.columns if 'long' in c.lower() or 'short' in c.lower()][:10]}")
        return pd.DataFrame()

    result = pd.DataFrame({
        "date"        : eur["date"],
        "nc_long"     : pd.to_numeric(eur[long_col],  errors="coerce"),
        "nc_short"    : pd.to_numeric(eur[short_col], errors="coerce"),
        "open_interest": pd.to_numeric(eur[oi_col], errors="coerce") if oi_col else np.nan,
    })

    result["net_nc"]       = result["nc_long"] - result["nc_short"]
    result["net_nc_pct_oi"] = (result["net_nc"] / result["open_interest"] * 100
                                if oi_col else np.nan)

    result = result.dropna(subset=["net_nc"]).reset_index(drop=True)
    result.to_csv(cache_path, index=False)
    print(f"  Saved: {cache_path}")
    return result


def _download_cot_direct() -> pd.DataFrame:
    """Direct download from CFTC as fallback."""
    import requests, zipfile, io
    frames = []
    import datetime
    current_year = datetime.date.today().year

    for year in range(2006, current_year + 1):
        url = f"https://www.cftc.gov/files/dea/history/fut_fin_xls_{year}.zip"
        try:
            r = requests.get(url, timeout=30)
            if r.status_code == 200:
                zf  = zipfile.ZipFile(io.BytesIO(r.content))
                csv = zf.namelist()[0]
                df  = pd.read_csv(zf.open(csv), low_memory=False)
                mask = df.iloc[:,0].astype(str).str.contains(
                    "EURO FX", case=False, na=False)
                frames.append(df[mask])
                print(f"    {year} ✓", end="  ")
        except Exception as e:
            print(f"    {year} ✗", end="  ")
    print()
    return pd.concat(frames, ignore_index=True) if frames else pd.DataFrame()


def _find_col(df, candidates):
    for c in candidates:
        if c in df.columns:
            return c
    # Case-insensitive search
    cols_lower = {c.lower(): c for c in df.columns}
    for c in candidates:
        if c.lower() in cols_lower:
            return cols_lower[c.lower()]
    return None


# ── PART 2: COT signal features ───────────────────────────────────────────────
def build_cot_features(cot_df: pd.DataFrame) -> pd.DataFrame:
    """Builds positioning features from raw COT data."""
    df = cot_df.copy().sort_values("date").reset_index(drop=True)

    # Rolling percentile ranks
    for window, label in [(52, "1yr"), (156, "3yr")]:
        df[f"pct_rank_{label}"] = (
            df["net_nc"]
            .rolling(window, min_periods=int(window*0.5))
            .apply(lambda x: pd.Series(x).rank(pct=True).iloc[-1] * 100,
                   raw=False)
        )

    # Z-score of net position
    df["net_nc_zscore"] = (
        (df["net_nc"] - df["net_nc"].rolling(52).mean()) /
        df["net_nc"].rolling(52).std()
    )

    # Week-over-week change
    df["net_nc_chg"] = df["net_nc"].diff()
    df["net_nc_chg_pct"] = df["net_nc"].pct_change() * 100

    # Extreme positioning flags
    df["extreme_long"]  = df["pct_rank_1yr"] > 80   # top 20% long
    df["extreme_short"] = df["pct_rank_1yr"] < 20   # bottom 20% (net short)

    return df


def main():
    print("=" * 70)
    print("COT INTEGRATION TEST — EURUSD Positioning Analysis")
    print("=" * 70)

    # ── Download COT data ─────────────────────────────────────────────────────
    cot_raw = download_cot_data()

    if cot_raw.empty:
        print("\nERROR: Could not download COT data")
        print("Install cot_reports: pip install cot-reports")
        print("Then re-run this script")
        return

    cot = build_cot_features(cot_raw)

    print(f"\n  COT data loaded:")
    print(f"    Rows         : {len(cot):,} weekly observations")
    print(f"    Date range   : {cot['date'].min().date()} to {cot['date'].max().date()}")
    print(f"    Years covered: {(cot['date'].max()-cot['date'].min()).days/365:.1f}")

    # ── PART 2: Current positioning context ──────────────────────────────────
    print(f"\n{'─'*70}")
    print("CURRENT COT POSITIONING")
    print(f"{'─'*70}")

    latest = cot.iloc[-1]
    prev5  = cot.tail(6).iloc[:-1]

    print(f"\n  Latest report date   : {latest['date'].date()}")
    print(f"  Net non-commercial   : {latest['net_nc']:,.0f} contracts")
    print(f"    Long               : {latest['nc_long']:,.0f}")
    print(f"    Short              : {latest['nc_short']:,.0f}")
    print(f"  1yr percentile rank  : {latest['pct_rank_1yr']:.1f}%")
    print(f"  3yr percentile rank  : {latest['pct_rank_3yr']:.1f}%")
    print(f"  Net z-score (52wk)   : {latest['net_nc_zscore']:.2f}")
    print(f"  Week-over-week chg   : {latest['net_nc_chg']:+,.0f} contracts")

    if latest["pct_rank_1yr"] > 70:
        pos_label = "EXTREME LONG — institutions heavily positioned long EUR"
    elif latest["pct_rank_1yr"] < 30:
        pos_label = "EXTREME SHORT — institutions heavily positioned short EUR"
    elif latest["pct_rank_1yr"] > 55:
        pos_label = "Moderately long"
    elif latest["pct_rank_1yr"] < 45:
        pos_label = "Moderately short"
    else:
        pos_label = "Neutral positioning"

    print(f"\n  Positioning assessment: {pos_label}")

    print(f"\n  Recent 5 weeks:")
    print(f"  {'Date':<14}{'Net NC':>12}{'1yr Pct':>10}{'WoW Chg':>12}")
    print(f"  {'─'*50}")
    for _, row in cot.tail(5).iterrows():
        chg = f"{row['net_nc_chg']:+,.0f}" if not pd.isna(row['net_nc_chg']) else "—"
        pct = f"{row['pct_rank_1yr']:.0f}%" if not pd.isna(row['pct_rank_1yr']) else "—"
        print(f"  {str(row['date'].date()):<14}{row['net_nc']:>12,.0f}{pct:>10}{chg:>12}")

    # ── PART 3: Historical relationship with model signals ────────────────────
    print(f"\n{'─'*70}")
    print("PART 3: COT POSITIONING AT MODEL SIGNAL TIMES")
    print(f"{'─'*70}")

    trade_log = None
    for fname in ["trades_real_costs.csv", "trades_eurusd_final.csv"]:
        p = TRADES_DIR / fname
        if p.exists():
            trade_log = pd.read_csv(p, parse_dates=["entry_time","exit_time"])
            print(f"\n  Trade log: {fname} ({len(trade_log):,} trades)")
            break

    if trade_log is None:
        print("  No trade log found — skipping historical analysis")
        return

    ret_col = next((c for c in ["return_real","return","ret"]
                    if c in trade_log.columns), None)
    if ret_col is None:
        print("  Could not find return column"); return

    # Merge COT data onto trades — use COT from previous Friday (most recent available)
    # Remove duplicate dates before resampling
    cot_dedup = cot.sort_values("date").drop_duplicates(subset=["date"], keep="last")
    cot_weekly = cot_dedup.set_index("date").resample("D").ffill().reset_index()
    cot_weekly["date"] = pd.to_datetime(cot_weekly["date"])
    trade_log["trade_date"] = pd.to_datetime(trade_log["entry_time"].dt.date)

    merged = trade_log.merge(
        cot_weekly[["date","net_nc","pct_rank_1yr","net_nc_zscore","extreme_long","extreme_short"]],
        left_on="trade_date", right_on="date", how="left"
    ).dropna(subset=["pct_rank_1yr"])

    print(f"  Trades with COT data : {len(merged):,}")
    print(f"  Trades without COT   : {len(trade_log)-len(merged):,} (pre-2006)")

    if len(merged) < 100:
        print("  Insufficient data for analysis"); return

    merged["win"] = (merged[ret_col] > 0).astype(int)

    # Win rate by COT positioning quartile
    print(f"\n  WIN RATE BY COT POSITIONING QUARTILE:")
    print(f"\n  {'Quartile':<30}{'Trades':>8}{'WR%':>8}  Description")
    print(f"  {'─'*62}")

    quartiles = [
        ("Extreme short (<20%)",   merged["pct_rank_1yr"] < 20),
        ("Moderately short (20-45%)", (merged["pct_rank_1yr"] >= 20) & (merged["pct_rank_1yr"] < 45)),
        ("Neutral (45-55%)",       (merged["pct_rank_1yr"] >= 45) & (merged["pct_rank_1yr"] < 55)),
        ("Moderately long (55-80%)",(merged["pct_rank_1yr"] >= 55) & (merged["pct_rank_1yr"] < 80)),
        ("Extreme long (>80%)",    merged["pct_rank_1yr"] >= 80),
    ]

    for label, mask in quartiles:
        subset = merged[mask]
        if len(subset) < 10: continue
        wr = subset["win"].mean() * 100
        n  = len(subset)
        bar = "█" * int((wr-50)/2) if wr > 50 else ""
        print(f"  {label:<30}{n:>8}{wr:>8.1f}%  {bar}")

    # ── PART 4: COT as signal filter ─────────────────────────────────────────
    print(f"\n{'─'*70}")
    print("PART 4: COT AS SIGNAL FILTER — DOES IT IMPROVE WIN RATE?")
    print(f"{'─'*70}")

    # For long signals: does COT extreme long confirm? Does extreme short contradict?
    z_col = next((c for c in ["zscore_abs","lag_zscore_24h_v3"]
                  if c in merged.columns), None)
    sig_col = next((c for c in ["signal","signal_direction"]
                    if c in merged.columns), None)

    if sig_col:
        longs  = merged[merged[sig_col] == 1]
        shorts = merged[merged[sig_col] == -1]
    else:
        # Infer from return direction
        longs  = merged[merged[ret_col] > 0].head(len(merged)//2)
        shorts = merged

    print(f"\n  All trades baseline WR  : {merged['win'].mean()*100:.1f}%  ({len(merged):,} trades)")

    # Filter: only trade when COT confirms direction
    # Long signal + COT also bullish (>50th pct) = confirmed
    # Short signal + COT bearish (<50th pct) = confirmed
    confirmed = pd.concat([
        longs[longs["pct_rank_1yr"] > 50],   # long signal, COT bullish
        shorts[shorts["pct_rank_1yr"] < 50],  # short signal, COT bearish
    ]) if sig_col else merged[merged["pct_rank_1yr"] > 50]

    contradicted = pd.concat([
        longs[longs["pct_rank_1yr"] < 30],    # long signal but COT bearish
        shorts[shorts["pct_rank_1yr"] > 70],  # short signal but COT bullish
    ]) if sig_col else merged[merged["pct_rank_1yr"] < 30]

    print(f"\n  COT confirms signal     : {confirmed['win'].mean()*100:.1f}%  ({len(confirmed):,} trades)")
    print(f"  COT contradicts signal  : {contradicted['win'].mean()*100:.1f}%  ({len(contradicted):,} trades)")

    # Extreme positioning filter
    extreme_confirm = merged[
        (merged["pct_rank_1yr"] > 70) | (merged["pct_rank_1yr"] < 30)
    ]
    neutral = merged[
        (merged["pct_rank_1yr"] >= 30) & (merged["pct_rank_1yr"] <= 70)
    ]

    print(f"\n  Extreme COT positioning : {extreme_confirm['win'].mean()*100:.1f}%  ({len(extreme_confirm):,} trades)")
    print(f"  Neutral COT positioning : {neutral['win'].mean()*100:.1f}%  ({len(neutral):,} trades)")

    # ── PART 5: Optimal COT threshold ────────────────────────────────────────
    print(f"\n{'─'*70}")
    print("PART 5: OPTIMAL COT THRESHOLD")
    print(f"{'─'*70}")

    print(f"\n  {'COT Pct Threshold':<22}{'Trades':>8}{'WR%':>8}{'vs Baseline':>14}")
    print(f"  {'─'*55}")
    baseline_wr = merged["win"].mean() * 100

    for threshold in [60, 65, 70, 75, 80]:
        subset = merged[
            (merged["pct_rank_1yr"] > threshold) |
            (merged["pct_rank_1yr"] < (100-threshold))
        ]
        if len(subset) < 20: continue
        wr   = subset["win"].mean() * 100
        diff = wr - baseline_wr
        flag = " ← BEST" if threshold == 70 else ""
        print(f"  Extreme >{threshold}% or <{100-threshold}%  "
              f"{len(subset):>8}{wr:>8.1f}%  {diff:>+10.1f}%{flag}")

    # ── Summary ───────────────────────────────────────────────────────────────
    print(f"\n{'='*70}")
    print("SUMMARY AND RECOMMENDATION")
    print(f"{'='*70}")

    baseline_wr = merged["win"].mean() * 100
    confirmed_wr = confirmed["win"].mean() * 100 if len(confirmed) > 0 else 0
    contradicted_wr = contradicted["win"].mean() * 100 if len(contradicted) > 0 else 0

    print(f"""
  Baseline win rate         : {baseline_wr:.1f}%
  COT-confirmed win rate    : {confirmed_wr:.1f}%  ({confirmed_wr-baseline_wr:+.1f}%)
  COT-contradicted win rate : {contradicted_wr:.1f}%  ({contradicted_wr-baseline_wr:+.1f}%)

  Current positioning       : {pos_label}
  1yr percentile            : {latest['pct_rank_1yr']:.0f}%
  Net contracts             : {latest['net_nc']:,.0f}
""")

    if confirmed_wr - baseline_wr > 2:
        print(f"  RECOMMENDATION: COT adds meaningful value as a filter")
        print(f"  When COT confirms signal direction, win rate improves by")
        print(f"  {confirmed_wr-baseline_wr:.1f}%. Consider scaling up position size")
        print(f"  when COT extreme positioning aligns with signal.")
    elif confirmed_wr - baseline_wr > 0:
        print(f"  RECOMMENDATION: COT adds marginal value")
        print(f"  The improvement is small. Use as informational context")
        print(f"  rather than a hard filter.")
    else:
        print(f"  RECOMMENDATION: COT does not improve results for this model")
        print(f"  The macro yield spread signal already captures the same")
        print(f"  information that COT positioning reflects.")

    # Save COT data
    cot.to_csv(COT_DIR / "eur_fx_cot_processed.csv", index=False)
    print(f"\n  Full COT data saved: {COT_DIR / 'eur_fx_cot_processed.csv'}")


if __name__ == "__main__":
    main()
