"""Value Area（市场轮廓）计算 + 叠加入烛龙回测
计算VAH/VAL/POC，作为信号质量加分的另一维度"""
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]
d = df_raw.copy()

O, H, L, C, V = d["open"].values, d["high"].values, d["low"].values, d["close"].values, d["volume"].values

# ── 1. 计算日级 Value Area ──
# 把15分钟K线汇总到日级，计算volume profile
# 用20个价格区间（bins）来算
N_BINS = 20

def calc_value_area(day_df):
    """计算一天的VAH, VAL, POC"""
    if len(day_df) < 10:
        return None
    
    day_high = day_df["high"].max()
    day_low = day_df["low"].min()
    price_range = day_high - day_low
    
    if price_range == 0:
        return None
    
    # 划分价格区间
    bin_size = price_range / N_BINS
    bins = [day_low + i * bin_size for i in range(N_BINS + 1)]
    
    # 计算每个价格区间的成交量
    vol_by_price = np.zeros(N_BINS)
    
    for _, row in day_df.iterrows():
        r_low = row["low"]
        r_high = row["high"]
        r_vol = row["volume"]
        
        # 这条K线的成交量按价格区间分配（简化：按中点归属）
        mid = (r_low + r_high) / 2
        bin_idx = min(N_BINS - 1, int((mid - day_low) / bin_size))
        vol_by_price[bin_idx] += r_vol
    
    # POC = 成交量最大的价格区间
    poc_idx = np.argmax(vol_by_price)
    poc = day_low + (poc_idx + 0.5) * bin_size
    
    # 从POC向两边扩展，直到累计70%成交量
    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:
        left_val = vol_by_price[left_idx] if left_idx >= 0 else 0
        right_val = vol_by_price[right_idx] if right_idx < N_BINS else 0
        
        if left_val >= right_val and left_idx >= 0:
            cum_vol += left_val
            left_idx -= 1
        elif right_idx < N_BINS:
            cum_vol += right_val
            right_idx += 1
        else:
            break
    
    val = day_low + (left_idx + 1) * bin_size
    vah = day_low + (right_idx) * bin_size
    
    return {"vah": vah, "val": val, "poc": poc}

# 逐日计算
daily = d.resample("1D").agg({"high":"max","low":"min","volume":"sum"}).dropna()
va_data = {}

for date in daily.index:
    date_str = date.date()
    day_df = d[d.index.date == date_str]
    result = calc_value_area(day_df)
    if result:
        va_data[date_str] = result

print(f"计算了 {len(va_data)} 天的 Value Area")

# 映射回15分钟级别
d["vah"] = np.nan
d["val"] = np.nan
d["poc"] = np.nan

for i in range(len(d)):
    date = d.index[i].date()
    if date in va_data:
        d.loc[d.index[i], "vah"] = va_data[date]["vah"]
        d.loc[d.index[i], "val"] = va_data[date]["val"]
        d.loc[d.index[i], "poc"] = va_data[date]["poc"]

# ── 2. 烛龙v1.6回测（含Value Area加分）──
dd = df_raw.resample("1D").agg({"close":"last"}).dropna()
dd["ma20"] = dd["close"].rolling(20).mean()
dd["trend_up"] = dd["close"] > dd["ma20"]

O1 = np.roll(O, 1); H1 = np.roll(H, 1); L1 = np.roll(L, 1); C1 = np.roll(C, 1)
O2 = np.roll(O, 2); C2 = np.roll(C, 2)

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

# 蜡烛形态
body = abs(C-O); body1 = abs(C1-O1); range_k = H - L
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

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_score(i, dirc):
    """共振评分 + Value Area加分"""
    score = 0
    
    # 放量加分
    if V[i] > vol_mean20[i] * 1.5: score += 15
    elif V[i] > vol_mean20[i] * 1.2: score += 8
    
    # Value Area加分：价格偏离价值区
    vah = d["vah"].iloc[i]
    val = d["val"].iloc[i]
    poc = d["poc"].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:  # 低于POC
            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):
    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_score(i, dirc)
                
                for j in range(1, MB + 1):
                    if i + j >= n: break
                    ret = (C[i+j] / entry - 1) * dirc
                    
                    if mode != "fixed":
                        eval_step = 3
                        if j % eval_step == 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":
                                current_ret = ret
                                
                                if res >= 30:
                                    base_max = 4; add_sz = 0.5
                                elif res >= 15:
                                    base_max = 3; add_sz = 0.35
                                else:
                                    base_max = 2; add_sz = 0.25
                                
                                if current_ret > 0.005: max_a = min(4, base_max + 1)
                                elif current_ret > 0: max_a = base_max
                                elif current_ret > -0.005: max_a = max(0, base_max - 1)
                                else: max_a = 0; add_sz = 0
                                
                                if add_count < max_a and pos_size < 2.0:
                                    pos_size = min(2.0, pos_size + add_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:
                    final_ret = (C[min(i+MB, n-1)] / entry - 1) * dirc
                    results[period].append(final_ret * pos_size - 0.001)
    
    return results, stats

def print_stats(results, stats, label):
    print(f"\n▶ {label}")
    if stats["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)
        sharpe = 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}% "
              f"周盈≈{pw:.1%} MDD={mdd:.1%} Sharpe={sharpe:.2f}")

print("=" * 80)
print("烛龙 v1.7：共振评分 + Value Area（市场轮廓理论）")
print("=" * 80)
print(f"SL={SL} TP={TP} MB={MB}h  数据: BTC 1H {n}根")

r1, s1 = backtest(lsig, ssig, "fixed")
print_stats(r1, s1, "固定止损（原版）")

r2, s2 = backtest(lsig, ssig, "addonly")
print_stats(r2, s2, "只加不减（v1.2）")

# 用同一份full逻辑，但用了含VA的resonance_score
# 需要单独跑v1.6（无VA）和v1.7（有VA）
# 先跑v1.7
r3, s3 = backtest(lsig, ssig, "full")
print_stats(r3, s3, "v1.7 共振+Value Area+自由能")