"""日线找底 + 周线确认融合策略"""
import pandas as pd, numpy as np

d = pd.read_parquet("/root/quant_pipeline/data/btc_daily.parquet")
w = pd.read_parquet("/root/quant_pipeline/data/btc_weekly.parquet")

DO, DH, DL, DC, DV = d["o"].values,d["h"].values,d["l"].values,d["c"].values,d["v"].values
nd = len(d)
d_idx = d.index

WO, WH, WL, WC = w["o"].values,w["h"].values,w["l"].values,w["c"].values
nw = len(w)
w_idx = w.index

# 将周线指标对齐到日线
w_sma20 = pd.Series(WC).rolling(20).mean().values
w_rsi14 = pd.Series(100 - 100/(1+pd.Series(WC).diff().clip(lower=0).rolling(14).mean()/(pd.Series(WC).diff().clip(upper=0).abs().rolling(14).mean()+1e-9))).values

# 每个日线对应的周线位置
# 找到每周的最后一天（周五），然后ffill到整个星期
week_map = {}
for i_w, date in enumerate(w_idx):
    week_end = date  # 周线index是周日
    week_map[week_end] = i_w

# 简化：直接看每周最后一个交易日
w_end_prices = WC
w_end_dates = w_idx

# 找到周线位置对应的日线index
# 取每个日线对应的周线index
d_week_end = d_idx.map(lambda x: x - pd.Timedelta(days=x.weekday()) + pd.Timedelta(days=6))  # 周日
# 使用merge_asof找最近周线
d["w_idx"] = pd.Series(d_idx).searchsorted(w_end_dates, side="right") - 1
d["w_idx"] = d["w_idx"].clip(0, nw-1)
d["w_sma20"] = w_sma20[d["w_idx"].values]
d["w_rsi"] = w_rsi14[d["w_idx"].values]

# ── 周线级别的状态判断 ──
d["in_bear_market"] = d["close"] < d["w_sma20"]  # 熊市区域
d["w_rsi_low"] = d["w_rsi"] < 30

# ── 日线找底信号 ──
# 1. 日线锤子线（长下影，收盘在顶部）
h_l = d["high"] - d["low"]
body = abs(d["close"] - d["open"])
d["hammer"] = (d["close"] > d["open"]) & ((d["low"] - d["open"]) / (h_l+1e-9) > 0.5) & (body/(h_l+1e-9) < 0.3)
# 2. 日线看涨吞没
d["engulf_bull"] = (d["close"].shift(1) < d["open"].shift(1)) & (d["close"] > d["open"]) & (d["open"] <= d["close"].shift(1)) & (d["close"] >= d["open"].shift(1))
# 3. 日线刺透
d["piercing"] = (d["close"].shift(1) < d["open"].shift(1)) & (d["close"] > d["open"]) & (d["open"] < d["low"].shift(1)) & (d["close"] > (d["open"].shift(1)+d["close"].shift(1))/2)
# 4. 放量反转（大跌后放量大涨）
ret_1 = d["close"].pct_change()
vol_ma20_i = d["volume"].rolling(20).mean().values
d["vol_surge"] = d["volume"] > vol_ma20_i * 1.5
d["reversal"] = (ret_1.shift(1) < -0.02) & (ret_1 > 0.02) & d["vol_surge"]

# ── 综合底部信号（日线形态+周线位置确认）──
d["buy_signal"] = d["in_bear_market"] & d["hammer"]
d["buy_signal2"] = d["in_bear_market"] & (d["engulf_bull"] | d["piercing"])
d["buy_signal3"] = d["in_bear_market"] & d["reversal"]
d["buy_all"] = d["buy_signal"] | d["buy_signal2"] | d["buy_signal3"]

# 只计算2020年之前的数据（避免未来数据污染）
train_mask = d_idx < "2023-01-01"
test_mask = d_idx >= "2023-01-01"

print("="*70)
print("🔥 日线找底 - 周线位置确认")
print("="*70)
print(f"数据: {nd}根日线, 训练集<2023, 测试集>=2023")

for name, col in [
    ("锤子线+熊市", "buy_signal"),
    ("吞没/刺透+熊市", "buy_signal2"),
    ("放量反转+熊市", "buy_signal3"),
    ("全部组合", "buy_all"),
]:
    for period_name, period_mask in [("训练", train_mask), ("测试", test_mask)]:
        mask = d[col].values & period_mask.values
        idxs = np.where(mask)[0]
        valid = [i for i in idxs if i+20 < nd]
        if len(valid) < 5: continue
        # 未来20天（约一个月）收益
        fwd = DC[np.array([min(i+20,nd-1) for i in valid])] / DC[np.array(valid)] - 1
        wr = sum(1 for r in fwd if r > 0.05) / len(valid)  # 涨超5%算成功
        avg = np.mean(fwd)
        print(f"  {name:<25s} {period_name}: {len(valid):>3d}笔 胜率{wr:>3.0%} 均收益{avg*100:+.0f}%")

# 最佳参数优化
print(f"\n{'='*70}")
print("🏆 测试集表现最好的组合")
print(f"{'='*70}")

# 只用测试集
for name, col in [
    ("锤子线+熊市", "buy_signal"),
    ("吞没+熊市", "buy_signal2"),
    ("放量反转", "buy_signal3"),
]:
    mask = d[col].values & test_mask.values
    idxs = np.where(mask)[0]
    valid = [i for i in idxs if i+20 < nd]
    if len(valid) < 3: continue
    print(f"\n  {name} ({len(valid)}笔):")
    for days in [5, 10, 20, 30]:
        fwd = DC[np.array([min(i+days,nd-1) for i in valid])] / DC[np.array(valid)] - 1
        wr = sum(1 for r in fwd if r > 0) / len(valid)
        avg = np.mean(fwd)
        print(f"    {days:>2d}天后: 胜率{wr:.0%} 均收益{avg*100:+.0f}%")
    
    # 模拟10x杠杆最佳参数
    best_params = None
    best_period = None
    best_cum = -99
    for days in [5, 10, 20, 30]:
        sl = 0.03
        tp = 0.06
        rets = []
        for i in valid:
            entry = DC[i]
            for j in range(1, min(days+1, nd-i)):
                r = DC[i+j]/entry - 1
                if r >= tp: rets.append(tp); break
                if r <= -sl: rets.append(-sl); break
            else:
                rets.append(DC[min(i+days,nd-1)]/entry - 1)
        if rets:
            cum = np.prod([1+r*10 for r in rets])
            wr = sum(1 for r in rets if r>0)/len(rets)
            if cum > best_cum:
                best_cum = cum
                best_params = (days, sl, tp, cum, wr, len(rets))
    print(f"    最佳: {best_params[0]}天 SL={best_params[1]*100:.0f}% TP={best_params[2]*100:.0f}% → 10x累计{best_params[3]:.1f}x 胜率{best_params[4]:.0%} ({best_params[5]}笔)")