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.
 
 

455 lines
14 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)
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
}