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.
 
 

522 lines
16 KiB

package trading
import (
"context"
"fmt"
"io"
"math"
"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/trade"
"sig-pub/pkg/types"
"sig-pub/pkg/utils/collect"
"sig-pub/pkg/utils/lang"
"sig-pub/pkg/utils/times"
"sig-pub/pkg/zlog"
"sort"
"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 *types.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 %s", plan.Interval)
return
}
// sigStrategy
var sigStrategyInput types.Input
if err = sonic.UnmarshalString(plan.SigStrategyParam, &sigStrategyInput); 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
}
_, _ = 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 {
// 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
}
// Indicators
func (svc *TradingService) Indicators() (indicatorNames []string) {
svc.indicatorReg.RangeIndicators(func(k string, v indicator.IIndicator) bool {
indicatorNames = append(indicatorNames, k)
return true
})
sort.Strings(indicatorNames)
return
}
// IndicatorMeta
func (svc *TradingService) IndicatorMeta(indicatorNames ...string) (indMetas map[string]indicator.IndicatorMeta, err error) {
indMetas = make(map[string]indicator.IndicatorMeta, len(indicatorNames))
for _, indicatorName := range indicatorNames {
ind, ok := svc.indicatorReg.Indicator(indicatorName)
if !ok {
err = fmt.Errorf("indicator %s not exists", indicatorName)
return nil, err
}
indMetas[indicatorName] = ind.Meta()
}
return
}
// IndicatorPlots
func (svc *TradingService) IndicatorPlots(indicatorNames ...string) (indPlots map[string][]indicator.Plot, err error) {
indPlots = make(map[string][]indicator.Plot, len(indicatorNames))
for _, indicatorName := range indicatorNames {
plots, e := svc.indicatorPlots(indicatorName)
if e != nil {
return nil, e
}
indPlots[indicatorName] = plots
}
return
}
func (svc *TradingService) indicatorPlots(indicatorName string) (plots []indicator.Plot, err error) {
ind, ok := svc.indicatorReg.Indicator(indicatorName)
if !ok {
err = fmt.Errorf("indicator %s not exists", indicatorName)
return
}
plots = ind.Meta().Plots
if len(plots) == 0 {
plots = append(plots, indicator.Plot{
State: "vector",
Type: indicator.PlotLine,
Props: indicator.PlotProps{"color": indicator.ColorBlue},
})
}
return
}
// IndicatorSeries 获取指标实时或历史序列数据, 闭区间
func (svc *TradingService) IndicatorSeries(ctx context.Context, indicatorName string, digit int32, input types.Input, sr *pb.SeriesRange) (matrix []float64, times []int64, states map[string][]float64, err error) {
appros := indicator.ApproCandles
indicator, ok := svc.indicatorReg.Indicator(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
}
// 查询历史指标数据
kSeries := types.NewKlineSeries(sr.Exchange, sr.InstId, types.Interval(sr.Interval))
indicatorContext := sig.NewIndicatorContext(indicator, input, sig.NewIndicatorStates(), kSeries, svc.indicatorReg)
candlePeriods := int(indicator.CandlePeriods(indicatorContext))
sr.Desc = false
// 计算第一个指标值需要多取candlePeriods - 1根k线; 计算ema等递归指标需要多拉取appros根k线逼近值
srRsp, err := svc.exchangeClient.SeriesRange(ctx, &pb.ReqSeriesRange{Series: sr})
if err != nil {
return
}
srBefore := srRsp.Before
sr.WindowExtra = uint32(max(0, candlePeriods-1)) + uint32(appros)
indicatorStates := indicator.Meta().State // 指标导出状态
digit = lang.Ternary(digit > 0 && digit <= 10, digit, 6) // 保留小数位数
pow := math.Pow(10, float64(digit))
matrix = make([]float64, 0, 200)
times = make([]int64, 0, 200)
states = make(map[string][]float64, len(indicatorStates))
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 k.Ts < srBefore {
return
}
if kSeries.Length() < candlePeriods {
return
}
vector := indicator.Calculate(indicatorContext)
matrix = append(matrix, math.Round(vector*pow)/pow)
times = append(times, indicatorContext.Get(0).Ts)
// 状态填充
for _, state := range indicatorStates {
sv, _ := indicatorContext.State().Get(state, 0)
states[state] = append(states[state], math.Round(sv*pow)/pow)
}
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
}
interval := types.Interval(req.Series.Interval)
_, ok = types.SupportedIntervals[interval]
if !ok {
err = fmt.Errorf("unsupport interval %s", interval)
return
}
// 信号策略参数
sigStrategyInput := types.Input(req.Input.AsMap())
// 使用回测器回测信号策略
driverInstId := req.Series.InstId
backtester := backtest.NewSigStrategyBacktester(sigStrategyType, sigStrategy, svc.indicatorReg, svc.exchangeClient)
err = backtester.Backtest(ctx, sigStrategyInput, req.Series, nil, func(instId string, sigSide types.Side, k types.Kline) (err error) {
if instId != driverInstId {
// todo 多币种回测
return
}
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)
return
})
return
}
// Backtest 回测交易计划
func (svc *TradingService) Backtest(ctx context.Context, 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)
if _, ok := types.SupportedIntervals[interval]; !ok {
err = fmt.Errorf("unsupport interval %s", interval)
return
}
sr := &pb.SeriesRange{
Exchange: exchange,
InstId: plan.InstId,
Interval: plan.Interval,
Before: stime,
After: etime,
Open: false,
Live: false,
Desc: false,
}
sigStrategyType, sigStrategy, ok := svc.strategyReg.NewSigStrategy(plan.SigStrategy)
if !ok {
err = fmt.Errorf("sig strategy not exists %s", plan.SigStrategy)
return
}
tester := backtest.NewTradingPlanBacktester(sigStrategyType, sigStrategy, svc.indicatorReg, svc.exchangeClient)
if err = tester.Init(10000, *plan, sr); err != nil {
return
}
w := times.NewWatch()
backtestTradingPlan, err := tester.Backtest(ctx)
if err != nil {
return
}
zlog.Infof("backtest use %s", w.ElapsedFmt(""))
w.Reset()
err = svc.tradingDataPersist.SaveBacktestTradingPlan(ctx, backtestTradingPlan)
zlog.Infof("insert backtest ret use %s", w.ElapsedFmt(""))
return
}
// BacktestLog 回测记录查询
func (svc *TradingService) BacktestLog(ctx context.Context, userId int64) (backtestLogs []*trade.BacktestTradingPlan, err error) {
// todo userid from ctx
backtestLogs, err = svc.tradingDataPersist.ListBacktestLogs(userId)
return
}
// BacktestRace 交易计划参数调试回测
func (svc *TradingService) BacktestRace(ctx context.Context, req *pb.ReqBacktestRace) (err error) {
// 参数组合
dbPlan, err := svc.tradingDataPersist.GetTradePlanById(req.PlanId)
if err != nil {
return
}
sigStrategyType, sigStrategy, ok := svc.strategyReg.NewSigStrategy(dbPlan.SigStrategy)
if !ok {
err = fmt.Errorf("sig strategy not exists %s", dbPlan.SigStrategy)
return
}
exchange := pb.ExchangeType(dbPlan.Exchange)
interval := types.Interval(dbPlan.Interval)
if _, ok := types.SupportedIntervals[interval]; !ok {
err = fmt.Errorf("unsupport interval %s", interval)
return
}
var seriesInputs, sigInputs, closeInputs [][]types.Input
if seriesInputs, err = parseProtoInputRanges(req.SeriesInputRange); err != nil {
return
}
if sigInputs, err = parseProtoInputRanges(req.SigInputRange); err != nil {
return
}
if closeInputs, err = parseProtoInputRanges(req.CloseInputRange); err != nil {
return
}
_, _, _ = seriesInputs, sigInputs, closeInputs
progress := len(seriesInputs) * len(sigInputs) * len(closeInputs)
_ = progress
var backtesters []*backtest.TradingPlanBacktester
// 输入参数组合
seriesInputGroups := collect.Mapping(collect.CartesianProduct(seriesInputs...), func(inputs []types.Input) types.Input {
return (types.Input{}).Assign(inputs...)
})
sigInputGroups := collect.Mapping(collect.CartesianProduct(sigInputs...), func(inputs []types.Input) types.Input {
return (types.Input{}).Assign(inputs...)
})
closeInputGroups := collect.Mapping(collect.CartesianProduct(closeInputs...), func(inputs []types.Input) types.Input {
return (types.Input{}).Assign(inputs...)
})
for _, seriesIn := range seriesInputGroups {
for _, sigIn := range sigInputGroups {
for _, closeIn := range closeInputGroups {
// 生成一个 tester
plan := *dbPlan
if plan.SigStrategyParam, err = sonic.MarshalString(sigIn); err != nil {
return
}
if plan.CloseStrategyParam, err = sonic.MarshalString(closeIn); err != nil {
return
}
before := seriesIn.Time("before")
after := seriesIn.Time("after")
sr := &pb.SeriesRange{
Exchange: exchange,
InstId: dbPlan.InstId,
Interval: dbPlan.Interval,
Before: before.UnixMilli(),
After: after.UnixMilli(),
Open: false,
Live: false,
Desc: false,
}
// todo 参数校验
// if err = sigStrategy.New().Init(sigIn); err != nil {
// return
// }
tester := backtest.NewTradingPlanBacktester(sigStrategyType, sigStrategy, svc.indicatorReg, svc.exchangeClient)
err = tester.Init(10000, plan, sr)
if err != nil {
return
}
backtesters = append(backtesters, tester)
}
}
}
var results []*trade.BacktestTradingPlan
for _, backtester := range backtesters {
r, e := backtester.Backtest(ctx)
if e != nil {
err = e
return
}
results = append(results, r)
}
for _, r := range results {
winRate := fmt.Sprintf(" %.2f%% ", float64(r.WinningTrades)/float64(r.TotalTrades)*100)
zlog.Info(r.Id, r.Cash, r.EndCash, r.Profit, r.Fee, r.Singals, r.TotalTrades, winRate, r.MaxDrawdown)
}
return
}
func parseProtoInputRanges(irs []*pb.InputRange) (inputs [][]types.Input, err error) {
for _, ir := range irs {
irt := &types.InputRange{
Name: ir.Name,
Type: ir.Type,
Value: ir.Value,
}
gen, e := irt.NewInputGen()
if e != nil {
err = e
return
}
values := gen.Values()
if len(values) == 0 {
err = fmt.Errorf("input %s no value", ir.Name)
return
}
inputs = append(inputs, values)
}
return
}