"""
《市场轮廓理论》: 价值区间(Value Area)叠加
只做价格偏离公平价值的反转
"""
import pandas as pd, numpy as np

df_raw = pd.read_parquet("data/btc_multidim.parquet")
d = df_raw.resample("1h").agg({"open":"first","high":"max","low":"min","close":"last","volume":"sum"}).dropna()
O,H,L,C,V = d["open"].values,d["high"].values,d["low"].values,d["close"].values,d["volume"].values
n = len(d)

# ── 价值区间计算 ──
# 用24小时(日级别)计算价值区间，映射到每小时
# 每个价格区间的成交量分布
FEE, SL, TP, MB = 0.001, 0.015, 0.045, 72

def calc_value_area(prices, volumes, num_bins=50):
    """计算日级别的价值区间: POC + VAH + VAL"""
    if len(prices)==0 or volumes.sum()==0:
        return None, None, None
    
    p_range = (prices.min(), prices.max())
    if p_range[1] - p_range[0] < 1: return None, None, None
    
    bins = np.linspace(p_range[0], p_range[1], num_bins)
    vol_profile = np.zeros(num_bins-1)
    
    for i in range(len(prices)):
        px = (prices.iloc[i] if hasattr(prices,'iloc') else prices[i])
        vol = (volumes.iloc[i] if hasattr(volumes,'iloc') else volumes[i])
        for j in range(len(bins)-1):
            if bins[j] <= px < bins[j+1]:
                vol_profile[j] += vol
                break
    
    if vol_profile.sum() == 0: return None, None, None
    
    # POC: 最大成交量价格
    poc_idx = np.argmax(vol_profile)
    poc = (bins[poc_idx] + bins[poc_idx+1]) / 2
    
    # 70%价值区间
    total_vol = vol_profile.sum()
    target = total_vol * 0.7
    best_start, best_end = 0, 0
    best_sum = 0
    
    for start in range(len(vol_profile)):
        current_sum = 0
        for end in range(start, len(vol_profile)):
            current_sum += vol_profile[end]
            if abs(current_sum - target) < abs(best_sum - target) or (best_start==0 and best_end==0):
                if current_sum > best_sum:
                    best_sum = current_sum
                    best_start, best_end = start, end
    
    vah = bins[best_end+1] if best_end+1 < len(bins) else bins[-1]
    val = bins[best_start]
    
    return poc, vah, val

# 逐日滑动计算
d_index = d.index
d_dates = pd.Series(d_index.date, index=d_index)

va_data = {}
for date in d_dates.unique():
    mask = d_dates == date
    prices = (H[mask.values] + L[mask.values] + C[mask.values]) / 3
    volumes = V[mask.values]
    poc, vah, val = calc_value_area(pd.Series(prices), pd.Series(volumes))
    if poc:
        va_data[date] = {"poc": poc, "vah": vah, "val": val}

# 映射到每小时
va_1h = np.array([va_data.get(d_index[i].date(), {}) for i in range(n)])
poc_1h = np.array([v.get("poc", C[i]) for i, v in enumerate(va_1h)])
vah_1h = np.array([v.get("vah", C[i]*1.02) for i, v in enumerate(va_1h)])
val_1h = np.array([v.get("val", C[i]*0.98) for i, v in enumerate(va_1h)])

# 价格相对价值区间的位置
above_va = C > vah_1h   # 高于价值区 → 偏贵
below_va = C < val_1h   # 低于价值区 → 便宜
in_va = (C >= val_1h) & (C <= vah_1h)  # 在价值区内

# 偏离幅度
deviation_pct = np.where(above_va, C/vah_1h-1, np.where(below_va, val_1h/C-1, 0)) * 100
strong_deviation = deviation_pct > 1.5  # 偏离超过1.5%

print(f"价值区间分布: 上方{above_va.sum()/n*100:.0f}% 区间内{in_va.sum()/n*100:.0f}% 下方{below_va.sum()/n*100:.0f}%")
print(f"强偏离(>1.5%): {strong_deviation.sum()/n*100:.0f}%")

# ── 蜡烛图信号 ──
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);body=abs(C-O);body1=abs(C1-O1)
nbull=np.roll(C,-1)>np.roll(O,-1);nbear=np.roll(C,-1)<np.roll(O,-1)
near_s=abs(L-pd.Series(L).shift(1).rolling(20).min().values)/(pd.Series(L).shift(1).rolling(20).min().values+1e-9)<0.008
near_r=abs(H-pd.Series(H).shift(1).rolling(20).max().values)/(pd.Series(H).shift(1).rolling(20).max().values+1e-9)<0.008

rev_long=((C1<O1)&(C>O)&(O<=C1)&(C>=O1)|(C2<O2)&(body1<(H-L)*0.3)&(C>O)&(C>(O2+C2)/2))&near_s&nbull
rev_short=((C1>O1)&(C<O)&(O>=C1)&(C<=O1)|(C2>O2)&(body1<(H-L)*0.3)&(C<O)&(C<(O2+C2)/2))&near_r&nbear

# 趋势 + ADX
dd=df_raw.resample("1D").agg({"close":"last"}).dropna();dd["ma20"]=dd["close"].rolling(20).mean()
du=dd["close"]>dd["ma20"];dts=d.index
trend_d=np.array([du.reindex([ts],method="ffill").values[0] if ts>=dd.index[0] else True for ts in dts])

tr=np.maximum(H-L,np.maximum(abs(H-np.roll(C,1)),abs(L-np.roll(C,1))))
atr14=pd.Series(tr).rolling(14).mean().values
pdi=100*pd.Series(np.where((H-np.roll(H,1)>np.roll(L,1)-L)&(H-np.roll(H,1)>0),H-np.roll(H,1),0)).rolling(14).mean().values/atr14
mdi=100*pd.Series(np.where((np.roll(L,1)-L>H-np.roll(H,1))&(np.roll(L,1)-L>0),np.roll(L,1)-L,0)).rolling(14).mean().values/atr14
adx=100*pd.Series(abs(pdi-mdi)/(pdi+mdi+1e-9)).rolling(14).mean().values
not_strong_trend=np.nan_to_num(adx,nan=0)<30

split=int(n*0.67)

# ── 逐层测试 ──
layers = [
    ("L0: 纯形态", rev_long&trend_d, rev_short&(~trend_d)),
    ("L1: +ADX<30", rev_long&trend_d&not_strong_trend, rev_short&(~trend_d)&not_strong_trend),
    ("L2: +偏离价值区", rev_long&trend_d&(below_va|above_va), rev_short&(~trend_d)&(below_va|above_va)),
    ("L3: +强偏离(>1.5%)", rev_long&trend_d&strong_deviation, rev_short&(~trend_d)&strong_deviation),
    ("L4: L1+L3全部", rev_long&trend_d&not_strong_trend&strong_deviation, rev_short&(~trend_d)&not_strong_trend&strong_deviation),
]

print(f"\n{'策略':20s} {'样本外':>20s}")
print(f"{'':20s} {'交易':>5s} {'胜率':>7s} {'累计':>8s} {'周盈':>7s}")
print("-"*58)

for label, ls, ss in layers:
    t=[]
    for mask,dirc in [(ls,1),(ss,-1)]:
        for i in range(split,min(n,len(mask))):
            if not mask[i] or i+MB>=n: continue
            entry=C[i];w=l=False
            for j in range(1,min(MB,n-i-1)):
                ret=(C[i+j]/entry-1)*dirc
                if ret<=-SL:t.append(-SL-FEE);l=True;break
                elif ret>=TP:t.append(TP-FEE);w=True;break
            if not w and not l:t.append((C[min(i+MB,n-1)]/entry-1)*dirc-FEE)
    if len(t)>5:
        wr=sum(1 for r in t if r>0)/len(t);cum=np.prod([1+r for r in t])
        ch=[t[i:i+5] for i in range(0,len(t),5)]
        pw=sum(1 for c in ch if sum(c)>0)/len(ch) if ch else 0
        print(f"{label:20s} {len(t):5d} {wr:6.1%} {cum:8.3f} {pw:6.1%}")
