You can not select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
354 lines
12 KiB
354 lines
12 KiB
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 |
|
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, 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()) |
|
|
|
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 |
|
} |
|
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, 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 |
|
} |
|
|
|
// 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{}) |
|
for _, interval := range otherIntervals { |
|
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) { |
|
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 |
|
close(stopCh) |
|
} |
|
}(interval, kSeries) |
|
} |
|
// 驱动周期数据拉取 |
|
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) |
|
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 { |
|
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) |
|
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 |
|
}
|
|
|