From 8d929e39b33981b62ca36ff0f5aff5079278f9fc Mon Sep 17 00:00:00 2001 From: strange Date: Thu, 6 Nov 2025 18:26:55 +0800 Subject: [PATCH] sig strategy backtester --- .../backtest/sig_strategy_backtester.go | 280 ++++++++++++++++++ ...backtest.go => trading_plan_backtester.go} | 134 +++++++-- internal/trading/trading_service.go | 185 +----------- pkg/data/entity/trade_plan.go | 8 +- pkg/trade/close_strategy.go | 22 +- pkg/trade/risk_strategy.go | 9 +- 6 files changed, 421 insertions(+), 217 deletions(-) create mode 100644 internal/trading/backtest/sig_strategy_backtester.go rename internal/trading/backtest/{backtest.go => trading_plan_backtester.go} (50%) diff --git a/internal/trading/backtest/sig_strategy_backtester.go b/internal/trading/backtest/sig_strategy_backtester.go new file mode 100644 index 0000000..7d94816 --- /dev/null +++ b/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 +} diff --git a/internal/trading/backtest/backtest.go b/internal/trading/backtest/trading_plan_backtester.go similarity index 50% rename from internal/trading/backtest/backtest.go rename to internal/trading/backtest/trading_plan_backtester.go index 66e5202..a1b2a97 100644 --- a/internal/trading/backtest/backtest.go +++ b/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) -} diff --git a/internal/trading/trading_service.go b/internal/trading/trading_service.go index f736e1b..465a97e 100644 --- a/internal/trading/trading_service.go +++ b/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 diff --git a/pkg/data/entity/trade_plan.go b/pkg/data/entity/trade_plan.go index 70f3f62..dd42f8a 100644 --- a/pkg/data/entity/trade_plan.go +++ b/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 { diff --git a/pkg/trade/close_strategy.go b/pkg/trade/close_strategy.go index c9cbbc2..f2682b2 100644 --- a/pkg/trade/close_strategy.go +++ b/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) } diff --git a/pkg/trade/risk_strategy.go b/pkg/trade/risk_strategy.go index 94077bf..f5281f5 100644 --- a/pkg/trade/risk_strategy.go +++ b/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 收到信号时进行评估, 返回过滤后的交易信号