package repository import ( "fmt" "math" "sig-pub/pkg/data" "sig-pub/pkg/storage/persist" "sig-pub/pkg/trade" "sig-pub/pkg/utils/conver" ) type BacktestRepository struct { db *persist.DB } func NewBacktestRepository(db *persist.DB) *BacktestRepository { return &BacktestRepository{ db: db, } } // GetBacktest 用户交易计划回测详情 func (p *BacktestRepository) GetBacktest(userId int64, backtestId int64) (test *trade.BacktestTradingPlan, err error) { test = new(trade.BacktestTradingPlan) err = p.db.Select(&test, ` select * from t_backtest_trading_plan where id = ? and user_id = ? `, backtestId, userId) return } // ListBacktest 用户交易计划回测记录查询 func (p *BacktestRepository) ListBacktest(userId int64, page data.Page) (total int64, backtestLogs []*trade.BacktestTradingPlan, err error) { sqlFrom := ` from t_backtest_trading_plan where user_id = ? ` err = p.db.Select(&total, fmt.Sprintf(` select count(*) %s `, sqlFrom), userId) if err != nil || total == 0 { return } err = p.db.Select(&backtestLogs, fmt.Sprintf(` select * %s order by id desc offset ? limit ? `, sqlFrom), userId, page.Offset, page.Limit) return } // ListBacktestTrades 交易计划回测交易单详情 func (p *BacktestRepository) ListBacktestTrades(userId, backtestId int64, page data.Page) (total int64, tradeOrders []*trade.TradeOrder, err error) { sqlFrom := ` from t_backtest_trading_order where backtest_id = (select id from t_backtest_trading_plan where id = ? and user_id = ?) ` err = p.db.Select(&total, fmt.Sprintf(` select count(*) %s `, sqlFrom), backtestId, userId) if err != nil || total == 0 { return } err = p.db.Select(&tradeOrders, fmt.Sprintf(` select * %s order by trade_id asc offset ? limit ? `, sqlFrom), backtestId, userId, page.Offset, page.Limit) return } // BacktestLogStats 交易计划回测结果统计信息 func (p *BacktestRepository) BacktestLogStats(userId, backtestId int64) (stats *trade.BacktestTradingPlan, err error) { stats = &trade.BacktestTradingPlan{} err = p.db.Select(stats, ` select * from t_backtest_trading_plan_stats where backtest_id = (select id from t_backtest_trading_plan where id = ? and user_id = ?) `, backtestId, userId) return } // BacktestEquities 回测记录资金曲线 func (p *BacktestRepository) BacktestEquities(userId, backtestId int64) (times []int64, equities []float64, tradeIds []string, err error) { var datas []map[string]any err = p.db.Select(&datas, ` select trade_id, ctime, equity from t_backtest_trading_order where backtest_id = (select id from t_backtest_trading_plan where id = ? and user_id = ?) and trade_type = 2 order by ctime, trade_id asc `, backtestId, userId) if err != nil { return } pow := math.Pow10(2) for _, data := range datas { trade_id := conver.ToInt64(data["trade_id"]) ctime := conver.ToInt64(data["ctime"]) equity := conver.ToFloat64(data["equity"]) equity = math.Round(equity*pow) / pow // 同一时间两笔单 if length := len(times); length > 0 && times[length-1] == ctime { equities[length-1] = equity tradeIds[length-1] += fmt.Sprintf(";%d", trade_id) } else { times = append(times, ctime) equities = append(equities, equity) tradeIds = append(tradeIds, fmt.Sprintf("%d", trade_id)) } } return }