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.
655 lines
20 KiB
655 lines
20 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/misc" |
|
"sig-pub/pkg/utils/progress" |
|
"sig-pub/pkg/utils/times" |
|
"sig-pub/pkg/zlog" |
|
"sort" |
|
"time" |
|
|
|
"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] // 运行中交易计划 |
|
backtestingTasks *progress.ProgressManager // 回测进度管理 |
|
} |
|
|
|
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](), |
|
backtestingTasks: progress.NewProgressManager(), |
|
} |
|
} |
|
|
|
// 初始化历史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 = misc.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 |
|
} |
|
|
|
func (svc *TradingService) IndicatorSummary(ctx context.Context, indicatorName string, digit int32, input types.Input, sr *pb.SeriesRange) (summary any, err error) { |
|
indicator, ok := svc.indicatorReg.SummaryIndicator(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) |
|
|
|
// 初始化指标 |
|
err = indicator.Init(input) |
|
if err != nil { |
|
return |
|
} |
|
|
|
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)) |
|
|
|
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 |
|
} |
|
indicator.Accumulate(indicatorContext) |
|
return |
|
}) |
|
if err != nil { |
|
return |
|
} |
|
|
|
// 计算指标汇总值 |
|
summary, ok = indicator.Summary(indicatorContext) |
|
if !ok { |
|
err = fmt.Errorf("indicator summary result error: %s", indicatorName) |
|
return |
|
} |
|
|
|
// 保留小数位数 |
|
digit = misc.Ternary(digit > 0 && digit <= 10, digit, 6) |
|
pow := math.Pow(10, float64(digit)) |
|
switch sm := summary.(type) { |
|
case *types.VRVPSummary: |
|
sm.Step = math.Round(sm.Step*pow) / pow |
|
for i := range len(sm.Buckets) { |
|
sm.Buckets[i].Price = math.Round(sm.Buckets[i].Price*pow) / pow |
|
sm.Buckets[i].Volume = math.Round(sm.Buckets[i].Volume*pow) / pow |
|
sm.Buckets[i].BuyVol = math.Round(sm.Buckets[i].BuyVol*pow) / pow |
|
sm.Buckets[i].SellVol = math.Round(sm.Buckets[i].SellVol*pow) / pow |
|
} |
|
} |
|
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) |
|
if _, ok = types.SupportedIntervals[interval]; !ok { |
|
err = fmt.Errorf("unsupport interval %s", interval) |
|
return |
|
} |
|
alignInterval := types.Interval(req.AlignInterval) |
|
if req.AlignInterval == "" { |
|
alignInterval = interval |
|
} |
|
if _, ok = types.SupportedIntervals[alignInterval]; !ok { |
|
err = fmt.Errorf("unsupport alignInterval %s", alignInterval) |
|
return |
|
} |
|
|
|
// 信号策略参数 |
|
sigStrategyInput := types.Input(req.Input.AsMap()) |
|
// 使用回测器回测信号策略 |
|
driverInstId := req.Series.InstId |
|
// 各周期k线数据 |
|
iiks := types.NewInstanceIntervalKlineSeries() |
|
// 额外拉取k线周期 |
|
intervalCandlePeriods := types.NewIntervalState[int16]() |
|
if alignInterval != interval { |
|
intervalCandlePeriods.Set(alignInterval, 1) |
|
} |
|
|
|
backtester := backtest.NewSigStrategyBacktester(sigStrategyType, sigStrategy, svc.indicatorReg, svc.exchangeClient) |
|
err = backtester.Backtest(ctx, sigStrategyInput, req.Series, iiks, intervalCandlePeriods, func(instId string, sigSide types.Side, k types.Kline) (err error) { |
|
if instId != driverInstId { |
|
// todo 多币种回测 |
|
return |
|
} |
|
side := misc.Ternary(sigSide == types.SideLong, pb.Side_BUY, pb.Side_SELL) |
|
rsp.Signal = append(rsp.Signal, int32(side)) |
|
rsp.Direction = append(rsp.Direction, misc.Ternary[int32](k.OpenF64() > k.CloseF64(), 1, 2)) |
|
rsp.Prices = append(rsp.Prices, k.CloseF64()) |
|
// 处理k线时间 |
|
ktime := k.Ts |
|
if alignInterval != interval { |
|
alignK, err := iiks.Get(driverInstId, alignInterval).Get(0) |
|
if err != nil { |
|
return err |
|
} |
|
ktime = alignK.Ts |
|
} |
|
rsp.Times = append(rsp.Times, ktime) |
|
return |
|
}) |
|
lastK := iiks.Get(driverInstId, interval).MustGet(0) |
|
fmt.Println("lastK: ", lastK, lastK.Ts, time.UnixMilli(lastK.Ts)) |
|
return |
|
} |
|
|
|
// Backtest 回测交易计划 |
|
func (svc *TradingService) Backtest(ctx context.Context, planId, stime, etime int64) (taskId string, 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 |
|
} |
|
|
|
// 生成异步任务, 获取任务进度 |
|
taskId = svc.backtestingTasks.StartTask(context.Background(), fmt.Sprintf("backtest %d", plan.Id), func(ctx context.Context, progress progress.IProgressReporter) error { |
|
w := times.NewWatch() |
|
backtestTradingPlan, err := tester.Backtest(ctx) |
|
if err != nil { |
|
return err |
|
} |
|
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 nil |
|
}) |
|
return |
|
} |
|
|
|
// BacktestLog 回测记录查询 |
|
func (svc *TradingService) BacktestLog(ctx context.Context, userId int64, paging data.Page, running bool) (backtestLogs []*trade.BacktestTradingPlan, err error) { |
|
// todo userid from ctx |
|
backtestLogs, err = svc.tradingDataPersist.ListBacktestLogs(userId) |
|
if err != nil { |
|
return |
|
} |
|
// 在首页加载运行中任务 |
|
if running && paging.Page == 1 { |
|
runningTasks := svc.backtestingTasks.ListRunningTask() |
|
var runningLogs []*trade.BacktestTradingPlan |
|
for _, task := range runningTasks { |
|
progress, total := task.Progress.Get() |
|
runningLogs = append(runningLogs, &trade.BacktestTradingPlan{ |
|
RunningStatus: int32(task.Status.Load()), |
|
RunningTaskId: task.ID, |
|
RunningProgress: progress, |
|
RunningTotal: total, |
|
RunningPct: fmt.Sprintf("%.2f", float64(progress)/float64(total)), |
|
}) |
|
} |
|
if len(runningLogs) > 0 { |
|
backtestLogs = append(runningLogs, backtestLogs...) |
|
} |
|
} |
|
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 |
|
}
|
|
|