Files
foxhunt/crates/backtesting/src/strategy_runner.rs
jgrusewski 9c3d741a08 refactor: restructure repo — crates/, bin/, testing/ layout
Move 17 library crates into crates/, CLI binary into bin/fxt,
consolidate 10 test crates into testing/, split config crate
from deployment config files.

Root directory reduced from 38+ to ~17 directories.
All Cargo.toml paths and build.rs proto refs updated.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-25 11:56:00 +01:00

1692 lines
58 KiB
Rust

#![allow(unsafe_code)] // Intentional unsafe for AVX2 vectorized backtesting performance
//! Adaptive Strategy Runner for Backtesting
//!
//! This module provides the bridge between the backtesting engine and the adaptive strategy
//! system, allowing historical data to flow through real ML models for validation.
use anyhow::Result;
use async_trait::async_trait;
use common::ml_strategy::ProductionFeatureExtractor225;
use common::Order;
use common::{OrderSide, OrderStatus, Position, Price, Quantity, Symbol};
use rust_decimal::prelude::ToPrimitive;
use rust_decimal::Decimal;
use trading_engine::types::events::MarketEvent;
// Use canonical types from ML module and real ML registry
use chrono::{DateTime, Utc};
use dashmap::DashMap;
use ml::features::ProductionFeatureExtractorAdapter;
use ml::{get_global_registry, Features, ModelPrediction};
use parking_lot::RwLock;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
use tracing::{debug, info, warn};
// SIMD optimizations for HFT performance
#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
use crate::strategy_tester::{
PerformanceSnapshot, SignalType, Strategy, StrategyConfig, StrategyContext, StrategyResult,
TradeRecord, TradingSignal,
};
/// Adaptive strategy runner that integrates ML models with backtesting
pub struct AdaptiveStrategyRunner {
/// Strategy configuration
config: AdaptiveStrategyConfig,
/// Current market state
market_state: Arc<RwLock<MarketState>>,
/// Model predictions cache (lock-free for HFT performance)
predictions_cache: Arc<DashMap<String, ModelPrediction>>,
/// Performance tracking
performance_tracker: Arc<RwLock<PerformanceTracker>>,
/// Feature extractor (legacy 3-5 dim fallback)
feature_extractor: Arc<FeatureExtractor>,
/// Production 51-dimension feature extractor from ml::features
production_extractor: Arc<RwLock<ProductionFeatureExtractorAdapter>>,
/// Risk manager
risk_manager: Arc<RiskManager>,
}
/// Configuration for adaptive strategy
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AdaptiveStrategyConfig {
/// Models to use for predictions
pub active_models: Vec<String>,
/// Minimum confidence threshold for trading
pub min_confidence: f64,
/// Maximum position size as fraction of portfolio
pub max_position_size: f64,
/// Lookback period for feature extraction
pub lookback_period: usize,
/// Model update frequency (in ticks)
pub model_update_frequency: u64,
/// Risk management settings
pub risk_settings: RiskSettings,
/// Feature extraction settings
pub feature_settings: FeatureSettings,
}
impl Default for AdaptiveStrategyConfig {
fn default() -> Self {
Self {
active_models: vec![
"TLOB".to_string(),
"MAMBA".to_string(),
"TFT".to_string(),
"DQN".to_string(),
"PPO".to_string(),
],
min_confidence: 0.65,
max_position_size: 0.05, // 5% max position
lookback_period: 100,
model_update_frequency: 1000,
risk_settings: RiskSettings::default(),
feature_settings: FeatureSettings::default(),
}
}
}
/// Risk management settings
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RiskSettings {
/// Maximum drawdown before stopping
pub max_drawdown: f64,
/// Stop loss percentage
pub stop_loss: f64,
/// Take profit percentage
pub take_profit: f64,
/// Kelly fraction multiplier
pub kelly_fraction: f64,
}
impl Default for RiskSettings {
fn default() -> Self {
Self {
max_drawdown: 0.10, // 10% max drawdown
stop_loss: 0.02, // 2% stop loss
take_profit: 0.04, // 4% take profit
kelly_fraction: 0.25, // Conservative 25% of Kelly
}
}
}
/// Feature extraction settings
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FeatureSettings {
/// Price features to extract
pub price_features: Vec<String>,
/// Volume features to extract
pub volume_features: Vec<String>,
/// Technical indicators to compute
pub technical_indicators: Vec<String>,
/// Microstructure features
pub microstructure_features: Vec<String>,
}
impl Default for FeatureSettings {
fn default() -> Self {
Self {
price_features: vec![
"returns".to_string(),
"log_returns".to_string(),
"volatility".to_string(),
"price_momentum".to_string(),
],
volume_features: vec![
"volume".to_string(),
"volume_momentum".to_string(),
"vwap".to_string(),
],
technical_indicators: vec![
"rsi".to_string(),
"macd".to_string(),
"bollinger_bands".to_string(),
],
microstructure_features: vec![
"bid_ask_spread".to_string(),
"order_flow_imbalance".to_string(),
"market_impact".to_string(),
],
}
}
}
/// Current market state
#[derive(Debug, Clone)]
struct MarketState {
/// Current timestamp
current_time: DateTime<Utc>,
/// Price history
price_history: Vec<(DateTime<Utc>, Decimal)>,
/// Volume history
volume_history: Vec<(DateTime<Utc>, Decimal)>,
/// Current position
current_position: Option<Position>,
}
impl Default for MarketState {
fn default() -> Self {
Self {
current_time: Utc::now(),
price_history: Vec::new(),
volume_history: Vec::new(),
current_position: None,
}
}
}
/// Tracks an open (or partially open) entry for PnL accounting.
#[derive(Debug, Clone)]
struct OpenEntry {
/// Side of the original entry
side: OrderSide,
/// Average entry price
entry_price: Price,
/// Remaining quantity in the entry
quantity: Decimal,
/// Timestamp when the position was opened
entry_time: DateTime<Utc>,
}
/// Performance tracking for the strategy
#[derive(Debug, Clone)]
struct PerformanceTracker {
/// Total trades executed
total_trades: u64,
/// Winning trades
winning_trades: u64,
/// Total PnL
total_pnl: Decimal,
/// Maximum drawdown
max_drawdown: Decimal,
/// Current drawdown
current_drawdown: Decimal,
/// Peak portfolio value
peak_value: Decimal,
/// Initial capital for equity curve baseline
initial_capital: Decimal,
/// Completed round-trip trade records
trade_records: Vec<TradeRecord>,
/// Equity snapshots: (timestamp, portfolio_value)
equity_snapshots: Vec<(DateTime<Utc>, Decimal)>,
/// Open position entries keyed by symbol for PnL accounting
open_entries: HashMap<String, OpenEntry>,
}
impl Default for PerformanceTracker {
fn default() -> Self {
Self {
total_trades: 0,
winning_trades: 0,
total_pnl: Decimal::ZERO,
max_drawdown: Decimal::ZERO,
current_drawdown: Decimal::ZERO,
peak_value: Decimal::ZERO,
initial_capital: Decimal::ZERO,
trade_records: Vec::new(),
equity_snapshots: Vec::new(),
open_entries: HashMap::new(),
}
}
}
/// Feature extractor for ML models with object pooling for performance
struct FeatureExtractor {
config: FeatureSettings,
}
impl FeatureExtractor {
/// Create new feature extractor with configuration
///
/// # Arguments
/// * `config` - Feature extraction configuration
///
/// # Returns
/// * `Self` - New feature extractor instance
fn new(config: FeatureSettings) -> Self {
Self {
config,
}
}
/// Extract features from market data
///
/// # Arguments
/// * `market_state` - Current market state containing price and volume history
///
/// # Returns
/// * `Result<Features>` - Extracted feature vector ready for ML model input
///
/// # Errors
///
/// Returns error if feature extraction fails or insufficient data
async fn extract_features(&self, market_state: &MarketState) -> Result<Features> {
let mut feature_values = Vec::new();
let mut feature_names = Vec::new();
// Extract price features
if let Some(features) = self.extract_price_features(market_state).await? {
feature_values.extend(features.0);
feature_names.extend(features.1);
}
// Extract volume features
if let Some(features) = self.extract_volume_features(market_state).await? {
feature_values.extend(features.0);
feature_names.extend(features.1);
}
// Extract technical indicators
if let Some(features) = self.extract_technical_features(market_state).await? {
feature_values.extend(features.0);
feature_names.extend(features.1);
}
Ok(Features {
values: feature_values,
names: feature_names,
timestamp: std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_micros() as u64,
symbol: None, // Symbol will be set by the calling context when available
})
}
/// Extract price-based features from market data
///
/// # Arguments
/// * `market_state` - Market state with price history
///
/// # Returns
/// * `Result<Option<(Vec<f64>, Vec<String>)>>` - Feature values and names, or None if insufficient data
async fn extract_price_features(
&self,
market_state: &MarketState,
) -> Result<Option<(Vec<f64>, Vec<String>)>> {
if market_state.price_history.len() < 2 {
return Ok(None);
}
let mut values = Vec::new();
let mut names = Vec::new();
let prices: Vec<f64> = market_state
.price_history
.iter()
.map(|(_, price)| price.to_f64().unwrap_or(0.0))
.collect();
// Calculate returns
if self.config.price_features.contains(&"returns".to_string()) {
let returns = self.calculate_returns(&prices);
values.extend(returns);
names.push("returns".to_string());
}
// Calculate volatility
if self
.config
.price_features
.contains(&"volatility".to_string())
{
let volatility = self.calculate_volatility(&prices);
values.push(volatility);
names.push("volatility".to_string());
}
Ok(Some((values, names)))
}
/// Extract volume-based features from market data
///
/// # Arguments
/// * `market_state` - Market state with volume history
///
/// # Returns
/// * `Result<Option<(Vec<f64>, Vec<String>)>>` - Feature values and names, or None if insufficient data
async fn extract_volume_features(
&self,
market_state: &MarketState,
) -> Result<Option<(Vec<f64>, Vec<String>)>> {
if market_state.volume_history.len() < 2 {
return Ok(None);
}
let mut values = Vec::new();
let mut names = Vec::new();
let volumes: Vec<f64> = market_state
.volume_history
.iter()
.map(|(_, volume)| volume.to_f64().unwrap_or(0.0))
.collect();
// Average volume
if self.config.volume_features.contains(&"volume".to_string()) {
let avg_volume = volumes.iter().sum::<f64>() / volumes.len() as f64;
values.push(avg_volume);
names.push("avg_volume".to_string());
}
Ok(Some((values, names)))
}
/// Extract technical indicator features from market data
///
/// # Arguments
/// * `market_state` - Market state with sufficient price history for indicators
///
/// # Returns
/// * `Result<Option<(Vec<f64>, Vec<String>)>>` - Technical indicator values and names, or None if insufficient data
async fn extract_technical_features(
&self,
market_state: &MarketState,
) -> Result<Option<(Vec<f64>, Vec<String>)>> {
if market_state.price_history.len() < 14 {
// Need minimum data for indicators
return Ok(None);
}
let mut values = Vec::new();
let mut names = Vec::new();
let prices: Vec<f64> = market_state
.price_history
.iter()
.map(|(_, price)| price.to_f64().unwrap_or(0.0))
.collect();
// RSI
if self
.config
.technical_indicators
.contains(&"rsi".to_string())
{
let rsi = self.calculate_rsi(&prices, 14);
values.push(rsi);
names.push("rsi_14".to_string());
}
Ok(Some((values, names)))
}
/// Calculate returns from price series with SIMD optimization
///
/// # Arguments
/// * `prices` - Array of price values
///
/// # Returns
/// * `Vec<f64>` - Vector of return percentages
///
/// # Note
///
/// Uses SIMD instructions on x86_64 for performance when available
fn calculate_returns(&self, prices: &[f64]) -> Vec<f64> {
// OPTIMIZATION: Use SIMD for vectorized return calculations
if prices.len() < 2 {
return Vec::new();
}
#[cfg(target_arch = "x86_64")]
{
if is_x86_feature_detected!("avx2") {
return self.calculate_returns_simd(prices);
}
}
// Fallback to scalar implementation
self.calculate_returns_scalar(prices)
}
/// Calculate returns using scalar operations (fallback)
///
/// # Arguments
/// * `prices` - Array of price values
///
/// # Returns
/// * `Vec<f64>` - Vector of return percentages
fn calculate_returns_scalar(&self, prices: &[f64]) -> Vec<f64> {
prices
.windows(2)
.map(|window| {
if window[0] != 0.0 {
(window[1] - window[0]) / window[0]
} else {
0.0
}
})
.collect()
}
/// Calculate returns using SIMD AVX2 instructions (x86_64 only)
///
/// # Arguments
/// * `prices` - Array of price values
///
/// # Returns
/// * `Vec<f64>` - Vector of return percentages
///
/// # Safety
///
/// Uses unsafe AVX2 intrinsics for vectorized computation
#[cfg(target_arch = "x86_64")]
fn calculate_returns_simd(&self, prices: &[f64]) -> Vec<f64> {
let mut returns = Vec::with_capacity(prices.len() - 1);
let len = prices.len() - 1;
// SAFETY: SIMD intrinsics validated with feature detection and proper data alignment
unsafe {
// Process 4 elements at a time with AVX2
let mut i = 0;
while i + 4 <= len {
let prev = _mm256_loadu_pd(prices.as_ptr().add(i));
let curr = _mm256_loadu_pd(prices.as_ptr().add(i + 1));
// Calculate (curr - prev) / prev
let diff = _mm256_sub_pd(curr, prev);
let result = _mm256_div_pd(diff, prev);
// Store results
let mut temp = [0.0; 4];
_mm256_storeu_pd(temp.as_mut_ptr(), result);
for j in 0..4 {
returns.push(if prices[i + j] != 0.0 { temp[j] } else { 0.0 });
}
i += 4;
}
// Handle remaining elements
for j in i..len {
let ret = if prices[j] != 0.0 {
(prices[j + 1] - prices[j]) / prices[j]
} else {
0.0
};
returns.push(ret);
}
}
returns
}
/// Calculate price volatility (standard deviation of returns)
///
/// # Arguments
/// * `prices` - Array of price values
///
/// # Returns
/// * `f64` - Volatility as standard deviation of returns
fn calculate_volatility(&self, prices: &[f64]) -> f64 {
let returns = self.calculate_returns(prices);
if returns.is_empty() {
return 0.0;
}
let mean = returns.iter().sum::<f64>() / returns.len() as f64;
let variance =
returns.iter().map(|r| (r - mean).powi(2)).sum::<f64>() / returns.len() as f64;
variance.sqrt()
}
/// Calculate Relative Strength Index (RSI)
///
/// # Arguments
/// * `prices` - Array of price values
///
/// * `period` - Period for RSI calculation (typically 14)
///
/// # Returns
/// * `f64` - RSI value between 0 and 100
fn calculate_rsi(&self, prices: &[f64], period: usize) -> f64 {
if prices.len() < period + 1 {
return 50.0; // Neutral RSI
}
let mut gains = Vec::new();
let mut losses = Vec::new();
for window in prices.windows(2) {
let change = window[1] - window[0];
if change > 0.0 {
gains.push(change);
losses.push(0.0);
} else {
gains.push(0.0);
losses.push(-change);
}
}
if gains.len() < period {
return 50.0;
}
let avg_gain = gains.iter().rev().take(period).sum::<f64>() / period as f64;
let avg_loss = losses.iter().rev().take(period).sum::<f64>() / period as f64;
if avg_loss == 0.0 {
return 100.0;
}
let rs = avg_gain / avg_loss;
100.0 - (100.0 / (1.0 + rs))
}
}
/// Risk manager for position sizing and risk controls
struct RiskManager {
config: RiskSettings,
}
impl RiskManager {
/// Create new risk manager with configuration
///
/// # Arguments
/// * `config` - Risk management settings
///
/// # Returns
/// * `Self` - New risk manager instance
fn new(config: RiskSettings) -> Self {
Self { config }
}
/// Calculate position size using Kelly criterion
///
/// # Arguments
/// * `prediction` - Model prediction with confidence score
///
/// * `account_value` - Current account value
/// * `current_price` - Current market price
///
/// # Returns
/// * `Result<Decimal>` - Position size in shares/units
///
/// # Note
///
/// Uses conservative Kelly fraction scaling for risk management
fn calculate_position_size(
&self,
prediction: &ModelPrediction,
account_value: Decimal,
current_price: Decimal,
) -> Result<Decimal> {
// Simple Kelly-based position sizing
let confidence = prediction.confidence;
let edge = (confidence - 0.5) * 2.0; // Convert to [-1, 1] range
if edge <= 0.0 {
return Ok(Decimal::ZERO);
}
// Kelly fraction with conservative scaling
let kelly_size =
Decimal::try_from(edge * self.config.kelly_fraction).unwrap_or(Decimal::ZERO);
let max_size = account_value
* Decimal::try_from(self.config.max_drawdown).unwrap_or(Decimal::new(5, 2)); // 5% fallback
let position_value = kelly_size * account_value;
let position_size = if current_price > Decimal::ZERO {
position_value / current_price
} else {
Decimal::ZERO
};
Ok(position_size.min(max_size / current_price))
}
/// Check if trade passes risk checks
///
/// # Arguments
/// * `signal` - Trading signal to validate
///
/// * `current_position` - Current position if any
/// * `account_value` - Current account value
///
/// # Returns
/// * `Result<bool>` - True if trade passes all risk checks
fn validate_trade(
&self,
signal: &TradingSignal,
_current_position: Option<&Position>,
account_value: Decimal,
) -> Result<bool> {
// Check position size limits
if let Some(price) = signal.target_price {
let trade_value = signal.quantity.to_decimal()? * price.to_decimal()?;
let position_fraction = trade_value / account_value;
if position_fraction
> Decimal::try_from(self.config.max_drawdown).unwrap_or(Decimal::new(10, 2))
{
debug!(
"Trade rejected: position size too large ({:.2}%)",
position_fraction * Decimal::from(100)
);
return Ok(false);
}
}
// Additional risk checks can be added here
Ok(true)
}
}
impl AdaptiveStrategyRunner {
/// Create new adaptive strategy runner
///
/// # Arguments
/// * `config` - Configuration for adaptive strategy
///
/// # Returns
/// * `Self` - New adaptive strategy runner with initialized components
pub fn new(config: AdaptiveStrategyConfig) -> Self {
let feature_extractor = Arc::new(FeatureExtractor::new(config.feature_settings.clone()));
let risk_manager = Arc::new(RiskManager::new(config.risk_settings.clone()));
let production_extractor = Arc::new(RwLock::new(ProductionFeatureExtractorAdapter::new()));
Self {
config,
market_state: Arc::new(RwLock::new(MarketState::default())),
predictions_cache: Arc::new(DashMap::new()),
performance_tracker: Arc::new(RwLock::new(PerformanceTracker::default())),
feature_extractor,
production_extractor,
risk_manager,
}
}
/// Get ensemble prediction from all active models (optimized for HFT performance)
///
/// # Arguments
/// * `features` - Feature vector for prediction
///
/// # Returns
/// * `Result<ModelPrediction>` - Ensemble prediction with confidence-weighted averaging
///
/// # Note
///
/// Uses lock-free caching and parallel model execution for low-latency performance
async fn get_ensemble_prediction(&self, features: &Features) -> Result<ModelPrediction> {
let registry = get_global_registry();
// OPTIMIZATION: Use predict_selected for parallel model execution
let predictions = registry
.predict_selected(&self.config.active_models, features)
.await;
// Filter successful predictions
let valid_predictions: Vec<ModelPrediction> = predictions
.into_iter()
.filter_map(|result| result.ok())
.collect();
if valid_predictions.is_empty() {
return Err(anyhow::anyhow!("No valid predictions from any model"));
}
// Ensemble using confidence-weighted average
let total_confidence: f64 = valid_predictions.iter().map(|p| p.confidence).sum();
if total_confidence == 0.0 {
return Err(anyhow::anyhow!("Zero total confidence in predictions"));
}
let weighted_value = valid_predictions
.iter()
.map(|p| p.value * p.confidence)
.sum::<f64>()
/ total_confidence;
let ensemble_confidence = valid_predictions.iter().map(|p| p.confidence).sum::<f64>()
/ valid_predictions.len() as f64;
// OPTIMIZATION: Use lock-free DashMap instead of async RwLock
for prediction in &valid_predictions {
self.predictions_cache
.insert(prediction.model_id.clone(), prediction.clone());
}
Ok(ModelPrediction::new(
"ensemble".to_string(),
weighted_value,
ensemble_confidence,
))
}
/// Generate trading signal from prediction
///
/// # Arguments
/// * `prediction` - Model prediction with confidence and direction
///
/// * `symbol` - Symbol to trade
/// * `current_price` - Current market price
///
/// * `account_value` - Current account value for position sizing
///
/// # Returns
/// * `Result<Option<TradingSignal>>` - Trading signal if confidence threshold is met
fn generate_signal(
&self,
prediction: &ModelPrediction,
symbol: Symbol,
current_price: Decimal,
account_value: Decimal,
) -> Result<Option<TradingSignal>> {
// Check confidence threshold
if prediction.confidence < self.config.min_confidence {
debug!(
"Prediction confidence {:.3} below threshold {:.3}",
prediction.confidence, self.config.min_confidence
);
return Ok(None);
}
// Determine trade direction
let side = if prediction.value > 0.5 {
OrderSide::Buy
} else if prediction.value < -0.5 {
OrderSide::Sell
} else {
return Ok(None); // Neutral signal
};
// Calculate position size
let quantity =
self.risk_manager
.calculate_position_size(prediction, account_value, current_price)?;
if quantity <= Decimal::ZERO {
return Ok(None);
}
let signal_type = match side {
OrderSide::Buy => SignalType::Buy,
OrderSide::Sell => SignalType::Sell,
};
let quantity_as_quantity = Quantity::from_f64(quantity.to_f64().unwrap_or(0.0))?;
let mut metadata = HashMap::new();
metadata.insert("strategy".to_string(), serde_json::json!("adaptive_ml"));
metadata.insert(
"ensemble_confidence".to_string(),
serde_json::json!(prediction.confidence),
);
metadata.insert(
"prediction_value".to_string(),
serde_json::json!(prediction.value),
);
metadata.insert(
"model_count".to_string(),
serde_json::json!(self.config.active_models.len()),
);
let signal = TradingSignal {
symbol,
signal_type,
quantity: quantity_as_quantity,
target_price: Some(Price::from_f64(current_price.to_f64().unwrap_or(0.0))?),
stop_loss: None,
take_profit: None,
confidence: Decimal::try_from(prediction.confidence).unwrap_or(Decimal::ZERO),
metadata,
};
Ok(Some(signal))
}
}
#[async_trait(?Send)]
impl Strategy for AdaptiveStrategyRunner {
fn name(&self) -> &str {
"adaptive_ml_strategy"
}
async fn initialize(
&mut self,
initial_capital: Decimal,
_config: StrategyConfig,
) -> Result<()> {
info!(
"Initializing Adaptive ML Strategy with capital: {}",
initial_capital
);
{
let mut tracker = self.performance_tracker.write();
tracker.peak_value = initial_capital;
tracker.initial_capital = initial_capital;
}
// Verify models are available
let registry = get_global_registry();
let available_models = registry.get_model_names();
for model_name in &self.config.active_models {
if !available_models.contains(model_name) {
warn!("Model {} not found in registry", model_name);
}
}
info!("Adaptive ML Strategy initialized successfully");
Ok(())
}
async fn on_market_event(
&mut self,
event: &MarketEvent,
context: &StrategyContext,
) -> Result<Vec<TradingSignal>> {
let mut signals = Vec::new();
match event {
MarketEvent::Trade {
symbol,
price,
timestamp,
..
} => {
// Update market state
{
let mut state = self.market_state.write();
state.current_time = *timestamp;
state.price_history.push((*timestamp, price.to_decimal()?));
// Note: volume not available in MarketEvent::Trade, using placeholder
// state.volume_history.push((*timestamp, Volume::ZERO));
// Keep only recent history
let max_history = self.config.lookback_period;
if state.price_history.len() > max_history {
let excess = state.price_history.len() - max_history;
state.price_history.drain(0..excess);
}
if state.volume_history.len() > max_history {
let excess = state.volume_history.len() - max_history;
state.volume_history.drain(0..excess);
}
}
// Update production extractor with new price data
let price_f64 = price.to_decimal()?.to_f64().unwrap_or(0.0);
// Volume not available in MarketEvent::Trade; use a default
let volume_f64 = 1000.0;
{
let mut prod_ext = self.production_extractor.write();
if let Err(e) = prod_ext.update(price_f64, volume_f64, *timestamp) {
debug!("Production extractor update failed: {}", e);
}
}
// Try production 51-dim feature extraction first
let production_features = {
let mut prod_ext = self.production_extractor.write();
prod_ext.extract_features()
};
match production_features {
Ok(feat_values) if feat_values.len() == 51 => {
// Build feature names for the 51-dim vector
let feature_names: Vec<String> = (0..51)
.map(|i| format!("prod_feature_{}", i))
.collect();
let ts_micros = timestamp
.timestamp_micros();
let features = Features {
values: feat_values,
names: feature_names,
timestamp: ts_micros as u64,
symbol: Some(symbol.to_string()),
};
match self.get_ensemble_prediction(&features).await {
Ok(prediction) => {
if let Some(signal) = self.generate_signal(
&prediction,
symbol.clone(),
price.to_decimal()?,
context.account_balance,
)? {
let current_position = context.positions.get(symbol);
if self.risk_manager.validate_trade(
&signal,
current_position,
context.account_balance,
)? {
signals.push(signal);
}
}
}
Err(e) => {
debug!("Prediction failed: {}", e);
}
}
}
Ok(_) | Err(_) => {
// Production extractor not warmed up yet, fall back to legacy extractor
let market_state = self.market_state.read().clone();
if market_state.price_history.len() >= 10 {
match self
.feature_extractor
.extract_features(&market_state)
.await
{
Ok(features) => {
match self.get_ensemble_prediction(&features).await {
Ok(prediction) => {
if let Some(signal) = self.generate_signal(
&prediction,
symbol.clone(),
price.to_decimal()?,
context.account_balance,
)? {
let current_position =
context.positions.get(symbol);
if self.risk_manager.validate_trade(
&signal,
current_position,
context.account_balance,
)? {
signals.push(signal);
}
}
}
Err(e) => {
debug!("Prediction failed: {}", e);
}
}
}
Err(e) => {
debug!("Feature extraction failed: {}", e);
}
}
}
}
}
},
_ => {
// Handle other event types if needed
},
}
Ok(signals)
}
async fn on_order_update(&mut self, order: &Order, context: &StrategyContext) -> Result<()> {
if order.status == OrderStatus::Filled {
let fill_price = order
.average_price
.unwrap_or_else(|| order.price.unwrap_or(Price::ZERO));
let fill_qty = order.quantity.to_decimal().unwrap_or(Decimal::ZERO);
let symbol_key = order.symbol.to_string();
let mut tracker = self.performance_tracker.write();
// Check if this fill closes (or reduces) an existing open entry
let mut closed = false;
if let Some(entry) = tracker.open_entries.get(&symbol_key).cloned() {
// A Buy closes a short entry; a Sell closes a long entry
let is_closing = matches!(
(&entry.side, &order.side),
(OrderSide::Buy, OrderSide::Sell) | (OrderSide::Sell, OrderSide::Buy)
);
if is_closing {
let closed_qty = entry.quantity.min(fill_qty);
let entry_dec =
entry.entry_price.to_decimal().unwrap_or(Decimal::ZERO);
let exit_dec = fill_price.to_decimal().unwrap_or(Decimal::ZERO);
// PnL: long close = (exit - entry) * qty
// short close = (entry - exit) * qty
let pnl = if entry.side == OrderSide::Buy {
(exit_dec - entry_dec) * closed_qty
} else {
(entry_dec - exit_dec) * closed_qty
};
let return_pct = if entry_dec != Decimal::ZERO {
pnl / (entry_dec * closed_qty) * Decimal::from(100)
} else {
Decimal::ZERO
};
let trade = TradeRecord {
trade_id: uuid::Uuid::new_v4().to_string(),
symbol: order.symbol.clone(),
side: entry.side.clone(),
entry_price: entry.entry_price,
exit_price: fill_price,
quantity: order.quantity,
entry_time: entry.entry_time,
exit_time: context.current_time,
pnl,
return_pct,
commission: Decimal::ZERO,
};
tracker.total_trades += 1;
tracker.total_pnl += pnl;
if pnl > Decimal::ZERO {
tracker.winning_trades += 1;
}
tracker.trade_records.push(trade);
// Update or remove the open entry
let remaining = entry.quantity - closed_qty;
if remaining <= Decimal::ZERO {
tracker.open_entries.remove(&symbol_key);
} else {
if let Some(e) = tracker.open_entries.get_mut(&symbol_key) {
e.quantity = remaining;
}
}
// If the fill is larger than the old entry, open a new entry
// for the remainder (position flip)
let excess = fill_qty - closed_qty;
if excess > Decimal::ZERO {
tracker.open_entries.insert(
symbol_key.clone(),
OpenEntry {
side: order.side.clone(),
entry_price: fill_price,
quantity: excess,
entry_time: context.current_time,
},
);
}
closed = true;
}
}
if !closed {
// Opening a new position or adding to existing
let entry = tracker
.open_entries
.entry(symbol_key)
.or_insert_with(|| OpenEntry {
side: order.side.clone(),
entry_price: fill_price,
quantity: Decimal::ZERO,
entry_time: context.current_time,
});
entry.quantity += fill_qty;
// Update average entry price (weighted average)
// This simple version just keeps the latest; for proper
// accounting the PositionTracker in strategy_tester handles it
}
// Take an equity snapshot
tracker
.equity_snapshots
.push((context.current_time, context.account_balance));
info!(
"Order filled: {} {} @ {} | total_pnl={} trades={} wins={}",
order.side,
order.quantity,
fill_price,
tracker.total_pnl,
tracker.total_trades,
tracker.winning_trades,
);
}
Ok(())
}
async fn on_position_update(
&mut self,
position: &Position,
context: &StrategyContext,
) -> Result<()> {
// Update market state with current position
{
let mut state = self.market_state.write();
state.current_position = Some(position.clone());
}
// Update performance tracking
{
let mut tracker = self.performance_tracker.write();
if context.account_balance > tracker.peak_value {
tracker.peak_value = context.account_balance;
tracker.current_drawdown = Decimal::ZERO;
} else if tracker.peak_value > Decimal::ZERO {
tracker.current_drawdown =
(tracker.peak_value - context.account_balance) / tracker.peak_value;
if tracker.current_drawdown > tracker.max_drawdown {
tracker.max_drawdown = tracker.current_drawdown;
}
}
// Record equity snapshot on every position update
tracker
.equity_snapshots
.push((context.current_time, context.account_balance));
}
Ok(())
}
async fn finalize(&mut self, context: &StrategyContext) -> Result<StrategyResult> {
let tracker = self.performance_tracker.read();
let baseline = if tracker.initial_capital > Decimal::ZERO {
tracker.initial_capital
} else {
tracker.peak_value
};
let total_return = if baseline > Decimal::ZERO {
(context.account_balance - baseline) / baseline
} else {
Decimal::ZERO
};
let win_rate = if tracker.total_trades > 0 {
Decimal::from(tracker.winning_trades) / Decimal::from(tracker.total_trades)
} else {
Decimal::ZERO
};
// Calculate Sharpe ratio (simplified)
let sharpe_ratio = if tracker.max_drawdown > Decimal::ZERO {
total_return / tracker.max_drawdown
} else {
Decimal::ZERO
};
// Build performance timeline from equity snapshots
let performance_timeline: Vec<PerformanceSnapshot> = tracker
.equity_snapshots
.iter()
.map(|(ts, value)| {
let realized = tracker.total_pnl;
let dd = if tracker.peak_value > Decimal::ZERO {
(tracker.peak_value - *value) / tracker.peak_value
} else {
Decimal::ZERO
};
PerformanceSnapshot {
timestamp: *ts,
portfolio_value: *value,
cash_balance: *value,
unrealized_pnl: Decimal::ZERO,
realized_pnl: realized,
open_positions: 0,
drawdown: if dd > Decimal::ZERO { dd } else { Decimal::ZERO },
}
})
.collect();
let trades = tracker.trade_records.clone();
Ok(StrategyResult {
strategy_name: "adaptive_ml_strategy".to_string(),
total_return,
annualized_return: total_return, // Simplified
max_drawdown: tracker.max_drawdown,
sharpe_ratio,
total_trades: tracker.total_trades,
win_rate,
avg_trade_return: if tracker.total_trades > 0 {
tracker.total_pnl / Decimal::from(tracker.total_trades)
} else {
Decimal::ZERO
},
final_value: context.account_balance,
trades,
performance_timeline,
})
}
async fn get_state(&self) -> Result<serde_json::Value> {
let market_state = self.market_state.read();
let tracker = self.performance_tracker.read();
let cache = &self.predictions_cache;
Ok(serde_json::json!({
"strategy_name": "adaptive_ml_strategy",
"config": self.config,
"current_time": market_state.current_time,
"price_history_length": market_state.price_history.len(),
"volume_history_length": market_state.volume_history.len(),
"current_position": market_state.current_position,
"total_trades": tracker.total_trades,
"max_drawdown": tracker.max_drawdown,
"cached_predictions": cache.len(),
"active_models": self.config.active_models,
}))
}
}
/// Create a configured adaptive strategy runner
///
/// # Returns
/// * `AdaptiveStrategyRunner` - Strategy runner with default configuration
pub fn create_adaptive_strategy() -> AdaptiveStrategyRunner {
AdaptiveStrategyRunner::new(AdaptiveStrategyConfig::default())
}
/// Create adaptive strategy with custom configuration
///
/// # Arguments
/// * `config` - Custom adaptive strategy configuration
///
/// # Returns
/// * `AdaptiveStrategyRunner` - Strategy runner with specified configuration
pub fn create_adaptive_strategy_with_config(
config: AdaptiveStrategyConfig,
) -> AdaptiveStrategyRunner {
AdaptiveStrategyRunner::new(config)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_adaptive_strategy_config_default() {
let config = AdaptiveStrategyConfig::default();
assert_eq!(config.active_models.len(), 5);
assert_eq!(config.min_confidence, 0.65);
assert!(config.max_position_size > 0.0);
}
#[test]
fn test_risk_settings_default() {
let risk = RiskSettings::default();
assert_eq!(risk.max_drawdown, 0.10);
assert!(risk.kelly_fraction > 0.0);
}
#[tokio::test]
async fn test_adaptive_strategy_creation() {
let strategy = create_adaptive_strategy();
assert_eq!(strategy.name(), "adaptive_ml_strategy");
}
#[tokio::test]
async fn test_feature_extractor() {
let config = FeatureSettings::default();
let extractor = FeatureExtractor::new(config);
let mut market_state = MarketState::default();
market_state.price_history = vec![
(Utc::now(), Decimal::from(100)),
(Utc::now(), Decimal::from(101)),
(Utc::now(), Decimal::from(102)),
];
let features = extractor.extract_features(&market_state).await;
assert!(features.is_ok());
}
#[test]
fn test_production_features_have_51_dimensions() -> Result<(), Box<dyn std::error::Error>> {
let mut extractor = ProductionFeatureExtractorAdapter::new();
// Feed 55 price updates (past the warmup period of 50)
for i in 0..55 {
let price = 100.0 + i as f64 * 0.1;
let volume = 1000.0;
let timestamp = Utc::now();
extractor.update(price, volume, timestamp)?;
}
let features = extractor.extract_features()?;
assert_eq!(
features.len(),
51,
"Production extractor should produce exactly 51 features, got {}",
features.len()
);
// Validate all features are finite
for (i, val) in features.iter().enumerate() {
assert!(
val.is_finite(),
"Feature {} should be finite, found {}",
i, val
);
}
Ok(())
}
#[test]
fn test_production_extractor_warmup_skips_gracefully() {
let mut extractor = ProductionFeatureExtractorAdapter::new();
// Feed only 5 price updates (below warmup threshold)
for i in 0..5 {
let price = 100.0 + i as f64 * 0.1;
let _ = extractor.update(price, 1000.0, Utc::now());
}
// Should either return an error or return a short vector
// (during warmup, extraction may fail)
let result = extractor.extract_features();
// We accept either an error (warmup not ready) or a valid 51-dim result
match result {
Ok(features) => {
assert_eq!(features.len(), 51);
}
Err(_) => {
// Expected during warmup - this is fine
}
}
}
#[test]
fn test_adaptive_strategy_has_production_extractor() {
let strategy = create_adaptive_strategy();
// Verify the production extractor is initialized
let ext = strategy.production_extractor.read();
// Should be able to access the extractor without issues
drop(ext);
}
#[test]
fn test_performance_tracker_records_trade() {
use crate::strategy_tester::TradeRecord;
let mut tracker = PerformanceTracker::default();
tracker.initial_capital = Decimal::from(100_000);
tracker.peak_value = Decimal::from(100_000);
// Simulate opening a long entry
tracker.open_entries.insert(
"AAPL".to_string(),
OpenEntry {
side: OrderSide::Buy,
entry_price: Price::from_f64(100.0).unwrap(),
quantity: Decimal::from(10),
entry_time: Utc::now(),
},
);
// Simulate a winning close: sold at 110
let pnl = (Decimal::from(110) - Decimal::from(100)) * Decimal::from(10); // +100
tracker.total_trades += 1;
tracker.total_pnl += pnl;
if pnl > Decimal::ZERO {
tracker.winning_trades += 1;
}
tracker.trade_records.push(TradeRecord {
trade_id: "test-1".to_string(),
symbol: Symbol::from("AAPL"),
side: OrderSide::Buy,
entry_price: Price::from_f64(100.0).unwrap(),
exit_price: Price::from_f64(110.0).unwrap(),
quantity: Quantity::from_f64(10.0).unwrap(),
entry_time: Utc::now(),
exit_time: Utc::now(),
pnl,
return_pct: Decimal::from(10),
commission: Decimal::ZERO,
});
tracker.open_entries.remove("AAPL");
assert_eq!(tracker.total_trades, 1);
assert_eq!(tracker.winning_trades, 1);
assert_eq!(tracker.total_pnl, Decimal::from(100));
assert_eq!(tracker.trade_records.len(), 1);
assert_eq!(tracker.trade_records.first().map(|t| &t.symbol), Some(&Symbol::from("AAPL")));
}
#[test]
fn test_performance_tracker_losing_trade() {
let mut tracker = PerformanceTracker::default();
tracker.initial_capital = Decimal::from(100_000);
tracker.peak_value = Decimal::from(100_000);
// Simulate opening a long entry
tracker.open_entries.insert(
"TSLA".to_string(),
OpenEntry {
side: OrderSide::Buy,
entry_price: Price::from_f64(200.0).unwrap(),
quantity: Decimal::from(5),
entry_time: Utc::now(),
},
);
// Simulate a losing close: sold at 180
let pnl = (Decimal::from(180) - Decimal::from(200)) * Decimal::from(5); // -100
tracker.total_trades += 1;
tracker.total_pnl += pnl;
if pnl > Decimal::ZERO {
tracker.winning_trades += 1;
}
tracker.trade_records.push(TradeRecord {
trade_id: "test-2".to_string(),
symbol: Symbol::from("TSLA"),
side: OrderSide::Buy,
entry_price: Price::from_f64(200.0).unwrap(),
exit_price: Price::from_f64(180.0).unwrap(),
quantity: Quantity::from_f64(5.0).unwrap(),
entry_time: Utc::now(),
exit_time: Utc::now(),
pnl,
return_pct: Decimal::from(-10),
commission: Decimal::ZERO,
});
tracker.open_entries.remove("TSLA");
assert_eq!(tracker.total_trades, 1);
assert_eq!(tracker.winning_trades, 0);
assert_eq!(tracker.total_pnl, Decimal::from(-100));
assert_eq!(tracker.trade_records.len(), 1);
}
#[tokio::test]
async fn test_finalize_populates_trades() {
let mut runner = create_adaptive_strategy();
let initial_capital = Decimal::from(100_000);
// Initialize the runner so peak_value / initial_capital are set
runner
.initialize(initial_capital, StrategyConfig::default())
.await
.unwrap();
// Manually inject trade records and equity snapshots into the tracker
{
let mut tracker = runner.performance_tracker.write();
tracker.total_trades = 2;
tracker.winning_trades = 1;
tracker.total_pnl = Decimal::from(50); // net +50
tracker.trade_records.push(TradeRecord {
trade_id: "finalize-1".to_string(),
symbol: Symbol::from("AAPL"),
side: OrderSide::Buy,
entry_price: Price::from_f64(100.0).unwrap(),
exit_price: Price::from_f64(110.0).unwrap(),
quantity: Quantity::from_f64(10.0).unwrap(),
entry_time: Utc::now(),
exit_time: Utc::now(),
pnl: Decimal::from(100),
return_pct: Decimal::from(10),
commission: Decimal::ZERO,
});
tracker.trade_records.push(TradeRecord {
trade_id: "finalize-2".to_string(),
symbol: Symbol::from("MSFT"),
side: OrderSide::Buy,
entry_price: Price::from_f64(300.0).unwrap(),
exit_price: Price::from_f64(290.0).unwrap(),
quantity: Quantity::from_f64(5.0).unwrap(),
entry_time: Utc::now(),
exit_time: Utc::now(),
pnl: Decimal::from(-50),
return_pct: Decimal::new(-333, 2),
commission: Decimal::ZERO,
});
// Add equity snapshots
let now = Utc::now();
tracker.equity_snapshots.push((now, Decimal::from(100_100)));
tracker.equity_snapshots.push((now, Decimal::from(100_050)));
}
// Build a context with the final balance
let context = StrategyContext {
current_time: Utc::now(),
account_balance: Decimal::from(100_050),
buying_power: Decimal::from(100_050),
positions: HashMap::new(),
open_orders: HashMap::new(),
market_prices: HashMap::new(),
performance: crate::strategy_tester::PerformanceMetrics::default(),
};
let result = runner.finalize(&context).await.unwrap();
// Trades should be populated
assert_eq!(result.trades.len(), 2, "finalize() must return trade records");
assert_eq!(result.trades.first().map(|t| t.trade_id.as_str()), Some("finalize-1"));
// Performance timeline should be populated from equity snapshots
assert_eq!(
result.performance_timeline.len(),
2,
"finalize() must return performance timeline"
);
// Verify aggregated stats
assert_eq!(result.total_trades, 2);
assert_eq!(result.win_rate, Decimal::from(1) / Decimal::from(2));
assert_eq!(result.avg_trade_return, Decimal::from(25)); // 50 / 2
assert_eq!(result.final_value, Decimal::from(100_050));
}
#[tokio::test]
async fn test_on_order_update_tracks_pnl() {
let mut runner = create_adaptive_strategy();
let initial_capital = Decimal::from(100_000);
runner
.initialize(initial_capital, StrategyConfig::default())
.await
.unwrap();
let now = Utc::now();
let context = StrategyContext {
current_time: now,
account_balance: initial_capital,
buying_power: initial_capital,
positions: HashMap::new(),
open_orders: HashMap::new(),
market_prices: HashMap::new(),
performance: crate::strategy_tester::PerformanceMetrics::default(),
};
// Step 1: simulate a Buy fill at $100 (opening)
let buy_order = Order {
id: "order_1".to_string().into(),
client_order_id: None,
broker_order_id: None,
account_id: None,
symbol: Symbol::from("AAPL"),
side: OrderSide::Buy,
order_type: common::OrderType::Market,
status: OrderStatus::Filled,
time_in_force: common::TimeInForce::Day,
quantity: Quantity::from_f64(10.0).unwrap(),
price: Some(Price::from_f64(100.0).unwrap()),
stop_price: None,
filled_quantity: Quantity::from_f64(10.0).unwrap(),
remaining_quantity: Quantity::ZERO,
average_price: Some(Price::from_f64(100.0).unwrap()),
avg_fill_price: None,
average_fill_price: None,
exchange_order_id: None,
parent_id: None,
execution_algorithm: None,
execution_params: serde_json::json!({}),
stop_loss: None,
take_profit: None,
created_at: common::HftTimestamp::now_or_zero(),
updated_at: None,
expires_at: None,
metadata: serde_json::json!({}),
};
runner.on_order_update(&buy_order, &context).await.unwrap();
// After a buy open, no round-trip trade yet
{
let tracker = runner.performance_tracker.read();
assert_eq!(tracker.total_trades, 0, "Opening buy should not count as a completed trade");
assert!(tracker.open_entries.contains_key("AAPL"));
}
// Step 2: simulate a Sell fill at $110 (closing the long)
let sell_order = Order {
id: "order_2".to_string().into(),
client_order_id: None,
broker_order_id: None,
account_id: None,
symbol: Symbol::from("AAPL"),
side: OrderSide::Sell,
order_type: common::OrderType::Market,
status: OrderStatus::Filled,
time_in_force: common::TimeInForce::Day,
quantity: Quantity::from_f64(10.0).unwrap(),
price: Some(Price::from_f64(110.0).unwrap()),
stop_price: None,
filled_quantity: Quantity::from_f64(10.0).unwrap(),
remaining_quantity: Quantity::ZERO,
average_price: Some(Price::from_f64(110.0).unwrap()),
avg_fill_price: None,
average_fill_price: None,
exchange_order_id: None,
parent_id: None,
execution_algorithm: None,
execution_params: serde_json::json!({}),
stop_loss: None,
take_profit: None,
created_at: common::HftTimestamp::now_or_zero(),
updated_at: None,
expires_at: None,
metadata: serde_json::json!({}),
};
let close_context = StrategyContext {
current_time: now,
account_balance: Decimal::from(100_100), // gained 100
buying_power: Decimal::from(100_100),
positions: HashMap::new(),
open_orders: HashMap::new(),
market_prices: HashMap::new(),
performance: crate::strategy_tester::PerformanceMetrics::default(),
};
runner
.on_order_update(&sell_order, &close_context)
.await
.unwrap();
// After close, we should have a completed trade with PnL = +100
{
let tracker = runner.performance_tracker.read();
assert_eq!(tracker.total_trades, 1, "Closing sell should create a completed trade");
assert_eq!(tracker.winning_trades, 1);
assert_eq!(tracker.total_pnl, Decimal::from(100));
assert_eq!(tracker.trade_records.len(), 1);
assert!(!tracker.open_entries.contains_key("AAPL"), "Entry should be removed after close");
// Equity snapshots recorded on each order update
assert!(tracker.equity_snapshots.len() >= 2);
}
}
}