"""v1.6 vs v1.7 公平对比（统一1H频率）"""
import pandas as pd, numpy as np

df_raw = pd.read_parquet("data/btc_multidim.parquet")
df_raw.columns = [c.lower() for c in df_raw.columns]

# ── 统一resample到1H ──
d = df_raw.resample("1h").agg({"open":"first","high":"max","low":"min","close":"last","volume":"sum"}).dropna()
print(f"1H数据: {len(d)}根K线 ({d.index[0].date()} ~ {d.index[-1].date()})")

# ── Value Area（基于日级，映射回1H）──
daily = df_raw.resample("1D").agg({"high":"max","low":"min","volume":"sum"}).dropna()
N_BINS = 20

va_map = {}
for date in daily.index:
    date_str = date.date()
    day_df = df_raw[df_raw.index.date == date_str]
    if len(day_df) < 10:
        continue
    
    day_high = day_df["high"].max()
    day_low = day_df["low"].min()
    price_range = day_high - day_low
    if price_range == 0:
        continue
    
    bin_size = price_range / N_BINS
    vol_by_price = np.zeros(N_BINS)
    
    for _, r in day_df.iterrows():
        mid = (r["low"] + r["high"]) / 2
        bin_idx = min(N_BINS - 1, int((mid - day_low) / bin_size))
        vol_by_price[bin_idx] += r["volume"]
    
    poc_idx = np.argmax(vol_by_price)
    poc = day_low + (poc_idx + 0.5) * bin_size
    
    total_vol = vol_by_price.sum()
    target_vol = total_vol * 0.7
    cum_vol = vol_by_price[poc_idx]
    left_idx = poc_idx - 1
    right_idx = poc_idx + 1
    
    while cum_vol < target_vol:
        lv = vol_by_price[left_idx] if left_idx >= 0 else 0
        rv = vol_by_price[right_idx] if right_idx < N_BINS else 0
        if lv >= rv and left_idx >= 0:
            cum_vol += lv; left_idx -= 1
        elif right_idx < N_BINS:
            cum_vol += rv; right_idx += 1
        else:
            break
    
    val = day_low + (left_idx + 1) * bin_size
    vah = day_low + right_idx * bin_size
    
    # 用当天的VA数据
    # 同一个自然日内的1H K线共享当天的VA
    for i in range(len(d)):
        if d.index[i].date() == date_str:
            d.loc[d.index[i], "vah"] = vah
            d.loc[d.index[i], "val"] = val
            d.loc[d.index[i], "poc"] = poc

# ── 通用参数 ──
O, H, L, C, V = d["open"].values, d["high"].values, d["low"].values, d["close"].values, d["volume"].values
O1, H1, L1, C1 = np.roll(O,1), np.roll(H,1), np.roll(L,1), np.roll(C,1)
n = len(d)

prev_l20 = pd.Series(L).shift(1).rolling(20).min().values
prev_h20 = pd.Series(H).shift(1).rolling(20).max().values
near_s = abs(L - prev_l20) / (prev_l20 + 1e-9) < 0.008
near_r = abs(H - prev_h20) / (prev_h20 + 1e-9) < 0.008

bull_engulf = (C1<O1) & (C>O) & (O<=C1) & (C>=O1)
nbull = np.roll(C,-1) > np.roll(O,-1)
long_raw = bull_engulf & near_s & nbull

swup = (H>pd.Series(H).shift(1).rolling(20).max().values) & (C<pd.Series(H).shift(1).rolling(20).max().values)
nbear = np.roll(C,-1) < np.roll(O,-1)
short_raw = swup & near_r & nbear

dd = df_raw.resample("1D").agg({"close":"last"}).dropna()
dd["ma20"] = dd["close"].rolling(20).mean()
dd["trend_up"] = dd["close"] > dd["ma20"]
d_ts = d.index
trend_1h = np.array([dd["trend_up"].reindex([ts], method="ffill").values[0] if ts >= dd.index[0] else True for ts in d_ts])

lsig = long_raw & trend_1h
ssig = short_raw & (~trend_1h)

vol_mean20 = pd.Series(V).rolling(20).mean().values
split = int(n * 0.67)
SL, TP, MB = 0.015, 0.045, 48

# ── 两种共振打分 ──
def resonance_v1(i, dirc):
    """v1.6 无VA"""
    score = 0
    if V[i] > vol_mean20[i] * 1.5: score += 15
    elif V[i] > vol_mean20[i] * 1.2: score += 8
    return score

def resonance_v2(i, dirc):
    """v1.7 含VA"""
    score = resonance_v1(i, dirc)
    poc = d["poc"].iloc[i]
    val = d["val"].iloc[i]
    vah = d["vah"].iloc[i]
    
    if not pd.isna(poc):
        price = C[i]
        if dirc == 1 and price < val * 0.995:
            score += 20
        elif dirc == 1 and price < poc * 0.998:
            score += 10
        elif dirc == -1 and price > vah * 1.005:
            score += 20
        elif dirc == -1 and price > poc * 1.002:
            score += 10
    return score

def backtest(lsig, ssig, mode, resonance_fn=None):
    results = {"in": [], "out": []}
    stats = {"add": 0}
    
    for period, st, en in [("in", 0, split), ("out", split, n)]:
        for mask, dirc in [(lsig, 1), (ssig, -1)]:
            for i in range(st, min(en, n)):
                if not mask[i]: continue
                if i + MB >= n: continue
                
                entry = C[i]; closed = False; pos_size = 1.0
                add_count = 0; last_eval = 0
                
                res = resonance_fn(i, dirc) if resonance_fn else 0
                
                for j in range(1, MB + 1):
                    if i + j >= n: break
                    ret = (C[i+j] / entry - 1) * dirc
                    
                    if mode != "fixed":
                        if j % 3 == 0 and j != last_eval:
                            last_eval = j
                            
                            if mode == "addonly":
                                if add_count < 3 and pos_size < 2.0:
                                    pos_size = min(2.0, pos_size + 0.3)
                                    add_count += 1; stats["add"] += 1
                            elif mode == "full":
                                cr = ret
                                if res >= 30: bm=4; sz=0.5
                                elif res >= 15: bm=3; sz=0.35
                                else: bm=2; sz=0.25
                                if cr > 0.005: mxa = min(4, bm+1)
                                elif cr > 0: mxa = bm
                                elif cr > -0.005: mxa = max(0, bm-1)
                                else: mxa = 0; sz = 0
                                if add_count < mxa and pos_size < 2.0:
                                    pos_size = min(2.0, pos_size + sz)
                                    add_count += 1; stats["add"] += 1
                    
                    if ret >= TP:
                        results[period].append(TP * pos_size - 0.001); closed = True; break
                    if ret <= -SL:
                        results[period].append(-SL * pos_size - 0.001); closed = True; break
                
                if not closed:
                    fr = (C[min(i+MB, n-1)] / entry - 1) * dirc
                    results[period].append(fr * pos_size - 0.001)
    
    return results, stats

def ps(results, stats, label):
    print(f"\n▶ {label}")
    if stats and stats.get("add"): print(f"  [统计] 加仓={stats['add']}")
    for nm, key in [("样本内", "in"), ("样本外", "out")]:
        tr = results[key]
        if len(tr) < 5: print(f"  {nm}: 仅{len(tr)}笔"); continue
        wr = sum(1 for r in tr if r>0)/len(tr)
        cum = np.prod([1+r for r in tr])
        ch = [tr[i:i+5] for i in range(0,len(tr),5)]
        pw = sum(1 for c in ch if sum(c)>0)/len(ch) if ch else 0
        avg = np.mean(tr)*100
        running=1.0;peak=1.0;mdd=0
        for r in tr:
            running*=(1+r);peak=max(peak,running)
            mdd=min(mdd,(running-peak)/peak)
        sp = np.mean(tr)/(np.std(tr)+1e-9)*np.sqrt(365*24/MB) if np.std(tr)>0 else 0
        print(f"  {nm}: {len(tr)}笔 wr={wr:.1%} cum={cum:.3f} avg={avg:+.2f}% 周盈≈{pw:.1%} MDD={mdd:.1%} Sharpe={sp:.2f}")

print("="*70)
print("烛龙 v1.6 vs v1.7（统一1H频率公平对比）")
print("="*70)

r0,_ = backtest(lsig, ssig, "fixed")
ps(r0, {}, "固定止损（原版）")

r1,s1 = backtest(lsig, ssig, "addonly")
ps(r1, s1, "只加不减（v1.2）")

r2,s2 = backtest(lsig, ssig, "full", resonance_v1)
ps(r2, s2, "v1.6 共振评分（无VA）")

r3,s3 = backtest(lsig, ssig, "full", resonance_v2)
ps(r3, s3, "v1.7 共振评分 + Value Area")

print(f"\n{'='*70}")
print("📊 样本外对比：")
print(f"  {'策略':12s} {'笔数':>4s} {'胜率':>7s} {'累计':>8s} {'周盈':>7s} {'MDD':>8s} {'Sharpe':>7s}")

# 从上面的结果手动填入
pairs = [("固定止损",r0), ("只加不减",r1), ("v1.6",r2), ("v1.7",r3)]
for name, res in pairs:
    tr = res["out"]
    if len(tr) < 5: continue
    wr = sum(1 for r in tr if r>0)/len(tr)
    cum = np.prod([1+r for r in tr])
    ch = [tr[i:i+5] for i in range(0,len(tr),5)]
    pw = sum(1 for c in ch if sum(c)>0)/len(ch) if ch else 0
    running=1.0;peak=1.0;mdd=0
    for r in tr:
        running*=(1+r);peak=max(peak,running)
        mdd=min(mdd,(running-peak)/peak)
    sp = np.mean(tr)/(np.std(tr)+1e-9)*np.sqrt(365*24/MB) if np.std(tr)>0 else 0
    print(f"  {name:12s}: {len(tr):>4d}笔 {wr:>6.1%} {cum:>8.3f} {pw:>6.1%} {mdd:>7.1%} {sp:>7.2f}")