Browse Source

sig strategy backtester

main
strange 10 months ago
parent
commit
8d929e39b3
  1. 280
      internal/trading/backtest/sig_strategy_backtester.go
  2. 134
      internal/trading/backtest/trading_plan_backtester.go
  3. 185
      internal/trading/trading_service.go
  4. 8
      pkg/data/entity/trade_plan.go
  5. 22
      pkg/trade/close_strategy.go
  6. 9
      pkg/trade/risk_strategy.go

280
internal/trading/backtest/sig_strategy_backtester.go

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

134
internal/trading/backtest/backtest.go → internal/trading/backtest/trading_plan_backtester.go

@ -6,6 +6,7 @@ import (
"io"
"sig-pub/api/pb"
"sig-pub/internal/trading/sig"
"sig-pub/pkg/data/entity"
"sig-pub/pkg/indicator"
"sig-pub/pkg/strategy"
"sig-pub/pkg/trade"
@ -15,25 +16,122 @@ import (
"sig-pub/pkg/utils/times"
"sig-pub/pkg/zlog"
"github.com/bytedance/sonic"
"google.golang.org/grpc"
)
type Backtest struct {
exchangeClient pb.ExchangeServiceClient
indReg *indicator.IndicatorRegistry
type BacktestStat struct {
// total_return 年化收益
// sharpe_ratio 夏普比率
// max_drawdown 最大回撤
// num_trades 单数
// win_rate 胜率
}
riskStrategy *trade.RiskStrategy
type TradeAccount struct {
trade.ITradeAccount
cash float64
}
func NewTradeAccount(cash float64) *TradeAccount {
return &TradeAccount{
cash: cash,
}
}
func NewBacktest(exchangeClient pb.ExchangeServiceClient, indReg *indicator.IndicatorRegistry, sigStrategyReg *strategy.SigStrategyRegistry) *Backtest {
return &Backtest{
type TradingPlanBacktester struct {
exchangeClient pb.ExchangeServiceClient
indicatorReg *indicator.IndicatorRegistry
sigStrategyReg *strategy.SigStrategyRegistry
}
func NewTradingPlanBacktester(exchangeClient pb.ExchangeServiceClient, indicatorReg *indicator.IndicatorRegistry, sigStrategyReg *strategy.SigStrategyRegistry) *TradingPlanBacktester {
return &TradingPlanBacktester{
exchangeClient: exchangeClient,
indReg: indReg,
riskStrategy: trade.NewRiskStrategy(),
indicatorReg: indicatorReg,
sigStrategyReg: sigStrategyReg,
}
}
// 核心引擎,模拟交易、持仓跟踪、费用计算
func (b *TradingPlanBacktester) Backtest(ctx context.Context, cash float64, plan entity.TradePlan, sr *pb.SeriesRange) (err error) {
// sig strategy
sigStrategyType, sigStrategy, ok := b.sigStrategyReg.NewSigStrategy(plan.SigStrategy)
if !ok {
err = fmt.Errorf("sig strategy not exists %s", plan.SigStrategy)
return
}
sigStrategyParam := make(strategy.StrategyParam)
if err = sonic.UnmarshalString(plan.SigStrategyParam, &sigStrategyParam); err != nil {
return
}
if err = sigStrategy.Init(sigStrategyParam); err != nil {
return
}
closeStrategyParam, tradeStrategyParam, riskStrategyParam := new(trade.CloseStrategyParam),
new(trade.TradeStrategyParam), new(trade.RiskStrategyParam)
if err = sonic.UnmarshalString(plan.CloseStrategyParam, closeStrategyParam); err != nil {
return
}
if err = sonic.UnmarshalString(plan.TradeStrategyParam, tradeStrategyParam); err != nil {
return
}
if err = sonic.UnmarshalString(plan.RiskStrategyParam, riskStrategyParam); err != nil {
return
}
// 平仓策略
closeStrategy, err := trade.NewCloseStrategy(*closeStrategyParam)
if err != nil {
return
}
// 风险管理策略
riskStrategy, err := trade.NewRiskStrategy(*riskStrategyParam)
if err != nil {
return
}
tradeAccount := NewTradeAccount(cash)
_ = tradeAccount
// 平仓管理器
closeManager := NewCloseManager(0.02, 0)
closeManager.SetDynamicParams(0.1, 0.02, 0)
// 回测账户
account := NewAccount(10000, NewSimulator(0.0005, 0.0008))
sigStrategyBacktester := NewSigStrategyBacktester(sigStrategyType, sigStrategy, b.indicatorReg, b.exchangeClient)
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)
})
return
}
// onSideSingal 出现买卖信号
func (b *TradingPlanBacktester) onSideSingal(sigSide types.Side, k types.Kline, closeStrategy *trade.CloseStrategy) (err error) {
return
}
// onSigSideSignalWithAccount 交易策略发出交易信号(使用指定的账户和平仓管理器)
func (b *TradingPlanBacktester) onSigSideSignalWithAccount(sigSide types.Side, k types.Kline, account *Account, closeManager *CloseManager, riskStrategy *trade.RiskStrategy) {
// risk check before executing
side := riskStrategy.SideAssess(sigSide)
if !side.IsValid() {
zlog.Debugf("risk strategy filter sig side: %s", sigSide.String())
return
}
// zlog.Debugf("apply market order: ts=%d, side=%s", k.Ts, side.String())
price := decimals.MustToFloat64(k.Close)
account.ApplyMarketOrder(side, 0.01, price, k.Ts)
// 根据信号方向平掉相反方向的仓位:如果信号是买入,平掉所有卖出仓位;如果信号是卖出,平掉所有买入仓位
closeManager.CloseBySignal(sigSide, account, k)
}
func (b *Backtest) RunTradingPlan(ctx context.Context, tradingPlan *sig.TradingPlan, stime, etime int64, sigKlineSeries *sig.KlineSeries) (err error) {
func (b *TradingPlanBacktester) RunTradingPlan(ctx context.Context, tradingPlan *sig.TradingPlan, stime, etime int64, sigKlineSeries *sig.KlineSeries) (err error) {
plan := tradingPlan.Plan
exchange := pb.ExchangeType(plan.Exchange)
interval := types.Interval(plan.Interval)
@ -108,7 +206,7 @@ func (b *Backtest) RunTradingPlan(ctx context.Context, tradingPlan *sig.TradingP
sigSide := tradingPlan.Update(strategy.StrategyTypeSig)
if sigSide.IsValid() {
b.onSigSideSignalWithAccount(sigSide, *kline, account, closeManager)
b.onSigSideSignalWithAccount(sigSide, *kline, account, closeManager, nil)
}
}
}
@ -120,19 +218,3 @@ func (b *Backtest) RunTradingPlan(ctx context.Context, tradingPlan *sig.TradingP
_ = exposure
return
}
// onSigSideSignalWithAccount 交易策略发出交易信号(使用指定的账户和平仓管理器)
func (b *Backtest) onSigSideSignalWithAccount(sigSide types.Side, k types.Kline, account *Account, closeManager *CloseManager) {
// risk check before executing
side := b.riskStrategy.SideAssess(sigSide)
if !side.IsValid() {
zlog.Debugf("risk strategy filter sig side: %s", sigSide.String())
return
}
// zlog.Debugf("apply market order: ts=%d, side=%s", k.Ts, side.String())
price := decimals.MustToFloat64(k.Close)
account.ApplyMarketOrder(side, 0.01, price, k.Ts)
// 根据信号方向平掉相反方向的仓位:如果信号是买入,平掉所有卖出仓位;如果信号是卖出,平掉所有买入仓位
closeManager.CloseBySignal(sigSide, account, k)
}

185
internal/trading/trading_service.go

@ -238,7 +238,7 @@ func (svc *TradingService) IndicatorSeries(ctx context.Context, indicatorName st
return
}
// StrategySeries 简单策略信号测
// StrategySeries 信号策略回
func (svc *TradingService) StrategySeries(ctx context.Context, req *pb.ReqStrategySeries, rsp *pb.RspStrategySeries) (err error) {
// sigStrategy
sigStrategyType, sigStrategy, ok := svc.strategyReg.NewSigStrategy(req.SigStrategy)
@ -258,183 +258,14 @@ func (svc *TradingService) StrategySeries(ctx context.Context, req *pb.ReqStrate
return
}
switch sigStrategyType {
case strategy.SigStrategyTypeSingle:
rsp.Signal, rsp.Times, err = svc.singleStrategySeries(ctx, sigStrategy.(strategy.ISingleSigStrategy), req.Series)
case strategy.SigStrategyTypeInterval:
rsp.Signal, rsp.Times, err = svc.intervalStrategySeries(ctx, sigStrategy.(strategy.IIntervalSigStrategy), req.Series)
default:
err = fmt.Errorf("unknown sig strategy type %v", sigStrategyType)
}
return
}
// singleStrategySeries 单周期策略
func (svc *TradingService) singleStrategySeries(ctx context.Context, sigStrategy strategy.ISingleSigStrategy, sr *pb.SeriesRange) (signals []pb.Side, times []int64, err error) {
interval := types.Interval(sr.Interval)
requiredSeries := int(sigStrategy.RequiredSeries())
sr.WindowExtra = uint32(requiredSeries - 1)
kSeries := sig.NewKlineSeries(sr.Exchange, sr.InstId, interval)
indicatorContext := sig.NewIndicatorContext(kSeries)
strategyContext := sig.NewStrategyContext(indicatorContext, svc.indicatorReg)
err = svc.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() {
side := lang.Ternary(sigSide == types.SideLong, pb.Side_BUY, pb.Side_SELL)
signals = append(signals, side)
times = append(times, k.Ts)
}
return
})
if err != nil {
// 使用回测器回测信号
backtester := backtest.NewSigStrategyBacktester(sigStrategyType, sigStrategy, svc.indicatorReg, svc.exchangeClient)
err = backtester.Backtest(ctx, req.Series, func(sigSide types.Side, k types.Kline) (err error) {
side := lang.Ternary(sigSide == types.SideLong, pb.Side_BUY, pb.Side_SELL)
rsp.Signal = append(rsp.Signal, side)
rsp.Times = append(rsp.Times, k.Ts)
return
}
return
}
// intervalStrategySeries 多周期策略
func (svc *TradingService) intervalStrategySeries(ctx context.Context, intervalSigStrategy strategy.IIntervalSigStrategy, sr *pb.SeriesRange) (signals []pb.Side, times []int64, 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)
}
})
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 := svc.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, svc.indicatorReg)
// 驱动周期数据拉取
driverSeries := sig.NewKlineSeries(sr.Exchange, sr.InstId, driverInterval)
intervalKlineSeries.Set(driverInterval, driverSeries)
sr.WindowExtra = uint32(requiredIntervalSeries.Get(driverInterval) - 1)
err = svc.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() {
side := lang.Ternary(sigSide == types.SideLong, pb.Side_BUY, pb.Side_SELL)
signals = append(signals, side)
times = append(times, k.Ts)
}
return
})
if err != io.EOF {
close(stopCh)
}
return
}
@ -453,7 +284,7 @@ func (svc *TradingService) Backtest(planId, stime, etime int64) (err error) {
return
}
test := backtest.NewBacktest(svc.exchangeClient, svc.indicatorReg, svc.strategyReg)
test := backtest.NewTradingPlanBacktester(svc.exchangeClient, svc.indicatorReg, svc.strategyReg)
err = test.RunTradingPlan(context.Background(), tradingPlan, stime, etime, sigKlineSeries)
if err != nil {
return

8
pkg/data/entity/trade_plan.go

@ -9,13 +9,15 @@ type TradePlan struct {
InstId string `gorm:"column:inst_id" json:"instId"` // 交易产品id
Interval string `gorm:"column:interval" json:"interval"` // 交易周期
SigStrategy string `gorm:"column:sig_strategy" json:"sigStrategy"` // 交易信号策略
ExitStrategy string `gorm:"column:exit_strategy" json:"exitStrategy"` // 退出策略
TradeStrategy string `gorm:"column:trade_strategy" json:"tradeStrategy"` // 下单仓位管理策略
SigStrategyParam string `gorm:"column:sig_strategy_param" json:"sigStrategyParam"` // 交易信号策略参数
ExitStrategyParam string `gorm:"column:exit_strategy_param" json:"exitStrategyParam"` // 退出策略名称参数
CloseStrategyParam string `gorm:"column:close_strategy_param" json:"closeStrategyParam"` // 退出策略名称参数
TradeStrategyParam string `gorm:"column:trade_strategy_param" json:"tradeStrategyParam"` // 下单仓位管理策略参数
RiskStrategyParam string `gorm:"column:risk_strategy_param" json:"riskStrategyParam"` // 风险管理策略参数
UpdateBy string `gorm:"column:update_by" json:"updateBy"` // 更新人
UpdateTime int64 `gorm:"column:update_time" json:"updateTime"` // 更新时间戳毫秒
// ExitStrategy string `gorm:"column:exit_strategy" json:"exitStrategy"` // 退出策略
// TradeStrategy string `gorm:"column:trade_strategy" json:"tradeStrategy"` // 下单仓位管理策略
}
func (TradePlan) TableName() string {

22
pkg/trade/close_strategy.go

@ -1,6 +1,7 @@
package trade
import (
"fmt"
"sig-pub/pkg/types"
"sig-pub/pkg/types/decimals"
)
@ -28,22 +29,25 @@ type CloseStrategy struct {
CloseStrategyParam
}
func NewCloseStrategy(param CloseStrategyParam) *CloseStrategy {
return &CloseStrategy{
CloseStrategyParam: param,
func NewCloseStrategy(param CloseStrategyParam) (cs *CloseStrategy, err error) {
if param.StopLossPct < 0 {
err = fmt.Errorf("stopLossPct can't less zero")
return
}
cs = &CloseStrategy{CloseStrategyParam: param}
return
}
// Update 当k线更新判断是否关闭仓位
func (s *CloseStrategy) OnKline(k types.Kline, pos *Position) (closePos bool, cause Cause) {
closePrice := decimals.MustToFloat64(k.Close)
// update peak px
if pos.Side == types.SideLong && closePrice > pos.PeakPx {
pos.PeakPx = closePrice
}
if pos.Side == types.SideShort && closePrice < pos.PeakPx {
pos.PeakPx = closePrice
}
// if pos.Side == types.SideLong && closePrice > pos.PeakPx {
// pos.PeakPx = closePrice
// }
// if pos.Side == types.SideShort && closePrice < pos.PeakPx {
// pos.PeakPx = closePrice
// }
return s.OnPrice(closePrice, pos)
}

9
pkg/trade/risk_strategy.go

@ -14,11 +14,16 @@ type RiskStrategyParam struct {
}
// RiskStrategy 风险管理策略
// Kelly准则优化方法
type RiskStrategy struct {
RiskStrategyParam
}
func NewRiskStrategy() *RiskStrategy {
return &RiskStrategy{}
func NewRiskStrategy(param RiskStrategyParam) (rs *RiskStrategy, err error) {
rs = &RiskStrategy{
RiskStrategyParam: param,
}
return
}
// SideAssess 收到信号时进行评估, 返回过滤后的交易信号

Loading…
Cancel
Save