From 614548e01814c377d89ca8ac451304e39d01765a Mon Sep 17 00:00:00 2001 From: strange Date: Wed, 12 Nov 2025 18:45:00 +0800 Subject: [PATCH] indicator/strategy input --- api/trading.proto | 4 +- cmd/test/test.go | 4 - .../backtest/sig_strategy_backtester.go | 30 +++--- .../backtest/trading_plan_backtester.go | 28 +++--- internal/trading/sig/indicator_context.go | 95 +++---------------- internal/trading/sig/kline_series.go | 40 ++++++-- internal/trading/sig/strategy_context.go | 88 ++++++++++------- internal/trading/sig/trading_plan.go | 4 +- internal/trading/trading_grpc_server.go | 5 +- internal/trading/trading_service.go | 31 +++--- pkg/indicator/atr.go | 7 +- pkg/indicator/ema.go | 35 +++++++ pkg/indicator/indicator.go | 4 +- pkg/indicator/indicator_registry.go | 2 + pkg/indicator/macd.go | 34 +++++++ pkg/indicator/rsi.go | 2 +- pkg/indicator/sam.go | 4 +- pkg/strategy/cross_star.go | 12 +-- pkg/strategy/gold_x.go | 12 +-- pkg/strategy/sig_strategy.go | 16 ++-- pkg/strategy/sig_strategy_params.go | 6 -- pkg/strategy/super_trend.go | 20 ++-- pkg/types/input.go | 85 +++++++++++++++++ pkg/types/kline.go | 14 +-- pkg/types/series/floats.go | 6 ++ 25 files changed, 360 insertions(+), 228 deletions(-) create mode 100644 pkg/indicator/ema.go create mode 100644 pkg/indicator/macd.go create mode 100644 pkg/types/input.go diff --git a/api/trading.proto b/api/trading.proto index 754a8fe..f3e2d47 100644 --- a/api/trading.proto +++ b/api/trading.proto @@ -1,5 +1,6 @@ syntax = "proto3"; +import "google/protobuf/struct.proto"; import "api/pub.proto"; option go_package = "./pb"; @@ -30,6 +31,7 @@ message ReqIndicatorSeries { string indicator = 1; uint32 window = 2; // 指标窗口大小 SeriesRange series = 9; + google.protobuf.Struct input = 10; // 指标参数 } message RspIndicatorSeries{ repeated double matrix = 1; @@ -39,7 +41,7 @@ message RspIndicatorSeries{ message ReqStrategySeries { SeriesRange series = 1; string sigStrategy = 2; - map sigParam = 3; // 策略参数 + google.protobuf.Struct input = 3; // 指标参数 } message RspStrategySeries { repeated Side signal = 1; // 0.sell,1.buy diff --git a/cmd/test/test.go b/cmd/test/test.go index fcf9f87..9f1fec1 100644 --- a/cmd/test/test.go +++ b/cmd/test/test.go @@ -3,14 +3,10 @@ package main import ( "sig-pub/pkg/zlog" - "github.com/VictoriaMetrics/metrics" "github.com/govalues/decimal" ) func main() { - open := metrics.NewCounter("open") - open.Set(1234) - // curl -H 'Content-Type: application/json' --data-binary "@vmdata.json" -X POST http://localhost:8428/api/v1/import testDecimalScale() diff --git a/internal/trading/backtest/sig_strategy_backtester.go b/internal/trading/backtest/sig_strategy_backtester.go index 472d9e0..4cfe9b0 100644 --- a/internal/trading/backtest/sig_strategy_backtester.go +++ b/internal/trading/backtest/sig_strategy_backtester.go @@ -49,15 +49,21 @@ func (b *SigStrategyBacktester) SubKline(interval types.Interval, recv func(inte } // Backtest 基于历史数据回测信号策略 -func (b *SigStrategyBacktester) Backtest(ctx context.Context, sr *pb.SeriesRange, cleanIntervalSeries *types.IntervalState[*sig.KlineSeries], recvSignal func(sigSide types.Side, k types.Kline) (err error)) (err error) { +func (b *SigStrategyBacktester) Backtest(ctx context.Context, sigStrategyInput types.Input, sr *pb.SeriesRange, cleanIntervalSeries *types.IntervalState[*sig.KlineSeries], + recvSignal func(sigSide types.Side, k types.Kline) (err error)) (err error) { if cleanIntervalSeries == nil { cleanIntervalSeries = types.NewIntervalState[*sig.KlineSeries]() } + // 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), sr, cleanIntervalSeries, recvSignal) + err = b.singleStrategySeries(ctx, b.sigStrategy.(strategy.ISingleSigStrategy), sigStrategyInput, sr, cleanIntervalSeries, recvSignal) case strategy.SigStrategyTypeInterval: - err = b.intervalStrategySeries(ctx, b.sigStrategy.(strategy.IIntervalSigStrategy), sr, cleanIntervalSeries, recvSignal) + err = b.intervalStrategySeries(ctx, b.sigStrategy.(strategy.IIntervalSigStrategy), sigStrategyInput, sr, cleanIntervalSeries, recvSignal) default: err = fmt.Errorf("unknown sig strategy type %v", b.sigStrategyType) } @@ -65,12 +71,11 @@ func (b *SigStrategyBacktester) Backtest(ctx context.Context, sr *pb.SeriesRange } // singleStrategySeries 单周期策略 -func (b *SigStrategyBacktester) singleStrategySeries(ctx context.Context, sigStrategy strategy.ISingleSigStrategy, sr *pb.SeriesRange, intervalKlineSeries *types.IntervalState[*sig.KlineSeries], 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, intervalKlineSeries *types.IntervalState[*sig.KlineSeries], recvSignal func(sigSide types.Side, k types.Kline) (err error)) (err error) { interval := types.Interval(sr.Interval) kSeries := intervalKlineSeries.ComputeIfAbsent(interval, func() *sig.KlineSeries { return sig.NewKlineSeries(sr.Exchange, sr.InstId, interval) }) - indicatorContext := sig.NewIndicatorContext(kSeries) - strategyContext := sig.NewStrategyContext(indicatorContext, b.indicatorReg) - requiredSeries := int(sigStrategy.RequiredSeries()) + strategyContext := sig.NewStrategyContext(sigStrategyInput, kSeries, b.indicatorReg) + requiredSeries := int(sigStrategy.RequiredSeries(sigStrategyInput)) requiredIntervalSeries := types.NewIntervalState[int16]() requiredIntervalSeries.Set(interval, int16(max(1, requiredSeries))) @@ -95,11 +100,11 @@ func (b *SigStrategyBacktester) singleStrategySeries(ctx context.Context, sigStr } // intervalStrategySeries 多周期策略 -func (b *SigStrategyBacktester) intervalStrategySeries(ctx context.Context, intervalSigStrategy strategy.IIntervalSigStrategy, sr *pb.SeriesRange, intervalKlineSeries *types.IntervalState[*sig.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, intervalKlineSeries *types.IntervalState[*sig.KlineSeries], recvSignal func(sigSide types.Side, k types.Kline) (err error)) (err error) { // 各周期所需k线数量 - requiredIntervalSeries := intervalSigStrategy.RequiredIntervalSeries() + requiredIntervalSeries := intervalSigStrategy.RequiredIntervalSeries(sigStrategyInput) // 策略上下文 - intervalStrategyContext := sig.NewIntervalStrategyContext(intervalKlineSeries, b.indicatorReg) + intervalStrategyContext := sig.NewIntervalStrategyContext(sigStrategyInput, intervalKlineSeries, b.indicatorReg) err = b.multiIntervalSeries(ctx, sr, requiredIntervalSeries, intervalKlineSeries, func(driver bool, interval types.Interval, k *types.Kline) (err error) { if !driver { return @@ -325,10 +330,9 @@ recvLoop: 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) + strategyContext := sig.NewStrategyContext(nil, kSeries, b.indicatorReg) - requiredSeries := int(sigStrategy.RequiredSeries()) + requiredSeries := int(sigStrategy.RequiredSeries(nil)) 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 { diff --git a/internal/trading/backtest/trading_plan_backtester.go b/internal/trading/backtest/trading_plan_backtester.go index 503f155..1991102 100644 --- a/internal/trading/backtest/trading_plan_backtester.go +++ b/internal/trading/backtest/trading_plan_backtester.go @@ -24,13 +24,14 @@ type TradingPlanBacktester struct { sigStrategyReg *strategy.SigStrategyRegistry exchangeClient pb.ExchangeServiceClient - plan entity.TradePlan - account *BacktestAccount - sigStrategyType strategy.SigStrategyType - sigStrategy strategy.ISigStrategy - closeStrategy *trade.CloseStrategy - riskStrategy *trade.RiskStrategy - tradeStrategy *trade.TradeStrategy + plan entity.TradePlan + account *BacktestAccount + sigStrategyType strategy.SigStrategyType + sigStrategy strategy.ISigStrategy + sigStrategyInput types.Input + closeStrategy *trade.CloseStrategy + riskStrategy *trade.RiskStrategy + tradeStrategy *trade.TradeStrategy } func NewTradingPlanBacktester(indicatorReg *indicator.IndicatorRegistry, sigStrategyReg *strategy.SigStrategyRegistry, exchangeClient pb.ExchangeServiceClient) *TradingPlanBacktester { @@ -53,13 +54,13 @@ func (b *TradingPlanBacktester) Init(cash float64, plan entity.TradePlan) (err e err = fmt.Errorf("sig strategy not exists %s", plan.SigStrategy) return } - sigStrategyParam := make(strategy.StrategyParam) - if err = sonic.UnmarshalString(plan.SigStrategyParam, &sigStrategyParam); err != nil { + if err = sonic.UnmarshalString(plan.SigStrategyParam, &b.sigStrategyInput); err != nil { return } - if err = b.sigStrategy.Init(sigStrategyParam); err != nil { + if err = b.sigStrategy.Init(b.sigStrategyInput); err != nil { return } + closeStrategyParam, tradeStrategyParam, riskStrategyParam := new(trade.CloseStrategyParam), new(trade.TradeStrategyParam), new(trade.RiskStrategyParam) if err = sonic.UnmarshalString(plan.CloseStrategyParam, closeStrategyParam); err != nil { @@ -112,7 +113,7 @@ func (b *TradingPlanBacktester) Backtest(ctx context.Context, sr *pb.SeriesRange }) intervalSeries := types.NewIntervalState[*sig.KlineSeries]() - err = sigStrategyBacktester.Backtest(ctx, sr, intervalSeries, func(sigSide types.Side, k types.Kline) (err error) { + err = sigStrategyBacktester.Backtest(ctx, b.sigStrategyInput, sr, intervalSeries, func(sigSide types.Side, k types.Kline) (err error) { test.Singals++ // 根据交易信号检查仓位平仓 if err = b.closeBySigSingal(sigSide, k); err != nil { @@ -127,9 +128,8 @@ func (b *TradingPlanBacktester) Backtest(ctx context.Context, sr *pb.SeriesRange // 读最新的k线 kSeries := intervalSeries.Get(types.Interval(sr.Interval)) - lastCandle, ok := kSeries.Get(0) - if !ok { - err = fmt.Errorf("get series last candle error") + lastCandle, err := kSeries.Get(0) + if err != nil { return } // 关闭所有未平仓仓位 diff --git a/internal/trading/sig/indicator_context.go b/internal/trading/sig/indicator_context.go index 1f3d8a2..d82966a 100644 --- a/internal/trading/sig/indicator_context.go +++ b/internal/trading/sig/indicator_context.go @@ -1,16 +1,10 @@ package sig import ( - "context" - "fmt" - "io" - "sig-pub/api/pb" "sig-pub/pkg/indicator" "sig-pub/pkg/types" "sig-pub/pkg/types/series" "sig-pub/pkg/zlog" - - "google.golang.org/grpc" ) type IOffsetIndicatorContext interface { @@ -23,12 +17,14 @@ type IOffsetIndicatorContext interface { // IndicatorContext 指标上下文, 提供k线序列给指标计算使用 type IndicatorContext struct { IOffsetIndicatorContext + input types.Input kSeries *KlineSeries offset int16 } -func NewIndicatorContext(kSeries *KlineSeries) *IndicatorContext { +func NewIndicatorContext(input types.Input, kSeries *KlineSeries) *IndicatorContext { return &IndicatorContext{ + input: input, kSeries: kSeries, } } @@ -45,92 +41,27 @@ func (c *IndicatorContext) GetOffset() (offset int16) { return c.offset } +func (c *IndicatorContext) Input() (in types.Input) { + return c.input +} + func (c *IndicatorContext) Get(offset int16) (kline types.Kline) { offset += c.offset - k, ok := c.kSeries.Get(offset) - if !ok { + k, err := c.kSeries.Get(offset) + if err != nil { lastTs := c.kSeries.LastTs() zlog.Warningf("get kline series offset out of range: offset=%d, length=%d, lastTs=%d", offset, c.kSeries.Length(), lastTs) - panic(fmt.Errorf("get kline series offset out of range: offset=%d", offset)) + panic(err) } return k } func (c *IndicatorContext) Series(offset, count int16) (klines series.Klines) { offset += c.offset - ks, ok := c.kSeries.Series(offset, count) - if !ok { + ks, err := c.kSeries.Series(offset, count) + if err != nil { zlog.Warningf("get kline series offset out of range: offset=%d, count=%d, length=%d, lastTs=%d", offset, count, c.kSeries.Length(), c.kSeries.LastTs()) - panic(fmt.Errorf("get kline series offset out of range: offset=%d, count=%d", offset, count)) + panic(err) } return ks } - -// Deprecated: 用流处理(exchange rpc stream) -type HistoryIndicatorContext struct { - IOffsetIndicatorContext - exchangeClient pb.ExchangeServiceClient - context *IndicatorContext -} - -func NewHistoryIndicatorContext(exchangeClient pb.ExchangeServiceClient) *HistoryIndicatorContext { - return &HistoryIndicatorContext{ - exchangeClient: exchangeClient, - } -} -func (c *HistoryIndicatorContext) Init(sr *pb.SeriesRange) (totalK int, err error) { - // fetch history series - req := &pb.ReqHistoryKlineStream{ - Series: sr, - } - stream, err := c.exchangeClient.HistoryKlineStream(context.Background(), req, grpc.UseCompressor("snappy")) - if err != nil { - zlog.Errorf("fetch history kline stream error: instId=%s(%s), interval=%s, %#v, err=%v", sr.InstId, sr.Exchange, sr.Interval, req, err) - return - } - interval := types.Interval(sr.Interval) - klineSeries := NewKlineSeries(sr.Exchange, sr.InstId, interval) - for { - msg, err0 := stream.Recv() - if err0 == io.EOF { - break - } - if err0 != nil { - err = err0 - zlog.Error("fetch kline stream recv error: ", err0) - return - } - // zlog.Debugf("recv: %s(%s), %s, branch=%d, ts=%d~%d", instId, exchange, interval, len(msg.Klines), msg.Klines[0].Ts, msg.Klines[len(msg.Klines)-1].Ts) - totalK += len(msg.Klines) - for _, k := range msg.Klines { - kline := new(types.Kline) - kline.ParsePBKline(sr.Exchange, k) - if lastTs, ok := klineSeries.Update(kline); !ok { - err = fmt.Errorf("history stream kline not series: last=%d", lastTs) - return - } - } - } - c.context = NewIndicatorContext(klineSeries) - return -} - -func (c *HistoryIndicatorContext) SetOffset(offset int16) { - c.context.SetOffset(offset) -} - -func (c *HistoryIndicatorContext) AddOffset(offset int16) { - c.context.AddOffset(offset) -} - -func (c *HistoryIndicatorContext) GetOffset() (offset int16) { - return c.context.GetOffset() -} - -func (c *HistoryIndicatorContext) Get(offset int16) (kline types.Kline) { - return c.context.Get(offset) -} - -func (c *HistoryIndicatorContext) Series(offset, count int16) (klines series.Klines) { - return c.context.Series(offset, count) -} diff --git a/internal/trading/sig/kline_series.go b/internal/trading/sig/kline_series.go index beb84e9..40829b7 100644 --- a/internal/trading/sig/kline_series.go +++ b/internal/trading/sig/kline_series.go @@ -58,8 +58,18 @@ func NewKlineSeries(exchange pb.ExchangeType, instId string, interval types.Inte } // Get [0]当前k线 -func (s *KlineSeries) Get(offset int16) (k types.Kline, ok bool) { - if ok = offset >= 0 && offset < MaxSeriesKlines; !ok { +func (s *KlineSeries) MustGet(offset int16) (k types.Kline) { + k, err := s.Get(offset) + if err != nil { + panic(err) + } + return +} + +// GetE [0]当前k线 +func (s *KlineSeries) Get(offset int16) (k types.Kline, err error) { + if ok := offset >= 0 && offset < MaxSeriesKlines; !ok { + err = fmt.Errorf("get kline series offset out of range: offset=%d", offset) return } s.mu.RLock() @@ -67,20 +77,31 @@ func (s *KlineSeries) Get(offset int16) (k types.Kline, ok bool) { length := len(s.klines) index := (length - 1) - int(offset) - if ok = index >= 0 && index < length; !ok { + if ok := index >= 0 && index < length; !ok { + err = fmt.Errorf("get kline series offset out of range: offset=%d", offset) return } - return *(s.klines[index]), true + return *(s.klines[index]), nil +} + +func (s *KlineSeries) MustSeries(offset, count int16) (klines series.Klines) { + klines, err := s.Series(offset, count) + if err != nil { + panic(err) + } + return } // Series 时间降序序列[count...offset] // offset: 从序列尾部开始偏移量 // count: 从offset位置开始向序列头部k线条数 -func (s *KlineSeries) Series(offset, count int16) (klines series.Klines, ok bool) { - if ok = offset >= 0 && offset < MaxSeriesKlines; !ok { +func (s *KlineSeries) Series(offset, count int16) (klines series.Klines, err error) { + if ok := offset >= 0 && offset < MaxSeriesKlines; !ok { + err = fmt.Errorf("get kline series offset out of range: offset=%d, count=%d", offset, count) return } - if ok = count > 0 && offset+count < MaxSeriesKlines; !ok { + if ok := count > 0 && offset+count < MaxSeriesKlines; !ok { + err = fmt.Errorf("get kline series offset out of range: offset=%d, count=%d", offset, count) return } @@ -90,7 +111,8 @@ func (s *KlineSeries) Series(offset, count int16) (klines series.Klines, ok bool length := len(s.klines) indexEnd := (length - 1) - int(offset) indexStart := (length - 1) - int(offset) - int(count) + 1 - if ok = indexEnd >= 0 && indexEnd < length && indexStart >= 0 && indexStart < length; !ok { + if ok := indexEnd >= 0 && indexEnd < length && indexStart >= 0 && indexStart < length; !ok { + err = fmt.Errorf("get kline series offset out of range: offset=%d, count=%d", offset, count) return } total := indexEnd - indexStart + 1 @@ -99,7 +121,7 @@ func (s *KlineSeries) Series(offset, count int16) (klines series.Klines, ok bool offset := total - 1 - (i - indexStart) klines[offset] = *(s.klines[i]) } - return klines, true + return klines, nil } func (s *KlineSeries) Length() int { diff --git a/internal/trading/sig/strategy_context.go b/internal/trading/sig/strategy_context.go index 47b8a26..389837f 100644 --- a/internal/trading/sig/strategy_context.go +++ b/internal/trading/sig/strategy_context.go @@ -11,81 +11,105 @@ import ( type StrategyContext struct { strategy.ISingleSigStrategyContext - indicatorContext IOffsetIndicatorContext - indicatorsReg *indicator.IndicatorRegistry + input types.Input + kSeries *KlineSeries + indicatorsReg *indicator.IndicatorRegistry } -func NewStrategyContext(indicatorContext IOffsetIndicatorContext, indicatorsReg *indicator.IndicatorRegistry) *StrategyContext { +func NewStrategyContext(input types.Input, kSeries *KlineSeries, indicatorsReg *indicator.IndicatorRegistry) *StrategyContext { return &StrategyContext{ - indicatorContext: indicatorContext, - indicatorsReg: indicatorsReg, + input: input, + kSeries: kSeries, + indicatorsReg: indicatorsReg, } } +// Input 获取输入参数 +func (c *StrategyContext) Input() (in types.Input) { + return c.input +} + func (c *StrategyContext) Get(offset int16) (kline types.Kline) { - return c.indicatorContext.Get(offset) + return c.kSeries.MustGet(offset) } func (c *StrategyContext) Series(offset, count int16) (klines series.Klines) { - return c.indicatorContext.Series(offset, count) + return c.kSeries.MustSeries(offset, count) } // 获取窗口类型指标 -func (c *StrategyContext) IndicatorW(name string, window int16) (s indicator.IIndicatorSeries) { +func (c *StrategyContext) IndicatorW(name string, window int16, args ...any) (s indicator.IIndicatorSeries) { indicator, ok := c.indicatorsReg.IndicatorW(name) if !ok { panic(fmt.Errorf("indicatorW %s not exists", name)) } - return NewWindowIndicatorSeries(window, indicator, c.indicatorContext) + var input types.Input + if len(args) > 0 { + if in, ok := args[0].(types.Input); ok { + input = in + } + } + indicatorContext := NewIndicatorContext(input, c.kSeries) + return NewWindowIndicatorSeries(window, indicator, indicatorContext) } // IntervalStrategyContext 周期策略上下文 type IntervalStrategyContext struct { strategy.IIntervalSigStrategyContext - intervalIndicatorContexts map[types.Interval]*IndicatorContext - intervalKlineSeries *types.IntervalState[*KlineSeries] - indicatorsReg *indicator.IndicatorRegistry + + input types.Input + intervalKlineSeries *types.IntervalState[*KlineSeries] + indicatorsReg *indicator.IndicatorRegistry } -func NewIntervalStrategyContext(intervalKlineSeries *types.IntervalState[*KlineSeries], indicatorsReg *indicator.IndicatorRegistry) *IntervalStrategyContext { +func NewIntervalStrategyContext(input types.Input, intervalKlineSeries *types.IntervalState[*KlineSeries], indicatorsReg *indicator.IndicatorRegistry) *IntervalStrategyContext { return &IntervalStrategyContext{ - intervalIndicatorContexts: make(map[types.Interval]*IndicatorContext), - intervalKlineSeries: intervalKlineSeries, - indicatorsReg: indicatorsReg, + input: input, + intervalKlineSeries: intervalKlineSeries, + indicatorsReg: indicatorsReg, } } -func (c *IntervalStrategyContext) getIndicatorContext(interval types.Interval) *IndicatorContext { - ctx, ok := c.intervalIndicatorContexts[interval] - if !ok { - klineSeries := c.intervalKlineSeries.Get(interval) - if klineSeries == nil { - panic(fmt.Errorf("interval %s kline series is nil", interval)) - } - ctx = NewIndicatorContext(klineSeries) - c.intervalIndicatorContexts[interval] = ctx +// Input 获取输入参数 +func (c *IntervalStrategyContext) Input() (in types.Input) { + return c.input +} + +func (c *IntervalStrategyContext) getCandleSeries(interval types.Interval) *KlineSeries { + klineSeries := c.intervalKlineSeries.Get(interval) + if klineSeries == nil { + panic(fmt.Errorf("interval %s kline series is nil", interval)) } - return ctx + return klineSeries } // Get [0]当前k线 func (c *IntervalStrategyContext) Get(interval types.Interval, offset int16) (kline types.Kline) { - ctx := c.getIndicatorContext(interval) - return ctx.Get(offset) + ks := c.getCandleSeries(interval) + return ks.MustGet(offset) } // Series [offset...end] func (c *IntervalStrategyContext) Series(interval types.Interval, offset, count int16) (klines series.Klines) { - ctx := c.getIndicatorContext(interval) - return ctx.Series(offset, count) + cs := c.getCandleSeries(interval) + return cs.MustSeries(offset, count) } // 获取窗口类型指标 -func (c *IntervalStrategyContext) IndicatorW(interval types.Interval, name string, window int16) (series indicator.IIndicatorSeries) { - indicatorContext := c.getIndicatorContext(interval) +func (c *IntervalStrategyContext) IndicatorW(interval types.Interval, name string, window int16, args ...any) (series indicator.IIndicatorSeries) { indicator, ok := c.indicatorsReg.IndicatorW(name) if !ok { panic(fmt.Errorf("indicatorW %s not exists", name)) } + + var input types.Input + if len(args) > 0 { + if in, ok := args[0].(types.Input); ok { + input = in + } + } + cs := c.getCandleSeries(interval) + indicatorContext := NewIndicatorContext(input, cs) + return NewWindowIndicatorSeries(window, indicator, indicatorContext) } diff --git a/internal/trading/sig/trading_plan.go b/internal/trading/sig/trading_plan.go index 6534813..8b58f41 100644 --- a/internal/trading/sig/trading_plan.go +++ b/internal/trading/sig/trading_plan.go @@ -42,8 +42,8 @@ func (r *TradingPlan) Init() (err error) { // initSigStrategy 初始化多空信号策略 // buy/sell -> 过滤/风控 -> tradeStrategy -> closeStrategy -func (r *TradingPlan) InitSigStrategy(sigStrategyType strategy.SigStrategyType, sigStrategy strategy.ISigStrategy, param strategy.StrategyParam, sigStrategyContext strategy.ISingleSigStrategyContext) (err error) { - if err = sigStrategy.Init(param); err != nil { +func (r *TradingPlan) InitSigStrategy(sigStrategyType strategy.SigStrategyType, sigStrategy strategy.ISigStrategy, sigStrategyInput types.Input, sigStrategyContext strategy.ISingleSigStrategyContext) (err error) { + if err = sigStrategy.Init(sigStrategyInput); err != nil { return } r.sigStrategyType = sigStrategyType diff --git a/internal/trading/trading_grpc_server.go b/internal/trading/trading_grpc_server.go index 89015c6..bba5f91 100644 --- a/internal/trading/trading_grpc_server.go +++ b/internal/trading/trading_grpc_server.go @@ -3,6 +3,7 @@ package trading import ( "context" "sig-pub/api/pb" + "sig-pub/pkg/types" "sig-pub/pkg/utils/times" ) @@ -22,7 +23,9 @@ func (svr *TradingGrpcServer) Init() (err error) { } func (svr *TradingGrpcServer) IndicatorSeries(ctx context.Context, req *pb.ReqIndicatorSeries) (rsp *pb.RspIndicatorSeries, err error) { - matrix, times, err := svr.tradingService.IndicatorSeries(ctx, req.Indicator, req.Window, req.Series) + // s, err := structpb.NewStruct(map[string]any{}) + input := req.Input.AsMap() + matrix, times, err := svr.tradingService.IndicatorSeries(ctx, req.Indicator, req.Window, types.Input(input), req.Series) if err != nil { return } diff --git a/internal/trading/trading_service.go b/internal/trading/trading_service.go index 2ddd4cd..9a6ba27 100644 --- a/internal/trading/trading_service.go +++ b/internal/trading/trading_service.go @@ -109,8 +109,8 @@ func (svc *TradingService) getTradingPlan(plan *entity.TradePlan, sigKlineSeries } // sigStrategy - sigStrategyParam := make(strategy.StrategyParam) - if err = sonic.UnmarshalString(plan.SigStrategyParam, &sigStrategyParam); err != nil { + var sigStrategyInput types.Input + if err = sonic.UnmarshalString(plan.SigStrategyParam, &sigStrategyInput); err != nil { return } sigStrategyType, sigStrategy, ok := svc.strategyReg.NewSigStrategy(plan.SigStrategy) @@ -130,11 +130,12 @@ func (svc *TradingService) getTradingPlan(plan *entity.TradePlan, sigKlineSeries if err = tradingPlan.Init(); err != nil { return } - sigIndicatorCtx := sig.NewIndicatorContext(sigKlineSeries) - sigStrategyCtx := sig.NewStrategyContext(sigIndicatorCtx, svc.indicatorReg) - if err = tradingPlan.InitSigStrategy(sigStrategyType, sigStrategy, sigStrategyParam, sigStrategyCtx); err != nil { - return - } + _, _ = sigStrategyType, sigStrategy + // sigIndicatorCtx := sig.NewIndicatorContext(sigKlineSeries) + // sigStrategyCtx := sig.NewStrategyContext(sigIndicatorCtx, svc.indicatorReg) + // if err = tradingPlan.InitSigStrategy(sigStrategyType, sigStrategy, sigStrategyParam, sigStrategyCtx); err != nil { + // return + // } return // if _, load := svc.tradingPlans.LoadOrStore(planId, tradingPlan); load { @@ -197,8 +198,7 @@ func (svc *TradingService) fetchHistoryKlineSeries(ctx context.Context, sr *pb.S } // IndicatorSeries 获取指标实时或历史序列数据, 闭区间 -func (svc *TradingService) IndicatorSeries(ctx context.Context, indicatorName string, window uint32, sr *pb.SeriesRange) (matrix []float64, times []int64, err error) { - // indicatorName string, exchange pb.ExchangeType, instId string, interval types.Interval, window int +func (svc *TradingService) IndicatorSeries(ctx context.Context, indicatorName string, window uint32, input types.Input, sr *pb.SeriesRange) (matrix []float64, times []int64, err error) { indicator, ok := svc.indicatorReg.IndicatorW(indicatorName) if !ok { err = fmt.Errorf("indicator %s not exists", indicatorName) @@ -212,10 +212,10 @@ func (svc *TradingService) IndicatorSeries(ctx context.Context, indicatorName st } // 查询历史指标数据 - requiredSeries := int(indicator.RequiredSeries(int16(window))) + requiredSeries := int(indicator.RequiredSeries(int16(window), input)) sr.WindowExtra = uint32(max(0, requiredSeries-1)) kSeries := sig.NewKlineSeries(sr.Exchange, sr.InstId, types.Interval(sr.Interval)) - indicatorContext := sig.NewIndicatorContext(kSeries) + indicatorContext := sig.NewIndicatorContext(input, kSeries) matrix = make([]float64, 0, 200) times = make([]int64, 0, 200) @@ -246,10 +246,6 @@ func (svc *TradingService) StrategySeries(ctx context.Context, req *pb.ReqStrate err = fmt.Errorf("strategy %s not exists", req.SigStrategy) return } - // init sigStrategy - if err = sigStrategy.Init(req.SigParam); err != nil { - return - } interval := types.Interval(req.Series.Interval) _, ok = types.SupportedIntervals[interval] @@ -258,9 +254,12 @@ func (svc *TradingService) StrategySeries(ctx context.Context, req *pb.ReqStrate return } + // 信号策略参数 + sigStrategyInput := types.Input(req.Input.AsMap()) + // 使用回测器回测信号 backtester := backtest.NewSigStrategyBacktester(sigStrategyType, sigStrategy, svc.indicatorReg, svc.exchangeClient) - err = backtester.Backtest(ctx, req.Series, nil, func(sigSide types.Side, k types.Kline) (err error) { + err = backtester.Backtest(ctx, sigStrategyInput, req.Series, nil, 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) diff --git a/pkg/indicator/atr.go b/pkg/indicator/atr.go index be30ecb..ad10318 100644 --- a/pkg/indicator/atr.go +++ b/pkg/indicator/atr.go @@ -1,6 +1,9 @@ package indicator -import "sig-pub/pkg/types/series" +import ( + "sig-pub/pkg/types" + "sig-pub/pkg/types/series" +) // ATR = SMA(TR, N) // 平均真实波幅 (ATR) atr define: https://www.investopedia.com/terms/a/atr.asp @@ -12,7 +15,7 @@ func (c *ATR) Name() string { return "atr" } -func (c *ATR) RequiredSeries(window int16) int16 { +func (c *ATR) RequiredSeries(window int16, in types.Input) int16 { return window + 1 } diff --git a/pkg/indicator/ema.go b/pkg/indicator/ema.go new file mode 100644 index 0000000..36d7a9d --- /dev/null +++ b/pkg/indicator/ema.go @@ -0,0 +1,35 @@ +package indicator + +import ( + "sig-pub/pkg/types" +) + +// EMA stateful indicator +type EMA struct { +} + +func (c *EMA) Name() string { + return "ema" +} + +func (c *EMA) RequiredSeries(window int16, in types.Input) int16 { + return window + 1 +} + +// Calculate 计算单根k线sma指标 +func (c *EMA) Calculate(ctx IIndicatorContext, window int16) (vector float64) { + alpha := 2.0 / float64(window+1) + + close := ctx.Get(0).CloseF64() + prevSMA := ctx.Series(1, window).Close().Avg() + + vector = ((close - prevSMA) * alpha) + prevSMA + + // ctx.GetSelf(0) // 自己计算的上一个值 + + // 计算eam + // closeSeries := ctx.Series(0, window).Close().Reverse() + // ema := talib.Ema(closeSeries, int(window)) + // vector = ema[len(ema)-1] + return +} diff --git a/pkg/indicator/indicator.go b/pkg/indicator/indicator.go index 4418cf1..1711fb9 100644 --- a/pkg/indicator/indicator.go +++ b/pkg/indicator/indicator.go @@ -20,13 +20,15 @@ type IWindowIndicator interface { // Name 指标名称 Name() string // RequiredSeries 计算窗口大小的指标值需要的K线数量 - RequiredSeries(window int16) int16 + RequiredSeries(window int16, in types.Input) int16 // Calculate 计算窗口大小的指标值 Calculate(ctx IIndicatorContext, window int16) (vector float64) } // IIndicatorContext k线序列, trading服务提供 type IIndicatorContext interface { + // Input 获取输入参数 + Input() types.Input Get(offset int16) (kline types.Kline) Series(offset, count int16) (klines series.Klines) } diff --git a/pkg/indicator/indicator_registry.go b/pkg/indicator/indicator_registry.go index 84be087..e55fbb8 100644 --- a/pkg/indicator/indicator_registry.go +++ b/pkg/indicator/indicator_registry.go @@ -21,6 +21,8 @@ func (r *IndicatorRegistry) Init() (err error) { r.MustRegistIndicatorW(&RSI{}) r.MustRegistIndicatorW(&SMA{}) r.MustRegistIndicatorW(&ATR{}) + r.MustRegistIndicatorW(&EMA{}) + r.MustRegistIndicatorW(&MACD{}) return } diff --git a/pkg/indicator/macd.go b/pkg/indicator/macd.go new file mode 100644 index 0000000..40f117f --- /dev/null +++ b/pkg/indicator/macd.go @@ -0,0 +1,34 @@ +package indicator + +import ( + "sig-pub/pkg/types" + + "github.com/markcheno/go-talib" +) + +// todo macdSignal(信号线) macdHist(柱状图) +type MACD struct { +} + +func (c *MACD) Name() string { + return "macd" +} + +func (c *MACD) RequiredSeries(window int16, in types.Input) int16 { + return window +} + +// Calculate 计算单根k线sma指标 +func (c *MACD) Calculate(ctx IIndicatorContext, window int16) (vector float64) { + fast := ctx.Input().Int("fast") + slow := ctx.Input().Int("slow") + + // 计算eam + closeSeries := ctx.Series(0, window).Close().Reverse() + + aa, bb, cc := talib.Macd(closeSeries, fast, slow, int(window)) + _, _, _ = aa, bb, cc + + vector = 0 + return +} diff --git a/pkg/indicator/rsi.go b/pkg/indicator/rsi.go index 6f6df55..086892c 100644 --- a/pkg/indicator/rsi.go +++ b/pkg/indicator/rsi.go @@ -15,7 +15,7 @@ func (c *RSI) Name() string { return "rsi" } -func (c *RSI) RequiredSeries(window int16) int16 { +func (c *RSI) RequiredSeries(window int16, in types.Input) int16 { return window } diff --git a/pkg/indicator/sam.go b/pkg/indicator/sam.go index 35ff432..510d95f 100644 --- a/pkg/indicator/sam.go +++ b/pkg/indicator/sam.go @@ -1,6 +1,8 @@ package indicator import ( + "sig-pub/pkg/types" + "github.com/markcheno/go-talib" ) @@ -14,7 +16,7 @@ func (c *SMA) Name() string { return "sma" } -func (c *SMA) RequiredSeries(window int16) int16 { +func (c *SMA) RequiredSeries(window int16, in types.Input) int16 { return window } diff --git a/pkg/strategy/cross_star.go b/pkg/strategy/cross_star.go index a734268..0d1bc17 100644 --- a/pkg/strategy/cross_star.go +++ b/pkg/strategy/cross_star.go @@ -28,17 +28,13 @@ func (s *CrossStar) Meta() StrategyMeta { } } -func (s *CrossStar) Init(param StrategyParam) (err error) { // 校验参数, 并根据参数初始化策略 - if s.rate, err = param.GetFloat64E("rate"); err != nil { - return - } - if s.rate2, err = param.GetFloat64E("rate2"); err != nil { - return - } +func (s *CrossStar) Init(input types.Input) (err error) { // 校验参数, 并根据参数初始化策略 + s.rate = input.Float("rate") + s.rate2 = input.Float("rate2") return } -func (s *CrossStar) RequiredIntervalSeries() (iss *types.IntervalState[int16]) { +func (s *CrossStar) RequiredIntervalSeries(input types.Input) (iss *types.IntervalState[int16]) { iss = types.NewIntervalState[int16]() iss.Set(types.Interval5m, 1) iss.Set(types.Interval15m, 2) diff --git a/pkg/strategy/gold_x.go b/pkg/strategy/gold_x.go index 7e8b05a..cd76263 100644 --- a/pkg/strategy/gold_x.go +++ b/pkg/strategy/gold_x.go @@ -26,13 +26,9 @@ func (s *GoldX) Meta() StrategyMeta { } } -func (s *GoldX) Init(param StrategyParam) (err error) { // 校验参数, 并根据参数初始化策略 - if s.short, err = param.GetInt16E("short"); err != nil { - return - } - if s.long, err = param.GetInt16E("long"); err != nil { - return - } +func (s *GoldX) Init(input types.Input) (err error) { // 校验参数, 并根据参数初始化策略 + s.short = input.Int16("short") + s.long = input.Int16("long") if s.long <= s.short { err = fmt.Errorf("param short should bigger then short") return @@ -40,7 +36,7 @@ func (s *GoldX) Init(param StrategyParam) (err error) { // 校验参数, 并根 return } -func (s *GoldX) RequiredSeries() int16 { +func (s *GoldX) RequiredSeries(input types.Input) int16 { return max(s.long, s.short) + 1 } diff --git a/pkg/strategy/sig_strategy.go b/pkg/strategy/sig_strategy.go index e6f697e..6b8495e 100644 --- a/pkg/strategy/sig_strategy.go +++ b/pkg/strategy/sig_strategy.go @@ -10,7 +10,7 @@ import ( type ISigStrategy interface { New() ISigStrategy Meta() StrategyMeta - Init(param StrategyParam) (err error) // 校验参数, 并根据参数初始化策略 + Init(input types.Input) (err error) // 校验参数, 并根据参数初始化策略 } type StrategyMeta struct { @@ -22,33 +22,37 @@ type StrategyMeta struct { // ISingleSigStrategy 单周期单交易所策略 type ISingleSigStrategy interface { ISigStrategy - RequiredSeries() int16 // 需要的最小数据k线数, 回测时用, 若不定义则取最大窗口值 + RequiredSeries(input types.Input) int16 // 需要的最小数据k线数, 回测时用, 若不定义则取最大窗口值 Update(ctx ISingleSigStrategyContext) (side types.Side) } // ISingleSigStrategyContext 策略上下文 type ISingleSigStrategyContext interface { + // Input 获取输入参数 + Input() types.Input // Get [0]当前k线 Get(offset int16) types.Kline // Series [offset...end] Series(offset, count int16) (klines series.Klines) - // 获取窗口类型指标 - IndicatorW(name string, window int16) indicator.IIndicatorSeries + // IndicatorW 获取窗口类型指标 + IndicatorW(name string, window int16, args ...any) indicator.IIndicatorSeries } // 多周期k线策略接口 type IIntervalSigStrategy interface { ISigStrategy - RequiredIntervalSeries() (iss *types.IntervalState[int16]) // 需要的各周期最小数据k线数, 回测时用, 若不定义则取最大窗口值 + RequiredIntervalSeries(input types.Input) (iss *types.IntervalState[int16]) // 需要的各周期最小数据k线数, 回测时用, 若不定义则取最大窗口值 Update(ctx IIntervalSigStrategyContext) (side types.Side) } // IIntervalSigStrategyContext 多周期策略上下文 type IIntervalSigStrategyContext interface { + // Input 获取输入参数 + Input() types.Input // Get [0]当前k线 Get(interval types.Interval, offset int16) types.Kline // Series [offset...end] Series(interval types.Interval, offset, count int16) (klines series.Klines) // 获取窗口类型指标 - IndicatorW(interval types.Interval, name string, window int16) indicator.IIndicatorSeries + IndicatorW(interval types.Interval, name string, window int16, args ...any) indicator.IIndicatorSeries } diff --git a/pkg/strategy/sig_strategy_params.go b/pkg/strategy/sig_strategy_params.go index 8b0720b..ca35abc 100644 --- a/pkg/strategy/sig_strategy_params.go +++ b/pkg/strategy/sig_strategy_params.go @@ -2,7 +2,6 @@ package strategy import ( "fmt" - "sig-pub/pkg/types" "github.com/spf13/cast" ) @@ -79,11 +78,6 @@ type ISigStrategyParamGenerator interface { NextParam(map[string]string) (map[string]string, bool) // 根据当前策略参数, 返回下一批策略参数(并行回测 stateless) } -type IntervalStrategyParam struct { - Interval types.Interval `json:"interval"` // 策略驱动周期 - Param StrategyParam `json:"param"` // 策略执行参数 -} - type StrategyParam map[string]string func (s *StrategyParam) Get(key string) (v string, ok bool) { diff --git a/pkg/strategy/super_trend.go b/pkg/strategy/super_trend.go index 6b6a590..d8aa83d 100644 --- a/pkg/strategy/super_trend.go +++ b/pkg/strategy/super_trend.go @@ -48,23 +48,15 @@ func (s *SupertrendBOSWaves) Meta() StrategyMeta { } } -func (s *SupertrendBOSWaves) Init(param StrategyParam) (err error) { // 校验参数, 并根据参数初始化策略 - if s.atrLength, err = param.GetInt16E("atrLength"); err != nil { - return - } - if s.atrMult, err = param.GetFloat64E("atrMult"); err != nil { - return - } - // if s.radiusStrength, err = param.GetFloat64E("radiusStrength"); err != nil { - // return - // } - // if s.smoothness, err = param.GetInt16E("smoothness"); err != nil { - // return - // } +func (s *SupertrendBOSWaves) Init(input types.Input) (err error) { // 校验参数, 并根据参数初始化策略 + s.atrLength = input.Int16("atrLength") + s.atrMult = input.Float("atrMult") + // s.radiusStrength = input.Float("radiusStrength") + // s.smoothness = input.Float("smoothness") return } -func (s *SupertrendBOSWaves) RequiredSeries() int16 { +func (s *SupertrendBOSWaves) RequiredSeries(input types.Input) int16 { return s.atrLength + 1 } diff --git a/pkg/types/input.go b/pkg/types/input.go new file mode 100644 index 0000000..8ef3215 --- /dev/null +++ b/pkg/types/input.go @@ -0,0 +1,85 @@ +package types + +import ( + "fmt" + + "github.com/spf13/cast" +) + +const ( + inputCacheKey = "__$cache__" +) + +type Input map[string]any + +// getCache 避免多线程读写cache map +func (in Input) getCache(k string) (r any, ok bool) { + if in == nil { + return + } + c, ok := in[inputCacheKey] + if !ok { + return + } + r, ok = c.(map[string]any)[k] + return +} + +func (in Input) setCache(k string, v any) { + if in == nil { + return + } + c, ok := in[inputCacheKey] + if !ok { + c = make(map[string]any, 4) + in[inputCacheKey] = c + } + c.(map[string]any)[k] = v +} + +func (in Input) get(k string, t string) (r any) { + if in == nil { + panic(fmt.Errorf("input %s type %s not provide", k, t)) + } + r, ok := in[k] + if !ok { + panic(fmt.Errorf("input %s type %s not provide", k, t)) + } + return +} + +func (in Input) Float(k string) (v float64) { + if r, ok := in.getCache(k); ok { + return r.(float64) + } + v, err := cast.ToFloat64E(in.get(k, "float")) + if err != nil { + panic(fmt.Errorf("input float parse error: %s", k)) + } + in.setCache(k, v) + return +} + +func (in Input) Int(k string) (v int) { + if r, ok := in.getCache(k); ok { + return r.(int) + } + v, err := cast.ToIntE(in.get(k, "int")) + if err != nil { + panic(fmt.Errorf("input int parse error: %s", k)) + } + in.setCache(k, v) + return +} + +func (in Input) Int16(k string) (v int16) { + if r, ok := in.getCache(k); ok { + return r.(int16) + } + v, err := cast.ToInt16E(in.get(k, "int16")) + if err != nil { + panic(fmt.Errorf("input int16 parse error: %s", k)) + } + in.setCache(k, v) + return +} diff --git a/pkg/types/kline.go b/pkg/types/kline.go index 607132b..3f7aece 100644 --- a/pkg/types/kline.go +++ b/pkg/types/kline.go @@ -60,31 +60,31 @@ func (k *Kline) ToPBKline() (kline *pb.Kline) { return } -func (k *Kline) OpenF64() float64 { +func (k Kline) OpenF64() float64 { return decimals.MustToFloat64(k.Open) } -func (k *Kline) CloseF64() float64 { +func (k Kline) CloseF64() float64 { return decimals.MustToFloat64(k.Close) } -func (k *Kline) HighF64() float64 { +func (k Kline) HighF64() float64 { return decimals.MustToFloat64(k.High) } -func (k *Kline) LowF64() float64 { +func (k Kline) LowF64() float64 { return decimals.MustToFloat64(k.Low) } -func (k *Kline) VolF64() float64 { +func (k Kline) VolF64() float64 { return decimals.MustToFloat64(k.Vol) } -func (k *Kline) VolQtyF64() float64 { +func (k Kline) VolQtyF64() float64 { return decimals.MustToFloat64(k.VolQuote) } -func (k *Kline) HL2() float64 { +func (k Kline) HL2() float64 { return (k.HighF64() + k.LowF64()) / 2 } diff --git a/pkg/types/series/floats.go b/pkg/types/series/floats.go index bb8677a..d86afa2 100644 --- a/pkg/types/series/floats.go +++ b/pkg/types/series/floats.go @@ -2,6 +2,7 @@ package series import ( "math" + "sig-pub/pkg/utils/collect" "gonum.org/v1/gonum/floats" ) @@ -25,6 +26,11 @@ func (s Floats) Length() int { return len(s) } +func (s Floats) Reverse() (r Floats) { + collect.Reverse(s) + return s +} + func (s Floats) Diff() (values Floats) { length := s.Length() for i, v := range s {