package backtest import ( "context" "errors" "fmt" "io" "sig-pub/api/pb" "sig-pub/internal/trading/sig" "sig-pub/pkg/indicator" "sig-pub/pkg/strategy" "sig-pub/pkg/types" "sig-pub/pkg/utils/collect" "sig-pub/pkg/utils/times" "sig-pub/pkg/zlog" "google.golang.org/grpc" ) // SigStrategyBacktester 信号策略回测 type SigStrategyBacktester struct { sigStrategyType strategy.SigStrategyType sigStrategy strategy.ISigStrategy indicatorReg *indicator.IndicatorRegistry exchangeClient pb.ExchangeServiceClient // 回测过程中订阅k线 intervalSubscribe map[types.Interval][]func(instId string, interval types.Interval, k types.Kline) (err error) } func NewSigStrategyBacktester( sigStrategyType strategy.SigStrategyType, sigStrategy strategy.ISigStrategy, indicatorReg *indicator.IndicatorRegistry, exchangeServiceClient pb.ExchangeServiceClient, ) *SigStrategyBacktester { return &SigStrategyBacktester{ sigStrategyType: sigStrategyType, sigStrategy: sigStrategy, indicatorReg: indicatorReg, exchangeClient: exchangeServiceClient, intervalSubscribe: make(map[types.Interval][]func(instId string, interval types.Interval, k types.Kline) (err error)), } } // SubKline 在回测过程中订阅k线 func (b *SigStrategyBacktester) SubKline(instId string, interval types.Interval, recv func(instId string, interval types.Interval, k types.Kline) (err error)) { b.intervalSubscribe[interval] = append(b.intervalSubscribe[interval], recv) } // Backtest 基于历史数据回测信号策略 func (b *SigStrategyBacktester) Backtest(ctx context.Context, sigStrategyInput types.Input, sr *pb.SeriesRange, iiks *types.InstanceIntervalKlineSeries, recvSignal func(instId string, sigSide types.Side, k types.Kline) (err error), ) (err error) { if iiks == nil { iiks = types.NewInstanceIntervalKlineSeries() } // init sig strategy if err = b.sigStrategy.Init(sigStrategyInput); err != nil { return } switch b.sigStrategyType { case strategy.SigStrategyTypeSingle: err = b.singleStrategySeries(ctx, b.sigStrategy.(strategy.ISingleSigStrategy), sigStrategyInput, sr, iiks, recvSignal) case strategy.SigStrategyTypeInterval: err = b.intervalStrategySeries(ctx, b.sigStrategy.(strategy.IIntervalSigStrategy), sigStrategyInput, sr, iiks, recvSignal) case strategy.SigStrategyTypeInstanceInterval: err = b.instanceIntervalStrategySeries(ctx, b.sigStrategy.(strategy.IInstanceIntervalSigStrategy), sigStrategyInput, sr, iiks, recvSignal) default: err = fmt.Errorf("unknown sig strategy type %v", b.sigStrategyType) } return } // singleStrategySeries 单周期策略 func (b *SigStrategyBacktester) singleStrategySeries(ctx context.Context, sigStrategy strategy.ISingleSigStrategy, sigStrategyInput types.Input, sr *pb.SeriesRange, iiks *types.InstanceIntervalKlineSeries, recvSignal func(instId string, sigSide types.Side, k types.Kline) (err error), ) (err error) { driverInstId := sr.InstId driverInterval := types.Interval(sr.Interval) driverSeries := iiks.Get(driverInstId, driverInterval) strategyContext := sig.NewStrategyContext(sigStrategyInput, driverSeries, b.indicatorReg) requiredPeriods := int(sigStrategy.CandlePeriods(strategyContext)) intervalCandlePeriods := types.NewIntervalState[int16]() intervalCandlePeriods.Set(driverInterval, int16(max(1, requiredPeriods))) err = b.multiInstanceIntervalSeries(ctx, sr, []string{sr.InstId}, intervalCandlePeriods, iiks, func(driver bool, instId string, interval types.Interval, k *types.Kline) (err error) { if !driver || instId != driverInstId || interval != driverInterval { return } if driverSeries.Length() < requiredPeriods { return } sigSide := sigStrategy.Update(strategyContext) if sigSide.IsValid() { if err = recvSignal(driverInstId, sigSide, *k); err != nil { return } } return }) return } // intervalStrategySeries 多周期策略 func (b *SigStrategyBacktester) intervalStrategySeries(ctx context.Context, intervalSigStrategy strategy.IIntervalSigStrategy, sigStrategyInput types.Input, sr *pb.SeriesRange, iiks *types.InstanceIntervalKlineSeries, recvSignal func(instId string, sigSide types.Side, k types.Kline) (err error), ) (err error) { driverInstId := sr.InstId driverInterval := types.Interval(sr.Interval) intervalKlineSeries := iiks.GetIntervalKlineSeries(driverInstId) // 策略上下文 intervalStrategyContext := sig.NewIntervalStrategyContext(sigStrategyInput, intervalKlineSeries, b.indicatorReg) // 各周期所需k线数量 intervalCandlePeriods := intervalSigStrategy.CandlePeriods(intervalStrategyContext) periodsChecked := false err = b.multiInstanceIntervalSeries(ctx, sr, []string{sr.InstId}, intervalCandlePeriods, iiks, func(driver bool, instId string, interval types.Interval, k *types.Kline) (err error) { if !driver || instId != driverInstId || interval != driverInterval { return } // 策略所需周期窗口是否满足检查 if !periodsChecked { update := true intervalCandlePeriods.Range(func(interval types.Interval, require int16) { if update && require > 0 { series := intervalKlineSeries.Get(interval) update = series.Length() >= int(require) } }) if !update { return } periodsChecked = true } sigSide := intervalSigStrategy.Update(intervalStrategyContext) if sigSide.IsValid() { if err = recvSignal(driverInstId, sigSide, *k); err != nil { return } } return }) return } // instanceIntervalStrategySeries 多币种多周期策略 func (b *SigStrategyBacktester) instanceIntervalStrategySeries(ctx context.Context, intervalSigStrategy strategy.IInstanceIntervalSigStrategy, sigStrategyInput types.Input, sr *pb.SeriesRange, iiks *types.InstanceIntervalKlineSeries, recvSignal func(instId string, sigSide types.Side, k types.Kline) (err error), ) (err error) { driverInstId := sr.InstId driverInterval := types.Interval(sr.Interval) // 策略上下文 strategyContext := sig.NewInstanceIntervalSigStrategyContext(sigStrategyInput, iiks, b.indicatorReg) // 各周期所需k线数量 tradeInsts, intervalCandlePeriods := intervalSigStrategy.CandlePeriods(strategyContext) periodsChecked := false err = b.multiInstanceIntervalSeries(ctx, sr, tradeInsts, intervalCandlePeriods, iiks, func(driver bool, instId string, interval types.Interval, k *types.Kline) (err error) { if !driver || instId != driverInstId || interval != driverInterval { return } // 策略所需周期窗口是否满足检查 if !periodsChecked { update := true for _, instId := range tradeInsts { intervalCandlePeriods.Range(func(interval types.Interval, periods int16) { if update && periods > 0 { series := iiks.Get(instId, interval) update = series.Length() >= int(periods) } }) if !update { return } } periodsChecked = true } sigSideInsts := intervalSigStrategy.Update(strategyContext) for _, si := range sigSideInsts { if err = recvSignal(si.InstId, si.Side, *k); err != nil { return } } return }) return } var errStop = errors.New("stop") // multiInstanceIntervalSeries 多币种多周期数据拉取 func (b *SigStrategyBacktester) multiInstanceIntervalSeries(ctx context.Context, sr *pb.SeriesRange, tradeInsts []string, intervalCandlePeriods *types.IntervalState[int16], iiks *types.InstanceIntervalKlineSeries, recvFn func(driver bool, instId string, interval types.Interval, k *types.Kline) (err error), ) (err error) { driverInstId := sr.InstId driverInterval := types.Interval(sr.Interval) driverIntervalAdder := types.SupportedIntervals[driverInterval] // 查询主周期时间范围 rsp, err := b.exchangeClient.SeriesRange(ctx, &pb.ReqSeriesRange{Series: sr}) if err != nil { return } driverBefore, driverAfter := rsp.Before, driverIntervalAdder(rsp.After, 1) // 运行时周期 fetchIntervals := []types.Interval{driverInterval} intervalCandlePeriods.Range(func(interval types.Interval, window int16) { if window > 0 { fetchIntervals = append(fetchIntervals, interval) } }) // 运行时订阅周期 for interval := range b.intervalSubscribe { fetchIntervals = append(fetchIntervals, interval) } fetchIntervals = collect.Uniq(fetchIntervals) // 交易产品 fetchInsts := append([]string{driverInstId}, tradeInsts...) fetchInsts = collect.Uniq(fetchInsts) // fetch kline series var otherSrs []*pb.SeriesRange for _, instId := range fetchInsts { for _, interval := range fetchIntervals { if instId == driverInstId && interval == driverInterval { continue } intervalAdder := types.SupportedIntervals[interval] isr := &pb.SeriesRange{Exchange: sr.Exchange, Open: false, Live: sr.Live, Desc: sr.Desc} isr.InstId = instId isr.Interval = string(interval) isr.Before = intervalAdder(driverBefore, -1) isr.After = driverAfter isr.WindowExtra = uint32(max(0, intervalCandlePeriods.Get(interval)-1)) + indicator.ApproCandles otherSrs = append(otherSrs, isr) } } stopCh := make(chan struct{}) syncChans := make([]chan int64, 0, len(otherSrs)) for _, isr := range otherSrs { kSeries := iiks.Get(isr.InstId, types.Interval(isr.Interval)) syncCh := make(chan int64) syncChans = append(syncChans, syncCh) go func(isr *pb.SeriesRange, kSeries *types.KlineSeries, syncCh chan int64) { interval := types.Interval(isr.Interval) intervalAdder := types.SupportedIntervals[interval] driverTs := int64(0) isr.WindowExtra = uint32(max(0, intervalCandlePeriods.Get(interval)-1)) + indicator.ApproCandles err1 := b.fetchHistoryKlineSeries(ctx, isr, func(k *types.Kline) (err error) { closeTs := intervalAdder(k.Ts, 1) // 与驱动周期series保持同步更新 if closeTs > driverTs { waitLoop: for { if driverTs != 0 { syncCh <- 0 // 响应主周期更新完毕 } select { case <-ctx.Done(): return errStop case <-stopCh: return errStop case driverTs = <-syncCh: // 等待主周期通知更新 if closeTs <= driverTs { break waitLoop } } } } if lastTs, serial := kSeries.Update(k); !serial { err = fmt.Errorf("kline not series: %s(%s), interval=%s, lastTs=%d", sr.InstId, sr.Exchange, interval, lastTs) return } // 回调周期订阅 for _, subFn := range b.intervalSubscribe[interval] { if err = subFn(driverInstId, interval, *k); err != nil { return } } // 运行时周期 if intervalCandlePeriods.Get(interval) > 0 { return recvFn(false, isr.InstId, interval, k) } return }) zlog.Debugf("other sr finish with: %s(%s), %v, last=%d", isr.InstId, isr.Interval, err1, intervalAdder(kSeries.MustGet(0).Ts, 1)) if err1 == nil { syncCh <- -1 // 通知更新完毕, 后续不再更新 } else if err1 != errStop { zlog.Errorf("fetch interval history error: inst=%s(%s) interval=%s, err=%v", sr.InstId, sr.Exchange, interval, err1) err = err1 close(stopCh) } }(isr, kSeries, syncCh) } // 驱动交易产品周期数据拉取 driverSeries := iiks.Get(driverInstId, driverInterval) sr.WindowExtra = max( sr.WindowExtra, uint32(max(0, intervalCandlePeriods.Get(driverInterval)-1)), ) + indicator.ApproCandles err0 := b.fetchHistoryKlineSeries(ctx, sr, func(k *types.Kline) (err error) { // update kline series if lastTs, serial := driverSeries.Update(k); !serial { err = fmt.Errorf("kline not series: %s(%s), interval=%s, lastTs=%d", sr.InstId, sr.Exchange, driverInterval, lastTs) return } driverTS := driverIntervalAdder(k.Ts, 1) for i, syncCh := range syncChans { if syncCh == nil { continue } syncCh <- driverTS // 通知其他周期更新到主周期时间 select { case sig := <-syncCh: // 等待该周期更新完毕 if sig == -1 { // 后续不再更新 syncChans[i] = nil } case <-ctx.Done(): return errStop case <-stopCh: return errStop } } if k.Ts < driverBefore { return } // 回调周期订阅 for _, subFn := range b.intervalSubscribe[driverInterval] { if err = subFn(driverInstId, driverInterval, *k); err != nil { return } } return recvFn(true, driverInstId, driverInterval, k) }) zlog.Debugf("driver sr finish with: %s(%s), %v, last=%d", sr.InstId, sr.Interval, err0, driverIntervalAdder(driverSeries.MustGet(0).Ts, 1)) if err0 == nil { close(stopCh) } else if err0 != errStop { zlog.Errorf("fetch driver interval history error: inst=%s(%s) interval=%s, err=%v", sr.InstId, sr.Exchange, driverInterval, err) err = err0 close(stopCh) return } return } // fetchHistoryKlineSeries 请求k线数据流式处理 func (b *SigStrategyBacktester) fetchHistoryKlineSeries(ctx context.Context, sr *pb.SeriesRange, recvFn func(k *types.Kline) error) (err error) { // fetch history klines via stream req := &pb.ReqHistoryKlineStream{Series: sr} stream, err := b.exchangeClient.HistoryKlineStream(ctx, req, grpc.UseCompressor("snappy")) if err != nil { return } var msg *pb.RspHistoryKlineStream recvTimes, recvTotal := 0, 0 watch := times.NewWatch() recvLoop: for { select { case <-ctx.Done(): err = ctx.Err() return default: } msg, err = stream.Recv() if err == io.EOF { err = nil break } if err != nil { break } recvTimes++ recvTotal += len(msg.Klines) for _, k := range msg.Klines { kline := new(types.Kline) kline.ParsePBKline(sr.Exchange, k) if err = recvFn(kline); err != nil { break recvLoop } } } zlog.Debugf("fetch history kline series: inst=%s(%s), interval=%s, recv=%d, total=%d, use %s", sr.InstId, sr.Exchange, sr.Interval, recvTimes, recvTotal, watch.ElapsedFmt(".")) return }