"""debug: 检查"只加不减"的真实交易明细"""
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()

dd = df_raw.resample("1D").agg({"close":"last"}).dropna()
dd["ma20"] = dd["close"].rolling(20).mean()
dd["trend_up"] = dd["close"] > dd["ma20"]

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

tr = np.maximum(H-L, np.maximum(abs(H-np.roll(C,1)), abs(L-np.roll(C,1))))
atr_pct = pd.Series(tr).rolling(14).mean().values / C * 100

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

bull = (C1<O1) & (C>O) & (O<=C1) & (C>=O1)
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.005
nbull = np.roll(C,-1) > np.roll(O,-1)
long_sig = bull & near_s & nbull

swup = (H>pd.Series(H).shift(1).rolling(20).max().values) & (C<pd.Series(H).shift(1).rolling(20).max().values)
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.005
nbear = np.roll(C,-1) < np.roll(O,-1)
short_sig = 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_sig & trend_1h
ssig = short_sig & (~trend_1h)

split = int(n * 0.67)
SL, TP, MB = 0.015, 0.045, 48

# 只加不减参数
add_thresh, reduce_thresh, add_size, reduce_size, max_add, eval_interval = (0.002, 1.0, 0.3, 0.0, 3, 4)

trades_detail = []

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
            exit_reason = ""
            max_pos = 1.0
            
            for j in range(1, MB + 1):
                if i + j >= n: break
                ret = (C[i+j] / entry - 1) * dirc
                
                if j % eval_interval == 0 and j != last_eval:
                    last_eval = j
                    atr_val = atr_pct[i+j] if not np.isnan(atr_pct[i+j]) else 0.3
                    expected = 0.0005 * j + atr_val * 0.002 * j
                    expected = min(expected, TP * 0.8)
                    error = ret - expected
                    
                    if error > -add_thresh:
                        if add_count < max_add and pos_size < 2.0:
                            pos_size = min(2.0, pos_size + add_size)
                            add_count += 1
                            max_pos = max(max_pos, pos_size)
                
                if ret >= TP:
                    trades_detail.append({
                        "period": period, "i": i, "dir": dirc,
                        "entry": entry, "exit": C[i+j],
                        "ret": ret, "pos_size": pos_size,
                        "pnl": TP * pos_size - 0.001,
                        "bars": j, "adds": add_count,
                        "max_pos": max_pos,
                        "reason": "TP"
                    })
                    closed = True; break
                if ret <= -SL:
                    trades_detail.append({
                        "period": period, "i": i, "dir": dirc,
                        "entry": entry, "exit": C[i+j],
                        "ret": ret, "pos_size": pos_size,
                        "pnl": -SL * pos_size - 0.001,
                        "bars": j, "adds": add_count,
                        "max_pos": max_pos,
                        "reason": "SL"
                    })
                    closed = True; break
            
            if not closed:
                final_ret = (C[min(i+MB, n-1)] / entry - 1) * dirc
                trades_detail.append({
                    "period": period, "i": i, "dir": dirc,
                    "entry": entry, "exit": C[min(i+MB, n-1)],
                    "ret": final_ret, "pos_size": pos_size,
                    "pnl": final_ret * pos_size - 0.001,
                    "bars": MB, "adds": add_count,
                    "max_pos": max_pos,
                    "reason": "timeout"
                })

df_trades = pd.DataFrame(trades_detail)

print("=" * 80)
print("只加不减 - 交易明细分析")
print("=" * 80)

print(f"\n总交易数: {len(df_trades)}")
print(f"\n按平仓原因:")
for reason, grp in df_trades.groupby("reason"):
    print(f"  {reason}: {len(grp)}笔, 平均PnL={grp['pnl'].mean()*100:+.2f}%, "
          f"平均仓位={grp['pos_size'].mean():.2f}x, 平均加仓次数={grp['adds'].mean():.1f}")

print(f"\n按加仓次数:")
for adds, grp in df_trades.groupby("adds"):
    print(f"  加仓{adds}次: {len(grp)}笔, 平均PnL={grp['pnl'].mean()*100:+.2f}%, "
          f"胜率={sum(grp['pnl']>0)/len(grp):.1%}, 平均仓位={grp['pos_size'].mean():.2f}x")

print(f"\n关键问题检查:")
print(f"  TP时平均仓位: {df_trades[df_trades['reason']=='TP']['pos_size'].mean():.2f}x")
print(f"  SL时平均仓位: {df_trades[df_trades['reason']=='SL']['pos_size'].mean():.2f}x")
print(f"  SL时平均PnL: {df_trades[df_trades['reason']=='SL']['pnl'].mean()*100:+.2f}%")
print(f"  TP时平均PnL: {df_trades[df_trades['reason']=='TP']['pnl'].mean()*100:+.2f}%")

# 检查SL时的亏损是否被放大
sl_trades = df_trades[df_trades['reason']=='SL']
print(f"\n  SL交易中仓位>1的占比: {sum(sl_trades['pos_size']>1)/len(sl_trades):.1%}")
print(f"  SL交易中最大仓位: {sl_trades['pos_size'].max():.2f}x")
print(f"  SL交易中平均亏损: {sl_trades['pnl'].mean()*100:+.2f}%")
print(f"  原版SL固定亏损: {-SL*1.0*100-0.1:.2f}%")

# 对比：如果不用加仓，同样这些交易的表现
print(f"\n对比（不加仓的PnL vs 加仓后PnL）:")
df_trades['pnl_no_add'] = df_trades['ret'] - 0.001  # 不加仓的PnL
print(f"  不加仓累计: {np.prod(1+df_trades['pnl_no_add']):.3f}")
print(f"  加仓后累计: {np.prod(1+df_trades['pnl']):.3f}")
print(f"  不加仓胜率: {sum(df_trades['pnl_no_add']>0)/len(df_trades):.1%}")
print(f"  加仓后胜率: {sum(df_trades['pnl']>0)/len(df_trades):.1%}")