diff --git a/api/trading.proto b/api/trading.proto index 7e11076..0cfddcc 100644 --- a/api/trading.proto +++ b/api/trading.proto @@ -8,6 +8,7 @@ service TradingService { rpc SubIndicator(IndicatorSubReq) returns (stream Indicator); // 订阅指标 rpc IndicatorSeries(ReqIndicatorSeries) returns (RspIndicatorSeries); // 获取指标序列 rpc StrategySeries(ReqStrategySeries) returns (RspStrategySeries); // 获取指标序列 + rpc Backtest(ReqBacktest) returns (RspBacktest); // 交易计划回测 } message IndicatorSubReq { @@ -46,3 +47,12 @@ message RspStrategySeries { repeated bool wins = 3; // 下一根k线价格方向是否正确 double winRate = 4; } + +message ReqBacktest { + int64 planId = 1; + string stime = 2; + string etime = 3; +} +message RspBacktest{ + +} diff --git a/cmd/market/main.go b/cmd/market/main.go index f88dc58..76068bc 100644 --- a/cmd/market/main.go +++ b/cmd/market/main.go @@ -8,7 +8,7 @@ import ( "sig-pub/pkg/config" "sig-pub/pkg/grpc/discovery" "sig-pub/pkg/grpc/interceptor" - "sig-pub/pkg/storage/rdb" + "sig-pub/pkg/storage/persist" "sig-pub/pkg/utils/exit" "sig-pub/pkg/zlog" @@ -41,7 +41,7 @@ func main() { if err != nil { panic(err) } - rdb := rdb.NewRDB(db) + rdb := persist.NewRDB(db) if err := rdb.Init(); err != nil { panic(err) } diff --git a/cmd/trading/main.go b/cmd/trading/main.go index 82344b3..a98122b 100644 --- a/cmd/trading/main.go +++ b/cmd/trading/main.go @@ -10,6 +10,7 @@ import ( "sig-pub/pkg/grpc/discovery" "sig-pub/pkg/grpc/interceptor" "sig-pub/pkg/mq" + "sig-pub/pkg/storage/persist" "sig-pub/pkg/utils/exit" "sig-pub/pkg/zlog" @@ -45,6 +46,17 @@ func main() { panic(err) } + // database + db, err := conf.Database.Postgres.NewGormDB() + if err != nil { + panic(err) + } + rdb := persist.NewRDB(db) + if err := rdb.Init(); err != nil { + panic(err) + } + tradingDataPersist := trading.NewTradingDataPersist(rdb) + // new market grpc client marketClient, err := client.NewMarketClient( grpc.WithResolvers(resolver), @@ -63,7 +75,7 @@ func main() { } // services - tradingService := trading.NewTradingService(tradeInstanceAside, exchangeClient) + tradingService := trading.NewTradingService(tradeInstanceAside, exchangeClient, tradingDataPersist) if err := tradingService.Init(); err != nil { panic(err) } diff --git a/config/exchange.toml b/config/exchange.toml index c0a42b7..4d567a5 100644 --- a/config/exchange.toml +++ b/config/exchange.toml @@ -18,8 +18,8 @@ receiveBuffer = 4096 marketSubscribeLimit = 16 consumeBatch = 1024 consumeLater = 2000 # 时间到达later或者数据累计到batch触发consume -httpProxy = "http://192.168.1.5:7890" -# httpProxy = "http://10.255.183.209:7890" +# httpProxy = "http://192.168.1.5:7890" +httpProxy = "http://10.255.183.209:7890" # 模拟盘API交易地址如下: # REST:https://www.okx.com diff --git a/internal/exchange/exchange_service.go b/internal/exchange/exchange_service.go index a95bcee..0047b2d 100644 --- a/internal/exchange/exchange_service.go +++ b/internal/exchange/exchange_service.go @@ -745,7 +745,7 @@ func (svc *ExchangeService) CalcSeriesRange(arg *pb.SeriesRange) (after, before, } // HistoryKline 获取交易产品历史k线 (before < klines... < after) -func (svc *ExchangeService) HistoryKline(arg *pb.SeriesRange, recvBranch int, recvKline func(klines []*pb.Kline) error) (live bool, err error) { +func (svc *ExchangeService) HistoryKline(arg *pb.SeriesRange, recvBranch int, recvKlineFn func(klines []*pb.Kline) error) (live bool, err error) { // 交易产品参数检查 exchange := svc.exchanges.Get(arg.Exchange) exchangeInstId, ok := exchange.TradeInstIds.Load(arg.InstId) @@ -787,7 +787,7 @@ func (svc *ExchangeService) HistoryKline(arg *pb.SeriesRange, recvBranch int, re // 分批查询 branch := int64(2000) before, after := beforeTs, afterTs - recvBuffer := make([]*pb.Kline, 0, recvBranch) + recvBuffer := make([]*pb.Kline, 0, min(recvBranch, int(branch))) var prevFirstK, prevLastK *types.Kline for range 10000 { if arg.Desc { @@ -861,7 +861,7 @@ func (svc *ExchangeService) HistoryKline(arg *pb.SeriesRange, recvBranch int, re if len(recvBuffer) < recvBranch && i < length-1 { continue } - if err = recvKline(recvBuffer); err != nil { + if err = recvKlineFn(recvBuffer); err != nil { break } recvBuffer = recvBuffer[:0] diff --git a/internal/market/trade_instance_service.go b/internal/market/trade_instance_service.go index 28a11aa..ae10429 100644 --- a/internal/market/trade_instance_service.go +++ b/internal/market/trade_instance_service.go @@ -6,17 +6,17 @@ import ( "sig-pub/pkg/data" "sig-pub/pkg/data/args" "sig-pub/pkg/data/entity" - "sig-pub/pkg/storage/rdb" + "sig-pub/pkg/storage/persist" "time" ) // TradeInstanceService 交易产品管理 // TODO cache type TradeInstanceService struct { - db *rdb.RDB + db *persist.RDB } -func NewTradeInstanceService(db *rdb.RDB) *TradeInstanceService { +func NewTradeInstanceService(db *persist.RDB) *TradeInstanceService { return &TradeInstanceService{ db: db, } diff --git a/internal/trading/backtest/backtest.go b/internal/trading/backtest/backtest.go index 150a332..337081c 100644 --- a/internal/trading/backtest/backtest.go +++ b/internal/trading/backtest/backtest.go @@ -2,12 +2,15 @@ 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/types/decimals" + "sig-pub/pkg/utils/times" "sig-pub/pkg/zlog" "google.golang.org/grpc" @@ -16,18 +19,91 @@ import ( type Backtest struct { exchangeClient pb.ExchangeServiceClient indReg *indicator.IndicatorRegistry - sim *Simulator + sigStrategyReg *strategy.SigStrategyRegistry } -func NewBacktest(exchangeClient pb.ExchangeServiceClient, indReg *indicator.IndicatorRegistry, sim *Simulator) *Backtest { - return &Backtest{exchangeClient: exchangeClient, indReg: indReg, sim: sim} +func NewBacktest(exchangeClient pb.ExchangeServiceClient, indReg *indicator.IndicatorRegistry, sigStrategyReg *strategy.SigStrategyRegistry) *Backtest { + return &Backtest{exchangeClient: exchangeClient, indReg: indReg} +} + +func (b *Backtest) RunTradingPlan(ctx context.Context, tradingPlan *sig.TradingPlan, stime, etime int64, sigKlineSeries *sig.KlineSeries) (err error) { + var sim *Simulator + _ = sim + plan := tradingPlan.Plan + exchange := pb.ExchangeType(plan.Exchange) + interval := types.Interval(plan.Interval) + instId := plan.InstId + + sigStrategy := tradingPlan.GetSigStrategy() + maxWindow := sigStrategy.MaxWindow() + if maxWindow < 0 || maxWindow > indicator.MaxWindow { + err = fmt.Errorf("invalid window %d 0-%d, planId=%d", maxWindow, indicator.MaxWindow, plan.Id) + return + } + + seriesRange := &pb.SeriesRange{ + Exchange: exchange, + InstId: instId, + Interval: string(interval), + Before: stime, + After: etime, + Open: false, + Live: false, + Desc: false, + Window: uint32(maxWindow), + } + // fetch history klines via stream + req := &pb.ReqHistoryKlineStream{Series: seriesRange} + stream, err := b.exchangeClient.HistoryKlineStream(context.Background(), req, grpc.UseCompressor("snappy")) + if err != nil { + return + } + + recvTimes, total := 0, 0 + watch := times.NewWatch() + for { + msg, err0 := stream.Recv() + if err0 == io.EOF { + break + } + if err0 != nil { + err = err0 + return + } + recvTimes++ + total += len(msg.Klines) + for _, k := range msg.Klines { + kline := new(types.Kline) + kline.ParsePBKline(seriesRange.Exchange, k) + + if lastTs, serial := sigKlineSeries.Update(kline); !serial { + err = fmt.Errorf("kline not series: %s(%s), interval=%s, lastTs=%d", instId, exchange, interval, lastTs) + return + } + length := sigKlineSeries.Length() + if length <= maxWindow { + continue + } + + signalSide := tradingPlan.Update(strategy.StrategyTypeSig) + switch signalSide { + case types.SideBuy, types.SideSell: + k, _ := sigKlineSeries.Get(0) + _ = k + // zlog.Infof("sigSide: ts=%d, %s", k.Ts, sigSide) + } + } + } + + zlog.Debugf("recv=%d, total=%d, use %s", recvTimes, total, watch.ElapsedFmt(".")) + return } // Run 执行回测 // seriesRange: 回测的交易产品/周期/时间区间 // sigStrategy: 已创建的策略实例(将调用 New() 并 Init) // params: 策略参数 -func (b *Backtest) Run(ctx context.Context, seriesRange *pb.SeriesRange, sigStrategy strategy.ISigStrategy, params strategy.SigStrategyParam, initialCash float64) (res BacktestResult, err error) { +func (b *Backtest) Run0(ctx context.Context, seriesRange *pb.SeriesRange, sigStrategy strategy.ISigStrategy, params strategy.StrategyParam, initialCash float64) (res BacktestResult, err error) { // prepare strategy strat := sigStrategy.New() if err = strat.Init(params); err != nil { @@ -64,7 +140,7 @@ func (b *Backtest) Run(ctx context.Context, seriesRange *pb.SeriesRange, sigStra } // prepare account and set risk limits from params if provided - acct := NewAccount(initialCash, b.sim) + acct := NewAccount(initialCash, nil) if v, ok := params.GetFloat64("max_pos_pct"); ok && v > 0 { acct.MaxPosPct = v } @@ -99,12 +175,12 @@ func (b *Backtest) Run(ctx context.Context, seriesRange *pb.SeriesRange, sigStra qty = 1 } - if side == pb.Side_BUY || side == pb.Side_SELL { + if side == types.SideBuy || side == types.SideSell { // signal-based close: close opposite positions first if cm != nil { - cm.CloseBySignal(side, acct, klines[i]) + cm.CloseBySignal(pb.Side_BUY, acct, klines[i]) } - if _, ok := acct.ApplyMarketOrder(side, qty, klines[i].Ts, klines[i]); ok { + if _, ok := acct.ApplyMarketOrder(pb.Side_BUY, qty, klines[i].Ts, klines[i]); ok { // trade recorded } else { zlog.Debugf("order rejected or insufficient cash at ts=%d", klines[i].Ts) diff --git a/internal/trading/backtest/risk_strategy.go b/internal/trading/backtest/risk_strategy.go new file mode 100644 index 0000000..93107aa --- /dev/null +++ b/internal/trading/backtest/risk_strategy.go @@ -0,0 +1,22 @@ +package backtest + +import ( + "sig-pub/api/pb" + "sig-pub/pkg/indicator" +) + +// RickStrategy 风险管理策略 +type RickStrategy struct { + indicatorCtx indicator.IIndicatorContext +} + +// OnSingle 收到信号时进行评估, 返回过滤后的交易信号 +// todo 对交易方向进行信心分数评估, 后续开仓仓位 +func (s *RickStrategy) OnSingle(signalSide pb.Side) (side pb.Side) { + k := s.indicatorCtx.Get(0) + price := k.Close + ts := k.Ts + _, _ = price, ts + + return +} diff --git a/internal/trading/sig/kline_series.go b/internal/trading/sig/kline_series.go index cb06a00..7bbe94b 100644 --- a/internal/trading/sig/kline_series.go +++ b/internal/trading/sig/kline_series.go @@ -34,7 +34,7 @@ func NewTradeInstanceKlineSeries(exchange pb.ExchangeType, instId string) *Trade } type KlineSeries struct { - sync.RWMutex + mu sync.RWMutex Exchange pb.ExchangeType InstId string Interval types.Interval @@ -62,8 +62,8 @@ func (s *KlineSeries) Get(offset int16) (k types.Kline, ok bool) { if ok = offset >= 0 && offset < MaxSeriesKlines; !ok { return } - s.RLock() - defer s.RUnlock() + s.mu.RLock() + defer s.mu.RUnlock() length := len(s.klines) index := (length - 1) - int(offset) @@ -84,8 +84,8 @@ func (s *KlineSeries) Series(offset, count int16) (klines series.Klines, ok bool return } - s.RLock() - defer s.RUnlock() + s.mu.RLock() + defer s.mu.RUnlock() length := len(s.klines) indexEnd := (length - 1) - int(offset) @@ -102,14 +102,20 @@ func (s *KlineSeries) Series(offset, count int16) (klines series.Klines, ok bool return klines, true } +func (s *KlineSeries) Length() int { + s.mu.RLock() + defer s.mu.RUnlock() + return len(s.klines) +} + func (s *KlineSeries) LastTs() int64 { return s.lastTs } // 检查k线序列完整 func (s *KlineSeries) Update(kline *types.Kline) (lastTs int64, serial bool) { - s.Lock() - defer s.Unlock() + s.mu.Lock() + defer s.mu.Unlock() serial = true lastTs = s.lastTs diff --git a/internal/trading/sig/trading_plan_runner.go b/internal/trading/sig/trading_plan.go similarity index 67% rename from internal/trading/sig/trading_plan_runner.go rename to internal/trading/sig/trading_plan.go index 86e1c27..ded1c78 100644 --- a/internal/trading/sig/trading_plan_runner.go +++ b/internal/trading/sig/trading_plan.go @@ -4,23 +4,24 @@ import ( "sig-pub/pkg/data/entity" "sig-pub/pkg/indicator" "sig-pub/pkg/strategy" + "sig-pub/pkg/types" "sync/atomic" ) type TradingPlan struct { Status atomic.Int32 - plan entity.TradePlan + Plan entity.TradePlan indicatorReg *indicator.IndicatorRegistry sigStrategy strategy.ISigStrategy - sigStrategyContext *StrategyContext + sigStrategyContext strategy.ISigStrategyContext // publisher publish.Publisher[int32, any] // signalKey map[string]int32 } func NewTradingPlan(plan entity.TradePlan, indicatorReg *indicator.IndicatorRegistry) *TradingPlan { return &TradingPlan{ - plan: plan, + Plan: plan, indicatorReg: indicatorReg, } } @@ -42,19 +43,25 @@ func (r *TradingPlan) Init() (err error) { // initSigStrategy 初始化多空信号策略 // buy/sell -> 过滤/风控 -> tradeStrategy -> closeStrategy -func (r *TradingPlan) InitSigStrategy(sigStrategy strategy.ISigStrategy, params strategy.SigStrategyParam, sigIndCtx IOffsetIndicatorContext) (err error) { - if err = sigStrategy.Init(params); err != nil { +func (r *TradingPlan) InitSigStrategy(sigStrategy strategy.ISigStrategy, param strategy.StrategyParam, sigStrategyContext strategy.ISigStrategyContext) (err error) { + if err = sigStrategy.Init(param); err != nil { return } r.sigStrategy = sigStrategy - r.sigStrategyContext = NewStrategyContext(sigIndCtx, r.indicatorReg) + r.sigStrategyContext = sigStrategyContext return } +func (r *TradingPlan) GetSigStrategy() strategy.ISigStrategy { + return r.sigStrategy +} + // Update 订阅k线更新 -func (r *TradingPlan) Update(signalType strategy.StrategyType) { +// todo filted sigle +func (r *TradingPlan) Update(signalType strategy.StrategyType) (side types.Side) { switch signalType { case strategy.StrategyTypeSig: - r.sigStrategy.Update(r.sigStrategyContext) + return r.sigStrategy.Update(r.sigStrategyContext) } + return } diff --git a/internal/trading/trading_data_persist.go b/internal/trading/trading_data_persist.go new file mode 100644 index 0000000..9e17f7b --- /dev/null +++ b/internal/trading/trading_data_persist.go @@ -0,0 +1,34 @@ +package trading + +import ( + "sig-pub/pkg/data" + "sig-pub/pkg/data/entity" + "sig-pub/pkg/storage/persist" +) + +type TradingDataPersist struct { + db *persist.RDB +} + +func NewTradingDataPersist(db *persist.RDB) *TradingDataPersist { + return &TradingDataPersist{ + db: db, + } +} + +func (p *TradingDataPersist) Init() (err error) { + return +} + +func (p *TradingDataPersist) GetTradePlanById(planId int64) (plan *entity.TradePlan, err error) { + plan = new(entity.TradePlan) + err = p.db.Select(plan, `select * from t_trade_plan where id = ?`, planId) + if err != nil { + return + } + if plan.InstId == "" { + err = data.ErrorNotExists + return + } + return +} diff --git a/internal/trading/trading_grpc_server.go b/internal/trading/trading_grpc_server.go index e3edc37..1ea4e17 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/utils/times" ) type TradingGrpcServer struct { @@ -36,3 +37,20 @@ func (svr *TradingGrpcServer) StrategySeries(ctx context.Context, req *pb.ReqStr err = svr.tradingService.StrategySeries(req, rsp) return } + +func (svr *TradingGrpcServer) Backtest(ctx context.Context, req *pb.ReqBacktest) (rsp *pb.RspBacktest, err error) { + stime, err := times.ParseFORMAT(req.Stime) + if err != nil { + return + } + etime, err := times.ParseFORMAT(req.Etime) + if err != nil { + return + } + err = svr.tradingService.Backtest(req.PlanId, stime.UnixMilli(), etime.UnixMilli()) + if err != nil { + return + } + rsp = new(pb.RspBacktest) + return +} diff --git a/internal/trading/trading_service.go b/internal/trading/trading_service.go index d745a4f..f046c6f 100644 --- a/internal/trading/trading_service.go +++ b/internal/trading/trading_service.go @@ -1,6 +1,7 @@ package trading import ( + "context" "fmt" "sig-pub/api/pb" "sig-pub/pkg/client" @@ -11,16 +12,19 @@ import ( "sig-pub/pkg/strategy" "sig-pub/pkg/types" "sig-pub/pkg/utils/collect" + "sig-pub/pkg/utils/lang" "sig-pub/pkg/zlog" + "sig-pub/internal/trading/backtest" "sig-pub/internal/trading/sig" "github.com/bytedance/sonic" ) type TradingService struct { - marketClientAside *client.TradeInstanceAside - exchangeClient pb.ExchangeServiceClient + marketClientAside *client.TradeInstanceAside + exchangeClient pb.ExchangeServiceClient + tradingDataPersist *TradingDataPersist klineSeriesStore *KlineSeriesStore indicatorReg *indicator.IndicatorRegistry // 注册窗口指标 @@ -32,15 +36,17 @@ type TradingService struct { func NewTradingService( marketClientAside *client.TradeInstanceAside, exchangeClient pb.ExchangeServiceClient, + tradingDataPersist *TradingDataPersist, ) *TradingService { return &TradingService{ - marketClientAside: marketClientAside, - exchangeClient: exchangeClient, - klineSeriesStore: NewKlineSeriesStore(exchangeClient), - indicatorReg: indicator.NewIndicatorRegistry(), - strategyReg: strategy.NewSigStrategyRegistry(), - signalPublisher: publish.NewPublisher[int64, strategy.StrategyType](16), - tradingPlans: collect.NewSyncMap[int64, *sig.TradingPlan](), + marketClientAside: marketClientAside, + exchangeClient: exchangeClient, + tradingDataPersist: tradingDataPersist, + klineSeriesStore: NewKlineSeriesStore(exchangeClient), + indicatorReg: indicator.NewIndicatorRegistry(), + strategyReg: strategy.NewSigStrategyRegistry(), + signalPublisher: publish.NewPublisher[int64, strategy.StrategyType](16), + tradingPlans: collect.NewSyncMap[int64, *sig.TradingPlan](), } } @@ -84,47 +90,24 @@ func (svc *TradingService) consumerKlineSignal() { } } -// runTradingPlan 运行交易计划 +// getTradingPlan 获取交易计划 // todo 止盈止损策略, 下单策略... -func (svc *TradingService) runTradingPlan(plan *entity.TradePlan) (err error) { - var planId = plan.Id - var instId = plan.InstId - var exchange pb.ExchangeType - var sigInterval types.Interval - - exchange = pb.ExchangeType(plan.Exchange) +func (svc *TradingService) getTradingPlan(plan *entity.TradePlan, sigKlineSeries *sig.KlineSeries) (tradingPlan *sig.TradingPlan, err error) { + exchange := pb.ExchangeType(plan.Exchange) + interval := types.Interval(plan.Interval) + instId := plan.InstId if !types.IsSupportExchange(exchange) { err = fmt.Errorf("unsupport exchange %d", plan.Exchange) return } - - tradingPlan := sig.NewTradingPlan(*plan, svc.indicatorReg) - if _, load := svc.tradingPlans.LoadOrStore(planId, tradingPlan); load { - err = fmt.Errorf("plan already running: planId=%d", planId) + if _, ok := types.SupportedIntervals[interval]; !ok { + err = fmt.Errorf("unsupport interval %d", plan.Exchange) return } - tradingPlan.Status.Store(int32(data.StatusProcessing)) - - defer func() { - if err != nil { - svc.tradingPlans.Delete(planId) - } else { - // 订阅交易信号策略k线周期 - sigSubKey := strategy.DriverIntervalKey(instId, exchange, false, sigInterval) - svc.signalPublisher.Subscribe(sigSubKey, planId, strategy.StrategyTypeSig) - - tradingPlan.Status.Store(int32(data.StatusOk)) - } - }() // sigStrategy - sigStrategyParam := new(strategy.SigStrategyParam) - if err = sonic.UnmarshalString(plan.SigStrategyParam, sigStrategyParam); err != nil { - return - } - sigInterval = types.Interval(sigStrategyParam.Interval) - if _, ok := types.SupportedIntervals[sigInterval]; !ok { - err = fmt.Errorf("unsupport interval %d", plan.Exchange) + sigStrategyParam := make(strategy.StrategyParam) + if err = sonic.UnmarshalString(plan.SigStrategyParam, &sigStrategyParam); err != nil { return } sigStrategy, ok := svc.strategyReg.NewSigStrategy(plan.SigStrategy) @@ -132,19 +115,42 @@ func (svc *TradingService) runTradingPlan(plan *entity.TradePlan) (err error) { err = fmt.Errorf("strategy %s not exists", plan.SigStrategy) return } - sigKlineSeries, err := svc.klineSeriesStore.GetKlineSeires(exchange, instId, sigInterval) - if err != nil { - return + + if sigKlineSeries == nil { + sigKlineSeries, err = svc.klineSeriesStore.GetKlineSeires(exchange, instId, interval) + if err != nil { + return + } } + tradingPlan = sig.NewTradingPlan(*plan, svc.indicatorReg) if err = tradingPlan.Init(); err != nil { return } - sigIndCtx := sig.NewIndicatorContext(sigKlineSeries) - if err = tradingPlan.InitSigStrategy(sigStrategy, *sigStrategyParam, sigIndCtx); err != nil { + sigIndicatorCtx := sig.NewIndicatorContext(sigKlineSeries) + sigStrategyCtx := sig.NewStrategyContext(sigIndicatorCtx, svc.indicatorReg) + if err = tradingPlan.InitSigStrategy(sigStrategy, sigStrategyParam, sigStrategyCtx); err != nil { return } return + + // if _, load := svc.tradingPlans.LoadOrStore(planId, tradingPlan); load { + // err = fmt.Errorf("plan already running: planId=%d", planId) + // return + // } + // tradingPlan.Status.Store(int32(data.StatusProcessing)) + + // defer func() { + // if err != nil { + // svc.tradingPlans.Delete(planId) + // } else { + // // 订阅交易信号策略k线周期 + // sigSubKey := strategy.DriverIntervalKey(instId, exchange, false, sigInterval) + // svc.signalPublisher.Subscribe(sigSubKey, planId, strategy.StrategyTypeSig) + + // tradingPlan.Status.Store(int32(data.StatusOk)) + // } + // }() } // IndicatorSeries 获取指标实时或历史序列数据, 闭区间 @@ -195,10 +201,6 @@ func (svc *TradingService) IndicatorSeries(indicatorName string, window uint32, return } -const ( - MaxIndicatorWindow = 128 -) - // StrategySeries 简单策略信号测试 func (svc *TradingService) StrategySeries(req *pb.ReqStrategySeries, rsp *pb.RspStrategySeries) (err error) { // sigStrategy @@ -207,11 +209,8 @@ func (svc *TradingService) StrategySeries(req *pb.ReqStrategySeries, rsp *pb.Rsp err = fmt.Errorf("strategy %s not exists", req.SigStrategy) return } - err = sigStrategy.Init(strategy.SigStrategyParam{ - Interval: types.Interval(req.Series.Interval), - Param: req.SigParam, - }) - if err != nil { + // init sigStrategy + if err = sigStrategy.Init(req.SigParam); err != nil { return } @@ -230,17 +229,18 @@ func (svc *TradingService) StrategySeries(req *pb.ReqStrategySeries, rsp *pb.Rsp // recover todo out of range count, totalK := 0, 0 indicatorContext := sig.NewHistoryIndicatorContext(svc.exchangeClient) - req.Series.Window += MaxIndicatorWindow + req.Series.Window += indicator.MaxWindow if totalK, err = indicatorContext.Init(req.Series); err != nil { return } - count = totalK - MaxIndicatorWindow + count = totalK - indicator.MaxWindow strategyContext := sig.NewStrategyContext(indicatorContext, svc.indicatorReg) for i := count - 1; i >= 0; i-- { strategyContext.SetOffset(int16(i)) - side := sigStrategy.Update(strategyContext) - if side == pb.Side_BUY || side == pb.Side_SELL { + sigSide := sigStrategy.Update(strategyContext) + if sigSide == types.SideBuy || sigSide == types.SideSell { + side := lang.Ternary(sigSide == types.SideBuy, pb.Side_BUY, pb.Side_SELL) signalK := strategyContext.Get(0) rsp.Signal = append(rsp.Signal, side) rsp.Times = append(rsp.Times, signalK.Ts) @@ -264,3 +264,26 @@ func (svc *TradingService) StrategySeries(req *pb.ReqStrategySeries, rsp *pb.Rsp rsp.WinRate = float64(len(wins)) / float64(len(rsp.Wins)) return } + +// Backtest 回测交易计划 +func (svc *TradingService) Backtest(planId, stime, etime int64) (err error) { + plan, err := svc.tradingDataPersist.GetTradePlanById(planId) + if err != nil { + return + } + exchange := pb.ExchangeType(plan.Exchange) + interval := types.Interval(plan.Interval) + + sigKlineSeries := sig.NewKlineSeries(exchange, plan.InstId, interval) + tradingPlan, err := svc.getTradingPlan(plan, sigKlineSeries) + if err != nil { + return + } + + test := backtest.NewBacktest(svc.exchangeClient, svc.indicatorReg, svc.strategyReg) + err = test.RunTradingPlan(context.Background(), tradingPlan, stime, etime, sigKlineSeries) + if err != nil { + return + } + return +} diff --git a/pkg/indicator/base.go b/pkg/indicator/base.go index ef50a3e..6172c67 100644 --- a/pkg/indicator/base.go +++ b/pkg/indicator/base.go @@ -5,6 +5,10 @@ import ( "sig-pub/pkg/types/series" ) +const ( + MaxWindow = 128 +) + // IIndicator 指标基础计算接口 type IIndicator interface { Name() string diff --git a/pkg/storage/rdb/rdb.go b/pkg/storage/persist/rdb.go similarity index 98% rename from pkg/storage/rdb/rdb.go rename to pkg/storage/persist/rdb.go index 49b89f7..22ac7ad 100644 --- a/pkg/storage/rdb/rdb.go +++ b/pkg/storage/persist/rdb.go @@ -1,4 +1,4 @@ -package rdb +package persist import ( "gorm.io/gorm" diff --git a/pkg/strategy/gold_x.go b/pkg/strategy/gold_x.go index 59ed395..07e711c 100644 --- a/pkg/strategy/gold_x.go +++ b/pkg/strategy/gold_x.go @@ -2,7 +2,6 @@ package strategy import ( "fmt" - "sig-pub/api/pb" "sig-pub/pkg/types" ) @@ -28,7 +27,7 @@ func (s *GoldX) Meta() StrategyMeta { } } -func (s *GoldX) Init(param SigStrategyParam) (err error) { // 校验参数, 并根据参数初始化策略 +func (s *GoldX) Init(param StrategyParam) (err error) { // 校验参数, 并根据参数初始化策略 if s.short, err = param.GetInt16E("short"); err != nil { return } @@ -42,7 +41,11 @@ func (s *GoldX) Init(param SigStrategyParam) (err error) { // 校验参数, 并 return } -func (s *GoldX) Update(ctx ISigStrategyContext) (side pb.Side) { +func (s *GoldX) MaxWindow() int { + return int(max(s.long, s.short)) +} + +func (s *GoldX) Update(ctx ISigStrategyContext) (side types.Side) { sma14 := ctx.IndicatorW("sma", s.short) sma28 := ctx.IndicatorW("sma", s.long) // 包装方法 crossover/crossunder @@ -51,15 +54,15 @@ func (s *GoldX) Update(ctx ISigStrategyContext) (side pb.Side) { crossover := s14[0] > s28[0] && s14[1] < s28[1] // 上穿 crossunder := s14[0] < s28[0] && s14[1] > s28[1] // 下穿 if crossover { - return pb.Side_BUY + return types.SideBuy } if crossunder { - return pb.Side_SELL + return types.SideSell } return } -func (s *GoldX) UpdateByIntervals(ctx IIntervalStrategyContext) (side pb.Side) { +func (s *GoldX) UpdateByIntervals(ctx IIntervalStrategyContext) (side types.Side) { series5m := ctx.Series(types.Interval5m, 0, 2) series5m.Close().Diff() return diff --git a/pkg/strategy/sig_strategy.go b/pkg/strategy/sig_strategy.go index cc85922..6c7f0ff 100644 --- a/pkg/strategy/sig_strategy.go +++ b/pkg/strategy/sig_strategy.go @@ -1,7 +1,6 @@ package strategy import ( - "sig-pub/api/pb" "sig-pub/pkg/indicator" "sig-pub/pkg/types" "sig-pub/pkg/types/series" @@ -11,15 +10,15 @@ import ( type ISigStrategy interface { New() ISigStrategy Meta() StrategyMeta - Init(param SigStrategyParam) (err error) // 校验参数, 并根据参数初始化策略 - Update(ctx ISigStrategyContext) (side pb.Side) + MaxWindow() int // 需要的最大数据窗口数, 回测时用, 若不定义则取最大窗口值 + Init(param StrategyParam) (err error) // 校验参数, 并根据参数初始化策略 + Update(ctx ISigStrategyContext) (side types.Side) } type StrategyMeta struct { - Name string `json:"name"` - Desc string `json:"desc"` - MaxWindow int `json:"maxWindow"` // 需要的最大窗口数, 回测时用, 若不定义则取最大窗口值 - Args []Param `json:"args"` // 参数定义 + Name string `json:"name"` + Desc string `json:"desc"` + Args []Param `json:"args"` // 参数定义 } // ISigStrategyContext 策略外部访问能力 diff --git a/pkg/strategy/sig_strategy_intervals.go b/pkg/strategy/sig_strategy_intervals.go index f97189b..1201700 100644 --- a/pkg/strategy/sig_strategy_intervals.go +++ b/pkg/strategy/sig_strategy_intervals.go @@ -1,7 +1,6 @@ package strategy import ( - "sig-pub/api/pb" "sig-pub/pkg/indicator" "sig-pub/pkg/types" "sig-pub/pkg/types/series" @@ -9,7 +8,7 @@ import ( // 多周期k线策略接口 type IIntervalSigStrategy interface { - UpdateByInterval(ctx IIntervalStrategyContext) (side pb.Side) + UpdateByInterval(ctx IIntervalStrategyContext) (side types.Side) } // IIntervalStrategyContext 多周期策略上下文 diff --git a/pkg/strategy/sig_strategy_params.go b/pkg/strategy/sig_strategy_params.go index 817c715..fa7361c 100644 --- a/pkg/strategy/sig_strategy_params.go +++ b/pkg/strategy/sig_strategy_params.go @@ -79,20 +79,22 @@ type ISigStrategyParamGenerator interface { NextParam(map[string]string) (map[string]string, bool) // 根据当前策略参数, 返回下一批策略参数(并行回测 stateless) } -type SigStrategyParam struct { - Interval types.Interval `json:"interval"` // 策略驱动周期 - Param map[string]string `json:"param"` // 策略执行参数 +type IntervalStrategyParam struct { + Interval types.Interval `json:"interval"` // 策略驱动周期 + Param StrategyParam `json:"param"` // 策略执行参数 } -func (s *SigStrategyParam) Get(key string) (v string, ok bool) { - if len(s.Param) == 0 { +type StrategyParam map[string]string + +func (s *StrategyParam) Get(key string) (v string, ok bool) { + if len(*s) == 0 { return } - v, ok = s.Param[key] + v, ok = (*s)[key] return } -func (s *SigStrategyParam) GetInt(key string) (r int, ok bool) { +func (s *StrategyParam) GetInt(key string) (r int, ok bool) { v, ok := s.Get(key) if !ok { return @@ -104,7 +106,7 @@ func (s *SigStrategyParam) GetInt(key string) (r int, ok bool) { return } -func (s *SigStrategyParam) GetInt16E(key string) (r int16, err error) { +func (s *StrategyParam) GetInt16E(key string) (r int16, err error) { v, ok := s.Get(key) if !ok { err = fmt.Errorf("param %s not provided", key) @@ -114,7 +116,7 @@ func (s *SigStrategyParam) GetInt16E(key string) (r int16, err error) { return } -func (s *SigStrategyParam) GetFloat64(key string) (r float64, ok bool) { +func (s *StrategyParam) GetFloat64(key string) (r float64, ok bool) { v, ok := s.Get(key) if !ok { return @@ -126,7 +128,7 @@ func (s *SigStrategyParam) GetFloat64(key string) (r float64, ok bool) { return } -func (s *SigStrategyParam) GetBool(key string) (r bool, ok bool) { +func (s *StrategyParam) GetBool(key string) (r bool, ok bool) { v, ok := s.Get(key) if !ok { return diff --git a/pkg/types/signal.go b/pkg/types/signal.go new file mode 100644 index 0000000..59f77bd --- /dev/null +++ b/pkg/types/signal.go @@ -0,0 +1,19 @@ +package types + +type Side int32 + +const ( + SideBuy Side = 1 // BUY + SideSell Side = 2 // SELL +) + +func (side Side) String() string { + switch side { + default: + return "" + case SideBuy: + return "BUY" + case SideSell: + return "SELL" + } +} diff --git a/pkg/utils/times/times.go b/pkg/utils/times/times.go index 9f291a1..4eaa4bd 100644 --- a/pkg/utils/times/times.go +++ b/pkg/utils/times/times.go @@ -1,5 +1,7 @@ package times +import "time" + const FORMAT string = "2006-01-02 15:04:05" const FORMAT2 string = "2006/01/02 15:04:05" const FORMAT_DATE string = "2006-01-02" @@ -7,3 +9,9 @@ const FORMAT_DATE2 string = "2006/01/02" const FORMAT_MONTH string = "2006-01" const FORMAT_TIME string = "15:04:05" const FORMAT_TIME_Minute string = "15:04" + +// ParseFORMAT 按FORMAT格式解析时间 +// todo in location +func ParseFORMAT(datetime string) (time.Time, error) { + return time.Parse(FORMAT, datetime) +}