diff --git a/api/trading.proto b/api/trading.proto index 7e310b7..196abcb 100644 --- a/api/trading.proto +++ b/api/trading.proto @@ -72,10 +72,12 @@ message ReqStrategySeries { SeriesRange series = 1; string sig_strategy = 2; google.protobuf.Struct input = 3; // 指标参数 + string alignInterval = 4; // 展示周期, 若和交易周期不一致time向展示周期靠拢 } message RspStrategySeries { - repeated Side signal = 1; // 0.sell,1.buy + repeated int32 signal = 1; // pub.Side: 1.buy,2.sell repeated int64 times = 2; + repeated double prices = 3; } message ReqBacktest { diff --git a/internal/admin/service/backtest_service.go b/internal/admin/service/backtest_service.go index 3527362..a19c308 100644 --- a/internal/admin/service/backtest_service.go +++ b/internal/admin/service/backtest_service.go @@ -23,10 +23,10 @@ func NewBacktestService(repo *repository.BacktestRepository) *BacktestService { } func (svc *BacktestService) Route(group *gin.RouterGroup) { - group.POST("listBacktest", svc.ListBacktest) // 回测记录分页 - group.GET("getBacktest", svc.GetBacktest) // 回测记录详情 - group.POST("listBacktestTrades", svc.ListBacktestTrades) // 回测记录交易订单分页 - group.GET("testEquities", svc.TestEquities) // 回测记录资金曲线 + group.POST("list", svc.ListBacktest) // 回测记录分页 + group.GET("get", svc.GetBacktest) // 回测记录详情 + group.POST("listTrades", svc.ListBacktestTrades) // 回测记录交易订单分页 + group.GET("equities", svc.BacktestEquities) // 回测记录资金曲线 } func (svc *BacktestService) ListBacktest(ctx *gin.Context) { @@ -81,7 +81,7 @@ func (svc *BacktestService) ListBacktestTrades(ctx *gin.Context) { } // TestEquityCurve 回测记录资金曲线 -func (svc *BacktestService) TestEquities(ctx *gin.Context) { +func (svc *BacktestService) BacktestEquities(ctx *gin.Context) { backtestId, err := cast.ToInt64E(ctx.Query("backtestId")) if backtestId == 0 || err != nil { ctx.JSON(http.StatusBadRequest, resp.Error("param backtestId format error")) diff --git a/internal/admin/service/market_service.go b/internal/admin/service/market_service.go index 2a3f590..c28d9af 100644 --- a/internal/admin/service/market_service.go +++ b/internal/admin/service/market_service.go @@ -4,6 +4,7 @@ import ( "net/http" repository "sig-pub/internal/admin/repoitory" "sig-pub/pkg/resp" + "sort" "github.com/gin-gonic/gin" ) @@ -29,6 +30,7 @@ func (svc *MarketService) ListInstanceId(ctx *gin.Context) { ctx.JSON(http.StatusInternalServerError, resp.Error(err.Error())) return } + sort.Strings(instIds) ctx.JSON(http.StatusOK, resp.Success(instIds)) } @@ -38,7 +40,7 @@ func (svc *MarketService) ListInstance(ctx *gin.Context) { ctx.JSON(http.StatusInternalServerError, resp.Error(err.Error())) return } - + sort.Slice(insts, func(i, j int) bool { return insts[i].InstId < insts[j].InstId }) ctx.JSON(http.StatusOK, resp.Success(resp.H{ "total": len(insts), "insts": insts, diff --git a/internal/trading/backtest/sig_strategy_backtester.go b/internal/trading/backtest/sig_strategy_backtester.go index 1272e26..c2c336a 100644 --- a/internal/trading/backtest/sig_strategy_backtester.go +++ b/internal/trading/backtest/sig_strategy_backtester.go @@ -53,11 +53,15 @@ func (b *SigStrategyBacktester) Backtest(ctx context.Context, sigStrategyInput types.Input, sr *pb.SeriesRange, iiks *types.InstanceIntervalKlineSeries, + intervalCandlePeriods *types.IntervalState[int16], recvSignal func(instId string, sigSide types.Side, k types.Kline) (err error), ) (err error) { if iiks == nil { iiks = types.NewInstanceIntervalKlineSeries() } + if intervalCandlePeriods == nil { + intervalCandlePeriods = types.NewIntervalState[int16]() + } // init sig strategy if err = b.sigStrategy.Init(sigStrategyInput); err != nil { return @@ -65,11 +69,11 @@ func (b *SigStrategyBacktester) Backtest(ctx context.Context, switch b.sigStrategyType { case strategy.SigStrategyTypeSingle: - err = b.singleStrategySeries(ctx, b.sigStrategy.(strategy.ISingleSigStrategy), sigStrategyInput, sr, iiks, recvSignal) + err = b.singleStrategySeries(ctx, b.sigStrategy.(strategy.ISingleSigStrategy), sigStrategyInput, sr, iiks, intervalCandlePeriods, recvSignal) case strategy.SigStrategyTypeInterval: - err = b.intervalStrategySeries(ctx, b.sigStrategy.(strategy.IIntervalSigStrategy), sigStrategyInput, sr, iiks, recvSignal) + err = b.intervalStrategySeries(ctx, b.sigStrategy.(strategy.IIntervalSigStrategy), sigStrategyInput, sr, iiks, intervalCandlePeriods, recvSignal) case strategy.SigStrategyTypeInstanceInterval: - err = b.instanceIntervalStrategySeries(ctx, b.sigStrategy.(strategy.IInstanceIntervalSigStrategy), sigStrategyInput, sr, iiks, recvSignal) + err = b.instanceIntervalStrategySeries(ctx, b.sigStrategy.(strategy.IInstanceIntervalSigStrategy), sigStrategyInput, sr, iiks, intervalCandlePeriods, recvSignal) default: err = fmt.Errorf("unknown sig strategy type %v", b.sigStrategyType) } @@ -82,16 +86,17 @@ func (b *SigStrategyBacktester) singleStrategySeries(ctx context.Context, sigStrategyInput types.Input, sr *pb.SeriesRange, iiks *types.InstanceIntervalKlineSeries, + intervalCandlePeriods *types.IntervalState[int16], recvSignal func(instId string, sigSide types.Side, k types.Kline) (err error), ) (err error) { driverInstId := sr.InstId driverInterval := types.Interval(sr.Interval) driverSeries := iiks.Get(driverInstId, driverInterval) strategyContext := sig.NewStrategyContext(sigStrategyInput, driverSeries, b.indicatorReg) - requiredPeriods := int(sigStrategy.CandlePeriods(strategyContext)) - - intervalCandlePeriods := types.NewIntervalState[int16]() - intervalCandlePeriods.Set(driverInterval, int16(max(1, requiredPeriods))) + requiredPeriods := max(1, int(sigStrategy.CandlePeriods(strategyContext))) + intervalCandlePeriods.SetIf(driverInterval, int16(requiredPeriods), func(old int16) bool { + return old < int16(requiredPeriods) + }) err = b.multiInstanceIntervalSeries(ctx, sr, []string{sr.InstId}, intervalCandlePeriods, iiks, func(driver bool, instId string, interval types.Interval, k *types.Kline) (err error) { if !driver || instId != driverInstId || interval != driverInterval { @@ -117,6 +122,7 @@ func (b *SigStrategyBacktester) intervalStrategySeries(ctx context.Context, sigStrategyInput types.Input, sr *pb.SeriesRange, iiks *types.InstanceIntervalKlineSeries, + intervalCandlePeriods *types.IntervalState[int16], recvSignal func(instId string, sigSide types.Side, k types.Kline) (err error), ) (err error) { driverInstId := sr.InstId @@ -125,7 +131,15 @@ func (b *SigStrategyBacktester) intervalStrategySeries(ctx context.Context, // 策略上下文 intervalStrategyContext := sig.NewIntervalStrategyContext(sigStrategyInput, intervalKlineSeries, b.indicatorReg) // 各周期所需k线数量 - intervalCandlePeriods := intervalSigStrategy.CandlePeriods(intervalStrategyContext) + strategyIntervalCandlePeriods := intervalSigStrategy.CandlePeriods(intervalStrategyContext) + // 合并所有需加载周期 + strategyIntervalCandlePeriods.Range(func(interval types.Interval, require int16) { + if require > 0 { + intervalCandlePeriods.SetIf(interval, require, func(old int16) bool { + return old < require + }) + } + }) periodsChecked := false err = b.multiInstanceIntervalSeries(ctx, sr, []string{sr.InstId}, intervalCandlePeriods, iiks, func(driver bool, instId string, interval types.Interval, k *types.Kline) (err error) { @@ -163,6 +177,7 @@ func (b *SigStrategyBacktester) instanceIntervalStrategySeries(ctx context.Conte sigStrategyInput types.Input, sr *pb.SeriesRange, iiks *types.InstanceIntervalKlineSeries, + intervalCandlePeriods *types.IntervalState[int16], recvSignal func(instId string, sigSide types.Side, k types.Kline) (err error), ) (err error) { driverInstId := sr.InstId @@ -171,7 +186,14 @@ func (b *SigStrategyBacktester) instanceIntervalStrategySeries(ctx context.Conte // 策略上下文 strategyContext := sig.NewInstanceIntervalSigStrategyContext(sigStrategyInput, iiks, b.indicatorReg) // 各周期所需k线数量 - tradeInsts, intervalCandlePeriods := intervalSigStrategy.CandlePeriods(strategyContext) + tradeInsts, strategyIntervalCandlePeriods := intervalSigStrategy.CandlePeriods(strategyContext) + strategyIntervalCandlePeriods.Range(func(interval types.Interval, require int16) { + if require > 0 { + intervalCandlePeriods.SetIf(interval, require, func(old int16) bool { + return old < require + }) + } + }) periodsChecked := false err = b.multiInstanceIntervalSeries(ctx, sr, tradeInsts, intervalCandlePeriods, iiks, func(driver bool, instId string, interval types.Interval, k *types.Kline) (err error) { diff --git a/internal/trading/backtest/trading_plan_backtester.go b/internal/trading/backtest/trading_plan_backtester.go index f30c5ab..32055df 100644 --- a/internal/trading/backtest/trading_plan_backtester.go +++ b/internal/trading/backtest/trading_plan_backtester.go @@ -118,7 +118,7 @@ func (b *TradingPlanBacktester) Backtest(ctx context.Context) (test *trade.Backt iiks := types.NewInstanceIntervalKlineSeries() b.instanceIntervalSigStrategyContext = sig.NewInstanceIntervalSigStrategyContext(b.tradeStrategyInput, iiks, b.indicatorReg) - err = sigStrategyBacktester.Backtest(ctx, b.sigStrategyInput, b.sr, iiks, func(instId string, sigSide types.Side, k types.Kline) (err error) { + err = sigStrategyBacktester.Backtest(ctx, b.sigStrategyInput, b.sr, iiks, nil, func(instId string, sigSide types.Side, k types.Kline) (err error) { test.Singals++ // 根据交易信号检查仓位平仓 if err = b.closeBySigSingal(instId, sigSide, k); err != nil { diff --git a/internal/trading/trading_service.go b/internal/trading/trading_service.go index 4f563e3..f9bcb36 100644 --- a/internal/trading/trading_service.go +++ b/internal/trading/trading_service.go @@ -327,25 +327,50 @@ func (svc *TradingService) StrategySeries(ctx context.Context, req *pb.ReqStrate } interval := types.Interval(req.Series.Interval) - _, ok = types.SupportedIntervals[interval] - if !ok { + 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, nil, func(instId string, sigSide types.Side, k types.Kline) (err error) { + 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 := lang.Ternary(sigSide == types.SideLong, pb.Side_BUY, pb.Side_SELL) - rsp.Signal = append(rsp.Signal, side) - rsp.Times = append(rsp.Times, k.Ts) + rsp.Signal = append(rsp.Signal, int32(side)) + 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 }) return diff --git a/pkg/indicator/boll.go b/pkg/indicator/boll.go index c2493a2..148b20f 100644 --- a/pkg/indicator/boll.go +++ b/pkg/indicator/boll.go @@ -21,7 +21,7 @@ func (c *Boll) Meta() IndicatorMeta { {Name: "中轨", State: "vector", Type: PlotHistogram, Props: PlotProps{"color": ColorOrange}}, {Name: "上轨", State: "ub", Type: PlotLine, Props: PlotProps{"color": ColorRed2}}, {Name: "下轨", State: "lb", Type: PlotLine, Props: PlotProps{"color": ColorRed2}}, - {Name: "布林带阴影", State: "ub,lb", Type: PlotShadow, Props: PlotProps{"color": "rgba(247, 169, 167, 0.3)"}}, + {Name: "布林带阴影", State: "ub,lb", Type: PlotShadow, Props: PlotProps{"shadowColor": "rgba(241, 169, 166, 0.2)"}}, }, } }