Browse Source

indicator/strategy input

main
strange 10 months ago
parent
commit
614548e018
  1. 4
      api/trading.proto
  2. 4
      cmd/test/test.go
  3. 30
      internal/trading/backtest/sig_strategy_backtester.go
  4. 28
      internal/trading/backtest/trading_plan_backtester.go
  5. 95
      internal/trading/sig/indicator_context.go
  6. 40
      internal/trading/sig/kline_series.go
  7. 88
      internal/trading/sig/strategy_context.go
  8. 4
      internal/trading/sig/trading_plan.go
  9. 5
      internal/trading/trading_grpc_server.go
  10. 31
      internal/trading/trading_service.go
  11. 7
      pkg/indicator/atr.go
  12. 35
      pkg/indicator/ema.go
  13. 4
      pkg/indicator/indicator.go
  14. 2
      pkg/indicator/indicator_registry.go
  15. 34
      pkg/indicator/macd.go
  16. 2
      pkg/indicator/rsi.go
  17. 4
      pkg/indicator/sam.go
  18. 12
      pkg/strategy/cross_star.go
  19. 12
      pkg/strategy/gold_x.go
  20. 16
      pkg/strategy/sig_strategy.go
  21. 6
      pkg/strategy/sig_strategy_params.go
  22. 20
      pkg/strategy/super_trend.go
  23. 85
      pkg/types/input.go
  24. 14
      pkg/types/kline.go
  25. 6
      pkg/types/series/floats.go

4
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<string,string> sigParam = 3; //
google.protobuf.Struct input = 3; //
}
message RspStrategySeries {
repeated Side signal = 1; // 0.sell,1.buy

4
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()

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

28
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
}
// 关闭所有未平仓仓位

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

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

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

4
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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

Loading…
Cancel
Save