206 lines
5.4 KiB
Go
206 lines
5.4 KiB
Go
package trader
|
|
|
|
import (
|
|
"context"
|
|
"fmt"
|
|
"log/slog"
|
|
"math/rand"
|
|
"time"
|
|
|
|
"github.com/pheinrich/aitrade/pkg/db"
|
|
"github.com/pheinrich/aitrade/pkg/model"
|
|
)
|
|
|
|
type DryRunExecutor struct {
|
|
tradeRepo *db.TradeRepository
|
|
balanceRepo *db.BalanceRepository
|
|
positionRepo *db.PositionRepository
|
|
logger *slog.Logger
|
|
balance float64
|
|
positionPrices map[string]float64 // Track entry prices for P&L calculation
|
|
}
|
|
|
|
func NewDryRunExecutor(
|
|
tradeRepo *db.TradeRepository,
|
|
balanceRepo *db.BalanceRepository,
|
|
positionRepo *db.PositionRepository,
|
|
startingBalance float64,
|
|
logger *slog.Logger,
|
|
) *DryRunExecutor {
|
|
return &DryRunExecutor{
|
|
tradeRepo: tradeRepo,
|
|
balanceRepo: balanceRepo,
|
|
positionRepo: positionRepo,
|
|
logger: logger,
|
|
balance: startingBalance,
|
|
positionPrices: make(map[string]float64),
|
|
}
|
|
}
|
|
|
|
func (e *DryRunExecutor) ExecuteTrade(ctx context.Context, trade *model.Trade) error {
|
|
e.logger.Info("executing dry run trade",
|
|
slog.Int64("trade_id", trade.ID),
|
|
slog.String("symbol", trade.Symbol),
|
|
slog.String("action", string(trade.Action)),
|
|
slog.Int("quantity", trade.Quantity),
|
|
slog.Float64("current_balance", e.balance),
|
|
)
|
|
|
|
// Simulate order execution with random price movement
|
|
simulatedPrice := e.simulatePrice(trade)
|
|
|
|
// Calculate cost
|
|
cost := simulatedPrice * float64(trade.Quantity)
|
|
|
|
// Check if we have enough balance for BUY
|
|
if trade.Action == model.ActionBuy {
|
|
if cost > e.balance {
|
|
return fmt.Errorf("insufficient dry run balance: need %.2f, have %.2f", cost, e.balance)
|
|
}
|
|
e.balance -= cost
|
|
e.positionPrices[trade.Symbol] = simulatedPrice
|
|
|
|
// Create position record
|
|
position := &model.Position{
|
|
Symbol: trade.Symbol,
|
|
Quantity: trade.Quantity,
|
|
EntryPrice: simulatedPrice,
|
|
EntryTradeID: trade.ID,
|
|
}
|
|
if err := e.positionRepo.Create(ctx, position); err != nil {
|
|
e.logger.Error("failed to create position", slog.Any("error", err))
|
|
}
|
|
|
|
e.logger.Info("dry run BUY executed",
|
|
slog.String("symbol", trade.Symbol),
|
|
slog.Float64("price", simulatedPrice),
|
|
slog.Float64("cost", cost),
|
|
slog.Float64("remaining_balance", e.balance),
|
|
)
|
|
} else {
|
|
// SELL - calculate P&L
|
|
entryPrice, exists := e.positionPrices[trade.Symbol]
|
|
if !exists {
|
|
entryPrice = simulatedPrice * 0.95 // Assume we bought 5% lower
|
|
}
|
|
|
|
pnl := (simulatedPrice - entryPrice) * float64(trade.Quantity)
|
|
e.balance += cost
|
|
trade.DryRunPnL = &pnl
|
|
|
|
delete(e.positionPrices, trade.Symbol)
|
|
|
|
// Remove position record
|
|
position, err := e.positionRepo.GetBySymbol(ctx, trade.Symbol)
|
|
if err == nil && position != nil {
|
|
if err := e.positionRepo.Delete(ctx, position.ID); err != nil {
|
|
e.logger.Error("failed to delete position", slog.Any("error", err))
|
|
}
|
|
}
|
|
|
|
e.logger.Info("dry run SELL executed",
|
|
slog.String("symbol", trade.Symbol),
|
|
slog.Float64("entry_price", entryPrice),
|
|
slog.Float64("exit_price", simulatedPrice),
|
|
slog.Float64("pnl", pnl),
|
|
slog.Float64("new_balance", e.balance),
|
|
)
|
|
}
|
|
|
|
// Update trade status
|
|
now := time.Now()
|
|
trade.Status = model.TradeSubmitted
|
|
trade.SubmittedAt = &now
|
|
trade.ExecutedPrice = &simulatedPrice
|
|
trade.IsDryRun = true
|
|
|
|
if err := e.tradeRepo.Update(ctx, trade); err != nil {
|
|
return fmt.Errorf("failed to update trade status: %w", err)
|
|
}
|
|
|
|
// Simulate fill after 1 second
|
|
go e.simulateFill(trade)
|
|
|
|
return nil
|
|
}
|
|
|
|
func (e *DryRunExecutor) simulatePrice(trade *model.Trade) float64 {
|
|
// Use target price if available, otherwise simulate
|
|
if trade.TargetPrice != nil {
|
|
// Add some random slippage (-0.5% to +0.5%)
|
|
slippage := (*trade.TargetPrice) * (rand.Float64() - 0.5) * 0.01
|
|
return *trade.TargetPrice + slippage
|
|
}
|
|
|
|
// Generate a random price based on symbol
|
|
// In production, you'd fetch real market data
|
|
basePrice := 150.0
|
|
|
|
// Add some randomness
|
|
movement := (rand.Float64() - 0.5) * 10.0
|
|
return basePrice + movement
|
|
}
|
|
|
|
func (e *DryRunExecutor) simulateFill(trade *model.Trade) {
|
|
time.Sleep(1 * time.Second)
|
|
|
|
ctx := context.Background()
|
|
now := time.Now()
|
|
|
|
trade.Status = model.TradeFilled
|
|
trade.FilledAt = &now
|
|
|
|
if err := e.tradeRepo.Update(ctx, trade); err != nil {
|
|
e.logger.Error("failed to update filled trade", slog.Any("error", err))
|
|
return
|
|
}
|
|
|
|
e.logger.Info("dry run trade filled",
|
|
slog.Int64("trade_id", trade.ID),
|
|
slog.Float64("executed_price", *trade.ExecutedPrice),
|
|
)
|
|
|
|
// Mark as completed (no stop-loss in dry run for simplicity)
|
|
time.Sleep(500 * time.Millisecond)
|
|
|
|
now = time.Now()
|
|
trade.Status = model.TradeCompleted
|
|
trade.CompletedAt = &now
|
|
|
|
if err := e.tradeRepo.Update(ctx, trade); err != nil {
|
|
e.logger.Error("failed to mark trade completed", slog.Any("error", err))
|
|
return
|
|
}
|
|
|
|
// Update balance in database
|
|
balance := &model.Balance{
|
|
Timestamp: time.Now(),
|
|
TotalValue: e.balance,
|
|
CashBalance: e.balance,
|
|
BuyingPower: e.balance,
|
|
UnrealizedPnL: nil,
|
|
RealizedPnL: trade.DryRunPnL,
|
|
}
|
|
|
|
if err := e.balanceRepo.Create(ctx, balance); err != nil {
|
|
e.logger.Error("failed to store dry run balance", slog.Any("error", err))
|
|
return
|
|
}
|
|
|
|
e.logger.Info("dry run trade completed",
|
|
slog.Int64("trade_id", trade.ID),
|
|
slog.Float64("balance", e.balance),
|
|
)
|
|
}
|
|
|
|
func (e *DryRunExecutor) GetBalance() float64 {
|
|
return e.balance
|
|
}
|
|
|
|
func (e *DryRunExecutor) CancelTrade(ctx context.Context, trade *model.Trade) error {
|
|
e.logger.Info("dry run trade cancelled",
|
|
slog.Int64("trade_id", trade.ID),
|
|
)
|
|
return nil
|
|
}
|