You can not select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
462 lines
15 KiB
462 lines
15 KiB
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 - 1) |
|
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) |
|
default: |
|
err = fmt.Errorf("unknown sig strategy type %v", sigStrategyType) |
|
} |
|
return |
|
} |
|
|
|
// singleStrategySeries 单周期策略 |
|
func (svc *TradingService) singleStrategySeries(ctx context.Context, sigStrategy strategy.ISingleSigStrategy, sr *pb.SeriesRange) (signals []pb.Side, times []int64, err error) { |
|
interval := types.Interval(sr.Interval) |
|
requiredSeries := int(sigStrategy.RequiredSeries()) |
|
sr.WindowExtra = uint32(requiredSeries - 1) |
|
kSeries := sig.NewKlineSeries(sr.Exchange, sr.InstId, interval) |
|
indicatorContext := sig.NewIndicatorContext(kSeries) |
|
strategyContext := sig.NewStrategyContext(indicatorContext, svc.indicatorReg) |
|
err = svc.fetchHistoryKlineSeries(ctx, sr, func(k *types.Kline) (err error) { |
|
if lastTs, serial := kSeries.Update(k); !serial { |
|
err = fmt.Errorf("kline not series: %s(%s), interval=%s, lastTs=%d", sr.InstId, sr.Exchange, interval, lastTs) |
|
return |
|
} |
|
if kSeries.Length() < requiredSeries { |
|
return |
|
} |
|
sigSide := sigStrategy.Update(strategyContext) |
|
if sigSide.IsValid() { |
|
side := lang.Ternary(sigSide == types.SideLong, pb.Side_BUY, pb.Side_SELL) |
|
signals = append(signals, side) |
|
times = append(times, k.Ts) |
|
} |
|
return |
|
}) |
|
if err != nil { |
|
return |
|
} |
|
return |
|
} |
|
|
|
// intervalStrategySeries 多周期策略 |
|
func (svc *TradingService) intervalStrategySeries(ctx context.Context, intervalSigStrategy strategy.IIntervalSigStrategy, sr *pb.SeriesRange) (signals []pb.Side, times []int64, err error) { |
|
driverInterval := types.Interval(sr.Interval) |
|
driverIntervalAdder := types.SupportedIntervals[driverInterval] |
|
// 各周期所需k线数量 |
|
requiredIntervalSeries := intervalSigStrategy.RequiredIntervalSeries() |
|
// 驱动周期外其他周期 |
|
otherIntervals := make([]types.Interval, 0, 3) |
|
requiredIntervalSeries.Range(func(interval types.Interval, series int16) { |
|
if series > 0 { |
|
otherIntervals = append(otherIntervals, interval) |
|
} |
|
}) |
|
otherIntervals = collect.Filter(otherIntervals, func(_ int, interval types.Interval) bool { return interval != driverInterval }) |
|
|
|
// 通知其他周期更新的channel |
|
otherIntervalCh := types.NewIntervalState[[]chan int64]() |
|
// otherIntervalDstCh := types.NewIntervalState[chan int64]() |
|
for _, interval := range otherIntervals { |
|
otherIntervalCh.Set(interval, []chan int64{make(chan int64), make(chan int64)}) |
|
} |
|
// 各周期 series |
|
intervalKlineSeries := types.NewIntervalState[*sig.KlineSeries]() |
|
// 其他周期数据拉取 |
|
stopCh := make(chan struct{}) |
|
for _, interval := range otherIntervals { |
|
kSeries := sig.NewKlineSeries(sr.Exchange, sr.InstId, interval) |
|
intervalKlineSeries.Set(interval, kSeries) |
|
go func(interval types.Interval, kSeries *sig.KlineSeries) { |
|
ch := otherIntervalCh.Get(interval) |
|
srcCh := ch[0] |
|
dstCh := ch[1] |
|
intervalAdder := types.SupportedIntervals[interval] |
|
driverTs := int64(0) |
|
isr := &pb.SeriesRange{Exchange: sr.Exchange, InstId: sr.InstId, Before: sr.Before, After: sr.After, Count: sr.Count, Open: sr.Open, Live: sr.Live, Desc: sr.Desc, Limit: sr.Limit} |
|
isr.Interval = string(interval) |
|
isr.WindowExtra = uint32(requiredIntervalSeries.Get(interval) - 1) |
|
err1 := svc.fetchHistoryKlineSeries(ctx, isr, func(k *types.Kline) (err error) { |
|
closeTs := intervalAdder(k.Ts, 1) |
|
// 与驱动周期series保持同步更新 |
|
if closeTs > driverTs { |
|
waitLoop: |
|
for { |
|
if driverTs != 0 { |
|
dstCh <- 0 // 通知更新完毕 |
|
} |
|
select { |
|
case <-stopCh: |
|
return io.EOF |
|
case driverTs = <-srcCh: |
|
if closeTs <= driverTs { |
|
break waitLoop |
|
} |
|
} |
|
} |
|
} |
|
if lastTs, serial := kSeries.Update(k); !serial { |
|
err = fmt.Errorf("kline not series: %s(%s), interval=%s, lastTs=%d", sr.InstId, sr.Exchange, interval, lastTs) |
|
return |
|
} |
|
return |
|
}) |
|
if err1 == nil { |
|
otherIntervalCh.Set(interval, nil) // 该周期数据拉取结束 |
|
dstCh <- 0 // 通知更新完毕 |
|
} else if err1 != io.EOF { |
|
zlog.Errorf("fetch history interval error: inst=%s(%s) interval=%s, err=%v", sr.InstId, sr.Exchange, interval, err1) |
|
err = err1 |
|
close(stopCh) |
|
} |
|
}(interval, kSeries) |
|
} |
|
|
|
// 策略上下文 |
|
intervalStrategyContext := sig.NewIntervalStrategyContext(intervalKlineSeries, svc.indicatorReg) |
|
|
|
// 驱动周期数据拉取 |
|
driverSeries := sig.NewKlineSeries(sr.Exchange, sr.InstId, driverInterval) |
|
intervalKlineSeries.Set(driverInterval, driverSeries) |
|
sr.WindowExtra = uint32(requiredIntervalSeries.Get(driverInterval) - 1) |
|
err = svc.fetchHistoryKlineSeries(ctx, sr, func(k *types.Kline) (err error) { |
|
driverTS := driverIntervalAdder(k.Ts, 1) |
|
otherIntervalCh.Range(func(interval types.Interval, ch []chan int64) { |
|
if len(ch) == 2 { |
|
// 通知其它周期先更新 |
|
select { |
|
case <-stopCh: |
|
err = io.EOF |
|
return |
|
case ch[0] <- driverTS: |
|
// 等待其它周期更新完毕 |
|
select { |
|
case <-stopCh: |
|
err = io.EOF |
|
return |
|
case <-ch[1]: |
|
} |
|
} |
|
} |
|
}) |
|
if err != nil { |
|
return |
|
} |
|
|
|
zlog.Debugf("driver series update: %s, %d", driverInterval, driverTS) |
|
if lastTs, serial := driverSeries.Update(k); !serial { |
|
err = fmt.Errorf("kline not series: %s(%s), interval=%s, lastTs=%d", sr.InstId, sr.Exchange, driverInterval, lastTs) |
|
return |
|
} |
|
// 检查满足策略执行条件 |
|
update := true |
|
requiredIntervalSeries.Range(func(interval types.Interval, require int16) { |
|
if update && require > 0 { |
|
series := intervalKlineSeries.Get(interval) |
|
update = series.Length() >= int(require) |
|
} |
|
}) |
|
if !update { |
|
return |
|
} |
|
intervalKlineSeries.Range(func(interval types.Interval, v *sig.KlineSeries) { |
|
if v != nil { |
|
zlog.Debugf("strategy update: interval series %s, %d", interval, v.Length()) |
|
} |
|
}) |
|
sigSide := intervalSigStrategy.Update(intervalStrategyContext) |
|
if sigSide.IsValid() { |
|
side := lang.Ternary(sigSide == types.SideLong, pb.Side_BUY, pb.Side_SELL) |
|
signals = append(signals, side) |
|
times = append(times, k.Ts) |
|
} |
|
return |
|
}) |
|
if err != io.EOF { |
|
close(stopCh) |
|
} |
|
return |
|
} |
|
|
|
// 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 |
|
}
|
|
|