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.
277 lines
8.0 KiB
277 lines
8.0 KiB
package backtest |
|
|
|
import ( |
|
"fmt" |
|
"sig-pub/pkg/trade" |
|
"sig-pub/pkg/types" |
|
"sig-pub/pkg/types/decimals" |
|
"sig-pub/pkg/utils/collect" |
|
"sig-pub/pkg/utils/conver" |
|
"sig-pub/pkg/utils/lang" |
|
"sig-pub/pkg/zlog" |
|
"time" |
|
|
|
"github.com/govalues/decimal" |
|
) |
|
|
|
type BacktestTradeAccount struct { |
|
trader *TradeSimulator // 下单器 |
|
traderOrderId int64 // 订单id递增 |
|
initialCash float64 // 初始资金 |
|
cash float64 // 可用资金 |
|
positions map[string]*trade.Position // 仓位信息,key: symbol, value: qty(不能持有同一交易产品反方向单, 交易中不能改变杠杆) |
|
orders map[int64]*trade.TradeOrder // 所有交易单 |
|
openTrades map[string][]int64 // 未平仓交易单,key: symbol, value: tradeId |
|
closeTrades []int64 // 平仓交易单 |
|
maxDrawdown [4]float64 // 最大回撤 |
|
} |
|
|
|
func NewBacktestTradeAccount(cash float64) *BacktestTradeAccount { |
|
return &BacktestTradeAccount{ |
|
trader: NewTradeSimulator(0.0005, 0.0008), |
|
traderOrderId: 1, |
|
initialCash: cash, |
|
cash: cash, |
|
positions: make(map[string]*trade.Position), |
|
orders: make(map[int64]*trade.TradeOrder), |
|
openTrades: make(map[string][]int64), |
|
maxDrawdown: [4]float64{cash, cash, cash, cash}, |
|
} |
|
} |
|
|
|
// 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 _, 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 |
|
} |
|
|
|
// CurrentExposure 返回当前仓位的名义总敞口(绝对值) |
|
func (a *BacktestTradeAccount) CurrentExposure() float64 { |
|
var sum float64 |
|
for _, p := range a.positions { |
|
sum += p.EntryPx * decimals.MustToFloat64(p.Qty) |
|
} |
|
return sum |
|
} |
|
|
|
// IsSymbolOpen 指定交易对是否有持仓 |
|
func (a *BacktestTradeAccount) IsSymbolOpen(symbol string) bool { |
|
return a.positions[symbol] != nil |
|
} |
|
|
|
func (a *BacktestTradeAccount) GetOpenTradeInsts() (instIds []string) { |
|
for instId := range a.openTrades { |
|
instIds = append(instIds, instId) |
|
} |
|
return |
|
} |
|
|
|
func (a *BacktestTradeAccount) GetOpenTrades(instId string) []*trade.TradeOrder { |
|
var trades []*trade.TradeOrder |
|
for _, tradeId := range a.openTrades[instId] { |
|
trades = append(trades, a.orders[tradeId]) |
|
} |
|
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] |
|
} |
|
|
|
// MarketOrder 市价下单 |
|
func (a *BacktestTradeAccount) MarketOrder(ticket trade.TradeTicket) (order *trade.TradeOrder, err error) { |
|
instId := ticket.InstId |
|
if pos, ok := a.positions[instId]; ok { |
|
// 检查反方向单 |
|
if pos.Side != ticket.Side { |
|
err = fmt.Errorf("cannot open opposite side position") |
|
return |
|
} |
|
// 同方向杠杆倍数 |
|
if pos.Leverage != ticket.Leverage { |
|
err = fmt.Errorf("cannot change leverage on existing position") |
|
return |
|
} |
|
} |
|
|
|
order = a.trader.ExecuteMarket(a.traderOrderId, instId, ticket) |
|
orderQty := decimals.MustToFloat64(order.Qty) |
|
cost := order.Price*orderQty/float64(order.Leverage) + order.Fee |
|
if cost > a.cash { |
|
// err = fmt.Errorf("insufficient cash") |
|
return |
|
} |
|
a.traderOrderId++ |
|
a.cash -= cost |
|
a.orders[order.TradeId] = order |
|
a.openTrades[instId] = append(a.openTrades[instId], order.TradeId) |
|
|
|
// open/add position |
|
pos, ok := a.positions[instId] |
|
if !ok { |
|
pos = &trade.Position{ |
|
InstId: order.InstId, |
|
Side: order.Side, |
|
Qty: order.Qty, |
|
Leverage: order.Leverage, |
|
EntryPx: order.Price, |
|
EntryTs: order.Ctime, |
|
PeakPx: order.Price, |
|
LastPx: order.Price, |
|
} |
|
a.positions[instId] = pos |
|
} else { |
|
posQty := decimals.MustToFloat64(pos.Qty) |
|
orderQty := decimals.MustToFloat64(order.Qty) |
|
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 |
|
} |
|
if pos.Side == types.SideShort && order.Price > pos.PeakPx { |
|
pos.PeakPx = order.Price |
|
} |
|
} |
|
return |
|
} |
|
|
|
// CloseTradeOrder 订单订单平仓 |
|
func (a *BacktestTradeAccount) CloseTradeOrder(ticket trade.TradeTicket) (err error) { |
|
instId := ticket.InstId |
|
// 检查仓位 |
|
pos, ok := a.positions[instId] |
|
if !ok { |
|
err = fmt.Errorf("no open position for symbol: %s", instId) |
|
return |
|
} |
|
if pos.Side == ticket.Side { |
|
err = fmt.Errorf("cannot close position with same side trade") |
|
return |
|
} |
|
if ticket.Qty.Cmp(pos.Qty) > 0 { |
|
// ticket.Qty = pos.Qty |
|
err = fmt.Errorf("close quantity exceeds position quantity") |
|
return |
|
} |
|
|
|
order := a.trader.ExecuteMarket(a.traderOrderId, instId, ticket) |
|
a.traderOrderId++ |
|
|
|
var insure, profit float64 |
|
var entryPxs, entryPxAvg, entryTrades float64 |
|
for _, tradeId := range ticket.TradesId { |
|
trade, ok := a.orders[tradeId] |
|
if !ok { |
|
continue |
|
} |
|
entryTrades++ |
|
entryPxs += trade.Price |
|
order.EntryFee += trade.Fee |
|
order.EntryTime = lang.Ternary(order.EntryTime == 0, trade.Ctime, min(order.EntryTime, trade.Ctime)) |
|
order.PeakPx = lang.Ternary( |
|
trade.Side == types.SideLong, |
|
max(order.PeakPx, trade.PeakPx), |
|
min(order.PeakPx, trade.PeakPx), |
|
) |
|
tradeQty := decimals.MustToFloat64(trade.Qty) |
|
|
|
// 计算盈利 |
|
if trade.Side == types.SideLong { |
|
profit += tradeQty * (order.Price - trade.Price) |
|
} else { |
|
profit += tradeQty * (trade.Price - order.Price) |
|
} |
|
} |
|
if entryTrades == 0 { |
|
err = fmt.Errorf("entry error") |
|
return |
|
} |
|
|
|
entryPxAvg = entryPxs / entryTrades // 开仓平均价 |
|
orderQty := decimals.MustToFloat64(order.Qty) // 平仓数量 |
|
insure = entryPxAvg * orderQty / float64(order.Leverage) // 开仓保证金 |
|
a.cash += (insure + profit - order.Fee) // 平仓收益加到账户 |
|
|
|
a.orders[order.TradeId] = order |
|
a.closeTrades = append(a.closeTrades, order.TradeId) |
|
if len(ticket.TradesId) > 0 { |
|
if trades, ok := a.openTrades[instId]; ok { |
|
// remove open trade order |
|
removes := collect.Remove(&trades, func(tradeId int64) bool { |
|
return collect.In(tradeId, ticket.TradesId...) |
|
}) |
|
a.openTrades[instId] = trades |
|
if removes != len(ticket.TradesId) { |
|
zlog.Warningf("close ticket tradeId not exists: instId=%s, tradeId=%#v", ticket.InstId, ticket.TradesId) |
|
} |
|
} |
|
if len(a.openTrades[instId]) == 0 { |
|
delete(a.openTrades, instId) |
|
} |
|
} |
|
|
|
// update position |
|
pos.Qty = decimals.MustSub(pos.Qty, ticket.Qty) |
|
if pos.Qty.Cmp(decimal.Zero) <= 0 { |
|
delete(a.positions, instId) |
|
if _, ok := a.openTrades[instId]; ok { |
|
zlog.Warningf("close all position trades: %s", instId) |
|
} |
|
} |
|
|
|
// 账户净值 |
|
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, ".") |
|
order.Profit = profit - order.Fee - order.EntryFee |
|
order.Trades = ticket.TradesId |
|
|
|
// 记录最大回撤 [high, low, highPrev, lowPrev] |
|
if currentEquity < a.maxDrawdown[1] { |
|
a.maxDrawdown[1] = currentEquity |
|
} |
|
if currentEquity > a.maxDrawdown[0] { |
|
if a.maxDrawdown[0]-a.maxDrawdown[1] > a.maxDrawdown[2]-a.maxDrawdown[3] { |
|
a.maxDrawdown[2], a.maxDrawdown[3] = a.maxDrawdown[0], a.maxDrawdown[1] |
|
} |
|
a.maxDrawdown[0] = currentEquity |
|
a.maxDrawdown[1] = currentEquity |
|
} |
|
return |
|
}
|
|
|