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

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
}