package backtest import ( "context" "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" "sync/atomic" "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(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(interval types.Interval, k types.Kline) (err error)), } } // SubKline 在回测过程中订阅k线 func (b *SigStrategyBacktester) SubKline(interval types.Interval, recv func(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, cleanIntervalSeries *types.IntervalState[*sig.KlineSeries], recvSignal func(sigSide types.Side, k types.Kline) (err error)) (err error) { if cleanIntervalSeries == nil { cleanIntervalSeries = types.NewIntervalState[*sig.KlineSeries]() } // 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, cleanIntervalSeries, recvSignal) case strategy.SigStrategyTypeInterval: err = b.intervalStrategySeries(ctx, b.sigStrategy.(strategy.IIntervalSigStrategy), sigStrategyInput, sr, cleanIntervalSeries, 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, intervalKlineSeries *types.IntervalState[*sig.KlineSeries], recvSignal func(sigSide types.Side, k types.Kline) (err error)) (err error) { interval := types.Interval(sr.Interval) kSeries := intervalKlineSeries.ComputeIfAbsent(interval, func() *sig.KlineSeries { return sig.NewKlineSeries(sr.Exchange, sr.InstId, interval) }) strategyContext := sig.NewStrategyContext(sigStrategyInput, kSeries, b.indicatorReg) requiredSeries := int(sigStrategy.CandlePeriods(strategyContext)) requiredIntervalSeries := types.NewIntervalState[int16]() requiredIntervalSeries.Set(interval, int16(max(1, requiredSeries))) err = b.multiIntervalSeries(ctx, sr, requiredIntervalSeries, intervalKlineSeries, func(driver bool, interval types.Interval, k *types.Kline) (err error) { if !driver { return } kSeries := intervalKlineSeries.Get(interval) if kSeries.Length() < requiredSeries { return } sigSide := sigStrategy.Update(strategyContext) if sigSide.IsValid() { if err = recvSignal(sigSide, *k); err != nil { return } } return }) return } // intervalStrategySeries 多周期策略 func (b *SigStrategyBacktester) intervalStrategySeries(ctx context.Context, intervalSigStrategy strategy.IIntervalSigStrategy, sigStrategyInput types.Input, sr *pb.SeriesRange, intervalKlineSeries *types.IntervalState[*sig.KlineSeries], recvSignal func(sigSide types.Side, k types.Kline) (err error)) (err error) { // 策略上下文 intervalStrategyContext := sig.NewIntervalStrategyContext(sigStrategyInput, intervalKlineSeries, b.indicatorReg) // 各周期所需k线数量 requiredIntervalSeries := intervalSigStrategy.CandlePeriods(intervalStrategyContext) err = b.multiIntervalSeries(ctx, sr, requiredIntervalSeries, intervalKlineSeries, func(driver bool, interval types.Interval, k *types.Kline) (err error) { if !driver { return } update := true requiredIntervalSeries.Range(func(interval types.Interval, require int16) { if update && require > 0 { series := intervalKlineSeries.Get(interval) update = series.Length() >= int(require) } }) if !update { return } sigSide := intervalSigStrategy.Update(intervalStrategyContext) if sigSide.IsValid() { if err = recvSignal(sigSide, *k); err != nil { return } } return }) return } // multiIntervalSeries 多周期k线数据拉取 // intervalKlineSeries: 各周期 series 从外部传入方便外部处理逻辑 func (b *SigStrategyBacktester) multiIntervalSeries(ctx context.Context, sr *pb.SeriesRange, requiredIntervalSeries *types.IntervalState[int16], intervalKlineSeries *types.IntervalState[*sig.KlineSeries], recvFn func(driver bool, interval types.Interval, k *types.Kline) (err error)) (err error) { // 查询主周期时间范围 rsp, err := b.exchangeClient.SeriesRange(ctx, &pb.ReqSeriesRange{Series: sr}) if err != nil { return } before, after := rsp.Before, rsp.After driverInterval := types.Interval(sr.Interval) driverIntervalAdder := types.SupportedIntervals[driverInterval] var otherIntervals []types.Interval // 运行时周期 requiredIntervalSeries.Range(func(interval types.Interval, window int16) { if interval != driverInterval && window > 0 { otherIntervals = append(otherIntervals, interval) } }) // 运行时订阅周期 for interval := range b.intervalSubscribe { if interval != driverInterval && collect.NotIn(interval, otherIntervals...) { otherIntervals = append(otherIntervals, interval) } } // otherIntervals = collect.Filter(otherIntervals, func(_ int, interval types.Interval) bool { return interval != driverInterval }) // 通知其他周期更新的channel otherIntervalSyncCh := types.NewIntervalState[chan int64]() for _, interval := range otherIntervals { otherIntervalSyncCh.Set(interval, make(chan int64)) } stopCh := make(chan struct{}) stopChClosed := atomic.Bool{} for _, interval := range otherIntervals { kSeries := intervalKlineSeries.ComputeIfAbsent(interval, func() *sig.KlineSeries { return sig.NewKlineSeries(sr.Exchange, sr.InstId, interval) }) go func(interval types.Interval, kSeries *sig.KlineSeries) { syncCh := otherIntervalSyncCh.Get(interval) intervalAdder := types.SupportedIntervals[interval] driverTs := int64(0) isr := &pb.SeriesRange{Exchange: sr.Exchange, InstId: sr.InstId, Open: false, Live: sr.Live, Desc: sr.Desc} isr.Before = intervalAdder(before, -1) isr.After = after isr.Interval = string(interval) isr.WindowExtra = uint32(max(0, requiredIntervalSeries.Get(interval)-1)) 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 fmt.Errorf("kline series canceled") case <-stopCh: return io.EOF 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 } // 回调周期订阅 if subs, ok := b.intervalSubscribe[interval]; ok { for _, subFn := range subs { if err = subFn(interval, *k); err != nil { return } } } // 运行时周期 if requiredIntervalSeries.Get(interval) > 0 { return recvFn(false, interval, k) } return }) if err1 == nil { otherIntervalSyncCh.Set(interval, nil) // 该周期数据拉取结束 syncCh <- 0 // 通知更新完毕 } else if err1 != io.EOF { zlog.Errorf("fetch interval history error: inst=%s(%s) interval=%s, err=%v", sr.InstId, sr.Exchange, interval, err1) err = err1 if stopChClosed.CompareAndSwap(false, true) { close(stopCh) } } }(interval, kSeries) } // 驱动周期数据拉取 driverSeries := intervalKlineSeries.ComputeIfAbsent(driverInterval, func() *sig.KlineSeries { return sig.NewKlineSeries(sr.Exchange, sr.InstId, driverInterval) }) sr.WindowExtra = max(sr.WindowExtra, uint32(max(0, requiredIntervalSeries.Get(driverInterval)-1))) err0 := b.fetchHistoryKlineSeries(ctx, sr, func(k *types.Kline) (err error) { 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) otherIntervalSyncCh.RangeBreak(func(interval types.Interval, syncCh chan int64) bool { if syncCh != nil { syncCh <- driverTS // 通知其他周期更新到主周期时间 select { case <-syncCh: // 等待其它周期更新完毕 case <-ctx.Done(): err = fmt.Errorf("kline series canceled") return false case <-stopCh: err = io.EOF return false } } return true }) if err != nil { return } // intervalKlineSeries.Range(func(interval types.Interval, v *sig.KlineSeries) { // if v != nil { // zlog.Debugf("interval series update: %s, %d", interval, v.Length()) // } // }) // 回调周期订阅 if subs, ok := b.intervalSubscribe[driverInterval]; ok { for _, subFn := range subs { if err = subFn(driverInterval, *k); err != nil { return } } } return recvFn(true, driverInterval, k) }) if err0 != io.EOF { if stopChClosed.CompareAndSwap(false, true) { close(stopCh) } if err0 != nil { err = err0 zlog.Errorf("fetch driver interval history error: inst=%s(%s) interval=%s, err=%v", sr.InstId, sr.Exchange, driverInterval, err0) } } 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 } // Deprecated // _singleStrategySeries 单周期策略 func (b *SigStrategyBacktester) _singleStrategySeries(ctx context.Context, sigStrategy strategy.ISingleSigStrategy, sr *pb.SeriesRange, recvSignal func(sigSide types.Side, k types.Kline) (err error)) (err error) { interval := types.Interval(sr.Interval) kSeries := sig.NewKlineSeries(sr.Exchange, sr.InstId, interval) strategyContext := sig.NewStrategyContext(nil, kSeries, b.indicatorReg) requiredSeries := int(sigStrategy.CandlePeriods(nil)) sr.WindowExtra = uint32(max(0, requiredSeries-1)) err = b.fetchHistoryKlineSeries(ctx, sr, func(k *types.Kline) (err error) { 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 } if kSeries.Length() < requiredSeries { return } sigSide := sigStrategy.Update(strategyContext) if sigSide.IsValid() { if err = recvSignal(sigSide, *k); err != nil { return } } return }) return }