package backtest import ( "context" "fmt" "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" "sig-pub/pkg/zlog" "time" "github.com/bytedance/sonic" ) // TradingPlanBacktester 交易计划回测 // todo 将backtest独立成单独服务横向扩展 type TradingPlanBacktester struct { indicatorReg *indicator.IndicatorRegistry sigStrategyReg *strategy.SigStrategyRegistry exchangeClient pb.ExchangeServiceClient plan entity.TradePlan account *BacktestTradeAccount // account *BacktestAccount sigStrategyType strategy.SigStrategyType sigStrategy strategy.ISigStrategy sigStrategyInput types.Input tradeStrategyInput types.Input closeStrategy trade.ICloseStrategy riskStrategy trade.IRiskStrategy tradeStrategy trade.ITradeStrategy instanceIntervalSigStrategyContext strategy.IInstanceIntervalSigStrategyContext } func NewTradingPlanBacktester( indicatorReg *indicator.IndicatorRegistry, sigStrategyReg *strategy.SigStrategyRegistry, exchangeClient pb.ExchangeServiceClient, ) *TradingPlanBacktester { return &TradingPlanBacktester{ indicatorReg: indicatorReg, sigStrategyReg: sigStrategyReg, exchangeClient: exchangeClient, } } func (b *TradingPlanBacktester) Init(cash float64, plan entity.TradePlan) (err error) { b.plan = plan // trade account simulator := NewTradeSimulator(0.0005, 0.0008) // b.account = NewBacktestAccount(cash, simulator) _ = simulator b.account = NewBacktestTradeAccount(cash) // sig strategy ok := false b.sigStrategyType, b.sigStrategy, ok = b.sigStrategyReg.NewSigStrategy(plan.SigStrategy) if !ok { err = fmt.Errorf("sig strategy not exists %s", plan.SigStrategy) return } 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, sr *pb.SeriesRange) (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: sr.Before, SeriesAfter: sr.After, Ctime: time.Now().UnixMilli(), Cash: b.account.cash, } // 交易信号回测器 sigStrategyBacktester := NewSigStrategyBacktester(b.sigStrategyType, b.sigStrategy, b.indicatorReg, b.exchangeClient) sigStrategyBacktester.SubKline(sr.InstId, 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, 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 } // 读最新的k线 kSeries := iiks.Get(sr.InstId, types.Interval(sr.Interval)) lastCandle, err := kSeries.Get(0) if err != nil { return } // 关闭所有未平仓仓位 b.forceCloseAllHoldingPosition(lastCandle) // 回测结果 test.Etime = time.Now().UnixMilli() test.EndCash = b.account.cash // test.Profit = b.account.profit // test.TotalTrades = len(b.account.trades) // test.WinningTrades = b.account.winningTrades // test.LosingTrades = b.account.losingTrades // test.Fee = b.account.fee // trades var trades []*Trade for _, trade := range b.account.trades { // trade.BacktestId = test.Id trade.Ctime = test.Ctime // trades = append(trades, trade) } collect.SortAsc(trades, func(t *Trade) int64 { return t.Id }) test.Trades = trades // 最大回撤 // drawdown := b.account.maxDrawdown // test.MaxDrawdown = max((drawdown[0]-drawdown[1])/drawdown[0], (drawdown[2]-drawdown[3])/drawdown[2]) return } // forceCloseAllHoldingPosition 关闭所有未平仓仓位 func (b *TradingPlanBacktester) forceCloseAllHoldingPosition(k types.Kline) (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) trades := b.account.GetOpenTrades(instId) for _, trd := range trades { closeTickets = append(closeTickets, trade.TradeTicket{ 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 { zlog.Errorf("close trade order error") } } 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, nil, instId, price) if err != nil { return } _ = closeTickets if len(closeTickets) == 0 { return } // var posErrs []error // positions := b.account.GetOpenTrades() // for _, pos := range positions { // closePos, cause := b.closeStrategy.OnKline(k, pos) // if closePos { // errc := b.account.ClosePosition(pos, k, cause) // if errc != nil { // posErrs = append(posErrs, errc) // zlog.Errorf("close position error: k=%#v, err=%v", k, errc) // } // } // } // if len(posErrs) > 0 { // err = errors.Join(posErrs...) // return // } return } // closeBySigSingal 交易信号出现时检查平仓 func (b *TradingPlanBacktester) closeBySigSingal(instId string, sigSide types.Side, kline types.Kline) (err error) { closeTickets, err := b.closeStrategy.CloseAssessOnSig(b.instanceIntervalSigStrategyContext, nil, instId, sigSide) if err != nil { return } _ = closeTickets if len(closeTickets) == 0 { return } // var posErrs []error // positions := b.account.GetOpenTrades() // for _, pos := range positions { // closePos, cause := b.closeStrategy.OnSigStrategySingal(sigSide, pos) // if closePos { // errc := b.account.ClosePosition(pos, kline, cause) // if errc != nil { // posErrs = append(posErrs, errc) // zlog.Error("close position error: ", errc) // } // } // } // if len(posErrs) > 0 { // err = errors.Join(posErrs...) // return // } return } // onSideSingal 出现买卖信号 func (b *TradingPlanBacktester) onSideSingal(instId string, sigSide types.Side, k types.Kline) (err error) { // 买卖信号交易风险分析 doTrade, causes, err := b.riskStrategy.RishAssess(b.instanceIntervalSigStrategyContext, nil, instId, sigSide) if err != nil { return } if !doTrade { _ = causes // todo 记录信号不交易原因分析 log db analyze return } ticket, err := b.tradeStrategy.TradeAssess(b.instanceIntervalSigStrategyContext, nil, instId, sigSide) if err != nil { return } order, err := b.account.MarketOrder(ticket) if err != nil { return } _ = order return }