Browse Source

strategy series align interval

main
strange 7 months ago
parent
commit
5e1c23e73e
  1. 4
      api/trading.proto
  2. 10
      internal/admin/service/backtest_service.go
  3. 4
      internal/admin/service/market_service.go
  4. 40
      internal/trading/backtest/sig_strategy_backtester.go
  5. 2
      internal/trading/backtest/trading_plan_backtester.go
  6. 35
      internal/trading/trading_service.go
  7. 2
      pkg/indicator/boll.go

4
api/trading.proto

@ -72,10 +72,12 @@ message ReqStrategySeries {
SeriesRange series = 1; SeriesRange series = 1;
string sig_strategy = 2; string sig_strategy = 2;
google.protobuf.Struct input = 3; // google.protobuf.Struct input = 3; //
string alignInterval = 4; // , time向展示周期靠拢
} }
message RspStrategySeries { 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 int64 times = 2;
repeated double prices = 3;
} }
message ReqBacktest { message ReqBacktest {

10
internal/admin/service/backtest_service.go

@ -23,10 +23,10 @@ func NewBacktestService(repo *repository.BacktestRepository) *BacktestService {
} }
func (svc *BacktestService) Route(group *gin.RouterGroup) { func (svc *BacktestService) Route(group *gin.RouterGroup) {
group.POST("listBacktest", svc.ListBacktest) // 回测记录分页 group.POST("list", svc.ListBacktest) // 回测记录分页
group.GET("getBacktest", svc.GetBacktest) // 回测记录详情 group.GET("get", svc.GetBacktest) // 回测记录详情
group.POST("listBacktestTrades", svc.ListBacktestTrades) // 回测记录交易订单分页 group.POST("listTrades", svc.ListBacktestTrades) // 回测记录交易订单分页
group.GET("testEquities", svc.TestEquities) // 回测记录资金曲线 group.GET("equities", svc.BacktestEquities) // 回测记录资金曲线
} }
func (svc *BacktestService) ListBacktest(ctx *gin.Context) { func (svc *BacktestService) ListBacktest(ctx *gin.Context) {
@ -81,7 +81,7 @@ func (svc *BacktestService) ListBacktestTrades(ctx *gin.Context) {
} }
// TestEquityCurve 回测记录资金曲线 // TestEquityCurve 回测记录资金曲线
func (svc *BacktestService) TestEquities(ctx *gin.Context) { func (svc *BacktestService) BacktestEquities(ctx *gin.Context) {
backtestId, err := cast.ToInt64E(ctx.Query("backtestId")) backtestId, err := cast.ToInt64E(ctx.Query("backtestId"))
if backtestId == 0 || err != nil { if backtestId == 0 || err != nil {
ctx.JSON(http.StatusBadRequest, resp.Error("param backtestId format error")) ctx.JSON(http.StatusBadRequest, resp.Error("param backtestId format error"))

4
internal/admin/service/market_service.go

@ -4,6 +4,7 @@ import (
"net/http" "net/http"
repository "sig-pub/internal/admin/repoitory" repository "sig-pub/internal/admin/repoitory"
"sig-pub/pkg/resp" "sig-pub/pkg/resp"
"sort"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
) )
@ -29,6 +30,7 @@ func (svc *MarketService) ListInstanceId(ctx *gin.Context) {
ctx.JSON(http.StatusInternalServerError, resp.Error(err.Error())) ctx.JSON(http.StatusInternalServerError, resp.Error(err.Error()))
return return
} }
sort.Strings(instIds)
ctx.JSON(http.StatusOK, resp.Success(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())) ctx.JSON(http.StatusInternalServerError, resp.Error(err.Error()))
return return
} }
sort.Slice(insts, func(i, j int) bool { return insts[i].InstId < insts[j].InstId })
ctx.JSON(http.StatusOK, resp.Success(resp.H{ ctx.JSON(http.StatusOK, resp.Success(resp.H{
"total": len(insts), "total": len(insts),
"insts": insts, "insts": insts,

40
internal/trading/backtest/sig_strategy_backtester.go

@ -53,11 +53,15 @@ func (b *SigStrategyBacktester) Backtest(ctx context.Context,
sigStrategyInput types.Input, sigStrategyInput types.Input,
sr *pb.SeriesRange, sr *pb.SeriesRange,
iiks *types.InstanceIntervalKlineSeries, iiks *types.InstanceIntervalKlineSeries,
intervalCandlePeriods *types.IntervalState[int16],
recvSignal func(instId string, sigSide types.Side, k types.Kline) (err error), recvSignal func(instId string, sigSide types.Side, k types.Kline) (err error),
) (err error) { ) (err error) {
if iiks == nil { if iiks == nil {
iiks = types.NewInstanceIntervalKlineSeries() iiks = types.NewInstanceIntervalKlineSeries()
} }
if intervalCandlePeriods == nil {
intervalCandlePeriods = types.NewIntervalState[int16]()
}
// init sig strategy // init sig strategy
if err = b.sigStrategy.Init(sigStrategyInput); err != nil { if err = b.sigStrategy.Init(sigStrategyInput); err != nil {
return return
@ -65,11 +69,11 @@ func (b *SigStrategyBacktester) Backtest(ctx context.Context,
switch b.sigStrategyType { switch b.sigStrategyType {
case strategy.SigStrategyTypeSingle: 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: 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: 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: default:
err = fmt.Errorf("unknown sig strategy type %v", b.sigStrategyType) err = fmt.Errorf("unknown sig strategy type %v", b.sigStrategyType)
} }
@ -82,16 +86,17 @@ func (b *SigStrategyBacktester) singleStrategySeries(ctx context.Context,
sigStrategyInput types.Input, sigStrategyInput types.Input,
sr *pb.SeriesRange, sr *pb.SeriesRange,
iiks *types.InstanceIntervalKlineSeries, iiks *types.InstanceIntervalKlineSeries,
intervalCandlePeriods *types.IntervalState[int16],
recvSignal func(instId string, sigSide types.Side, k types.Kline) (err error), recvSignal func(instId string, sigSide types.Side, k types.Kline) (err error),
) (err error) { ) (err error) {
driverInstId := sr.InstId driverInstId := sr.InstId
driverInterval := types.Interval(sr.Interval) driverInterval := types.Interval(sr.Interval)
driverSeries := iiks.Get(driverInstId, driverInterval) driverSeries := iiks.Get(driverInstId, driverInterval)
strategyContext := sig.NewStrategyContext(sigStrategyInput, driverSeries, b.indicatorReg) strategyContext := sig.NewStrategyContext(sigStrategyInput, driverSeries, b.indicatorReg)
requiredPeriods := int(sigStrategy.CandlePeriods(strategyContext)) requiredPeriods := max(1, int(sigStrategy.CandlePeriods(strategyContext)))
intervalCandlePeriods.SetIf(driverInterval, int16(requiredPeriods), func(old int16) bool {
intervalCandlePeriods := types.NewIntervalState[int16]() return old < int16(requiredPeriods)
intervalCandlePeriods.Set(driverInterval, int16(max(1, requiredPeriods))) })
err = b.multiInstanceIntervalSeries(ctx, sr, []string{sr.InstId}, intervalCandlePeriods, iiks, func(driver bool, instId string, interval types.Interval, k *types.Kline) (err error) { 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 { if !driver || instId != driverInstId || interval != driverInterval {
@ -117,6 +122,7 @@ func (b *SigStrategyBacktester) intervalStrategySeries(ctx context.Context,
sigStrategyInput types.Input, sigStrategyInput types.Input,
sr *pb.SeriesRange, sr *pb.SeriesRange,
iiks *types.InstanceIntervalKlineSeries, iiks *types.InstanceIntervalKlineSeries,
intervalCandlePeriods *types.IntervalState[int16],
recvSignal func(instId string, sigSide types.Side, k types.Kline) (err error), recvSignal func(instId string, sigSide types.Side, k types.Kline) (err error),
) (err error) { ) (err error) {
driverInstId := sr.InstId driverInstId := sr.InstId
@ -125,7 +131,15 @@ func (b *SigStrategyBacktester) intervalStrategySeries(ctx context.Context,
// 策略上下文 // 策略上下文
intervalStrategyContext := sig.NewIntervalStrategyContext(sigStrategyInput, intervalKlineSeries, b.indicatorReg) intervalStrategyContext := sig.NewIntervalStrategyContext(sigStrategyInput, intervalKlineSeries, b.indicatorReg)
// 各周期所需k线数量 // 各周期所需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 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) { 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, sigStrategyInput types.Input,
sr *pb.SeriesRange, sr *pb.SeriesRange,
iiks *types.InstanceIntervalKlineSeries, iiks *types.InstanceIntervalKlineSeries,
intervalCandlePeriods *types.IntervalState[int16],
recvSignal func(instId string, sigSide types.Side, k types.Kline) (err error), recvSignal func(instId string, sigSide types.Side, k types.Kline) (err error),
) (err error) { ) (err error) {
driverInstId := sr.InstId driverInstId := sr.InstId
@ -171,7 +186,14 @@ func (b *SigStrategyBacktester) instanceIntervalStrategySeries(ctx context.Conte
// 策略上下文 // 策略上下文
strategyContext := sig.NewInstanceIntervalSigStrategyContext(sigStrategyInput, iiks, b.indicatorReg) strategyContext := sig.NewInstanceIntervalSigStrategyContext(sigStrategyInput, iiks, b.indicatorReg)
// 各周期所需k线数量 // 各周期所需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 periodsChecked := false
err = b.multiInstanceIntervalSeries(ctx, sr, tradeInsts, intervalCandlePeriods, iiks, func(driver bool, instId string, interval types.Interval, k *types.Kline) (err error) { err = b.multiInstanceIntervalSeries(ctx, sr, tradeInsts, intervalCandlePeriods, iiks, func(driver bool, instId string, interval types.Interval, k *types.Kline) (err error) {

2
internal/trading/backtest/trading_plan_backtester.go

@ -118,7 +118,7 @@ func (b *TradingPlanBacktester) Backtest(ctx context.Context) (test *trade.Backt
iiks := types.NewInstanceIntervalKlineSeries() iiks := types.NewInstanceIntervalKlineSeries()
b.instanceIntervalSigStrategyContext = sig.NewInstanceIntervalSigStrategyContext(b.tradeStrategyInput, iiks, b.indicatorReg) 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++ test.Singals++
// 根据交易信号检查仓位平仓 // 根据交易信号检查仓位平仓
if err = b.closeBySigSingal(instId, sigSide, k); err != nil { if err = b.closeBySigSingal(instId, sigSide, k); err != nil {

35
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) interval := types.Interval(req.Series.Interval)
_, ok = types.SupportedIntervals[interval] if _, ok = types.SupportedIntervals[interval]; !ok {
if !ok {
err = fmt.Errorf("unsupport interval %s", interval) err = fmt.Errorf("unsupport interval %s", interval)
return 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()) sigStrategyInput := types.Input(req.Input.AsMap())
// 使用回测器回测信号策略 // 使用回测器回测信号策略
driverInstId := req.Series.InstId 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) 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 { if instId != driverInstId {
// todo 多币种回测 // todo 多币种回测
return return
} }
side := lang.Ternary(sigSide == types.SideLong, pb.Side_BUY, pb.Side_SELL) side := lang.Ternary(sigSide == types.SideLong, pb.Side_BUY, pb.Side_SELL)
rsp.Signal = append(rsp.Signal, side) rsp.Signal = append(rsp.Signal, int32(side))
rsp.Times = append(rsp.Times, k.Ts) 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
}) })
return return

2
pkg/indicator/boll.go

@ -21,7 +21,7 @@ func (c *Boll) Meta() IndicatorMeta {
{Name: "中轨", State: "vector", Type: PlotHistogram, Props: PlotProps{"color": ColorOrange}}, {Name: "中轨", State: "vector", Type: PlotHistogram, Props: PlotProps{"color": ColorOrange}},
{Name: "上轨", State: "ub", Type: PlotLine, Props: PlotProps{"color": ColorRed2}}, {Name: "上轨", State: "ub", Type: PlotLine, Props: PlotProps{"color": ColorRed2}},
{Name: "下轨", State: "lb", 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)"}},
}, },
} }
} }

Loading…
Cancel
Save