diff --git a/README.md b/README.md index a8d89a8..b372950 100644 --- a/README.md +++ b/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} diff --git a/api/exchange.proto b/api/exchange.proto index 6b445bd..e5f607b 100644 --- a/api/exchange.proto +++ b/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继续请求 diff --git a/api/market.proto b/api/market.proto index 1edb347..bc459a3 100644 --- a/api/market.proto +++ b/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; } diff --git a/api/pub.proto b/api/pub.proto index bd79a17..702af13 100644 --- a/api/pub.proto +++ b/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"]] +} diff --git a/api/trading.proto b/api/trading.proto index 7a6ece4..98015f3 100644 --- a/api/trading.proto +++ b/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; // 买卖点描述 +} diff --git a/cmd/trading/test_exchange_subscribe.go b/cmd/trading/test_exchange_subscribe.go index 64a7993..a8188c7 100644 --- a/cmd/trading/test_exchange_subscribe.go +++ b/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{ diff --git a/config/exchange.toml b/config/exchange.toml index 823264d..e65fb93 100644 --- a/config/exchange.toml +++ b/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 diff --git a/internal/exchange/exchange_grpc_server.go b/internal/exchange/exchange_grpc_server.go index 44af003..60d808a 100644 --- a/internal/exchange/exchange_grpc_server.go +++ b/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) } } diff --git a/internal/sig/sig_server.go b/internal/sig/sig_server.go index 1fd7e05..0f00224 100644 --- a/internal/sig/sig_server.go +++ b/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 diff --git a/internal/trading/backtest/trade_account.go b/internal/trading/backtest/trade_account.go index 3036008..7353822 100644 --- a/internal/trading/backtest/trade_account.go +++ b/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, ".") diff --git a/internal/trading/backtest/trade_simulator.go b/internal/trading/backtest/trade_simulator.go index 40f7b6f..43226f4 100644 --- a/internal/trading/backtest/trade_simulator.go +++ b/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 diff --git a/internal/trading/backtest/trading_plan_backtester.go b/internal/trading/backtest/trading_plan_backtester.go index 7012381..f7b21e1 100644 --- a/internal/trading/backtest/trading_plan_backtester.go +++ b/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, diff --git a/internal/trading/backtest/trading_plan_input_racer.go b/internal/trading/backtest/trading_plan_input_racer.go new file mode 100644 index 0000000..d604f5b --- /dev/null +++ b/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 +} diff --git a/internal/trading/backtest/types.go b/internal/trading/backtest/types.go index 6c88e2f..6739d5a 100644 --- a/internal/trading/backtest/types.go +++ b/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" } + +// 参数迭代 diff --git a/internal/trading/kline_series_store.go b/internal/trading/kline_series_store.go index 4110901..d169369 100644 --- a/internal/trading/kline_series_store.go +++ b/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, diff --git a/internal/trading/trading_data_persist.go b/internal/trading/trading_data_persist.go index c8fd69c..a9994f1 100644 --- a/internal/trading/trading_data_persist.go +++ b/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 +} diff --git a/internal/trading/trading_grpc_server.go b/internal/trading/trading_grpc_server.go index ee91b23..04fbf69 100644 --- a/internal/trading/trading_grpc_server.go +++ b/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 +} diff --git a/internal/trading/trading_service.go b/internal/trading/trading_service.go index f86f641..3f34113 100644 --- a/internal/trading/trading_service.go +++ b/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 +} diff --git a/pkg/strategy/super_trend_macd_rsi.go b/pkg/strategy/super_trend_macd_rsi.go index 1cd07bf..ac603ab 100644 --- a/pkg/strategy/super_trend_macd_rsi.go +++ b/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) diff --git a/pkg/trade/sig_close_strategy.go b/pkg/trade/sig_close_strategy.go index aab8b86..2884841 100644 --- a/pkg/trade/sig_close_strategy.go +++ b/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, diff --git a/pkg/trade/sig_trade_strategy.go b/pkg/trade/sig_trade_strategy.go index 80d9c22..9e6dedb 100644 --- a/pkg/trade/sig_trade_strategy.go +++ b/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{ diff --git a/pkg/trade/trade_account.go b/pkg/trade/trade_account.go index 6c26896..94631cf 100644 --- a/pkg/trade/trade_account.go +++ b/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) // 获取未平仓交易单数 diff --git a/pkg/trade/types.go b/pkg/trade/types.go index d9d9d8d..525914b 100644 --- a/pkg/trade/types.go +++ b/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"` // 开仓时间(最早) diff --git a/pkg/types/input.go b/pkg/types/input.go index ff8a5bc..f136a9b 100644 --- a/pkg/types/input.go +++ b/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 // 单选 diff --git a/pkg/types/input_gen.go b/pkg/types/input_gen.go new file mode 100644 index 0000000..3327fd8 --- /dev/null +++ b/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 +} diff --git a/pkg/types/input_test.go b/pkg/types/input_test.go index 24d503e..ea33732 100644 --- a/pkg/types/input_test.go +++ b/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) +} diff --git a/pkg/types/ring_series_test.go b/pkg/types/ring_series_test.go deleted file mode 100644 index 4c328b7..0000000 --- a/pkg/types/ring_series_test.go +++ /dev/null @@ -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)) -} diff --git a/pkg/types/kline_series_test.go b/pkg/types/types_test.go similarity index 60% rename from pkg/types/kline_series_test.go rename to pkg/types/types_test.go index 381cf6f..244e468 100644 --- a/pkg/types/kline_series_test.go +++ b/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] diff --git a/pkg/utils/codec/codec_test.go b/pkg/utils/codec/codec_test.go new file mode 100644 index 0000000..e1d314d --- /dev/null +++ b/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) +} diff --git a/pkg/utils/codec/mapstructure.go b/pkg/utils/codec/mapstructure.go new file mode 100644 index 0000000..fd6c0f1 --- /dev/null +++ b/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 + }) +} diff --git a/pkg/utils/collect/combine.go b/pkg/utils/collect/combine.go new file mode 100644 index 0000000..f868014 --- /dev/null +++ b/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 +} diff --git a/pkg/utils/collect/combine_test.go b/pkg/utils/collect/combine_test.go new file mode 100644 index 0000000..0ea5935 --- /dev/null +++ b/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 + }) +}