"""
Quick session comparison — london vs all vs tokyo
Final locked params except session.
"""
import pandas as pd
import numpy as np
from pathlib import Path
import sys
import warnings
warnings.filterwarnings('ignore')

BASE_PATH  = Path(__file__).resolve().parents[2]
RATES_PATH = BASE_PATH / "data" / "raw" / "rates"
if str(BASE_PATH / "src") not in sys.path:
    sys.path.insert(0, str(BASE_PATH / "src"))

NOTIONAL=300_000; SPREAD_COST=0.0001
WF_SPLIT=pd.Timestamp('2020-01-01')
Z_WINDOW=30; THRESHOLD=2.0; TP=0.0070; SL=0.0040; HOLD=24; FIB=0.786


def load_data():
    us2y = pd.read_csv(RATES_PATH/"us2y.csv")
    dc=[c for c in us2y.columns if 'date' in c.lower()][0]
    vc=[c for c in us2y.columns if c!=dc][0]
    us2y[dc]=pd.to_datetime(us2y[dc])
    us2y=us2y.rename(columns={dc:'date',vc:'us2y'})[['date','us2y']].dropna()
    jp2y=pd.read_csv(RATES_PATH/"jp2y.csv")
    jp2y['date']=pd.to_datetime(jp2y['date'])
    jp2y=jp2y[['date','jp2y']].dropna()
    rates=pd.merge(us2y,jp2y,on='date',how='outer').sort_values('date')
    rates=rates.set_index('date')
    rates=rates.reindex(pd.date_range(rates.index.min(),rates.index.max(),freq='D')).ffill().dropna()
    rates['spread']=rates['us2y']-rates['jp2y']
    s=rates['spread']
    z=(s-s.rolling(Z_WINDOW).mean())/s.rolling(Z_WINDOW).std()
    return z.reindex(pd.date_range(z.index.min(),z.index.max(),freq='h')).ffill()


def load_prices():
    fpath=next(BASE_PATH.rglob("USDJPY_15M*.csv"),None)
    df=pd.read_csv(fpath); df.columns=df.columns.str.strip().str.lower()
    dc=next(c for c in df.columns if 'date' in c or 'time' in c)
    df[dc]=pd.to_datetime(df[dc])
    return df.rename(columns={dc:'datetime'}).set_index('datetime').sort_index()


def in_session(dt, session):
    h=dt.hour
    if session=='london': return 7<=(h+2)%24<=17
    if session=='tokyo':  return 0<=h<=9
    return True


def run(z_hourly, m15, session):
    trades=[]; last_exit=None
    for dt in pd.date_range('2003-01-01','2026-04-01',freq='h'):
        if dt not in z_hourly.index: continue
        z=float(z_hourly.at[dt])
        if abs(z)<THRESHOLD: continue
        if not in_session(dt, session): continue
        if last_exit is not None and dt<=last_exit: continue
        direction=1 if z>=THRESHOLD else -1
        z_abs=abs(z); mult=1.0 if z_abs<2.5 else(1.5 if z_abs<3.5 else 2.0)
        try: bar_slice=m15.loc[dt:dt+pd.Timedelta(hours=1)-pd.Timedelta(minutes=1)]
        except: continue
        if bar_slice.empty: continue
        bh=float(bar_slice['high'].max()); bl=float(bar_slice['low'].min()); bc=float(bar_slice['close'].iloc[-1])
        if direction==1:
            pull=bc-bl
            if pull<=0.05: continue
            target=bc-FIB*pull
        else:
            pull=bh-bc
            if pull<=0.05: continue
            target=bc+FIB*pull
        entry_price=entry_time=None
        for edt,ebar in m15.loc[dt+pd.Timedelta(minutes=1):dt+pd.Timedelta(hours=6)].iterrows():
            if direction==1 and float(ebar['low'])<=target: entry_price=target; entry_time=edt; break
            elif direction==-1 and float(ebar['high'])>=target: entry_price=target; entry_time=edt; break
        if entry_price is None: continue
        tp_px=entry_price*(1+TP) if direction==1 else entry_price*(1-TP)
        sl_px=entry_price*(1-SL) if direction==1 else entry_price*(1+SL)
        hold_bars=m15.loc[entry_time+pd.Timedelta(minutes=1):entry_time+pd.Timedelta(hours=HOLD)]
        exit_price=None
        for hdt,hbar in hold_bars.iterrows():
            if direction==1:
                if float(hbar['high'])>=tp_px: exit_price=tp_px; break
                if float(hbar['low'])<=sl_px:  exit_price=sl_px; break
            else:
                if float(hbar['low'])<=tp_px:  exit_price=tp_px; break
                if float(hbar['high'])>=sl_px: exit_price=sl_px; break
        if exit_price is None:
            exit_price=float(hold_bars['close'].iloc[-1]) if not hold_bars.empty else entry_price
        pnl=round(((exit_price-entry_price)/entry_price*direction-SPREAD_COST)*NOTIONAL*mult,2)
        trades.append({'entry_time':entry_time,'pnl':pnl})
        last_exit=entry_time+pd.Timedelta(hours=HOLD)
    return pd.DataFrame(trades)


def score(df, years):
    if len(df)<10: return 0,0,0,0
    pnl=df['pnl'].values; wr=(pnl>0).mean()*100
    eq=100_000+np.cumsum(pnl); pk=np.maximum.accumulate(eq)
    dd=((eq-pk)/pk*100).min(); pyr=len(pnl)/years
    exc=pnl/100_000-0.04/pyr
    sh=exc.mean()/exc.std()*np.sqrt(pyr) if exc.std()>0 else 0
    return round(sh,2),round(wr,1),round(dd,2),len(pnl)


def main():
    print("SESSION COMPARISON — TP0.70% SL0.40% (final locked params)")
    print("="*60)
    z_hourly=load_data(); m15=load_prices()
    years=(m15.index.max()-m15.index.min()).days/365.25
    yrs_oos=(m15.index.max()-WF_SPLIT).days/365.25

    print(f"\n  {'Session':<10}{'N':>5}{'N/yr':>6}{'WR%':>6}{'Sh_full':>9}{'Sh_OOS':>8}{'DD%':>7}")
    print(f"  {'─'*52}")

    for session in ['london','all','tokyo']:
        print(f"  Running {session}...", end="\r")
        trades=run(z_hourly,m15,session)
        oos=trades[trades['entry_time']>=WF_SPLIT] if len(trades) else pd.DataFrame()
        sh,wr,dd,n=score(trades,years)
        sh_oos,wr_oos,dd_oos,n_oos=score(oos,yrs_oos)
        pyr=n/years
        print(f"  {session:<10}{n:>5}{pyr:>6.0f}{wr:>6.1f}{sh:>9.2f}{sh_oos:>8.2f}{dd:>7.2f}")

    print(f"\n  Note: 'all' gives more trades at minimal Sharpe cost.")
    print(f"  Pick london for Sharpe, all for trade frequency.")

if __name__ == "__main__":
    main()
