From a084afe44e502eed5e70d449a277ea869a522cf1 Mon Sep 17 00:00:00 2001 From: strange Date: Fri, 7 Nov 2025 16:51:52 +0800 Subject: [PATCH] sig strategy multi interval backtest --- api/exchange.proto | 12 + internal/exchange/exchange_grpc_server.go | 7 + internal/exchange/exchange_service.go | 5 + .../backtest/sig_strategy_backtester.go | 262 +++++++++++------- .../backtest/trading_plan_backtester.go | 7 +- internal/trading/trading_service.go | 2 +- pkg/types/decimals/decimal.go | 12 + pkg/types/interval.go | 10 + 8 files changed, 221 insertions(+), 96 deletions(-) diff --git a/api/exchange.proto b/api/exchange.proto index 68715d5..6b445bd 100644 --- a/api/exchange.proto +++ b/api/exchange.proto @@ -15,6 +15,9 @@ service ExchangeService { // 获取交易所交易产品状态 rpc ExchangeInstanceState(ReqExchangeInstanceState) returns (RspExchangeInstanceState); + // 查询SeriesRange时间范围 + rpc SeriesRange(ReqSeriesRange) returns (RspSeriesRange); + // 查询历史k线 rpc HistoryKline(ReqHistoryKline) returns (RspHistoryKline); @@ -56,6 +59,15 @@ message RspExchangeInstanceState { repeated TradeInstanceState instsState = 1; // 交易产品列表 } +message ReqSeriesRange { + SeriesRange series = 1; +} +message RspSeriesRange { + int64 after = 1; + int64 before = 2; + int64 total = 3; +} + message ReqHistoryKline { SeriesRange series = 1; } diff --git a/internal/exchange/exchange_grpc_server.go b/internal/exchange/exchange_grpc_server.go index c3cbeab..44af003 100644 --- a/internal/exchange/exchange_grpc_server.go +++ b/internal/exchange/exchange_grpc_server.go @@ -117,6 +117,13 @@ func (svr *ExchangeGrpcServer) ExchangeInstanceState(ctx context.Context, req *p return &pb.RspExchangeInstanceState{InstsState: states}, nil } +// 查询SeriesRange时间范围 +func (svr *ExchangeGrpcServer) SeriesRange(ctx context.Context, req *pb.ReqSeriesRange) (rsp *pb.RspSeriesRange, err error) { + rsp = new(pb.RspSeriesRange) + rsp.After, rsp.Before, rsp.Total, err = svr.exchangeService.CalcSeriesRange(req.Series) + return +} + // HistoryKline 获取交易产品历史k线 (before < klines... < after) func (svr *ExchangeGrpcServer) HistoryKline(ctx context.Context, req *pb.ReqHistoryKline) (rsp *pb.RspHistoryKline, err error) { if req.Series == nil { diff --git a/internal/exchange/exchange_service.go b/internal/exchange/exchange_service.go index 1a15641..7f4a432 100644 --- a/internal/exchange/exchange_service.go +++ b/internal/exchange/exchange_service.go @@ -8,6 +8,7 @@ import ( "sig-pub/api/pb" "sig-pub/pkg/client" "sig-pub/pkg/data" + "sig-pub/pkg/indicator" "sig-pub/pkg/mq" "sig-pub/pkg/publish" "sig-pub/pkg/types" @@ -733,6 +734,10 @@ func (svc *ExchangeService) CalcSeriesRange(arg *pb.SeriesRange) (after, before, } // 额外拉取 if arg.WindowExtra > 0 { + if arg.WindowExtra > indicator.MaxWindow { + err = fmt.Errorf("series range window extra %d big then %d", arg.WindowExtra, indicator.MaxWindow) + return + } before = max(intervalAdder(before, -int64(arg.WindowExtra)), KlineBefore0) } if before > after { diff --git a/internal/trading/backtest/sig_strategy_backtester.go b/internal/trading/backtest/sig_strategy_backtester.go index 7d94816..d271ac9 100644 --- a/internal/trading/backtest/sig_strategy_backtester.go +++ b/internal/trading/backtest/sig_strategy_backtester.go @@ -18,13 +18,13 @@ import ( // SigStrategyBacktester 信号策略回测 type SigStrategyBacktester struct { - sigStrategyType strategy.SigStrategyType - sigStrategy strategy.ISigStrategy - indicatorReg *indicator.IndicatorRegistry - exchangeServiceClient pb.ExchangeServiceClient + sigStrategyType strategy.SigStrategyType + sigStrategy strategy.ISigStrategy + indicatorReg *indicator.IndicatorRegistry + exchangeClient pb.ExchangeServiceClient // 回测过程中订阅k线 - intervalSubscribe map[types.Interval][]func(k types.Kline) + intervalSubscribe map[types.Interval][]func(interval types.Interval, k *types.Kline) (err error) } func NewSigStrategyBacktester( @@ -34,16 +34,16 @@ func NewSigStrategyBacktester( exchangeServiceClient pb.ExchangeServiceClient, ) *SigStrategyBacktester { return &SigStrategyBacktester{ - sigStrategyType: sigStrategyType, - sigStrategy: sigStrategy, - indicatorReg: indicatorReg, - exchangeServiceClient: exchangeServiceClient, - intervalSubscribe: make(map[types.Interval][]func(k types.Kline)), + 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(k types.Kline)) { +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) } @@ -66,14 +66,19 @@ func (b *SigStrategyBacktester) singleStrategySeries(ctx context.Context, sigStr kSeries := sig.NewKlineSeries(sr.Exchange, sr.InstId, interval) indicatorContext := sig.NewIndicatorContext(kSeries) strategyContext := sig.NewStrategyContext(indicatorContext, b.indicatorReg) - requiredSeries := int(sigStrategy.RequiredSeries()) - sr.WindowExtra = uint32(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) + + requiredIntervalSeries := types.NewIntervalState[int16]() + requiredIntervalSeries.Set(interval, int16(max(1, requiredSeries))) + + intervalKlineSeries := types.NewIntervalState[*sig.KlineSeries]() + intervalKlineSeries.Set(interval, kSeries) + + 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 } @@ -88,55 +93,91 @@ func (b *SigStrategyBacktester) singleStrategySeries(ctx context.Context, sigStr return } -// multiIntervalSeries 多周期k线数据拉取 -func (b *SigStrategyBacktester) multiIntervalSeries(sr *pb.SeriesRange, otherInterval []types.Interval, - revcFn func(driver bool, interval types.Interval, k *types.Kline) (err error)) (err error) { +// intervalStrategySeries 多周期策略 +func (b *SigStrategyBacktester) intervalStrategySeries(ctx context.Context, intervalSigStrategy strategy.IIntervalSigStrategy, sr *pb.SeriesRange, recvSignal func(sigSide types.Side, k types.Kline) (err error)) (err error) { + // 各周期所需k线数量 + requiredIntervalSeries := intervalSigStrategy.RequiredIntervalSeries() + // 各周期 series + intervalKlineSeries := types.NewIntervalState[*sig.KlineSeries]() + // 策略上下文 + intervalStrategyContext := sig.NewIntervalStrategyContext(intervalKlineSeries, b.indicatorReg) + 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 } -// intervalStrategySeries 多周期策略 -func (b *SigStrategyBacktester) intervalStrategySeries(ctx context.Context, intervalSigStrategy strategy.IIntervalSigStrategy, sr *pb.SeriesRange, recvSignal func(sigSide types.Side, k types.Kline) (err error)) (err error) { +// 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] - // 各周期所需k线数量 - requiredIntervalSeries := intervalSigStrategy.RequiredIntervalSeries() - // 驱动周期外其他周期 - otherIntervals := make([]types.Interval, 0, 3) - requiredIntervalSeries.Range(func(interval types.Interval, series int16) { - if series > 0 { + + 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 collect.NotIn(interval, otherIntervals...) { + if interval != driverInterval && collect.NotIn(interval, otherIntervals...) { otherIntervals = append(otherIntervals, interval) } } - otherIntervals = collect.Filter(otherIntervals, func(_ int, interval types.Interval) bool { return interval != driverInterval }) + // otherIntervals = collect.Filter(otherIntervals, func(_ int, interval types.Interval) bool { return interval != driverInterval }) // 通知其他周期更新的channel - otherIntervalCh := types.NewIntervalState[[]chan int64]() - // otherIntervalDstCh := types.NewIntervalState[chan int64]() + otherIntervalSyncCh := types.NewIntervalState[chan int64]() for _, interval := range otherIntervals { - otherIntervalCh.Set(interval, []chan int64{make(chan int64), make(chan int64)}) + otherIntervalSyncCh.Set(interval, make(chan int64)) } - // 各周期 series - intervalKlineSeries := types.NewIntervalState[*sig.KlineSeries]() - // 其他周期数据拉取 stopCh := make(chan struct{}) for _, interval := range otherIntervals { - kSeries := sig.NewKlineSeries(sr.Exchange, sr.InstId, interval) - intervalKlineSeries.Set(interval, kSeries) + kSeries := intervalKlineSeries.Get(interval) + if kSeries == nil { + kSeries = sig.NewKlineSeries(sr.Exchange, sr.InstId, interval) + intervalKlineSeries.Set(interval, kSeries) + } go func(interval types.Interval, kSeries *sig.KlineSeries) { - ch := otherIntervalCh.Get(interval) - srcCh := ch[0] - dstCh := ch[1] + syncCh := otherIntervalSyncCh.Get(interval) intervalAdder := types.SupportedIntervals[interval] driverTs := int64(0) - isr := &pb.SeriesRange{Exchange: sr.Exchange, InstId: sr.InstId, Before: sr.Before, After: sr.After, Count: sr.Count, Open: sr.Open, Live: sr.Live, Desc: sr.Desc, Limit: sr.Limit} + 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(requiredIntervalSeries.Get(interval) - 1) + 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保持同步更新 @@ -144,12 +185,14 @@ func (b *SigStrategyBacktester) intervalStrategySeries(ctx context.Context, inte waitLoop: for { if driverTs != 0 { - dstCh <- 0 // 通知更新完毕 + syncCh <- 0 // 通知更新完毕 } select { + case <-ctx.Done(): + return fmt.Errorf("kline series canceled") case <-stopCh: return io.EOF - case driverTs = <-srcCh: + case driverTs = <-syncCh: // 等待主周期通知更新 if closeTs <= driverTs { break waitLoop } @@ -160,81 +203,82 @@ func (b *SigStrategyBacktester) intervalStrategySeries(ctx context.Context, inte 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 { - otherIntervalCh.Set(interval, nil) // 该周期数据拉取结束 - dstCh <- 0 // 通知更新完毕 + otherIntervalSyncCh.Set(interval, nil) // 该周期数据拉取结束 + syncCh <- 0 // 通知更新完毕 } else if err1 != io.EOF { - zlog.Errorf("fetch history interval error: inst=%s(%s) interval=%s, err=%v", sr.InstId, sr.Exchange, interval, err1) + zlog.Errorf("fetch interval history error: inst=%s(%s) interval=%s, err=%v", sr.InstId, sr.Exchange, interval, err1) err = err1 close(stopCh) } }(interval, kSeries) } - - // 策略上下文 - intervalStrategyContext := sig.NewIntervalStrategyContext(intervalKlineSeries, b.indicatorReg) - // 驱动周期数据拉取 - driverSeries := sig.NewKlineSeries(sr.Exchange, sr.InstId, driverInterval) - intervalKlineSeries.Set(driverInterval, driverSeries) - sr.WindowExtra = uint32(requiredIntervalSeries.Get(driverInterval) - 1) - err = b.fetchHistoryKlineSeries(ctx, sr, func(k *types.Kline) (err error) { + driverSeries := intervalKlineSeries.Get(driverInterval) + if driverSeries == nil { + driverSeries = sig.NewKlineSeries(sr.Exchange, sr.InstId, driverInterval) + intervalKlineSeries.Set(driverInterval, driverSeries) + } + 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) - otherIntervalCh.Range(func(interval types.Interval, ch []chan int64) { - if len(ch) == 2 { - // 通知其它周期先更新 + 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 - case ch[0] <- driverTS: - // 等待其它周期更新完毕 - select { - case <-stopCh: - err = io.EOF - return - case <-ch[1]: - } + return false } } + return true }) if err != nil { return } - - // zlog.Debugf("driver series update: %s, %d", driverInterval, driverTS) - 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 - } - // 检查满足策略执行条件 - 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 - } // intervalKlineSeries.Range(func(interval types.Interval, v *sig.KlineSeries) { // if v != nil { - // zlog.Debugf("strategy update: interval series %s, %d", interval, v.Length()) + // zlog.Debugf("interval series update: %s, %d", interval, v.Length()) // } // }) - sigSide := intervalSigStrategy.Update(intervalStrategyContext) - if sigSide.IsValid() { - if err = recvSignal(sigSide, *k); err != nil { - return + // 回调周期订阅 + if subs, ok := b.intervalSubscribe[driverInterval]; ok { + for _, subFn := range subs { + if err = subFn(driverInterval, k); err != nil { + return + } } } - return + return recvFn(true, driverInterval, k) }) - if err != io.EOF { + if err0 != io.EOF { 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 } @@ -243,13 +287,14 @@ func (b *SigStrategyBacktester) intervalStrategySeries(ctx context.Context, inte 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.exchangeServiceClient.HistoryKlineStream(ctx, req, grpc.UseCompressor("snappy")) + 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(): @@ -263,7 +308,7 @@ func (b *SigStrategyBacktester) fetchHistoryKlineSeries(ctx context.Context, sr break } if err != nil { - return + break } recvTimes++ recvTotal += len(msg.Klines) @@ -271,10 +316,39 @@ func (b *SigStrategyBacktester) fetchHistoryKlineSeries(ctx context.Context, sr kline := new(types.Kline) kline.ParsePBKline(sr.Exchange, k) if err = recvFn(kline); err != nil { - return + 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) + indicatorContext := sig.NewIndicatorContext(kSeries) + strategyContext := sig.NewStrategyContext(indicatorContext, b.indicatorReg) + + requiredSeries := int(sigStrategy.RequiredSeries()) + 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 +} diff --git a/internal/trading/backtest/trading_plan_backtester.go b/internal/trading/backtest/trading_plan_backtester.go index a1b2a97..da3ad0c 100644 --- a/internal/trading/backtest/trading_plan_backtester.go +++ b/internal/trading/backtest/trading_plan_backtester.go @@ -101,8 +101,13 @@ func (b *TradingPlanBacktester) Backtest(ctx context.Context, cash float64, plan account := NewAccount(10000, NewSimulator(0.0005, 0.0008)) sigStrategyBacktester := NewSigStrategyBacktester(sigStrategyType, sigStrategy, b.indicatorReg, b.exchangeClient) + // 平仓策略订阅1分钟曲线 + sigStrategyBacktester.SubKline(types.Interval1m, func(interval types.Interval, k *types.Kline) (err error) { + closeManager.OnKline(*k, account) + return + }) + // 下单 err = sigStrategyBacktester.Backtest(ctx, sr, func(sigSide types.Side, k types.Kline) (err error) { - closeManager.OnKline(k, account) b.onSigSideSignalWithAccount(sigSide, k, account, closeManager, riskStrategy) return b.onSideSingal(sigSide, k, closeStrategy) }) diff --git a/internal/trading/trading_service.go b/internal/trading/trading_service.go index 465a97e..a0d01c3 100644 --- a/internal/trading/trading_service.go +++ b/internal/trading/trading_service.go @@ -213,7 +213,7 @@ func (svc *TradingService) IndicatorSeries(ctx context.Context, indicatorName st // 查询历史指标数据 requiredSeries := int(indicator.RequiredSeries(int16(window))) - sr.WindowExtra = uint32(requiredSeries - 1) + sr.WindowExtra = uint32(max(0, requiredSeries-1)) kSeries := sig.NewKlineSeries(sr.Exchange, sr.InstId, types.Interval(sr.Interval)) indicatorContext := sig.NewIndicatorContext(kSeries) diff --git a/pkg/types/decimals/decimal.go b/pkg/types/decimals/decimal.go index e4c29ce..84c481a 100644 --- a/pkg/types/decimals/decimal.go +++ b/pkg/types/decimals/decimal.go @@ -6,6 +6,12 @@ import ( "github.com/govalues/decimal" ) +func panicE(err error) { + if err != nil { + panic(err) + } +} + func MustToFloat64(v decimal.Decimal) float64 { f, ok := v.Float64() if !ok { @@ -21,3 +27,9 @@ func MustFromFloat64(f float64) (v decimal.Decimal) { } return } + +func MustAdd(a, b decimal.Decimal) (r decimal.Decimal) { + r, err := a.Add(b) + panicE(err) + return +} diff --git a/pkg/types/interval.go b/pkg/types/interval.go index 38ba50f..a2937a5 100644 --- a/pkg/types/interval.go +++ b/pkg/types/interval.go @@ -154,6 +154,16 @@ func (s *IntervalState[T]) Range(f func(interval Interval, v T)) { } } +func (s *IntervalState[T]) RangeBreak(f func(interval Interval, v T) bool) { + for i, interval := range iotasIntervals { + index := i + 1 // 0保留 + v := s.state[index] + if !f(interval, v) { + break + } + } +} + func (s *IntervalState[T]) SetIf(interval Interval, v T, cond func(old T) bool) { i := intervalIotas[interval] old := s.state[i]