package trading import ( "context" "fmt" "io" "sig-pub/api/pb" "sig-pub/pkg/client" "sig-pub/pkg/data" "sig-pub/pkg/data/entity" "sig-pub/pkg/indicator" "sig-pub/pkg/publish" "sig-pub/pkg/strategy" "sig-pub/pkg/types" "sig-pub/pkg/utils/collect" "sig-pub/pkg/utils/lang" "sig-pub/pkg/utils/times" "sig-pub/pkg/zlog" "sig-pub/internal/trading/backtest" "sig-pub/internal/trading/sig" "github.com/bytedance/sonic" "google.golang.org/grpc" ) type TradingService struct { marketClientAside *client.TradeInstanceAside exchangeClient pb.ExchangeServiceClient tradingDataPersist *TradingDataPersist klineSeriesStore *KlineSeriesStore indicatorReg *indicator.IndicatorRegistry // 注册窗口指标 strategyReg *strategy.SigStrategyRegistry // 注册信号策略 signalPublisher *publish.Publisher[int64, strategy.StrategyType] // planId -> strategyType tradingPlans *collect.SyncMap[int64, *sig.TradingPlan] // 运行中交易计划 } func NewTradingService( marketClientAside *client.TradeInstanceAside, exchangeClient pb.ExchangeServiceClient, tradingDataPersist *TradingDataPersist, ) *TradingService { return &TradingService{ 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](), } } // 初始化历史k线, 订阅实时k线 func (svc *TradingService) Init() (err error) { if err = svc.indicatorReg.Init(); err != nil { return } if err = svc.strategyReg.Init(); err != nil { return } if err = svc.klineSeriesStore.Init(); err != nil { return } go svc.consumerKlineSignal() // todo loading trading plan return } // consumerKlineSignal 订阅k线更新 func (svc *TradingService) consumerKlineSignal() { c := svc.klineSeriesStore.ConsumerKlineSignel() for { signalKey := <-c zlog.Debugf("signal: %s", signalKey) planIds, strategyTypes := svc.signalPublisher.Publisher(signalKey) for i, strategyType := range strategyTypes { planId := planIds[i] plan, ok := svc.tradingPlans.Load(planId) if !ok { zlog.Warningf("plan not running: id=%d", planId) continue } if plan.Status.Load() == int32(data.StatusOk) { plan.Update(strategyType) } } } } // getTradingPlan 获取交易计划 // todo 止盈止损策略, 下单策略... 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 } if _, ok := types.SupportedIntervals[interval]; !ok { err = fmt.Errorf("unsupport interval %d", plan.Exchange) return } // sigStrategy sigStrategyParam := make(strategy.StrategyParam) if err = sonic.UnmarshalString(plan.SigStrategyParam, &sigStrategyParam); err != nil { return } sigStrategyType, sigStrategy, ok := svc.strategyReg.NewSigStrategy(plan.SigStrategy) if !ok { err = fmt.Errorf("strategy %s not exists", plan.SigStrategy) 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 } 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 { // 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)) // } // }() } // fetchHistoryKlineSeries 请求k线数据流式处理 func (svc *TradingService) fetchHistoryKlineSeries(ctx context.Context, sr *pb.SeriesRange, recvFn func(k *types.Kline) error) (err error) { // fetch history klines via stream req := &pb.ReqHistoryKlineStream{Series: sr} stream, err := svc.exchangeClient.HistoryKlineStream(ctx, req, grpc.UseCompressor("snappy")) if err != nil { return } var msg *pb.RspHistoryKlineStream recvTimes, recvTotal := 0, 0 watch := times.NewWatch() for { select { case <-ctx.Done(): err = ctx.Err() return default: } msg, err = stream.Recv() if err == io.EOF { err = nil break } if err != nil { return } recvTimes++ recvTotal += len(msg.Klines) for _, k := range msg.Klines { kline := new(types.Kline) kline.ParsePBKline(sr.Exchange, k) if err = recvFn(kline); err != nil { return } } } zlog.Debugf("fetch history kline series: inst=%s(%s), interval=%s, recv=%d, total=%d, use %s", sr.InstId, sr.Exchange, sr.Interval, recvTimes, recvTotal, watch.ElapsedFmt(".")) return } // 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 indicator, ok := svc.indicatorReg.IndicatorW(indicatorName) if !ok { err = fmt.Errorf("indicator %s not exists", indicatorName) return } interval := types.Interval(sr.Interval) _, ok = types.SupportedIntervals[interval] if !ok { err = fmt.Errorf("unsupport interval %s", interval) return } // 查询历史指标数据 requiredSeries := int(indicator.RequiredSeries(int16(window))) sr.WindowExtra = uint32(requiredSeries) kSeries := sig.NewKlineSeries(sr.Exchange, sr.InstId, types.Interval(sr.Interval)) indicatorContext := sig.NewIndicatorContext(kSeries) matrix = make([]float64, 0, 200) times = make([]int64, 0, 200) err = svc.fetchHistoryKlineSeries(ctx, sr, func(k *types.Kline) (err error) { if lastTs, serial := kSeries.Update(k); !serial { err = fmt.Errorf("kline not series: %s(%s), interval=%s, lastTs=%d", sr.InstId, sr.Exchange, interval, lastTs) return } if kSeries.Length() < requiredSeries { return } vector := indicator.Calculate(indicatorContext, int16(window)) matrix = append(matrix, vector) times = append(times, indicatorContext.Get(0).Ts) return }) if err != nil { return } return } // StrategySeries 简单策略信号测试 func (svc *TradingService) StrategySeries(ctx context.Context, req *pb.ReqStrategySeries, rsp *pb.RspStrategySeries) (err error) { // sigStrategy sigStrategyType, sigStrategy, ok := svc.strategyReg.NewSigStrategy(req.SigStrategy) if !ok { 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] if !ok { err = fmt.Errorf("unsupport interval %s", interval) return } switch sigStrategyType { case strategy.SigStrategyTypeSingle: rsp.Signal, rsp.Times, err = svc.singleStrategySeries(ctx, sigStrategy.(strategy.ISingleSigStrategy), req.Series) case strategy.SigStrategyTypeInterval: rsp.Signal, rsp.Times, err = svc.intervalStrategySeries(ctx, sigStrategy.(strategy.IIntervalSigStrategy), req.Series) } return } // intervalStrategySeries 多周期策略 func (svc *TradingService) intervalStrategySeries(ctx context.Context, intervalSigStrategy strategy.IIntervalSigStrategy, sr *pb.SeriesRange) (signals []pb.Side, times []int64, err error) { driverInterval := types.Interval(sr.Interval) driverIntervalAdder := types.SupportedIntervals[driverInterval] // 各周期所需k线数量 requiredIntervalSeries := intervalSigStrategy.RequiredIntervalSeries() // 驱动周期外其他周期 otherIntervals := make([]types.Interval, 0, 3) requiredIntervalSeries.Range(func(interval types.Interval, series int16) { if series > 0 { otherIntervals = append(otherIntervals, interval) } }) otherIntervals = collect.Filter(otherIntervals, func(_ int, interval types.Interval) bool { return interval != driverInterval }) // 通知其他周期更新的channel otherIntervalCh := types.NewIntervalState[[]chan int64]() // otherIntervalDstCh := types.NewIntervalState[chan int64]() for _, interval := range otherIntervals { otherIntervalCh.Set(interval, []chan int64{make(chan int64), make(chan int64)}) } // 各周期 series intervalKlineSeries := types.NewIntervalState[*sig.KlineSeries]() // 其他周期数据拉取 stopCh := make(chan struct{}) for _, interval := range otherIntervals { kSeries := sig.NewKlineSeries(sr.Exchange, sr.InstId, interval) intervalKlineSeries.Set(interval, kSeries) go func(interval types.Interval, kSeries *sig.KlineSeries) { ch := otherIntervalCh.Get(interval) srcCh := ch[0] dstCh := ch[1] intervalAdder := types.SupportedIntervals[interval] driverTs := int64(0) isr := &pb.SeriesRange{Exchange: sr.Exchange, InstId: sr.InstId, Before: sr.Before, After: sr.After, Count: sr.Count, Open: sr.Open, Live: sr.Live, Desc: sr.Desc, WindowExtra: sr.WindowExtra, Limit: sr.Limit, Interval: string(interval)} err1 := svc.fetchHistoryKlineSeries(ctx, isr, func(k *types.Kline) (err error) { closeTs := intervalAdder(k.Ts, 1) // 与驱动周期series保持同步更新 if closeTs > driverTs { if driverTs != 0 { dstCh <- 0 // 通知更新完毕 } waitLoop: for { select { case <-stopCh: return io.EOF case driverTs = <-srcCh: if closeTs <= driverTs { break waitLoop } else { dstCh <- 0 // 通知更新完毕 } } } } if lastTs, serial := kSeries.Update(k); !serial { err = fmt.Errorf("kline not series: %s(%s), interval=%s, lastTs=%d", sr.InstId, sr.Exchange, interval, lastTs) return } return }) if err1 == nil { otherIntervalCh.Set(interval, nil) } else if err1 != io.EOF { zlog.Errorf("fetch history interval error: inst=%s(%s) interval=%s, err=%v", sr.InstId, sr.Exchange, interval, err1) err = err1 close(stopCh) } }(interval, kSeries) } // 策略上下文 intervalStrategyContext := sig.NewIntervalStrategyContext(intervalKlineSeries, svc.indicatorReg) // 驱动周期数据拉取 driverSeries := sig.NewKlineSeries(sr.Exchange, sr.InstId, driverInterval) intervalKlineSeries.Set(driverInterval, driverSeries) err = svc.fetchHistoryKlineSeries(ctx, sr, func(k *types.Kline) (err error) { driverTS := driverIntervalAdder(k.Ts, 1) // 通知其他周期先更新 otherIntervalCh.Range(func(interval types.Interval, ch []chan int64) { if len(ch) == 2 { select { case <-stopCh: err = io.EOF return case ch[0] <- driverTS: <-ch[1] // 等待更新完毕 } } }) if err != nil { return } if lastTs, serial := driverSeries.Update(k); !serial { err = fmt.Errorf("kline not series: %s(%s), interval=%s, lastTs=%d", sr.InstId, sr.Exchange, driverInterval, lastTs) return } // 检查满足策略执行条件 update := true requiredIntervalSeries.Range(func(interval types.Interval, require int16) { if update && require > 0 { series := intervalKlineSeries.Get(interval) update = series.Length() >= int(require) } }) if !update { return } intervalKlineSeries.Range(func(interval types.Interval, v *sig.KlineSeries) { if v != nil { zlog.Debugf("strategy update: interval series %s, %d", interval, v.Length()) } }) sigSide := intervalSigStrategy.Update(intervalStrategyContext) if sigSide.IsValid() { side := lang.Ternary(sigSide == types.SideLong, pb.Side_BUY, pb.Side_SELL) signals = append(signals, side) times = append(times, k.Ts) } zlog.Debugf("strategy update finish ------------------------------------") return }) if err != io.EOF { close(stopCh) } return } // singleStrategySeries 单周期策略 func (svc *TradingService) singleStrategySeries(ctx context.Context, sigStrategy strategy.ISingleSigStrategy, sr *pb.SeriesRange) (signals []pb.Side, times []int64, err error) { interval := types.Interval(sr.Interval) requiredSeries := int(sigStrategy.RequiredSeries()) sr.WindowExtra = uint32(requiredSeries) kSeries := sig.NewKlineSeries(sr.Exchange, sr.InstId, interval) indicatorContext := sig.NewIndicatorContext(kSeries) strategyContext := sig.NewStrategyContext(indicatorContext, svc.indicatorReg) err = svc.fetchHistoryKlineSeries(ctx, sr, func(k *types.Kline) (err error) { if lastTs, serial := kSeries.Update(k); !serial { err = fmt.Errorf("kline not series: %s(%s), interval=%s, lastTs=%d", sr.InstId, sr.Exchange, interval, lastTs) return } if kSeries.Length() < requiredSeries { return } sigSide := sigStrategy.Update(strategyContext) if sigSide.IsValid() { side := lang.Ternary(sigSide == types.SideLong, pb.Side_BUY, pb.Side_SELL) signals = append(signals, side) times = append(times, k.Ts) } return }) if err != nil { return } 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 }