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.
371 lines
12 KiB
371 lines
12 KiB
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, |
|
// cleanIntervalSeries *types.IntervalState[*types.KlineSeries], |
|
iiks *types.InstanceIntervalKlineSeries, |
|
recvSignal func(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: |
|
// todo |
|
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(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(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(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) |
|
|
|
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 |
|
} |
|
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 |
|
} |
|
sigSide := intervalSigStrategy.Update(intervalStrategyContext) |
|
if sigSide.IsValid() { |
|
if err = recvSignal(sigSide, *k); err != nil { |
|
return |
|
} |
|
} |
|
return |
|
}) |
|
|
|
// err = b.multiIntervalSeries(ctx, sr, intervalCandlePeriods, intervalKlineSeries, func(driver bool, interval types.Interval, k *types.Kline) (err error) { |
|
// if !driver { |
|
// return |
|
// } |
|
// 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 |
|
// } |
|
// sigSide := intervalSigStrategy.Update(intervalStrategyContext) |
|
// if sigSide.IsValid() { |
|
// if err = recvSignal(sigSide, *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", isr.InstId, isr.Interval, err1) |
|
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,", sr.InstId, sr.Interval, err0) |
|
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 |
|
}
|
|
|