Browse Source

trading backtestrace

main
strange 8 months ago
parent
commit
0794c674ff
  1. 2
      README.md
  2. 20
      api/exchange.proto
  3. 12
      api/market.proto
  4. 65
      api/pub.proto
  5. 73
      api/trading.proto
  6. 4
      cmd/trading/test_exchange_subscribe.go
  7. 4
      config/exchange.toml
  8. 6
      internal/exchange/exchange_grpc_server.go
  9. 7
      internal/sig/sig_server.go
  10. 54
      internal/trading/backtest/trade_account.go
  11. 1
      internal/trading/backtest/trade_simulator.go
  12. 50
      internal/trading/backtest/trading_plan_backtester.go
  13. 46
      internal/trading/backtest/trading_plan_input_racer.go
  14. 4
      internal/trading/backtest/types.go
  15. 2
      internal/trading/kline_series_store.go
  16. 7
      internal/trading/trading_data_persist.go
  17. 17
      internal/trading/trading_grpc_server.go
  18. 130
      internal/trading/trading_service.go
  19. 23
      pkg/strategy/super_trend_macd_rsi.go
  20. 4
      pkg/trade/sig_close_strategy.go
  21. 7
      pkg/trade/sig_trade_strategy.go
  22. 9
      pkg/trade/trade_account.go
  23. 3
      pkg/trade/types.go
  24. 19
      pkg/types/input.go
  25. 236
      pkg/types/input_gen.go
  26. 31
      pkg/types/input_test.go
  27. 30
      pkg/types/ring_series_test.go
  28. 24
      pkg/types/types_test.go
  29. 23
      pkg/utils/codec/codec_test.go
  30. 60
      pkg/utils/codec/mapstructure.go
  31. 99
      pkg/utils/collect/combine.go
  32. 17
      pkg/utils/collect/combine_test.go

2
README.md

@ -137,4 +137,4 @@ RSI[1,2,3,4] -> RSI[0]
因子挖掘 -> 策略挖掘
exchange_service.go: 100 task/一批, 批量成功后mark, 再发布下一波
viceAccount 负账户对冲交易
viceAccount 负账户对冲交易 {Long: mainAccount, Short: viceAccount}

20
api/exchange.proto

@ -26,19 +26,19 @@ service ExchangeService {
}
message ReqStreamSubscribeKline {
SubscribeType subType = 1;
SubscribeType sub_type = 1;
repeated ExchangeType exchanges = 2; //
repeated string instIds = 3;
repeated string inst_ids = 3;
repeated string intervals = 4;
bool onlyConfirm = 5; // confirm的k线
bool only_confirm = 5; // confirm的k线
}
message RspStreamSubscribeKline { StreamKline kline = 1; }
message StreamKline {
repeated Kline klines = 2;
ExchangeType exchange = 3; //
string instId = 4; // id
int64 streamId = 5; // subscribe stream id
string inst_id = 4; // id
int64 stream_id = 5; // subscribe stream id
}
message ReqExchanges {
@ -48,15 +48,15 @@ message RspExchanges {
}
message ReqExchangeInstanceState {
bool allExchange = 1; //
bool allInsts = 2; //
bool allStatus = 3; //
bool all_exchange = 1; //
bool all_insts = 2; //
bool all_status = 3; //
repeated ExchangeType exchanges = 4; //
repeated string insts = 5; //
repeated int32 status = 6; //
}
message RspExchangeInstanceState {
repeated TradeInstanceState instsState = 1; //
repeated TradeInstanceState insts_state = 1; //
}
message ReqSeriesRange {
@ -73,7 +73,7 @@ message ReqHistoryKline {
}
message RspHistoryKline {
ExchangeType exchange = 1; //
string instId = 2;
string inst_id = 2;
string interval = 3;
bool live = 4; // k线
// bool next = 5; // : true时, k线的ts作为before继续请求

12
api/market.proto

@ -16,7 +16,7 @@ service MarketService {
}
message ReqGetTradeInstance {
string instId = 1;
string inst_id = 1;
}
message RspGetTradeInstance {
TradeInstance inst = 1;
@ -37,7 +37,7 @@ message RspUpdateTradeInstance {
}
message ReqStatusTradeInstance {
string instId = 1;
string inst_id = 1;
}
message RspStatusTradeInstance {
TradeInstance inst = 1;
@ -45,20 +45,20 @@ message RspStatusTradeInstance {
message ReqListTradeInstance {
int32 page = 1;
int32 pageSize = 2;
int32 page_size = 2;
ExchangeType exchange = 3; //
string topic = 4;
string instId = 5;
string inst_id = 5;
}
message RspListTradeInstance {
repeated Kline klines = 1;
ExchangeType exchange = 2; //
string instId = 3; // id
string inst_id = 3; // id
}
message ReqListMarketTradeInstance {
ExchangeType exchange = 1;
}
message RspListMarketTradeInstance {
repeated MarketTradeInstance exchangeInsts = 1;
repeated MarketTradeInstance exchange_insts = 1;
}

65
api/pub.proto

@ -5,9 +5,9 @@ import "google/protobuf/struct.proto";
option go_package = "./pb";
enum SubscribeType {
Subscribe = 0;
Unsubscribe = 1;
UnsubscribeAll = 2;
SUB = 0;
UNSUB = 1;
UNSUB_ALL = 2;
}
enum ExchangeType {
@ -17,9 +17,9 @@ enum ExchangeType {
}
enum TradeInstanceType {
Unknow = 0;
Spot = 1; // 1
PerpetualContract = 2; // 2
UNKNOW = 0;
SPOT = 1; // 1
SWAP = 2; // 2
}
enum Event {
@ -42,7 +42,7 @@ enum Channel {
}
enum Side {
None = 0;
NONE = 0;
BUY = 1;
SELL = 2;
}
@ -63,47 +63,47 @@ message Error {
//
message TradeInstance {
string instId = 1;
string instPair = 2;
string instCoin = 3;
TradeInstanceType instType = 4;
string inst_id = 1;
string inst_pair = 2;
string inst_coin = 3;
TradeInstanceType inst_type = 4;
int32 status = 5;
int32 priceSz = 6;
int32 quantitySz = 7;
int32 price_sz = 6;
int32 quantity_sz = 7;
string icon = 8;
string updateBy = 9;
int64 updateTime = 10;
string update_by = 9;
int64 update_time = 10;
repeated int32 leverages = 11;
repeated TradeInstanceExchange exchanges = 15; //
}
//
message TradeInstanceExchange {
string exchangeInstId = 1;
string instId = 2;
string exchange_inst_id = 1;
string inst_id = 2;
int32 status = 3;
ExchangeType exchange = 4;
string updateBy = 5;
int64 updateTime = 6;
string update_by = 5;
int64 update_time = 6;
}
//
message MarketTradeInstance {
ExchangeType exchange = 1;
string instId = 2;
string exchangeInstId = 3;
TradeInstanceType instType = 4;
string instCoin = 5;
string inst_id = 2;
string exchange_inst_id = 3;
TradeInstanceType inst_type = 4;
string inst_coin = 5;
int32 status = 6;
int32 priceSz = 9;
int32 quantitySz = 10;
int32 price_sz = 9;
int32 quantity_sz = 10;
repeated int32 leverages = 11;
}
//
message TradeInstanceState {
ExchangeType exchange = 1;
string instId = 2;
string inst_id = 2;
int32 status = 3;
string last = 4; //
}
@ -116,7 +116,7 @@ message Kline {
double low = 6;
double close = 7;
double vol = 8; //
double volQuote = 9; //
double vol_quote = 9; //
bool confirm = 10; // k线是否完结
}
@ -139,7 +139,7 @@ message Order {
message SeriesRange {
ExchangeType exchange = 1; //
string instId = 2;
string inst_id = 2;
string interval = 3;
int64 before = 4;
int64 after = 5;
@ -148,7 +148,7 @@ message SeriesRange {
bool live = 8; // k线, before和after为0时是否追加实时k线
bool desc = 9; // ,
uint32 windowExtra = 10; // k线条数
uint32 window_extra = 10; // k线条数
uint32 limit = 11; // 0, limit则返回错误
}
@ -174,3 +174,10 @@ message IndicatorPlotExp {
string exp = 1;
google.protobuf.Struct props = 3;
}
//
message InputRange {
string name = 1; // names
int32 type = 2; // 0.fixed, 1.range, 2.enum, 3.simple group
string value = 3; // range => [10,20,1](min,max,step); enum => [1,2,3,4,5]; simple group => [["2025-10-01","2025-12-31"],["2025-01-01","2025-12-31"]]
}

73
api/trading.proto

@ -5,14 +5,33 @@ import "api/pub.proto";
option go_package = "./pb";
//
service TradingService {
rpc IndicatorPlots(ReqIndicatorPlots) returns (RspIndicatorPlots); //
rpc IndicatorSeries(ReqIndicatorSeries) returns (RspIndicatorSeries); //
rpc StrategySeries(ReqStrategySeries) returns (RspStrategySeries); //
rpc Backtest(ReqBacktest) returns (RspBacktest); //
rpc BacktestLog(ReqBacktestLog) returns (RspBacktestLog); //
rpc BacktestLogTrades(ReqBacktestLogTrades) returns (RspBacktestLogTrades); //
//
rpc IndicatorPlots(ReqIndicatorPlots) returns (RspIndicatorPlots);
//
rpc IndicatorSeries(ReqIndicatorSeries) returns (RspIndicatorSeries);
//
rpc StrategySeries(ReqStrategySeries) returns (RspStrategySeries);
//
rpc Backtest(ReqBacktest) returns (RspBacktest);
//
rpc BacktestRace(ReqBacktestRace) returns (RspBacktestRace);
//
rpc BacktestLog(ReqBacktestLog) returns (RspBacktestLog);
//
rpc BacktestLogTrades(ReqBacktestLogTrades) returns (RspBacktestLogTrades);
//
rpc BacktestLogStats(ReqBacktestLogStats) returns (RspBacktestLogStats);
}
message ReqIndicatorPlots {
repeated string indicators = 1;
}
@ -38,7 +57,7 @@ message IndicatorState {
message ReqStrategySeries {
SeriesRange series = 1;
string sigStrategy = 2;
string sig_strategy = 2;
google.protobuf.Struct input = 3; //
}
message RspStrategySeries {
@ -47,17 +66,17 @@ message RspStrategySeries {
}
message ReqBacktest {
int64 planId = 1;
int64 plan_id = 1;
string stime = 2;
string etime = 3;
}
message RspBacktest {
int64 backtestId = 1;
int64 backtest_id = 1;
}
message BacktestLog {
int64 backtestId = 1;
int64 planId = 2;
int64 backtest_id = 1;
int64 plan_id = 2;
int64 stime = 3;
int64 etime = 4;
string interval = 5; //
@ -73,11 +92,41 @@ message RspBacktestLog {
message BacktestTrade {
int64 id = 1;
int64 ctime = 2; //
Side side = 3; //
}
message ReqBacktestLogTrades {
Paging paging = 1;
int64 backtestId = 2;
int64 backtest_id = 2;
}
message RspBacktestLogTrades {
repeated BacktestTrade trades = 1;
}
message ReqBacktestRace {
int64 plan_id = 1; //
repeated InputRange series_input_range = 5; //
repeated InputRange sig_input_range = 6; //
repeated InputRange close_input_range = 7; //
repeated InputRange trade_input_range = 8; //
repeated InputRange risk_input_range = 9; //
}
message RspBacktestRace {
}
// (line charts)
message ReqBacktestLogStats {
int64 plan_id = 1;
int64 backtest_id = 2;
}
message RspBacktestLogStats {
repeated int64 times = 1; // k线时间
repeated double equitys = 2; //
repeated BacktestLogBuySellPoint buy_sell = 9; // charts data
}
// {time: 1763890200000, value: 86400, text: 'BUY:86400', direction: 'up'}
message BacktestLogBuySellPoint {
double value = 2; // k线收盘价
Side side = 3; //
string text = 4; //
}

4
cmd/trading/test_exchange_subscribe.go

@ -80,7 +80,7 @@ func subscribeEvents(client pb.ExchangeServiceClient) {
// 发送消息的goroutine
instIds := []string{"BTC-USDT", "DOGE-USDT-SWAP"}
msg := &pb.ReqStreamSubscribeKline{
SubType: pb.SubscribeType_Subscribe,
SubType: pb.SubscribeType_SUB,
Exchanges: []pb.ExchangeType{pb.ExchangeType_OKX},
InstIds: instIds,
Intervals: []string{
@ -98,7 +98,7 @@ func subscribeEvents(client pb.ExchangeServiceClient) {
go func() {
<-time.After(10 * time.Second)
msg := &pb.ReqStreamSubscribeKline{
SubType: pb.SubscribeType_Unsubscribe,
SubType: pb.SubscribeType_UNSUB,
Exchanges: []pb.ExchangeType{pb.ExchangeType_OKX},
InstIds: []string{"BTC-USDT"},
Intervals: []string{

4
config/exchange.toml

@ -18,9 +18,9 @@ receiveBuffer = 4096
marketSubscribeLimit = 16
consumeBatch = 1024
consumeLater = 2000 # 时间到达later或者数据累计到batch触发consume
# httpProxy = ""
httpProxy = ""
# httpProxy = "http://192.168.1.5:7890"
httpProxy = "http://10.255.183.209:7890"
# httpProxy = "http://10.255.183.209:7890"
# 模拟盘API交易地址如下:
# REST:https://www.okx.com

6
internal/exchange/exchange_grpc_server.go

@ -71,7 +71,7 @@ func (svr *ExchangeGrpcServer) SubscribeKline(stream grpc.BidiStreamingServer[pb
return
}
if msg.SubType == pb.SubscribeType_UnsubscribeAll {
if msg.SubType == pb.SubscribeType_UNSUB_ALL {
svr.klineSubscriber.UnsubscribeAll(streamId)
continue
}
@ -86,9 +86,9 @@ func (svr *ExchangeGrpcServer) SubscribeKline(stream grpc.BidiStreamingServer[pb
subKey := fmt.Sprintf("/kline/%s/%s/%s/%d", exchange.String(), instId, interval, confirm)
zlog.Debugf("stream: id=%d, sub %s", streamId, subKey)
switch msg.SubType {
case pb.SubscribeType_Subscribe:
case pb.SubscribeType_SUB:
svr.klineSubscriber.Subscribe(subKey, streamId, stream)
case pb.SubscribeType_Unsubscribe:
case pb.SubscribeType_UNSUB:
svr.klineSubscriber.Unsubscribe(subKey, streamId)
}
}

7
internal/sig/sig_server.go

@ -113,8 +113,11 @@ func (s *SigServer) handleGrpcGenericCall(c *gin.Context) {
return
}
// decode response
bytes, err := protojson.Marshal(rsp)
// encode response
bytes, err := protojson.MarshalOptions{
UseProtoNames: false, // false:lowerCamelCase, true:snake_case
EmitUnpopulated: false, // 是否包含默认值
}.Marshal(rsp)
if err != nil {
c.JSON(http.StatusInternalServerError, resp.Error(err.Error()))
return

54
internal/trading/backtest/trade_account.go

@ -39,34 +39,33 @@ func NewBacktestTradeAccount(cash float64) *BacktestTradeAccount {
}
}
// Equity 账户净值
// GetCash 账户可交易资金
func (a *BacktestTradeAccount) GetCash() float64 {
return a.cash
}
// OnPrice 交易产品价格更新
func (a *BacktestTradeAccount) OnPrice(instId string, price float64) {
if pos, ok := a.positions[instId]; ok {
pos.LastPx = price
}
}
// Equity 账户净值
func (a *BacktestTradeAccount) CurrentEquity() float64 {
equity := a.cash
for _, trades := range a.openTrades {
for _, tradeId := range trades {
if order, ok := a.orders[tradeId]; ok {
orderQty := decimals.MustToFloat64(order.Qty)
// 保证金
cost := order.Price * orderQty / float64(order.Leverage) // + order.Fee
// 浮盈
var pnl float64
if order.LastPx == 0 {
zlog.Warningf("order lastPx not updated")
} else {
if order.Side == types.SideLong {
pnl = (order.LastPx - order.Price) * orderQty
} else {
pnl = (order.Price - order.LastPx) * orderQty
}
}
equity += cost + pnl
}
for _, pos := range a.positions {
posQty := decimals.MustToFloat64(pos.Qty)
// 保证金
insure := pos.EntryPx * posQty / float64(pos.Leverage)
// 浮盈
var pnl float64
if pos.Side == types.SideLong {
pnl = (pos.LastPx - pos.EntryPx) * posQty
} else {
pnl = (pos.EntryPx - pos.LastPx) * posQty
}
equity += insure + pnl
}
return equity
}
@ -100,6 +99,11 @@ func (a *BacktestTradeAccount) GetOpenTrades(instId string) []*trade.TradeOrder
return trades
}
// 获取交易产品仓位
func (a *BacktestTradeAccount) GetPosition(instId string) *trade.Position {
return a.positions[instId]
}
// GetSymbolOpenPosition 获取指定交易对的仓位信息
func (a *BacktestTradeAccount) GetSymbolOpenPosition(symbol string) *trade.Position {
return a.positions[symbol]
@ -144,6 +148,7 @@ func (a *BacktestTradeAccount) MarketOrder(ticket trade.TradeTicket) (order *tra
EntryPx: order.Price,
EntryTs: order.Ctime,
PeakPx: order.Price,
LastPx: order.Price,
}
a.positions[instId] = pos
} else {
@ -152,6 +157,7 @@ func (a *BacktestTradeAccount) MarketOrder(ticket trade.TradeTicket) (order *tra
totalQty := decimals.MustAdd(pos.Qty, order.Qty)
pos.EntryPx = (pos.EntryPx*posQty + order.Price*orderQty) / decimals.MustToFloat64(totalQty)
pos.Qty = totalQty
pos.LastPx = order.Price
if pos.Side == types.SideLong && order.Price < pos.PeakPx {
pos.PeakPx = order.Price
}
@ -184,7 +190,6 @@ func (a *BacktestTradeAccount) CloseTradeOrder(ticket trade.TradeTicket) (err er
order := a.trader.ExecuteMarket(a.traderOrderId, instId, ticket)
a.traderOrderId++
// insure := order.Price * orderQty / float64(order.Leverage)
var insure, profit float64
var entryPxs, entryPxAvg, entryTrades float64
for _, tradeId := range ticket.TradesId {
@ -196,7 +201,11 @@ func (a *BacktestTradeAccount) CloseTradeOrder(ticket trade.TradeTicket) (err er
entryPxs += trade.Price
order.EntryFee += trade.Fee
order.EntryTime = lang.Ternary(order.EntryTime == 0, trade.Ctime, min(order.EntryTime, trade.Ctime))
order.PeakPx = trade.PeakPx
order.PeakPx = lang.Ternary(
trade.Side == types.SideLong,
max(order.PeakPx, trade.PeakPx),
min(order.PeakPx, trade.PeakPx),
)
tradeQty := decimals.MustToFloat64(trade.Qty)
// 计算盈利
@ -246,6 +255,7 @@ func (a *BacktestTradeAccount) CloseTradeOrder(ticket trade.TradeTicket) (err er
// 账户净值
currentEquity := a.CurrentEquity()
// 平仓单统计
order.LastPx = ticket.Price
order.EntryPx = entryPxAvg
order.Equity = currentEquity
order.HoldTime = conver.TimeDurationFormat(time.Duration(order.Ctime-order.EntryTime)*time.Millisecond, ".")

1
internal/trading/backtest/trade_simulator.go

@ -46,7 +46,6 @@ func (s *TradeSimulator) ExecuteMarket(tradeId int64, symbol string, ticket trad
Ctime: ticket.Ktime,
Status: data.StatusOk,
PeakPx: ticket.Price,
LastPx: ticket.Price,
CloseCause: ticket.Cause,
}
return

50
internal/trading/backtest/trading_plan_backtester.go

@ -2,7 +2,6 @@ package backtest
import (
"context"
"fmt"
"math"
"sig-pub/api/pb"
"sig-pub/internal/trading/sig"
@ -21,15 +20,14 @@ import (
// TradingPlanBacktester 交易计划回测
// todo 将backtest独立成单独服务横向扩展
type TradingPlanBacktester struct {
indicatorReg *indicator.IndicatorRegistry
sigStrategyReg *strategy.SigStrategyRegistry
exchangeClient pb.ExchangeServiceClient
sigStrategyType strategy.SigStrategyType
sigStrategy strategy.ISigStrategy
indicatorReg *indicator.IndicatorRegistry
exchangeClient pb.ExchangeServiceClient
plan entity.TradePlan
account *BacktestTradeAccount
// account *BacktestAccount
sigStrategyType strategy.SigStrategyType
sigStrategy strategy.ISigStrategy
plan entity.TradePlan
sr *pb.SeriesRange
account *BacktestTradeAccount
sigStrategyInput types.Input
tradeStrategyInput types.Input
@ -40,32 +38,26 @@ type TradingPlanBacktester struct {
}
func NewTradingPlanBacktester(
sigType strategy.SigStrategyType,
strategy strategy.ISigStrategy,
indicatorReg *indicator.IndicatorRegistry,
sigStrategyReg *strategy.SigStrategyRegistry,
exchangeClient pb.ExchangeServiceClient,
) *TradingPlanBacktester {
return &TradingPlanBacktester{
indicatorReg: indicatorReg,
sigStrategyReg: sigStrategyReg,
exchangeClient: exchangeClient,
sigStrategyType: sigType,
sigStrategy: strategy,
indicatorReg: indicatorReg,
exchangeClient: exchangeClient,
}
}
func (b *TradingPlanBacktester) Init(cash float64, plan entity.TradePlan) (err error) {
func (b *TradingPlanBacktester) Init(cash float64, plan entity.TradePlan, sr *pb.SeriesRange) (err error) {
b.plan = plan
b.sr = sr
// trade account
simulator := NewTradeSimulator(0.0005, 0.0008)
// b.account = NewBacktestAccount(cash, simulator)
_ = simulator
b.account = NewBacktestTradeAccount(cash)
// sig strategy
ok := false
b.sigStrategyType, b.sigStrategy, ok = b.sigStrategyReg.NewSigStrategy(plan.SigStrategy)
if !ok {
err = fmt.Errorf("sig strategy not exists %s", plan.SigStrategy)
return
}
// sig strategy param
if err = sonic.UnmarshalString(plan.SigStrategyParam, &b.sigStrategyInput); err != nil {
return
}
@ -102,7 +94,7 @@ func (b *TradingPlanBacktester) Init(cash float64, plan entity.TradePlan) (err e
}
// 核心引擎,模拟交易、持仓跟踪、费用计算
func (b *TradingPlanBacktester) Backtest(ctx context.Context, sr *pb.SeriesRange) (test *BacktestTradingPlan, err error) {
func (b *TradingPlanBacktester) Backtest(ctx context.Context) (test *BacktestTradingPlan, err error) {
test = &BacktestTradingPlan{
Id: time.Now().Unix(),
UserId: 10001,
@ -110,8 +102,8 @@ func (b *TradingPlanBacktester) Backtest(ctx context.Context, sr *pb.SeriesRange
InstId: b.plan.InstId,
Exchange: pb.ExchangeType(b.plan.Exchange),
Interval: b.plan.Interval,
SeriesBefore: sr.Before,
SeriesAfter: sr.After,
SeriesBefore: b.sr.Before,
SeriesAfter: b.sr.After,
Ctime: time.Now().UnixMilli(),
Cash: b.account.cash,
}
@ -126,7 +118,7 @@ func (b *TradingPlanBacktester) Backtest(ctx context.Context, sr *pb.SeriesRange
iiks := types.NewInstanceIntervalKlineSeries()
b.instanceIntervalSigStrategyContext = sig.NewInstanceIntervalSigStrategyContext(b.tradeStrategyInput, iiks, b.indicatorReg)
err = sigStrategyBacktester.Backtest(ctx, b.sigStrategyInput, sr, iiks, func(instId string, sigSide types.Side, k types.Kline) (err error) {
err = sigStrategyBacktester.Backtest(ctx, b.sigStrategyInput, b.sr, iiks, func(instId string, sigSide types.Side, k types.Kline) (err error) {
test.Singals++
// 根据交易信号检查仓位平仓
if err = b.closeBySigSingal(instId, sigSide, k); err != nil {
@ -198,9 +190,9 @@ func (b *TradingPlanBacktester) forceCloseAllHoldingPosition() (err error) {
for _, instId := range b.account.GetOpenTradeInsts() {
k := b.instanceIntervalSigStrategyContext.Get(instId, trade.PriceDriverInterval, 0)
price := decimals.MustToFloat64(k.Close)
b.account.OnPrice(instId, price)
trades := b.account.GetOpenTrades(instId)
for _, trd := range trades {
trd.LastPx = price
closeTickets = append(closeTickets, trade.TradeTicket{
TradeType: trade.TradeTypeClose,
InstId: instId,

46
internal/trading/backtest/trading_plan_input_racer.go

@ -0,0 +1,46 @@
package backtest
import (
"sig-pub/api/pb"
"sig-pub/pkg/indicator"
"sig-pub/pkg/strategy"
"sig-pub/pkg/types"
)
type TradingPlanInputRacer struct {
sigStrategyType strategy.SigStrategyType
sigStrategy strategy.ISigStrategy
indicatorReg *indicator.IndicatorRegistry
exchangeClient pb.ExchangeServiceClient
}
func NewTradingPlanInputRacer(
sigType strategy.SigStrategyType,
strategy strategy.ISigStrategy,
indicatorReg *indicator.IndicatorRegistry,
exchangeClient pb.ExchangeServiceClient,
) *TradingPlanInputRacer {
return &TradingPlanInputRacer{
sigStrategyType: sigType,
sigStrategy: strategy,
indicatorReg: indicatorReg,
exchangeClient: exchangeClient,
}
}
func (r *TradingPlanInputRacer) Init() (err error) {
a := types.InputRange{}
_ = a
// plan := &entity.TradePlan{}
// plan.SigStrategy
// SigStrategyParam
// CloseStrategyParam
// TradeStrategyParam
// RiskStrategyParam
// tester := NewTradingPlanBacktester(1, nil, nil, nil)
// tester.Init(10000, *plan)
// result, err := tester.Backtest(nil, nil)
// _, _ = result, err
return
}

4
internal/trading/backtest/types.go

@ -103,7 +103,7 @@ type BacktestTradingPlan struct {
LosingTrades int `json:"losingTrades" gorm:"column:losing_trades"` // 亏损单数
Fee float64 `json:"fee" gorm:"column:fee"` // 总手续费
MaxDrawdown float64 `json:"maxDrawdown" gorm:"column:max_drawdown"` // 最大回撤
SharpeRatio float64 `json:"sharpeRatio" gorm:"column:sharpe_ratio"` // 最大回撤
SharpeRatio float64 `json:"sharpeRatio" gorm:"column:sharpe_ratio"` // 夏普比率
Trades []*trade.TradeOrder `json:"-" gorm:"-"` // 回测交易单
}
@ -135,3 +135,5 @@ type Trade struct {
func (Trade) TableName() string {
return "t_backtest_trading_trade"
}
// 参数迭代

2
internal/trading/kline_series_store.go

@ -159,7 +159,7 @@ func (s *KlineSeriesStore) sendSubscribeKline(save bool, exchange pb.ExchangeTyp
// 发送订阅消息
subMsg := &pb.ReqStreamSubscribeKline{
SubType: pb.SubscribeType_Subscribe,
SubType: pb.SubscribeType_SUB,
Exchanges: []pb.ExchangeType{exchange},
InstIds: instIds,
Intervals: s.subKlineIntervals,

7
internal/trading/trading_data_persist.go

@ -66,3 +66,10 @@ func (p *TradingDataPersist) ListBacktestLogs(userId int64) (backtestLogs []*bac
`, userId)
return
}
func (p *TradingDataPersist) ListBacktestTrades(backtestId int64) (trades []*backtest.Trade, err error) {
err = p.db.Select(&trades, `
select * from t_backtest_trading_trade where backtest_id = ? order by id asc
`, backtestId)
return
}

17
internal/trading/trading_grpc_server.go

@ -117,3 +117,20 @@ func (svr *TradingGrpcServer) BacktestLog(ctx context.Context, req *pb.ReqBackte
}
return
}
// BacktestRace 交易计划参数调试回测
func (svr *TradingGrpcServer) BacktestRace(ctx context.Context, req *pb.ReqBacktestRace) (rsp *pb.RspBacktestRace, err error) {
rsp = new(pb.RspBacktestRace)
err = svr.tradingService.BacktestRace(ctx, req)
return
}
func (svr *TradingGrpcServer) BacktestLogTrades(ctx context.Context, req *pb.ReqBacktestLogTrades) (rsp *pb.RspBacktestLogTrades, err error) {
return
}
func (svr *TradingGrpcServer) BacktestLogStats(ctx context.Context, req *pb.ReqBacktestLogStats) (rsp *pb.RspBacktestLogStats, err error) {
return
}

130
internal/trading/trading_service.go

@ -340,13 +340,17 @@ func (svc *TradingService) Backtest(ctx context.Context, planId, stime, etime in
Live: false,
Desc: false,
}
tester := backtest.NewTradingPlanBacktester(svc.indicatorReg, svc.strategyReg, svc.exchangeClient)
if err = tester.Init(10000, *plan); err != nil {
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, sr)
backtestTradingPlan, err := tester.Backtest(ctx)
if err != nil {
return
}
@ -363,3 +367,121 @@ func (svc *TradingService) BacktestLog(ctx context.Context, userId int64) (backt
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 []*backtest.BacktestTradingPlan
for _, backtester := range backtesters {
r, e := backtester.Backtest(ctx)
if e != nil {
err = e
return
}
results = append(results, r)
}
for _, r := range results {
zlog.Info(r.Id, r.Cash, r.EndCash, r.Profit, r.Singals, r.TotalTrades, r.WinningTrades, r.LosingTrades, 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
}

23
pkg/strategy/super_trend_macd_rsi.go

@ -31,9 +31,10 @@ func (s *SuperTrendMacdRSI) CandlePeriods(ctx ISingleSigStrategyContext) int16 {
return max(
ctx.Indicator("SuperTrend", types.Input{"window": 10, "mul": 3}).CandlePeriods(),
ctx.Indicator("RSI", 14).CandlePeriods(),
ctx.Indicator("MacdDEA", types.Input{"fast": 12, "slow": 26, "singal": 9}).CandlePeriods(),
ctx.Indicator("MacdDEA", types.Input{"fast": 12, "slow": 26, "singal": 9}).CandlePeriods(),
ctx.Indicator("MacdDIF", types.Input{"fast": 12, "slow": 26, "singal": 9}).CandlePeriods(),
ctx.Indicator("MACD", types.Input{"fast": 12, "slow": 26, "singal": 9}).CandlePeriods(),
// ctx.Indicator("MacdDEA", types.Input{"fast": 12, "slow": 26, "singal": 9}).CandlePeriods(),
// ctx.Indicator("MacdDEA", types.Input{"fast": 12, "slow": 26, "singal": 9}).CandlePeriods(),
// ctx.Indicator("MacdDIF", types.Input{"fast": 12, "slow": 26, "singal": 9}).CandlePeriods(),
21,
)
}
@ -49,8 +50,8 @@ func (s *SuperTrendMacdRSI) Update(ctx ISingleSigStrategyContext) (side types.Si
// macdDea := ctx.Indicator("MacdDEA", types.Input{"fast": 12, "slow": 26, "singal": 9}).Series(0, 2) // macd_dea信号线
// macdDif := ctx.Indicator("MacdDIF", types.Input{"fast": 12, "slow": 26, "singal": 9}).Series(0, 2) // macd_dif线
crossover := macdDif[0] > macdDea[0] && macdDif[1] < macdDea[1] // 金叉
// crossunder := macdDif[0] < macdDea[0] && macdDif[1] > macdDea[1] // 死叉
crossover := macdDif[0] > macdDea[0] && macdDif[1] < macdDea[1] // 金叉
crossunder := macdDif[0] < macdDea[0] && macdDif[1] > macdDea[1] // 死叉
closeP := ctx.Get(0).CloseF64()
volAvg := ctx.Series(1, 20).Vol().Avg()
@ -73,6 +74,18 @@ func (s *SuperTrendMacdRSI) Update(ctx ISingleSigStrategyContext) (side types.Si
}
}
if crossunder && macdHist < 0 {
if rsi < 50 {
// SuperTrend 趋势确认
if closeP < trend && trendDirection == -1 {
// 成交量过滤
if vol > volAvg*1.5 {
return types.SideShort
}
}
}
}
_ = `
// 策略算子脚本DST, 优化golang底层不影响策略语法
st1 = sig.SuperTrend(window=10, mul=3)

4
pkg/trade/sig_close_strategy.go

@ -17,6 +17,7 @@ func (s *SigTradeStrategy) CloseAssessOnPrice(ctx strategy.IInstanceIntervalSigS
for _, trd := range openTrades {
k := ctx.Get(trd.InstId, PriceDriverInterval, 0)
price := decimals.MustToFloat64(k.Close)
account.OnPrice(instId, price)
closeTrade, cause := s.closeTradeOnPrice(price, trd)
if closeTrade {
closeTickets = append(closeTickets, TradeTicket{
@ -41,7 +42,6 @@ func (s *SigTradeStrategy) closeTradeOnPrice(price float64, trd *TradeOrder) (cl
if !trd.Side.IsValid() {
return
}
trd.LastPx = price
// update peak px
if trd.Side == types.SideLong && (price > trd.PeakPx) {
trd.PeakPx = price
@ -131,8 +131,10 @@ func (s *SigTradeStrategy) CloseAssessOnSig(ctx strategy.IInstanceIntervalSigStr
if trade.Side == oppositeSide {
k := ctx.Get(trade.InstId, PriceDriverInterval, 0)
price := decimals.MustToFloat64(k.Close)
account.OnPrice(instId, price)
closeTickets = append(closeTickets, TradeTicket{
TradeType: TradeTypeClose,
InstId: instId,
TradesId: []int64{trade.TradeId},
Side: trade.Side.Opposite(),
Price: price,

7
pkg/trade/sig_trade_strategy.go

@ -74,6 +74,13 @@ func (s *SigTradeStrategy) RishAssess(ctx strategy.IInstanceIntervalSigStrategyC
// 控制滑点, 仓位管理
// 持仓中币种不能改变杠杆
func (s *SigTradeStrategy) TradeAssess(ctx strategy.IInstanceIntervalSigStrategyContext, account ITradeAccount, sigInstId string, sigSide types.Side) (tickets []TradeTicket, err error) {
if pos := account.GetPosition(sigInstId); pos != nil {
// 已持仓不能下反方向单, todo 副账户做反方向单,对冲(viceAccount)
if pos.Side != sigSide {
return
}
}
k := ctx.Get(sigInstId, PriceDriverInterval, 0)
price := decimals.MustToFloat64(k.Close)
ta := TradeTicket{

9
pkg/trade/trade_account.go

@ -4,9 +4,14 @@ package trade
// sig -> risk strategy -> trade strategy -> tarde account
// position, trades
type ITradeAccount interface {
// 获取未平仓交易单
// 副账户
// ViceAccount() ITradeAccount
// OnPrice 交易产品价格更新
OnPrice(instId string, price float64)
// 获取交易产品未平仓交易单
GetOpenTrades(instId string) []*TradeOrder
// 获取交易产品仓位
GetPosition(instId string) *Position
// MarketOrder(symbol string, ticket TradeTicket) (order *TradeOrder, err error)
// CloseTradeOrder(order *TradeOrder) (err error)
// 获取未平仓交易单数

3
pkg/trade/types.go

@ -51,6 +51,7 @@ type Position struct {
EntryPx float64 // 入场价格(均价)
EntryTs int64 // 入场时间
PeakPx float64 // 持仓最高价格(空单最低价格)
LastPx float64 // 最后更新价格
}
// TradeOrder 交易订单
@ -67,12 +68,12 @@ type TradeOrder struct {
Ctime int64 `jsno:"ctime" gorm:"column:ctime"` // 交易时间
Status data.Status `jsno:"status" gorm:"column:status"` // 1.交易成功 2.交易中 4.交易失败
PeakPx float64 `jsno:"peakPx" gorm:"column:peak_px"` // 持仓最高价格(空单最低价格)
LastPx float64 `jsno:"lastPx" gorm:"column:last_px"` // 最后更新价格
// ----------------- 平仓单信息
CloseCause Cause `json:"closeCause" gorm:"column:close_cause"` // 平仓原因 ["stoploss", "takeprofit", "trailing", "retrace", "signal"](“止损”、“止盈”、“动态跟踪”、“回撤”、“信号”)
Equity float64 `json:"equity" gorm:"column:equity"` // 平仓后账户净值
HoldTime string `json:"holdTime" gorm:"column:hold_time"` // 持仓时间
Profit float64 `json:"profit" gorm:"column:profit"` // 利润
LastPx float64 `jsno:"lastPx" gorm:"column:last_px"` // 最后更新价格
EntryPx float64 `json:"entryPx" gorm:"column:entry_px"` // 开仓价格(均价)
EntryFee float64 `json:"entryFee" gorm:"column:entry_fee"` // 开仓手续费(总计)
EntryTime int64 `json:"entryTime" gorm:"column:entry_time"` // 开仓时间(最早)

19
pkg/types/input.go

@ -3,8 +3,9 @@ package types
import (
"fmt"
"maps"
"sig-pub/pkg/utils/codec"
"time"
"github.com/go-viper/mapstructure/v2"
"github.com/spf13/cast"
)
@ -100,25 +101,34 @@ func (in Input) String(k string) (v string) {
return
}
func (in Input) Time(k string) (v time.Time) {
v, err := cast.ToTimeInDefaultLocationE(in.get(k, "time"), time.Local)
if err != nil {
panic(fmt.Errorf("input time parse error: %s", k))
}
return
}
func (in Input) Decode(k string, point any) {
t := fmt.Sprintf("%T", point)
v := in.get(k, t)
if err := mapstructure.Decode(v, point); err != nil {
if err := codec.MapDecode(v, point); err != nil {
panic(fmt.Errorf("input %s decode error: %s, %v", t, k, err))
}
}
func (in Input) DecodeInput(point any) {
if err := mapstructure.Decode(in, point); err != nil {
if err := codec.MapDecode(in, point); err != nil {
panic(fmt.Errorf("input decode to %T error: %v", point, err))
}
}
// Assign 将other中的值赋给当前Input
func (in Input) Assign(others ...Input) {
func (in Input) Assign(others ...Input) Input {
for _, other := range others {
maps.Copy(in, other)
}
return in
}
// 参数类型
@ -132,6 +142,7 @@ const (
InputTypeUFloat
InputTypeInt
InputTypeUInt
InputTypeTime
InputTypeUFloats // float数组
InputTypeUFloats2D // float二维数组
InputTypeSelect // 单选

236
pkg/types/input_gen.go

@ -0,0 +1,236 @@
package types
import (
"errors"
"fmt"
"strconv"
"strings"
"github.com/bytedance/sonic"
)
// InputGroup 参数输入组合, 梯度下降
// type InputGroup struct {
// Inputs []InputRange `json:"inputs"`
// }
// InputRange 参数输入范围
// SuperTrend fast 1 [10,20,1]
// SuperTrend slow 1 [10,20,1]
// SuperTrend singal 1 [10,20,1]
type InputRange struct {
Name string `json:"name"` // 参数名 names
Type int32 `json:"type"` // 1.range, 2.enum, 3.simple group
Value string `json:"value"` // range => [10,20,1](min,max,step); enum => [1,2,3,4,5]; simple group => [[2025-10-01,2025-12-31],[2025-01-01,2025-12-31]]
}
func (ir InputRange) NewInputGen() (gen InputGen, err error) {
switch ir.Type {
case 0:
gen = new(inputFixedGen)
case 1:
gen = new(inputStepGen)
case 2:
gen = new(inputEnumGen)
case 3:
gen = new(simpleInputGroupGen)
}
if gen != nil {
err = gen.init(ir.Name, ir.Value)
}
return
}
type InputGen interface {
init(name, value string) (err error)
Next() (in Input, ok bool) // 生成下一个值
Values() []Input
}
// inputFixedGen 单个固定值
type inputFixedGen struct {
name, r string
got bool
}
func (c *inputFixedGen) init(name, r string) (err error) {
c.name = name
c.r = r
return
}
func (c *inputFixedGen) Next() (in Input, ok bool) {
if c.got {
return
}
return Input{c.name: c.r}, true
}
func (c *inputFixedGen) Values() (vs []Input) {
in, ok := c.Next()
if !ok {
return
}
vs = append(vs, in)
return
}
// inputStepGen 数值范围生成器 "min,max,step" -> "0,10,2"
type inputStepGen struct {
name string
genType int8 // 1.float64 2.int64
fmin, fmax, fstep float64
imin, imax, istep int64
}
func (c *inputStepGen) init(name, r string) (err error) {
c.name = name
r = strings.TrimSuffix(strings.TrimPrefix(r, "["), "]")
vs := strings.Split(r, ",")
if len(vs) != 3 {
return fmt.Errorf("step generator format error")
}
min, max, step := vs[0], vs[1], vs[2]
// any float use float
if strings.Contains(min, ".") || strings.Contains(max, ".") || strings.Contains(step, ".") {
c.genType = 1
if c.fmin, err = strconv.ParseFloat(min, 64); err != nil {
return
}
if c.fmax, err = strconv.ParseFloat(max, 64); err != nil {
return
}
if c.fstep, err = strconv.ParseFloat(step, 64); err != nil {
return
}
} else {
c.genType = 2
if c.imin, err = strconv.ParseInt(min, 10, 64); err != nil {
return
}
if c.imax, err = strconv.ParseInt(max, 10, 64); err != nil {
return
}
if c.istep, err = strconv.ParseInt(step, 10, 64); err != nil {
return
}
}
return
}
// Next 生成下一个值
func (c *inputStepGen) Next() (in Input, ok bool) {
switch c.genType {
case 1:
if c.fmin <= c.fmax {
v := c.fmin
c.fmin += c.fstep
return Input{c.name: v}, true
}
case 2:
if c.imin <= c.imax {
v := c.imin
c.imin += c.istep
return Input{c.name: v}, true
}
}
return
}
func (c *inputStepGen) Values() (vs []Input) {
for {
v, ok := c.Next()
if !ok {
break
}
vs = append(vs, v)
}
return
}
type inputEnumGen struct {
name string
enums []string
index int
}
func (c *inputEnumGen) init(name, r string) (err error) {
c.name = name
if r == "" {
return fmt.Errorf("enum value empty")
}
r = strings.TrimSuffix(strings.TrimPrefix(r, "["), "]")
c.enums = strings.Split(r, ",")
if len(c.enums) == 0 {
return
}
return
}
// Next 生成下一个值
func (c *inputEnumGen) Next() (in Input, ok bool) {
if c.index >= len(c.enums) {
return
}
v := c.enums[c.index]
c.index++
return Input{c.name: v}, true
}
func (c *inputEnumGen) Values() (vs []Input) {
vs = make([]Input, 0, len(c.enums))
for _, v := range c.enums {
vs = append(vs, Input{c.name: v})
}
return
}
type simpleInputGroupGen struct {
names []string
enums [][]any
index int
}
func (c *simpleInputGroupGen) init(name, r string) (err error) {
c.names = strings.Split(name, ",")
if r == "" {
return errors.New("empty values")
}
if err = sonic.UnmarshalString(r, &c.enums); err != nil {
return
}
if len(c.enums) == 0 {
return errors.New("empty values")
}
for i, enums := range c.enums {
if len(enums) != len(c.names) {
err = fmt.Errorf("simple input group enums %d, length %d not match names %#v", i, len(enums), c.names)
}
}
return
}
// Next 生成下一个值
func (c *simpleInputGroupGen) Next() (in Input, ok bool) {
if c.index >= len(c.enums) {
return
}
enums := c.enums[c.index]
c.index++
in = make(Input, len(c.names))
for i, name := range c.names {
in[name] = enums[i]
}
return in, true
}
func (c *simpleInputGroupGen) Values() (vs []Input) {
for {
v, ok := c.Next()
if !ok {
break
}
vs = append(vs, v)
}
return
}

31
pkg/types/input_test.go

@ -1,6 +1,7 @@
package types
import (
"fmt"
"testing"
"github.com/bytedance/sonic"
@ -37,3 +38,33 @@ func TestInput(t *testing.T) {
in.DecodeInput(csp)
t.Logf("%#v", csp)
}
func TestInputRange(t *testing.T) {
r1 := &InputRange{Name: "fast", Type: 1, Value: "3,30,2"}
g1, err := r1.NewInputGen()
if err != nil {
panic(err)
}
for {
v, ok := g1.Next()
if !ok {
break
}
t.Log(v)
}
g11, err := r1.NewInputGen()
if err != nil {
panic(err)
}
t.Log(g11.Values())
r3 := &InputRange{Name: "stime,etime", Type: 3, Value: `[["2025-10-01","2025-12-31"],["2025-01-01","2025-12-31"]]`}
g3, err := r3.NewInputGen()
if err != nil {
panic(err)
}
vs := g3.Values()
fmt.Println(vs)
stime := vs[0].Time("stime")
fmt.Println(stime)
}

30
pkg/types/ring_series_test.go

@ -1,30 +0,0 @@
package types
import (
"fmt"
"testing"
)
func TestRingSeries(t *testing.T) {
rs1 := NewRingSeries[int](1, 1)
rs1.Push(1)
r1, ok := rs1.Get(0)
if !ok {
t.Error(ok)
return
}
if r1 != 1 {
t.Error(r1)
return
}
rs2 := NewRingSeries[int](10, 3)
for i := range 11 {
rs2.Push(i)
}
for i := range 10 {
fmt.Println(rs2.Get(i))
}
fmt.Println("--------------------")
fmt.Println(rs2.Series(0, 1))
}

24
pkg/types/kline_series_test.go → pkg/types/types_test.go

@ -7,6 +7,30 @@ import (
"testing"
)
func TestRingSeries(t *testing.T) {
rs1 := NewRingSeries[int](1, 1)
rs1.Push(1)
r1, ok := rs1.Get(0)
if !ok {
t.Error(ok)
return
}
if r1 != 1 {
t.Error(r1)
return
}
rs2 := NewRingSeries[int](10, 3)
for i := range 11 {
rs2.Push(i)
}
for i := range 10 {
fmt.Println(rs2.Get(i))
}
fmt.Println("--------------------")
fmt.Println(rs2.Series(0, 1))
}
func TestKlineSeries(t *testing.T) {
ks := NewKlineSeries(pb.ExchangeType_SIG, "TEST_USDT", Interval5m)
intervalAdder := SupportedIntervals[Interval5m]

23
pkg/utils/codec/codec_test.go

@ -0,0 +1,23 @@
package codec
import (
"fmt"
"testing"
)
type Config struct {
Fm float64
}
func TestMapstructure(t *testing.T) {
// var arr1 [][]float64
// arr2T := reflect.SliceOf(reflect.SliceOf(reflect.TypeOf(1.0)))
// fmt.Println(reflect.TypeOf(arr1) == arr2T, arr2T.Kind())
cfg := &Config{}
err := MapDecode(map[string]any{"fm": "3.1415"}, cfg)
if err != nil {
panic(err)
}
fmt.Printf("%#v\n", cfg)
}

60
pkg/utils/codec/mapstructure.go

@ -0,0 +1,60 @@
package codec
import (
"fmt"
"reflect"
"strconv"
"strings"
"github.com/bytedance/sonic"
"github.com/go-viper/mapstructure/v2"
)
func MapDecode(input, output any) (err error) {
decoder, err := mapstructure.NewDecoder(&mapstructure.DecoderConfig{
Result: output,
WeaklyTypedInput: true, // 开启弱类型转换
DecodeHook: mapstructure.ComposeDecodeHookFunc(StringToNumberHook()),
})
if err != nil {
return
}
err = decoder.Decode(input)
return
}
// 方案1:最推荐 - 字符串 → 数字(int/uint/float)全覆盖
func StringToNumberHook() mapstructure.DecodeHookFunc {
return mapstructure.DecodeHookFuncType(func(from reflect.Type, to reflect.Type, data interface{}) (interface{}, error) {
fmt.Println("hook1")
if from.Kind() == reflect.String {
str := data.(string)
str = strings.TrimSpace(str)
switch to.Kind() {
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
return strconv.ParseInt(str, 10, 64)
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
return strconv.ParseUint(str, 10, 64)
case reflect.Float32, reflect.Float64:
return strconv.ParseFloat(str, 64)
}
// string to []float64, [][]float64
if to.Kind() == reflect.Slice {
switch to {
case reflect.SliceOf(reflect.TypeOf(float64(0.0))):
var v []float64
if err := sonic.UnmarshalString(str, &v); err == nil {
return v, nil
}
case reflect.SliceOf(reflect.SliceOf(reflect.TypeOf(float64(0.0)))):
var v [][]float64
if err := sonic.UnmarshalString(str, &v); err == nil {
return v, nil
}
}
}
}
return data, nil
})
}

99
pkg/utils/collect/combine.go

@ -0,0 +1,99 @@
package collect
// CartesianCount 笛卡尔积组合数
func CartesianCount[T any](sets ...[]T) (total int) {
total = 1
for _, set := range sets {
total *= len(set)
}
return
}
// CartesianProduct 笛卡尔积(返回所有可能的组合)
// 输入: [][]T 类型的二维切片
// 输出: []T 类型的切片组成的切片
func CartesianProduct[T any](sets ...[]T) [][]T {
if len(sets) == 0 {
return [][]T{{}}
}
// 初始化为第一个集合的所有单元素组合
result := make([][]T, len(sets[0]))
for i, v := range sets[0] {
result[i] = []T{v}
}
// 依次处理后续每一组
for _, set := range sets[1:] {
if len(set) == 0 {
return [][]T{}
}
newResult := make([][]T, 0, len(result)*len(set))
for _, prev := range result {
for _, curr := range set {
// 预分配空间,避免频繁扩容
combined := make([]T, 0, len(prev)+1)
combined = append(combined, prev...)
combined = append(combined, curr)
newResult = append(newResult, combined)
}
}
result = newResult
}
return result
}
// CartesianYield 生成笛卡尔积的所有组合,并为每个组合调用回调函数
// callback: func(comb []T) bool - 处理当前组合,返回 true 继续生成,返回 false 停止
// return: 生成的组合总数
func CartesianYield[T any](sets [][]T, callback func([]T) bool) int {
if len(sets) == 0 {
return 0
}
// 检查是否有空集
for _, s := range sets {
if len(s) == 0 {
return 0
}
}
// 初始化索引计数器
indices := make([]int, len(sets))
comb := make([]T, len(sets)) // 复用缓冲区,避免每次分配
count := 0
for {
// 构建当前组合
for i, idx := range indices {
comb[i] = sets[i][idx]
}
// 调用回调
count++
if !callback(comb) {
break // 停止生成
}
// 递增索引,像计数器一样
i := len(indices) - 1
for i >= 0 {
indices[i]++
if indices[i] < len(sets[i]) {
break
}
indices[i] = 0
i--
}
if i < 0 {
break // 所有组合已生成
}
}
return count
}

17
pkg/utils/collect/combine_test.go

@ -0,0 +1,17 @@
package collect
import (
"fmt"
"testing"
)
func TestCartesianYield(t *testing.T) {
arrs := [][]int{
{1, 2},
{3, 4, 5},
}
CartesianYield(arrs, func(t []int) bool {
fmt.Println(t)
return true
})
}
Loading…
Cancel
Save