You can not select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
108 lines
3.3 KiB
108 lines
3.3 KiB
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 |
|
}
|
|
|