Browse Source

sig strategy multi interval backtest

main
strange 10 months ago
parent
commit
a084afe44e
  1. 12
      api/exchange.proto
  2. 7
      internal/exchange/exchange_grpc_server.go
  3. 5
      internal/exchange/exchange_service.go
  4. 242
      internal/trading/backtest/sig_strategy_backtester.go
  5. 7
      internal/trading/backtest/trading_plan_backtester.go
  6. 2
      internal/trading/trading_service.go
  7. 12
      pkg/types/decimals/decimal.go
  8. 10
      pkg/types/interval.go

12
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;
}

7
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 {

5
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 {

242
internal/trading/backtest/sig_strategy_backtester.go

@ -21,10 +21,10 @@ type SigStrategyBacktester struct {
sigStrategyType strategy.SigStrategyType
sigStrategy strategy.ISigStrategy
indicatorReg *indicator.IndicatorRegistry
exchangeServiceClient pb.ExchangeServiceClient
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(
@ -37,13 +37,13 @@ func NewSigStrategyBacktester(
sigStrategyType: sigStrategyType,
sigStrategy: sigStrategy,
indicatorReg: indicatorReg,
exchangeServiceClient: exchangeServiceClient,
intervalSubscribe: make(map[types.Interval][]func(k types.Kline)),
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
}
// 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线数据拉取
func (b *SigStrategyBacktester) multiIntervalSeries(sr *pb.SeriesRange, otherInterval []types.Interval,
revcFn func(driver bool, interval types.Interval, k *types.Kline) (err error)) (err error) {
// 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
// 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 {
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)
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)
driverSeries := intervalKlineSeries.Get(driverInterval)
if driverSeries == nil {
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
}
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
case ch[0] <- driverTS:
// 等待其它周期更新完毕
}
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
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 {
// 回调周期订阅
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
}

7
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)
})

2
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)

12
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
}

10
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]

Loading…
Cancel
Save