"""烛龙 v1.8：信息熵（Shannon Entropy）市场状态过滤
熵低→趋势清晰→放心加仓  熵高→市场随机→保守操作"""
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]

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

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)

# ── 1. 信息熵计算 ──
def calc_market_entropy(closes, window=48):
    """用收益率方向持续性计算市场熵
    核心思想：连续同向 = 低熵(有序)  频繁反转 = 高熵(随机)
    返回值: 0(完全有序/趋势) ~ 1(完全随机/震荡)
    """
    rets = np.diff(closes) / closes[:-1]
    
    # 用符号判断方向：+1涨, -1跌, 阈值过滤平盘
    thresh = 0.0005
    signs = np.zeros(len(rets))
    signs[rets > thresh] = 1
    signs[rets < -thresh] = -1
    
    entropies = np.full(len(closes), 0.5)
    
    for i in range(window, len(rets)):
        seg = signs[i-window:i]
        
        # 排除平盘
        seg = seg[seg != 0]
        if len(seg) < 10:
            continue
        
        # 计算"持续性"：连续同方向的占比
        # 如果>70%的相邻对是同方向→低熵(趋势)
        # 如果~50%→高熵(随机)
        same_dir = np.mean(seg[1:] == seg[:-1])
        
        # 映射到0-1熵值
        # same_dir=0.5(完全随机)→熵=1
        # same_dir=1(完全趋势)→熵=0
        # same_dir=0(完全反转)→熵=0(也是有序的!)
        # 用正态化: 用|same_dir-0.5|*2 作为有序度
        order = abs(same_dir - 0.5) * 2  # 0=随机, 1=有序
        entropies[i+1] = 1 - order  # 0=有序, 1=随机
    
    # 熵值平滑
    entropies = pd.Series(entropies).rolling(5, min_periods=1).mean().values    
    return entropies

entropy = calc_market_entropy(C, window=48)

# ── 打印几个样本 ──
valid_entropy = entropy[~np.isnan(entropy)]
print(f"熵值范围: {valid_entropy.min():.3f} ~ {valid_entropy.max():.3f}")
print(f"熵值均值: {valid_entropy.mean():.3f}")
for q in [0.1, 0.25, 0.5, 0.75, 0.9]:
    print(f"  {q*100:.0f}%分位: {np.percentile(valid_entropy, q*100):.3f}")

q10 = np.percentile(valid_entropy, 10)
q25 = np.percentile(valid_entropy, 25)
q75 = np.percentile(valid_entropy, 75)
q90 = np.percentile(valid_entropy, 90)
print(f"\n用分位数做阈值:")
print(f"  低熵(<10%分位={q10:.3f}): 信号加分")
print(f"  高熵(>90%分位={q90:.3f}): 信号减分或跳过")

# ── 看看低熵和高熵时期的样子 ──
low_ent_idx = entropy < q10
high_ent_idx = entropy > q90
print(f"\n低熵时段(<{q10:.3f}): {low_ent_idx.sum()}根K线")
print(f"高熵时段(>{q90:.3f}): {high_ent_idx.sum()}根K线")

# ── 2. Value Area（复用）──
daily = df_raw.resample("1D").agg({"high":"max","low":"min","volume":"sum"}).dropna()
N_BINS = 20
va_map = {}
for date in daily.index:
    day_df = df_raw[df_raw.index.date == date.date()]
    if len(day_df) < 10: continue
    dh = day_df["high"].max(); dl = day_df["low"].min()
    pr = dh - dl
    if pr == 0: continue
    bs = pr / N_BINS
    vbp = np.zeros(N_BINS)
    for _, r in day_df.iterrows():
        mid = (r["low"] + r["high"]) / 2
        bi = min(N_BINS-1, int((mid - dl) / bs))
        vbp[bi] += r["volume"]
    poci = np.argmax(vbp)
    poc = dl + (poci + 0.5) * bs
    tv = vbp.sum(); tgt = tv * 0.7; cv = vbp[poci]
    li = poci - 1; ri = poci + 1
    while cv < tgt:
        lv = vbp[li] if li >= 0 else 0
        rv = vbp[ri] if ri < N_BINS else 0
        if lv >= rv and li >= 0: cv += lv; li -= 1
        elif ri < N_BINS: cv += rv; ri += 1
        else: break
    val = dl + (li + 1) * bs
    vah = dl + ri * bs
    for i in range(len(d)):
        if d.index[i].date() == date.date():
            d.loc[d.index[i], "vah"] = vah
            d.loc[d.index[i], "val"] = val
            d.loc[d.index[i], "poc"] = poc

# ── 3. 烛龙信号 ──
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

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

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

def resonance_v8(i, dirc, q10, q90):
    """v1.8 含VA+信息熵（动态阈值）"""
    s = resonance_v7(i, dirc)
    ent = entropy[i] if i < len(entropy) and not np.isnan(entropy[i]) else 0.5
    
    if ent < q10:
        s += 20  # 低熵→有序市场→信号可靠
    elif ent < q25:
        s += 10
    elif ent > q90:
        s -= 15  # 高熵→混乱市场→信号可疑
    elif ent > q75:
        s -= 5
    
    return s

def backtest(lsig, ssig, mode, resonance_fn=None, q10=None, q90=None):
    results = {"in": [], "out": []}
    stats = {"add": 0, "signals_skipped": 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
                
                # 信息熵极端高→直接跳过这个信号（v1.8专属）
                if mode == "full_v8" and q90 is not None and entropy[i] > q90:
                    stats["signals_skipped"] += 1
                    continue
                
                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
                            else:
                                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 mode == "full_v8" and q75 is not None and entropy[i] > q75:
                                    bm = max(0, bm - 1)
                                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}")
    extra = []
    if stats.get("add"): extra.append(f"加仓={stats['add']}")
    if stats.get("signals_skipped"): extra.append(f"跳过={stats['signals_skipped']}")
    if extra: print(f"  [统计] {' '.join(extra)}")
    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.8：信息熵市场状态过滤")
print("="*70)

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

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

r6,s6 = backtest(lsig, ssig, "full", resonance_v7)
ps(r6, s6, "v1.7 共振+Value Area")

r8,s8 = backtest(lsig, ssig, "full_v8", lambda i,d: resonance_v8(i,d,q10,q90), q10, q90)
ps(r8, s8, "v1.8 共振+VA+信息熵")

print(f"\n{'='*70}")
print("📊 样本外对比：")
for name, res, s in [("固定止损", r0, {}), ("v1.2", r1, s1), ("v1.7", r6, s6), ("v1.8", r8, s8)]:
    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
    sk = s.get("signals_skipped",0)
    sk_str = f" 跳过{sk}" if sk else ""
    print(f"  {name:12s}: {len(tr):>3d}笔 wr={wr:>5.1%} cum={cum:>7.3f} 周盈≈{pw:>5.1%} MDD={mdd:>6.1%} Sharpe={sp:>5.2f}{sk_str}")