from dataclasses import dataclass
from .engine import Candle,generate

def walk_forward(symbol_id:int,symbol:str,horizon:str,candles:list[Candle],spread_bps:float=2,slippage_bps:float=2,fee_bps:float=1,funding_bps:float=0)->dict:
    returns=[]; wins=losses=liquidations=0; regimes={}; calibration=[]
    for end in range(210,len(candles)-1):
        signal=generate(symbol_id,symbol,horizon,candles[:end]);
        if signal.direction=="NO_TRADE": continue
        entry=signal.entry*(1+(spread_bps+slippage_bps)/10000 if signal.direction=="LONG" else 1-(spread_bps+slippage_bps)/10000)
        future=candles[end:min(len(candles),end+20)]; exit_price=future[-1].close; liquidated=False
        for bar in future:
            liquidation_move=1/max(1,signal.recommended_leverage)
            if (signal.direction=="LONG" and bar.low<=entry*(1-liquidation_move)) or (signal.direction=="SHORT" and bar.high>=entry*(1+liquidation_move)):
                exit_price=entry*(1-liquidation_move if signal.direction=="LONG" else 1+liquidation_move);liquidated=True;liquidations+=1;break
            if signal.direction=="LONG" and bar.low<=signal.stop_loss: exit_price=signal.stop_loss;break
            if signal.direction=="SHORT" and bar.high>=signal.stop_loss: exit_price=signal.stop_loss;break
            if signal.direction=="LONG" and bar.high>=signal.take_profit: exit_price=signal.take_profit;break
            if signal.direction=="SHORT" and bar.low<=signal.take_profit: exit_price=signal.take_profit;break
        ret=((exit_price-entry)/entry)*(1 if signal.direction=="LONG" else -1)-(fee_bps+funding_bps)/10000;returns.append(ret);wins+=ret>0;losses+=ret<=0
        bucket=regimes.setdefault(signal.regime,{"trades":0,"wins":0,"return":0});bucket["trades"]+=1;bucket["wins"]+=ret>0;bucket["return"]+=ret
        calibration.append((signal.confidence/100,1 if ret>0 else 0))
    equity=1;peak=1;max_dd=0
    for value in returns: equity*=1+value;peak=max(peak,equity);max_dd=max(max_dd,(peak-equity)/peak)
    avg=sum(returns)/len(returns) if returns else 0
    variance=sum((x-avg)**2 for x in returns)/len(returns) if returns else 0
    regime_metrics={name:{"trades":x["trades"],"win_rate":x["wins"]/x["trades"],"expectancy":x["return"]/x["trades"]} for name,x in regimes.items()}
    calibration_error=sum(abs(predicted-actual) for predicted,actual in calibration)/len(calibration) if calibration else 0
    return {"symbol_id":symbol_id,"symbol":symbol,"horizon":horizon,"engine_version":"rules-v1","trades":len(returns),"win_rate":wins/len(returns) if returns else 0,"expectancy":avg,"max_drawdown":max_dd,"risk_adjusted_return":avg/(variance**.5) if variance else 0,"calibration_error":calibration_error,"turnover":len(returns),"liquidation_rate":liquidations/len(returns) if returns else 0,"by_regime":regime_metrics}
