@ -2,6 +2,7 @@ package backtest
import (
import (
"context"
"context"
"errors"
"fmt"
"fmt"
"io"
"io"
"sig-pub/api/pb"
"sig-pub/api/pb"
@ -12,7 +13,6 @@ import (
"sig-pub/pkg/utils/collect"
"sig-pub/pkg/utils/collect"
"sig-pub/pkg/utils/times"
"sig-pub/pkg/utils/times"
"sig-pub/pkg/zlog"
"sig-pub/pkg/zlog"
"sync/atomic"
"google.golang.org/grpc"
"google.golang.org/grpc"
)
)
@ -66,7 +66,7 @@ func (b *SigStrategyBacktester) Backtest(ctx context.Context, sigStrategyInput t
case strategy . SigStrategyTypeSingle :
case strategy . SigStrategyTypeSingle :
err = b . singleStrategySeries ( ctx , b . sigStrategy . ( strategy . ISingleSigStrategy ) , sigStrategyInput , sr , iiks , recvSignal )
err = b . singleStrategySeries ( ctx , b . sigStrategy . ( strategy . ISingleSigStrategy ) , sigStrategyInput , sr , iiks , recvSignal )
case strategy . SigStrategyTypeInterval :
case strategy . SigStrategyTypeInterval :
err = b . intervalStrategySeries ( ctx , b . sigStrategy . ( strategy . IIntervalSigStrategy ) , sigStrategyInput , sr , iiks . GetIntervalKlineSeries ( sr . InstId ) , recvSignal )
err = b . intervalStrategySeries ( ctx , b . sigStrategy . ( strategy . IIntervalSigStrategy ) , sigStrategyInput , sr , iiks , recvSignal )
case strategy . SigStrategyTypeInstanceInterval :
case strategy . SigStrategyTypeInstanceInterval :
// todo
// todo
default :
default :
@ -77,21 +77,20 @@ func (b *SigStrategyBacktester) Backtest(ctx context.Context, sigStrategyInput t
// singleStrategySeries 单周期策略
// 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 ) {
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 ) {
interval := types . Interval ( sr . Interval )
driverInstId := sr . InstId
intervalKlineSeries := iiks . GetIntervalKlineSeries ( sr . InstId )
dr iverI nterval := types . Interval ( sr . Interval )
kSeries := intervalKlineSeries . ComputeIfAbsent ( interval , func ( ) * types . KlineSeries { return types . NewKlineSeries ( sr . Exchange , sr . InstId , interval ) } )
driverSeries := iiks . Get ( driver InstId, dr iverI nterval)
strategyContext := sig . NewStrategyContext ( sigStrategyInput , k Series, b . indicatorReg )
strategyContext := sig . NewStrategyContext ( sigStrategyInput , driver Series, b . indicatorReg )
requiredSerie s := int ( sigStrategy . CandlePeriods ( strategyContext ) )
requiredPeriod s := int ( sigStrategy . CandlePeriods ( strategyContext ) )
requiredIntervalSerie s := types . NewIntervalState [ int16 ] ( )
intervalCandlePeriod s := types . NewIntervalState [ int16 ] ( )
requiredIntervalSerie s. Set ( interval , int16 ( max ( 1 , requiredSerie s ) ) )
intervalCandlePeriod s. Set ( dr iverI nterval, int16 ( max ( 1 , requiredPeriod s ) ) )
err = b . multiIntervalSeries ( ctx , sr , requiredIntervalSeries , intervalKlineSerie s, func ( driver bool , interval types . Interval , k * types . Kline ) ( err error ) {
err = b . multiInstanceIn tervalSeries ( ctx , sr , [ ] string { sr . InstId } , intervalCandlePeriods , iik s, func ( driver bool , instId string , interval types . Interval , k * types . Kline ) ( err error ) {
if ! driver {
if ! driver || instId != driverInstId || interval != driverInterval {
return
return
}
}
kSeries := intervalKlineSeries . Get ( interval )
if driverSeries . Length ( ) < requiredPeriods {
if kSeries . Length ( ) < requiredSeries {
return
return
}
}
sigSide := sigStrategy . Update ( strategyContext )
sigSide := sigStrategy . Update ( strategyContext )
@ -106,17 +105,21 @@ func (b *SigStrategyBacktester) singleStrategySeries(ctx context.Context, sigStr
}
}
// intervalStrategySeries 多周期策略
// intervalStrategySeries 多周期策略
func ( b * SigStrategyBacktester ) intervalStrategySeries ( ctx context . Context , intervalSigStrategy strategy . IIntervalSigStrategy , sigStrategyInput types . Input , sr * pb . SeriesRange , intervalKlineSeries * types . IntervalState [ * types . KlineSeries ] , recvSignal func ( sigSide types . Side , k types . Kline ) ( err error ) ) ( err error ) {
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 )
intervalStrategyContext := sig . NewIntervalStrategyContext ( sigStrategyInput , intervalKlineSeries , b . indicatorReg )
// 各周期所需k线数量
// 各周期所需k线数量
requiredIntervalSeries := intervalSigStrategy . CandlePeriods ( intervalStrategyContext )
intervalCandlePeriods := intervalSigStrategy . CandlePeriods ( intervalStrategyContext )
err = b . multiIntervalSeries ( ctx , sr , requiredIntervalSeries , intervalKlineSeries , func ( driver bool , interval types . Interval , k * types . Kline ) ( err error ) {
if ! driver {
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
return
}
}
update := true
update := true
requiredIntervalSerie s. Range ( func ( interval types . Interval , require int16 ) {
intervalCandlePeriod s. Range ( func ( interval types . Interval , require int16 ) {
if update && require > 0 {
if update && require > 0 {
series := intervalKlineSeries . Get ( interval )
series := intervalKlineSeries . Get ( interval )
update = series . Length ( ) >= int ( require )
update = series . Length ( ) >= int ( require )
@ -133,200 +136,33 @@ func (b *SigStrategyBacktester) intervalStrategySeries(ctx context.Context, inte
}
}
return
return
} )
} )
return
}
// multiIntervalSeries 多周期k线数据拉取
// intervalKlineSeries: 各周期 series 从外部传入方便外部处理逻辑
func ( b * SigStrategyBacktester ) multiIntervalSeries ( ctx context . Context , sr * pb . SeriesRange ,
requiredIntervalSeries * types . IntervalState [ int16 ] ,
intervalKlineSeries * types . IntervalState [ * types . 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
}
driverBefore , driverAfter := rsp . Before , rsp . After
driverInstId := sr . InstId
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 { } )
stopChClosed := atomic . Bool { }
for _ , interval := range otherIntervals {
kSeries := intervalKlineSeries . ComputeIfAbsent ( interval , func ( ) * types . KlineSeries { return types . NewKlineSeries ( sr . Exchange , sr . InstId , interval ) } )
go func ( interval types . Interval , kSeries * types . 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 ( driverBefore , - 1 )
isr . After = driverAfter
isr . Interval = string ( interval )
isr . WindowExtra = uint32 ( max ( 0 , requiredIntervalSeries . 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 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
}
// 回调周期订阅
for _ , subFn := range b . intervalSubscribe [ interval ] {
if err = subFn ( driverInstId , 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
if stopChClosed . CompareAndSwap ( false , true ) {
close ( stopCh )
}
}
} ( interval , kSeries )
}
// 驱动周期数据拉取
driverSeries := intervalKlineSeries . ComputeIfAbsent ( driverInterval , func ( ) * types . KlineSeries { return types . NewKlineSeries ( sr . Exchange , sr . InstId , driverInterval ) } )
sr . WindowExtra = max ( sr . WindowExtra , uint32 ( max ( 0 , requiredIntervalSeries . Get ( driverInterval ) - 1 ) ) ) + indicator . ApproCandles
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 ( _ 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
}
if k . Ts < driverBefore {
return
}
// 回调周期订阅
// err = b.multiIntervalSeries(ctx, sr, intervalCandlePeriods, intervalKlineSeries, func(driver bool, interval types.Interval, k *types.Kline) (err error) {
for _ , subFn := range b . intervalSubscribe [ driverInterval ] {
// if !driver {
if err = subFn ( driverInstId , driverInterval , * k ) ; err != nil {
// return
return
// }
}
// update := true
}
// intervalCandlePeriods.Range(func(interval types.Interval, require int16) {
return recvFn ( true , driverInterval , k )
// if update && require > 0 {
} )
// series := intervalKlineSeries.Get(interval)
if err0 != io . EOF {
// update = series.Length() >= int(require)
if stopChClosed . CompareAndSwap ( false , true ) {
// }
close ( stopCh )
// })
}
// if !update {
if err0 != nil {
// return
err = err0
// }
zlog . Errorf ( "fetch driver interval history error: inst=%s(%s) interval=%s, err=%v" , sr . InstId , sr . Exchange , driverInterval , err0 )
// sigSide := intervalSigStrategy.Update(intervalStrategyContext)
}
// if sigSide.IsValid() {
}
// if err = recvSignal(sigSide, *k); err != nil {
// return
// }
// }
// return
// })
return
return
}
}
// fetchHistoryKlineSeries 请求k线数据流式处理
var errStop = errors . New ( "stop" )
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
}
// multiInstanceIntervalSeries 多币种多周期数据拉取
// multiInstanceIntervalSeries 多币种多周期数据拉取
func ( b * SigStrategyBacktester ) multiInstanceIntervalSeries ( ctx context . Context , sr * pb . SeriesRange ,
func ( b * SigStrategyBacktester ) multiInstanceIntervalSeries ( ctx context . Context , sr * pb . SeriesRange ,
@ -334,16 +170,16 @@ func (b *SigStrategyBacktester) multiInstanceIntervalSeries(ctx context.Context,
iiks * types . InstanceIntervalKlineSeries ,
iiks * types . InstanceIntervalKlineSeries ,
recvFn func ( driver bool , instId string , interval types . Interval , k * types . Kline ) ( err error ) ,
recvFn func ( driver bool , instId string , interval types . Interval , k * types . Kline ) ( err error ) ,
) ( 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 } )
rsp , err := b . exchangeClient . SeriesRange ( ctx , & pb . ReqSeriesRange { Series : sr } )
if err != nil {
if err != nil {
return
return
}
}
driverBefore , driverAfter := rsp . Before , rsp . After
driverBefore , driverAfter := rsp . Before , driverIntervalAdder ( rsp . After , 1 )
driverInstId := sr . InstId
driverInterval := types . Interval ( sr . Interval )
driverIntervalAdder := types . SupportedIntervals [ driverInterval ]
// 运行时周期
// 运行时周期
fetchIntervals := [ ] types . Interval { driverInterval }
fetchIntervals := [ ] types . Interval { driverInterval }
@ -362,7 +198,7 @@ func (b *SigStrategyBacktester) multiInstanceIntervalSeries(ctx context.Context,
fetchInsts := append ( [ ] string { driverInstId } , tradeInsts ... )
fetchInsts := append ( [ ] string { driverInstId } , tradeInsts ... )
fetchInsts = collect . Uniq ( fetchInsts )
fetchInsts = collect . Uniq ( fetchInsts )
// channel s
// fetch kline serie s
var otherSrs [ ] * pb . SeriesRange
var otherSrs [ ] * pb . SeriesRange
for _ , instId := range fetchInsts {
for _ , instId := range fetchInsts {
for _ , interval := range fetchIntervals {
for _ , interval := range fetchIntervals {
@ -403,9 +239,9 @@ func (b *SigStrategyBacktester) multiInstanceIntervalSeries(ctx context.Context,
}
}
select {
select {
case <- ctx . Done ( ) :
case <- ctx . Done ( ) :
return io . EOF
return errStop
case <- stopCh :
case <- stopCh :
return io . EOF
return errStop
case driverTs = <- syncCh : // 等待主周期通知更新
case driverTs = <- syncCh : // 等待主周期通知更新
if closeTs <= driverTs {
if closeTs <= driverTs {
break waitLoop
break waitLoop
@ -425,15 +261,15 @@ func (b *SigStrategyBacktester) multiInstanceIntervalSeries(ctx context.Context,
}
}
}
}
// 运行时周期
// 运行时周期
// if requiredIntervalSeries.Get(interval) > 0 {
if intervalCandlePeriods . Get ( interval ) > 0 {
// return recvFn(false, isr.InstId, interval, k)
return recvFn ( false , isr . InstId , interval , k )
// }
}
return
return
} )
} )
zlog . Debugf ( "other sr finish with: %s(%s), %v" , isr . InstId , isr . Interval , err1 )
if err1 == nil {
if err1 == nil {
syncCh <- 0 // 通知更新完毕
syncCh <- - 1 // 通知更新完毕, 后续不再更新
// todo 后续不再更新
} else if err1 != errStop {
} else if err1 != io . EOF {
zlog . Errorf ( "fetch interval history 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
err = err1
close ( stopCh )
close ( stopCh )
@ -453,14 +289,21 @@ func (b *SigStrategyBacktester) multiInstanceIntervalSeries(ctx context.Context,
return
return
}
}
driverTS := driverIntervalAdder ( k . Ts , 1 )
driverTS := driverIntervalAdder ( k . Ts , 1 )
for _ , syncCh := range syncChans {
for i , syncCh := range syncChans {
if syncCh == nil {
continue
}
syncCh <- driverTS // 通知其他周期更新到主周期时间
syncCh <- driverTS // 通知其他周期更新到主周期时间
select {
select {
case <- syncCh : // 等待该周期更新完毕
case sig := <- syncCh : // 等待该周期更新完毕
if sig == - 1 {
// 后续不再更新
syncChans [ i ] = nil
}
case <- ctx . Done ( ) :
case <- ctx . Done ( ) :
return io . EOF
return errStop
case <- stopCh :
case <- stopCh :
return io . EOF
return errStop
}
}
}
}
if k . Ts < driverBefore {
if k . Ts < driverBefore {
@ -474,9 +317,10 @@ func (b *SigStrategyBacktester) multiInstanceIntervalSeries(ctx context.Context,
}
}
return recvFn ( true , driverInstId , driverInterval , k )
return recvFn ( true , driverInstId , driverInterval , k )
} )
} )
zlog . Debugf ( "driver sr finish with: %s(%s), %v," , sr . InstId , sr . Interval , err0 )
if err0 == nil {
if err0 == nil {
close ( stopCh )
close ( stopCh )
} else if err0 != io . EOF {
} else if err0 != errStop {
zlog . Errorf ( "fetch driver interval history error: inst=%s(%s) interval=%s, err=%v" , sr . InstId , sr . Exchange , driverInterval , err )
zlog . Errorf ( "fetch driver interval history error: inst=%s(%s) interval=%s, err=%v" , sr . InstId , sr . Exchange , driverInterval , err )
err = err0
err = err0
close ( stopCh )
close ( stopCh )
@ -484,3 +328,44 @@ func (b *SigStrategyBacktester) multiInstanceIntervalSeries(ctx context.Context,
}
}
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
}