6 changed files with 421 additions and 217 deletions
@ -0,0 +1,280 @@ |
|||||||
|
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" |
||||||
|
|
||||||
|
"google.golang.org/grpc" |
||||||
|
) |
||||||
|
|
||||||
|
// SigStrategyBacktester 信号策略回测
|
||||||
|
type SigStrategyBacktester struct { |
||||||
|
sigStrategyType strategy.SigStrategyType |
||||||
|
sigStrategy strategy.ISigStrategy |
||||||
|
indicatorReg *indicator.IndicatorRegistry |
||||||
|
exchangeServiceClient pb.ExchangeServiceClient |
||||||
|
|
||||||
|
// 回测过程中订阅k线
|
||||||
|
intervalSubscribe map[types.Interval][]func(k types.Kline) |
||||||
|
} |
||||||
|
|
||||||
|
func NewSigStrategyBacktester( |
||||||
|
sigStrategyType strategy.SigStrategyType, |
||||||
|
sigStrategy strategy.ISigStrategy, |
||||||
|
indicatorReg *indicator.IndicatorRegistry, |
||||||
|
exchangeServiceClient pb.ExchangeServiceClient, |
||||||
|
) *SigStrategyBacktester { |
||||||
|
return &SigStrategyBacktester{ |
||||||
|
sigStrategyType: sigStrategyType, |
||||||
|
sigStrategy: sigStrategy, |
||||||
|
indicatorReg: indicatorReg, |
||||||
|
exchangeServiceClient: exchangeServiceClient, |
||||||
|
intervalSubscribe: make(map[types.Interval][]func(k types.Kline)), |
||||||
|
} |
||||||
|
} |
||||||
|
|
||||||
|
// SubKline 在回测过程中订阅k线
|
||||||
|
func (b *SigStrategyBacktester) SubKline(interval types.Interval, recv func(k types.Kline)) { |
||||||
|
b.intervalSubscribe[interval] = append(b.intervalSubscribe[interval], recv) |
||||||
|
} |
||||||
|
|
||||||
|
// Backtest 基于历史数据回测信号策略
|
||||||
|
func (b *SigStrategyBacktester) Backtest(ctx context.Context, sr *pb.SeriesRange, recvSignal func(sigSide types.Side, k types.Kline) (err error)) (err error) { |
||||||
|
switch b.sigStrategyType { |
||||||
|
case strategy.SigStrategyTypeSingle: |
||||||
|
err = b.singleStrategySeries(ctx, b.sigStrategy.(strategy.ISingleSigStrategy), sr, recvSignal) |
||||||
|
case strategy.SigStrategyTypeInterval: |
||||||
|
err = b.intervalStrategySeries(ctx, b.sigStrategy.(strategy.IIntervalSigStrategy), sr, recvSignal) |
||||||
|
default: |
||||||
|
err = fmt.Errorf("unknown sig strategy type %v", b.sigStrategyType) |
||||||
|
} |
||||||
|
return |
||||||
|
} |
||||||
|
|
||||||
|
// 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(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 |
||||||
|
} |
||||||
|
|
||||||
|
// 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) { |
||||||
|
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) { |
||||||
|
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 { |
||||||
|
otherIntervals = append(otherIntervals, interval) |
||||||
|
} |
||||||
|
}) |
||||||
|
// 运行时订阅周期
|
||||||
|
for interval := range b.intervalSubscribe { |
||||||
|
if collect.NotIn(interval, otherIntervals...) { |
||||||
|
otherIntervals = append(otherIntervals, interval) |
||||||
|
} |
||||||
|
} |
||||||
|
otherIntervals = collect.Filter(otherIntervals, func(_ int, interval types.Interval) bool { return interval != driverInterval }) |
||||||
|
|
||||||
|
// 通知其他周期更新的channel
|
||||||
|
otherIntervalCh := types.NewIntervalState[[]chan int64]() |
||||||
|
// otherIntervalDstCh := types.NewIntervalState[chan int64]()
|
||||||
|
for _, interval := range otherIntervals { |
||||||
|
otherIntervalCh.Set(interval, []chan int64{make(chan int64), 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) |
||||||
|
go func(interval types.Interval, kSeries *sig.KlineSeries) { |
||||||
|
ch := otherIntervalCh.Get(interval) |
||||||
|
srcCh := ch[0] |
||||||
|
dstCh := ch[1] |
||||||
|
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.Interval = string(interval) |
||||||
|
isr.WindowExtra = uint32(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 { |
||||||
|
dstCh <- 0 // 通知更新完毕
|
||||||
|
} |
||||||
|
select { |
||||||
|
case <-stopCh: |
||||||
|
return io.EOF |
||||||
|
case driverTs = <-srcCh: |
||||||
|
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 |
||||||
|
} |
||||||
|
return |
||||||
|
}) |
||||||
|
if err1 == nil { |
||||||
|
otherIntervalCh.Set(interval, nil) // 该周期数据拉取结束
|
||||||
|
dstCh <- 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) |
||||||
|
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) { |
||||||
|
driverTS := driverIntervalAdder(k.Ts, 1) |
||||||
|
otherIntervalCh.Range(func(interval types.Interval, ch []chan int64) { |
||||||
|
if len(ch) == 2 { |
||||||
|
// 通知其它周期先更新
|
||||||
|
select { |
||||||
|
case <-stopCh: |
||||||
|
err = io.EOF |
||||||
|
return |
||||||
|
case ch[0] <- driverTS: |
||||||
|
// 等待其它周期更新完毕
|
||||||
|
select { |
||||||
|
case <-stopCh: |
||||||
|
err = io.EOF |
||||||
|
return |
||||||
|
case <-ch[1]: |
||||||
|
} |
||||||
|
} |
||||||
|
} |
||||||
|
}) |
||||||
|
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())
|
||||||
|
// }
|
||||||
|
// })
|
||||||
|
sigSide := intervalSigStrategy.Update(intervalStrategyContext) |
||||||
|
if sigSide.IsValid() { |
||||||
|
if err = recvSignal(sigSide, *k); err != nil { |
||||||
|
return |
||||||
|
} |
||||||
|
} |
||||||
|
return |
||||||
|
}) |
||||||
|
if err != io.EOF { |
||||||
|
close(stopCh) |
||||||
|
} |
||||||
|
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.exchangeServiceClient.HistoryKlineStream(ctx, req, grpc.UseCompressor("snappy")) |
||||||
|
if err != nil { |
||||||
|
return |
||||||
|
} |
||||||
|
var msg *pb.RspHistoryKlineStream |
||||||
|
recvTimes, recvTotal := 0, 0 |
||||||
|
watch := times.NewWatch() |
||||||
|
for { |
||||||
|
select { |
||||||
|
case <-ctx.Done(): |
||||||
|
err = ctx.Err() |
||||||
|
return |
||||||
|
default: |
||||||
|
} |
||||||
|
msg, err = stream.Recv() |
||||||
|
if err == io.EOF { |
||||||
|
err = nil |
||||||
|
break |
||||||
|
} |
||||||
|
if err != nil { |
||||||
|
return |
||||||
|
} |
||||||
|
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 { |
||||||
|
return |
||||||
|
} |
||||||
|
} |
||||||
|
} |
||||||
|
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 |
||||||
|
} |
||||||
Loading…
Reference in new issue