"""趋势腿加空头测试（crypto-multi-sleeve 报告附录）
引擎/数据/成本/资金费与研究第 1-2 轮完全相同：bt.run（4h bar，收盘信号、下一根开盘成交），taker 5bp + 分级滑点，
资金费按实际结算向持仓方收付（空头在正费率时收钱）。信号 = final.py 的 8 个子信号平均 S∈[-1,1]，
目标 = S_part × min(0.5/vol,3) / 有效币数，25% 带宽作用在未缩放的多空合并目标上（与实盘/famB 一致），之后再分别乘多头、空头乘数。
BTC200：BTC 日收盘 > 200 日 SMA（前一天判定，次日生效），与 famB 完全相同。
"""
import sys, json, numpy as np, pandas as pd
C0="/workspace/iq/strategies/crypto/"; sys.path.insert(0,C0); sys.path.insert(0,C0+"round2/")
import os; os.chdir(C0)
from bt import *
from lib2 import base_sleeves, daily
U6=["BTC","ETH","SOL","BNB","XRP","DOGE"]
def band(tgt, rel):
    out=np.zeros_like(tgt); cur=np.zeros(tgt.shape[1])
    for t in range(tgt.shape[0]):
        x=tgt[t]; chg=(np.sign(x)!=np.sign(cur))|((np.abs(x-cur)>rel*np.abs(x)+1e-12)&(np.abs(x-cur)>rel*np.abs(cur)))
        cur=np.where(chg,x,cur); out[t]=cur
    return out
raw={s:resample(load(s,"1h"),"4h") for s in U6}
P=lambda f: pd.DataFrame({s:raw[s][f] for s in U6}).loc[START:]
O,H,L,C=P("o"),P("h"),P("l"),P("c")
vol=C.pct_change().rolling(180,min_periods=60).std()*np.sqrt(365*6)
valid=C.notna()&vol.notna()&(C.notna().cumsum()>180)
sigs=[]
for N in [48,96,192]:
    hh=H.rolling(N).max().shift(1); ll=L.rolling(N).min().shift(1); hh2=H.rolling(N//2).max().shift(1); ll2=L.rolling(N//2).min().shift(1)
    c=C.values;a=hh.values;b=ll.values;x=hh2.values;y=ll2.values; sa=np.zeros(C.shape); st=np.zeros(C.shape[1])
    for t in range(C.shape[0]):
        ct=c[t]; st=np.where((st==0)&(ct>a[t]),1,np.where((st==0)&(ct<b[t]),-1,st)); st=np.where((st==1)&(ct<y[t]),0,np.where((st==-1)&(ct>x[t]),0,st)); sa[t]=st
    sigs.append(pd.DataFrame(sa,index=C.index,columns=C.columns))
for f,s in [(12,48),(24,96),(48,192)]: sigs.append(np.sign(C.ewm(span=f).mean()-C.ewm(span=s).mean()).fillna(0))
for N in [96,192]: sigs.append(np.sign(C/C.shift(N)-1).fillna(0))
S=(sum(sigs)/len(sigs)).where(valid,0)
scale=(0.5/vol).clip(upper=3); nval=valid.sum(1).replace(0,np.nan)
def tgt(Sx): return (Sx*scale).div(nval,axis=0).fillna(0)
W_LO=pd.DataFrame(band(tgt(S.clip(lower=0)).values,0.25),index=C.index,columns=U6)
W_LS=pd.DataFrame(band(tgt(S).values,0.25),index=C.index,columns=U6)
W_SO=pd.DataFrame(band(tgt(S.clip(upper=0)).values,0.25),index=C.index,columns=U6)
ref_w=pd.read_parquet("res/final_trend_w.parquet").reindex(C.index).fillna(0)
assert np.allclose(W_LO.values,ref_w.values,atol=1e-12), "LO weights mismatch"
btc=C["BTC"].resample("1D").last()
on=(btc>btc.rolling(200).mean()).astype(float).shift(1).reindex(O.index,method="ffill").fillna(0)  # 1=risk-on
F4=fund_panel(U6,O.index)
def side(W,lm,sm):
    lm=pd.Series(lm,index=W.index) if np.isscalar(lm) else lm; sm=pd.Series(sm,index=W.index) if np.isscalar(sm) else sm
    return W.clip(lower=0).mul(lm,axis=0)+W.clip(upper=0).mul(sm,axis=0)
half=on.replace(0,0.5)          # long multiplier: 1 risk-on, 0.5 risk-off (现行)
off=1-on                        # 1 when BTC < SMA200
V={
 "0 现行：只做多 + BTC200 减半":side(W_LO,half,0),
 "a 对称多空（无过滤）":side(W_LS,1,1),
 "a′ 对称多空，多头按 BTC200 减半":side(W_LS,half,1),
 "b 空头只在 BTC<SMA200 时开，多头照旧减半":side(W_LS,half,off),
 "b′ 同 b，空头 0.5×":side(W_LS,half,0.5*off),
 "c 空头始终 0.5×，多头照旧减半":side(W_LS,half,0.5),
 "d 只做空（诊断，无过滤）":side(W_SO,0,1),
 "d′ 只做空，仅 BTC<SMA200":side(W_SO,0,off),
}
# note: variants b/b′/c/a′ use W_LS (band on combined target); in risk-on b zeroes shorts after band.
B=base_sleeves(); A=pd.read_pickle(C0+"round2/res/all_daily.pkl")
xs=B["xs"].fillna(0); ls=A["Dr_LSENS_top+glb_L137_H3"].reindex(B.index).fillna(0)
ev=pd.read_pickle(C0+"round3_events/ev/fa_sleeves.pkl")[("PRE=MON+PDL+AIR","corrected funding + pre-reg delist exit (NEW HEADLINE)")].reindex(B.index).fillna(0)
famB=pd.read_pickle(C0+"round2/res/famB.pkl")
UTC=lambda s: pd.Timestamp(s,tz="UTC"); END=UTC("2026-08-31")
def sst(x,lo,hi):
    x=x[(x.index>=UTC(lo))&(x.index<UTC(hi))]; eq=(1+x).cumprod(); yrs=len(x)/365.25
    return dict(cagr=float(eq.iloc[-1]**(1/yrs)-1),sharpe=float(x.mean()/x.std()*np.sqrt(365)),mdd=float((eq/eq.cummax()-1).min()))
def split(r,lo,hi):
    r=r[(r.index>=UTC(lo))&(r.index<UTC(hi))]; yrs=(UTC(hi)-UTC(lo)).days/365.25
    return dict(gross=r.gross.sum()/yrs,fees=-r.fees.sum()/yrs,funding=-r.fund.sum()/yrs,net=r.net.sum()/yrs,turn=r.turn.sum()/yrs)
PER=[("IS","2020-01-01","2024-01-01"),("IS22","2022-01-01","2024-01-01"),("OOS","2024-01-01","2026-09-01")]
rows=[]; daily_out={}
for name,W in V.items():
    r=run(W,O,fund=F4)
    if name.startswith("0"): assert np.allclose(r.net.values,famB["BTCabove200_else0.5"].net.reindex(r.index).values,atol=1e-12),"baseline mismatch"
    d=daily(r).reindex(B.index).fillna(0); d=d[d.index<=END]; daily_out[name]=d
    held=W.shift(1).fillna(0)
    rec=dict(variant=name)
    for p,lo,hi in PER:
        rec.update({f"{p}_{k}":v for k,v in sst(d,lo,hi).items()}); rec.update({f"{p}_{k}":v for k,v in split(r,lo,hi).items()})
        m=(held.index>=UTC(lo))&(held.index<UTC(hi)); h=held[m]
        rec[f"{p}_tim"]=float((h.abs().sum(1)>1e-9).mean()); rec[f"{p}_tim_long"]=float((h.clip(lower=0).sum(1)>1e-9).mean()); rec[f"{p}_tim_short"]=float((h.clip(upper=0).sum(1)<-1e-9).mean())
        rec[f"{p}_gexp"]=float(h.abs().sum(1).mean()); rec[f"{p}_net_exp"]=float(h.sum(1).mean())
    o=pd.DataFrame({"t":d,"xs":xs,"ls":ls,"ev":ev}).dropna(); oo=o[o.index>=UTC("2024-01-01")]
    rec["corr_xs_OOS"]=float(oo["t"].corr(oo["xs"])); rec["corr_ls_OOS"]=float(oo["t"].corr(oo["ls"])); rec["corr_base_OOS"]=None
    rec["yearly"]={int(y):float(v) for y,v in ((1+d).groupby(d.index.year).prod()-1).items()}
    t=d; x_=xs.reindex(t.index).fillna(0); l_=ls.reindex(t.index).fillna(0); e_=ev.reindex(t.index).fillna(0)
    for nm,combo in [("p",0.4*t+0.25*x_+0.35*l_),("pe",0.4*t+0.25*x_+0.35*l_+0.1*e_)]:
        for p,lo,hi in PER[1:]:
            rec.update({f"{nm}_{p}_{k}":v for k,v in sst(combo,lo,hi).items()})
        rec[f"{nm}_yearly"]={int(y):float(v) for y,v in ((1+combo).groupby(combo.index.year).prod()-1).items()}
    rows.append(rec)
base=daily_out[list(V)[0]]
for rec in rows:
    d=daily_out[rec["variant"]]; m=d.index>=UTC("2024-01-01"); rec["corr_base_OOS"]=float(d[m].corr(base[m]))
OUT=sys.argv[1] if len(sys.argv)>1 else "/workspace/report_cms/trendls"
json.dump(rows,open(OUT+"/trend_ls.json","w"),ensure_ascii=False,indent=1)
pd.DataFrame(daily_out).to_csv(OUT+"/trend_ls_daily.csv",index_label="date_utc")
for rec in rows:
    print(rec["variant"],"| IS %.1f/%.2f/%.1f | IS22 %.1f/%.2f/%.1f | OOS %.1f/%.2f/%.1f | split OOS %.1f/%.1f/%.1f/%.1f | tim %.2f L%.2f S%.2f | corr xs %.2f ls %.2f base %.2f | combo %.1f/%.2f/%.1f ev %.1f/%.2f | IS22 combo %.2f"%(
      rec["IS_cagr"]*100,rec["IS_sharpe"],rec["IS_mdd"]*100,rec["IS22_cagr"]*100,rec["IS22_sharpe"],rec["IS22_mdd"]*100,rec["OOS_cagr"]*100,rec["OOS_sharpe"],rec["OOS_mdd"]*100,
      rec["OOS_gross"]*100,rec["OOS_fees"]*100,rec["OOS_funding"]*100,rec["OOS_net"]*100,rec["OOS_tim"],rec["OOS_tim_long"],rec["OOS_tim_short"],rec["corr_xs_OOS"],rec["corr_ls_OOS"],rec["corr_base_OOS"],
      rec["p_OOS_cagr"]*100,rec["p_OOS_sharpe"],rec["p_OOS_mdd"]*100,rec["pe_OOS_cagr"]*100,rec["pe_OOS_sharpe"],rec["p_IS22_sharpe"]))
# ---- flat CSV + chart ----
flat=[{k:v for k,v in r.items() if not isinstance(v,dict)}|{f"y{y}":v for y,v in r["yearly"].items()} for r in rows]
pd.DataFrame(flat).to_csv(OUT+"/trend_ls.csv",index=False,float_format="%.6f")
import matplotlib; matplotlib.use("Agg"); import matplotlib.pyplot as plt; from matplotlib import font_manager as fm
for f in ["/usr/share/fonts/opentype/noto/NotoSansCJK-Regular.ttc","/usr/share/fonts/opentype/noto/NotoSansCJK-Bold.ttc"]: fm.fontManager.addfont(f)
plt.rcParams["font.family"]="Noto Sans CJK JP"; plt.rcParams["axes.unicode_minus"]=False
names=list(V); pick=[(names[0],"现行：只做多 + BTC200","#e23d3d",1.8),(names[1],"a 对称多空","#6b3fc4",1.1),(names[3],"b 空头仅 BTC<SMA200","#2563c4",1.1),(names[5],"c 空头 0.5×","#c9941f",1.1),(names[6],"d 只做空（诊断）","#555",1.1)]
fig,axs=plt.subplots(1,2,figsize=(12,4.6),dpi=130)
for n,lab,col,lw in pick:
    t=daily_out[n]; axs[0].plot((1+t).cumprod(),color=col,lw=lw,label=lab)
    cb=0.4*t+0.25*xs.reindex(t.index).fillna(0)+0.35*ls.reindex(t.index).fillna(0); axs[1].plot((1+cb).cumprod(),color=col,lw=lw,label=lab)
for ax,ttl in zip(axs,["趋势腿变体（100% 资金，对数坐标）","40/25/35 组合中替换趋势腿（不含事件腿）"]):
    ax.set_yscale("log"); ax.grid(alpha=.25); ax.set_title(ttl,fontsize=10.5); ax.axvspan(UTC("2024-01-01"),END,color="#76b900",alpha=0.07)
axs[0].legend(fontsize=8.5,frameon=False,loc="upper left")
fig.tight_layout(); fig.savefig(OUT+"/trend_ls.png"); plt.close(fig)
