package backtest import ( "context" "math" "sig-pub/api/pb" "sig-pub/internal/trading/sig" "sig-pub/pkg/data/entity" "sig-pub/pkg/indicator" "sig-pub/pkg/strategy" "sig-pub/pkg/trade" "sig-pub/pkg/types" "sig-pub/pkg/types/decimals" "sig-pub/pkg/utils/collect" "time" "github.com/bytedance/sonic" ) // TradingPlanBacktester 交易计划回测 // todo 将backtest独立成单独服务横向扩展 type TradingPlanBacktester struct { sigStrategyType strategy.SigStrategyType sigStrategy strategy.ISigStrategy indicatorReg *indicator.IndicatorRegistry exchangeClient pb.ExchangeServiceClient plan entity.TradePlan sr *pb.SeriesRange account *BacktestTradeAccount sigStrategyInput types.Input tradeStrategyInput types.Input closeStrategy trade.ICloseStrategy riskStrategy trade.IRiskStrategy tradeStrategy trade.ITradeStrategy instanceIntervalSigStrategyContext strategy.IInstanceIntervalSigStrategyContext } func NewTradingPlanBacktester( sigType strategy.SigStrategyType, strategy strategy.ISigStrategy, indicatorReg *indicator.IndicatorRegistry, exchangeClient pb.ExchangeServiceClient, ) *TradingPlanBacktester { return &TradingPlanBacktester{ sigStrategyType: sigType, sigStrategy: strategy, indicatorReg: indicatorReg, exchangeClient: exchangeClient, } } func (b *TradingPlanBacktester) Init(cash float64, plan entity.TradePlan, sr *pb.SeriesRange) (err error) { b.plan = plan b.sr = sr // trade account b.account = NewBacktestTradeAccount(cash) // sig strategy param if err = sonic.UnmarshalString(plan.SigStrategyParam, &b.sigStrategyInput); err != nil { return } if err = b.sigStrategy.Init(b.sigStrategyInput); err != nil { return } // 交易策略参数 var tradeStrategyInput, closeStrategyInput, riskStrategyInput types.Input if err = sonic.UnmarshalString(plan.TradeStrategyParam, &tradeStrategyInput); err != nil { return } if err = sonic.UnmarshalString(plan.CloseStrategyParam, &closeStrategyInput); err != nil { return } if err = sonic.UnmarshalString(plan.RiskStrategyParam, &riskStrategyInput); err != nil { return } b.tradeStrategyInput = make(types.Input) b.tradeStrategyInput.Assign(tradeStrategyInput, closeStrategyInput, riskStrategyInput) // 交易策略 tradeStrate := trade.NewSigTradeStrategy() if err = tradeStrate.Init(b.tradeStrategyInput); err != nil { return } // 平仓策略 b.closeStrategy = tradeStrate // 风险管理策略 b.riskStrategy = tradeStrate // 交易策略 b.tradeStrategy = tradeStrate return } // 核心引擎,模拟交易、持仓跟踪、费用计算 func (b *TradingPlanBacktester) Backtest(ctx context.Context) (test *BacktestTradingPlan, err error) { test = &BacktestTradingPlan{ Id: time.Now().Unix(), UserId: 10001, PlanId: b.plan.Id, InstId: b.plan.InstId, Exchange: pb.ExchangeType(b.plan.Exchange), Interval: b.plan.Interval, SeriesBefore: b.sr.Before, SeriesAfter: b.sr.After, Ctime: time.Now().UnixMilli(), Cash: b.account.cash, } // 交易信号回测器 sigStrategyBacktester := NewSigStrategyBacktester(b.sigStrategyType, b.sigStrategy, b.indicatorReg, b.exchangeClient) sigStrategyBacktester.SubKline(types.Interval1m, func(instId string, interval types.Interval, k types.Kline) (err error) { // 检查仓位平仓 return b.closeByKlineInterval1m(instId, k) }) iiks := types.NewInstanceIntervalKlineSeries() b.instanceIntervalSigStrategyContext = sig.NewInstanceIntervalSigStrategyContext(b.tradeStrategyInput, iiks, b.indicatorReg) err = sigStrategyBacktester.Backtest(ctx, b.sigStrategyInput, b.sr, iiks, func(instId string, sigSide types.Side, k types.Kline) (err error) { test.Singals++ // 根据交易信号检查仓位平仓 if err = b.closeBySigSingal(instId, sigSide, k); err != nil { return } // 交易下单 return b.onSideSingal(instId, sigSide, k) }) if err != nil { return } // 关闭所有未平仓仓位 err = b.forceCloseAllHoldingPosition() if err != nil { return } // 统计回测结果 test.Etime = time.Now().UnixMilli() test.EndCash = b.account.cash var profit float64 var orders []*trade.TradeOrder var snapshots []*EquitySnapshot // collect equity snapshots once per timestamp for _, order := range b.account.orders { switch order.TradeType { case trade.TradeTypeOpen: test.TotalTrades++ case trade.TradeTypeClose: if order.Profit > 0 { test.WinningTrades += len(order.Trades) } else { test.LosingTrades += len(order.Trades) } snapshots = append(snapshots, &EquitySnapshot{TradeId: order.TradeId, Ts: order.Ctime, Equity: order.Equity}) } test.Profit += order.Profit test.Fee += order.Fee order.BacktestId = test.Id profit += order.Profit orders = append(orders, order) } collect.SortAsc(orders, func(t *trade.TradeOrder) int64 { return t.TradeId }) test.Trades = orders // 最大回撤 drawdown := b.account.maxDrawdown test.MaxDrawdown = max((drawdown[0]-drawdown[1])/drawdown[0], (drawdown[2]-drawdown[3])/drawdown[2]) // 夏普比率: 使用 equity-curve method (periodic equity snapshots) collect.SortAsc(snapshots, func(sna *EquitySnapshot) int64 { return sna.TradeId }) var snaps []*EquitySnapshot var prevTs int64 for _, snap := range snapshots { if snap.Ts != prevTs { prevTs = snap.Ts snaps = append(snaps, snap) } else if len(snaps) > 0 { snaps[len(snaps)-1] = snap } } sharpeRatio := sharpeFromEquitySnapshots(snaps, 0.01) pow := math.Pow(10, float64(6)) test.SharpeRatio = math.Round(sharpeRatio*pow) / pow return } // forceCloseAllHoldingPosition 关闭所有未平仓仓位 func (b *TradingPlanBacktester) forceCloseAllHoldingPosition() (err error) { var closeTickets []trade.TradeTicket for _, instId := range b.account.GetOpenTradeInsts() { k := b.instanceIntervalSigStrategyContext.Get(instId, trade.PriceDriverInterval, 0) price := decimals.MustToFloat64(k.Close) b.account.OnPrice(instId, price) trades := b.account.GetOpenTrades(instId) for _, trd := range trades { closeTickets = append(closeTickets, trade.TradeTicket{ TradeType: trade.TradeTypeClose, InstId: instId, TradesId: []int64{trd.TradeId}, Side: trd.Side.Opposite(), Price: price, Leverage: trd.Leverage, Qty: trd.Qty, Interval: string(k.Interval), Ktime: k.Interval.MustAddMul(k.Ts, 1), Ctime: time.Now().UnixMilli(), Cause: trade.CauseCloseForced, }) } } for _, ticket := range closeTickets { err = b.account.CloseTradeOrder(ticket) if err != nil { return } } return } // closeByKlineInterval1m k线更新时检查平仓 func (b *TradingPlanBacktester) closeByKlineInterval1m(instId string, k types.Kline) (err error) { price := decimals.MustToFloat64(k.Close) closeTickets, err := b.closeStrategy.CloseAssessOnPrice(b.instanceIntervalSigStrategyContext, b.account, instId, price) if err != nil { return } if len(closeTickets) == 0 { return } for _, ticket := range closeTickets { err = b.account.CloseTradeOrder(ticket) if err != nil { return } } return } // closeBySigSingal 交易信号出现时检查平仓 func (b *TradingPlanBacktester) closeBySigSingal(instId string, sigSide types.Side, _ types.Kline) (err error) { closeTickets, err := b.closeStrategy.CloseAssessOnSig(b.instanceIntervalSigStrategyContext, b.account, instId, sigSide) if err != nil { return } if len(closeTickets) == 0 { return } for _, ticket := range closeTickets { err = b.account.CloseTradeOrder(ticket) if err != nil { return } } return } // onSideSingal 出现买卖信号 func (b *TradingPlanBacktester) onSideSingal(instId string, sigSide types.Side, _ types.Kline) (err error) { // 买卖信号交易风险分析 doTrade, causes, err := b.riskStrategy.RishAssess(b.instanceIntervalSigStrategyContext, b.account, instId, sigSide) if err != nil { return } if !doTrade { _ = causes // todo 记录信号不交易原因分析 log db analyze return } tickets, err := b.tradeStrategy.TradeAssess(b.instanceIntervalSigStrategyContext, b.account, instId, sigSide) if err != nil || len(tickets) == 0 { return } for _, ticket := range tickets { _, err = b.account.MarketOrder(ticket) if err != nil { return } } return }