Files
foxhunt/services/backtesting_service/tests/helpers.rs
jgrusewski 52630a77d3 perf: eliminate heap-alloc Decimal→float casts across 19 files (36 instances)
Replace all `.to_string().parse::<f32/f64>()` patterns with
`num_traits::ToPrimitive` methods (`.to_f32()`, `.to_f64()`).
Each string roundtrip heap-allocated per conversion — fatal in
DQN hot loop (300K+ bars × epochs). Decimal stays as canonical
financial type; conversions happen at GPU/float boundaries only.

Also fixes blocking_read() in async context (risk_integration.rs).

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-02-28 02:29:04 +01:00

654 lines
18 KiB
Rust

//! Test Helper Utilities for Data Validation
//!
//! Provides reusable assertion and validation functions for backtesting tests.
//! All functions are designed to give clear, actionable error messages.
//!
//! # Categories
//!
//! - **OHLCV Validation**: Price relationship checks
//! - **Time Series Validation**: Chronological ordering, continuity
//! - **Statistical Validation**: Price ranges, volatility bounds
//! - **Trade Validation**: PnL calculations, execution logic
//!
//! # Usage
//!
//! ```rust
//! use helpers::{assert_valid_ohlcv, assert_chronological, assert_price_range};
//!
//! #[test]
//! fn test_market_data_quality() {
//! let bars = load_test_data();
//! assert_valid_ohlcv(&bars);
//! assert_chronological(&bars);
//! assert_price_range(&bars, "ES.FUT");
//! }
//! ```
use backtesting_service::strategy_engine::{BacktestTrade, MarketData, TradeSide};
use num_traits::ToPrimitive;
use rust_decimal::Decimal;
// ============================================================================
// OHLCV Validation
// ============================================================================
/// Assert that OHLCV bars have valid price relationships
///
/// # Validates
///
/// - High >= Low (always)
/// - High >= Open, Close
/// - Low <= Open, Close
/// - All prices > 0
/// - Volume >= 0
///
/// # Panics
///
/// On first invalid bar with detailed error message
///
/// # Example
///
/// ```rust
/// assert_valid_ohlcv(&bars);
/// ```
pub fn assert_valid_ohlcv(bars: &[MarketData]) {
for (i, bar) in bars.iter().enumerate() {
// Basic positivity checks
assert!(
bar.open > Decimal::ZERO,
"Bar {}: Open must be positive, got {}",
i,
bar.open
);
assert!(
bar.high > Decimal::ZERO,
"Bar {}: High must be positive, got {}",
i,
bar.high
);
assert!(
bar.low > Decimal::ZERO,
"Bar {}: Low must be positive, got {}",
i,
bar.low
);
assert!(
bar.close > Decimal::ZERO,
"Bar {}: Close must be positive, got {}",
i,
bar.close
);
assert!(
bar.volume >= Decimal::ZERO,
"Bar {}: Volume must be non-negative, got {}",
i,
bar.volume
);
// OHLCV relationship checks
assert!(
bar.high >= bar.low,
"Bar {}: High ({}) must be >= Low ({})",
i,
bar.high,
bar.low
);
assert!(
bar.high >= bar.open,
"Bar {}: High ({}) must be >= Open ({})",
i,
bar.high,
bar.open
);
assert!(
bar.high >= bar.close,
"Bar {}: High ({}) must be >= Close ({})",
i,
bar.high,
bar.close
);
assert!(
bar.low <= bar.open,
"Bar {}: Low ({}) must be <= Open ({})",
i,
bar.low,
bar.open
);
assert!(
bar.low <= bar.close,
"Bar {}: Low ({}) must be <= Close ({})",
i,
bar.low,
bar.close
);
// Open and Close must be within [Low, High]
assert!(
bar.open >= bar.low && bar.open <= bar.high,
"Bar {}: Open ({}) must be within [{}, {}]",
i,
bar.open,
bar.low,
bar.high
);
assert!(
bar.close >= bar.low && bar.close <= bar.high,
"Bar {}: Close ({}) must be within [{}, {}]",
i,
bar.close,
bar.low,
bar.high
);
}
}
// ============================================================================
// Time Series Validation
// ============================================================================
/// Assert that bars are sorted chronologically
///
/// # Panics
///
/// If any bar has timestamp <= previous bar
///
/// # Example
///
/// ```rust
/// assert_chronological(&bars);
/// ```
pub fn assert_chronological(bars: &[MarketData]) {
for i in 1..bars.len() {
assert!(
bars[i].timestamp >= bars[i - 1].timestamp,
"Bar {}: Timestamps not chronological. Bar[{}]={:?}, Bar[{}]={:?}",
i,
i - 1,
bars[i - 1].timestamp,
i,
bars[i].timestamp
);
}
}
/// Assert that bars have no gaps larger than specified duration
///
/// # Arguments
///
/// * `bars` - Market data bars
/// * `max_gap_minutes` - Maximum allowed gap in minutes
///
/// # Example
///
/// ```rust
/// // Assert no gaps > 5 minutes (for 1-minute data)
/// assert_no_large_gaps(&bars, 5);
/// ```
pub fn assert_no_large_gaps(bars: &[MarketData], max_gap_minutes: i64) {
use chrono::Duration;
let max_gap = Duration::minutes(max_gap_minutes);
for i in 1..bars.len() {
let gap = bars[i].timestamp - bars[i - 1].timestamp;
assert!(
gap <= max_gap,
"Bar {}: Gap too large ({:?} > {:?}). Bar[{}]={:?}, Bar[{}]={:?}",
i,
gap,
max_gap,
i - 1,
bars[i - 1].timestamp,
i,
bars[i].timestamp
);
}
}
// ============================================================================
// Statistical Validation
// ============================================================================
/// Assert that prices are within realistic range for a symbol
///
/// # Symbol Ranges (2024 typical)
///
/// - ES.FUT: 3000 - 6000
/// - NQ.FUT: 12000 - 20000
/// - CL.FUT: 50 - 100
///
/// # Panics
///
/// If any price is outside realistic range
///
/// # Example
///
/// ```rust
/// assert_price_range(&bars, "ES.FUT");
/// ```
pub fn assert_price_range(bars: &[MarketData], symbol: &str) {
let (min_price, max_price) = match symbol {
"ES.FUT" => (3000.0, 6000.0),
"NQ.FUT" => (12000.0, 20000.0),
"CL.FUT" => (50.0, 100.0),
_ => (0.1, 1_000_000.0), // Very permissive for unknown symbols
};
for (i, bar) in bars.iter().enumerate() {
let close = bar.close.to_f64().unwrap_or(0.0);
assert!(
close >= min_price && close <= max_price,
"Bar {}: {} price {} outside realistic range [{}, {}]",
i,
symbol,
close,
min_price,
max_price
);
}
}
/// Assert that volatility is within realistic bounds
///
/// Calculates return volatility and checks against expected ranges.
///
/// # Arguments
///
/// * `bars` - Market data (minimum 20 bars required)
/// * `max_volatility_pct` - Maximum expected annualized volatility (%)
///
/// # Example
///
/// ```rust
/// // Assert volatility < 100% annualized
/// assert_volatility_bounds(&bars, 100.0);
/// ```
pub fn assert_volatility_bounds(bars: &[MarketData], max_volatility_pct: f64) {
if bars.len() < 20 {
return; // Need sufficient data for volatility calculation
}
// Calculate returns
let mut returns = Vec::new();
for i in 1..bars.len() {
let prev_close = bars[i - 1].close.to_f64().unwrap_or(1.0);
let curr_close = bars[i].close.to_f64().unwrap_or(1.0);
let ret = (curr_close - prev_close) / prev_close;
returns.push(ret);
}
// Calculate standard deviation
let mean = returns.iter().sum::<f64>() / returns.len() as f64;
let variance: f64 =
returns.iter().map(|r| (r - mean).powi(2)).sum::<f64>() / returns.len() as f64;
let std_dev = variance.sqrt();
// Annualize (assume 252 trading days, 390 1-minute bars per day)
let bars_per_year: f64 = 252.0 * 390.0;
let annualized_volatility = std_dev * (bars_per_year).sqrt() * 100.0;
assert!(
annualized_volatility <= max_volatility_pct,
"Volatility ({:.2}%) exceeds maximum ({:.2}%)",
annualized_volatility,
max_volatility_pct
);
}
/// Calculate and return actual volatility (for informational purposes)
///
/// # Returns
///
/// Annualized volatility as percentage
pub fn calculate_volatility(bars: &[MarketData]) -> f64 {
if bars.len() < 2 {
return 0.0;
}
let mut returns = Vec::new();
for i in 1..bars.len() {
let prev_close = bars[i - 1].close.to_f64().unwrap_or(1.0);
let curr_close = bars[i].close.to_f64().unwrap_or(1.0);
let ret = (curr_close - prev_close) / prev_close;
returns.push(ret);
}
let mean = returns.iter().sum::<f64>() / returns.len() as f64;
let variance: f64 =
returns.iter().map(|r| (r - mean).powi(2)).sum::<f64>() / returns.len() as f64;
let std_dev = variance.sqrt();
// Annualize
let bars_per_year: f64 = 252.0 * 390.0;
std_dev * (bars_per_year).sqrt() * 100.0
}
// ============================================================================
// Trade Validation
// ============================================================================
/// Assert that a trade has valid structure
///
/// # Validates
///
/// - Exit time > entry time
/// - Exit price > 0
/// - Entry price > 0
/// - Quantity > 0
/// - PnL calculation correct for side
///
/// # Example
///
/// ```rust
/// assert_valid_trade(&trade);
/// ```
pub fn assert_valid_trade(trade: &BacktestTrade) {
// Time validation
assert!(
trade.exit_time > trade.entry_time,
"Trade {}: Exit time must be after entry time",
trade.trade_id
);
// Price validation
assert!(
trade.entry_price > Decimal::ZERO,
"Trade {}: Entry price must be positive",
trade.trade_id
);
assert!(
trade.exit_price > Decimal::ZERO,
"Trade {}: Exit price must be positive",
trade.trade_id
);
// Quantity validation
assert!(
trade.quantity > Decimal::ZERO,
"Trade {}: Quantity must be positive",
trade.trade_id
);
// PnL calculation validation
let entry = trade.entry_price.to_f64().unwrap_or(0.0);
let exit = trade.exit_price.to_f64().unwrap_or(0.0);
let qty = trade.quantity.to_f64().unwrap_or(0.0);
let expected_pnl = match trade.side {
TradeSide::Buy => (exit - entry) * qty,
TradeSide::Sell => (entry - exit) * qty,
};
let actual_pnl = trade.pnl.to_f64().unwrap_or(0.0);
// Allow small floating point errors
let pnl_diff = (actual_pnl - expected_pnl).abs();
assert!(
pnl_diff < 0.01,
"Trade {}: PnL mismatch. Expected {:.4}, got {:.4} (diff: {:.6})",
trade.trade_id,
expected_pnl,
actual_pnl,
pnl_diff
);
}
/// Assert that a list of trades has valid sequence
///
/// # Validates
///
/// - No overlapping trades (same symbol)
/// - Chronological order
/// - All trades individually valid
///
/// # Example
///
/// ```rust
/// assert_valid_trade_sequence(&trades);
/// ```
pub fn assert_valid_trade_sequence(trades: &[BacktestTrade]) {
// Validate each trade
for trade in trades {
assert_valid_trade(trade);
}
// Check chronological order
for i in 1..trades.len() {
assert!(
trades[i].entry_time >= trades[i - 1].entry_time,
"Trade {}: Trades not in chronological order",
i
);
}
// Check for overlapping trades (same symbol)
for i in 0..trades.len() {
for j in (i + 1)..trades.len() {
if trades[i].symbol == trades[j].symbol {
// If same symbol, ensure no time overlap
let overlap = trades[j].entry_time < trades[i].exit_time;
assert!(
!overlap,
"Trades {} and {}: Overlapping trades for symbol {}",
i, j, trades[i].symbol
);
}
}
}
}
// ============================================================================
// Performance Metrics Validation
// ============================================================================
/// Assert that Sharpe ratio is within realistic bounds
///
/// # Arguments
///
/// * `sharpe` - Sharpe ratio value
/// * `min` - Minimum realistic value (typically -3.0)
/// * `max` - Maximum realistic value (typically 5.0)
///
/// # Example
///
/// ```rust
/// assert_sharpe_bounds(sharpe_ratio, -3.0, 5.0);
/// ```
pub fn assert_sharpe_bounds(sharpe: f64, min: f64, max: f64) {
assert!(
sharpe >= min && sharpe <= max,
"Sharpe ratio {:.2} outside realistic bounds [{:.2}, {:.2}]",
sharpe,
min,
max
);
}
/// Assert that drawdown is positive and within bounds
///
/// # Arguments
///
/// * `drawdown` - Max drawdown as positive percentage (e.g., 15.5 for 15.5%)
/// * `max_drawdown` - Maximum acceptable drawdown percentage
///
/// # Example
///
/// ```rust
/// assert_drawdown_bounds(max_dd, 50.0); // Max 50% drawdown
/// ```
pub fn assert_drawdown_bounds(drawdown: f64, max_drawdown: f64) {
assert!(
drawdown >= 0.0,
"Drawdown must be non-negative, got {:.2}%",
drawdown
);
assert!(
drawdown <= max_drawdown,
"Drawdown {:.2}% exceeds maximum {:.2}%",
drawdown,
max_drawdown
);
}
/// Assert that win rate is between 0% and 100%
pub fn assert_win_rate_valid(win_rate: f64) {
assert!(
win_rate >= 0.0 && win_rate <= 100.0,
"Win rate must be between 0% and 100%, got {:.2}%",
win_rate
);
}
// ============================================================================
// Data Quality Reports
// ============================================================================
/// Generate a comprehensive data quality report
///
/// # Returns
///
/// String with detailed quality metrics
///
/// # Example
///
/// ```rust
/// let report = generate_quality_report(&bars);
/// println!("{}", report);
/// ```
pub fn generate_quality_report(bars: &[MarketData]) -> String {
if bars.is_empty() {
return "No data to analyze".to_string();
}
let mut report = String::new();
report.push_str("=== Data Quality Report ===\n\n");
// Basic stats
report.push_str(&format!("Total bars: {}\n", bars.len()));
report.push_str(&format!("Symbol: {}\n", bars[0].symbol));
report.push_str(&format!(
"Date range: {} to {}\n",
bars[0].timestamp,
bars[bars.len() - 1].timestamp
));
// Price stats
let prices: Vec<f64> = bars
.iter()
.map(|b| b.close.to_f64().unwrap_or(0.0))
.collect();
let min_price = prices.iter().cloned().fold(f64::INFINITY, f64::min);
let max_price = prices.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
let avg_price = prices.iter().sum::<f64>() / prices.len() as f64;
report.push_str(&format!("\nPrice Statistics:\n"));
report.push_str(&format!(" Min: {:.2}\n", min_price));
report.push_str(&format!(" Max: {:.2}\n", max_price));
report.push_str(&format!(" Avg: {:.2}\n", avg_price));
report.push_str(&format!(
" Range: {:.2}%\n",
((max_price - min_price) / avg_price) * 100.0
));
// Volatility
let volatility = calculate_volatility(bars);
report.push_str(&format!("\nVolatility:\n"));
report.push_str(&format!(" Annualized: {:.2}%\n", volatility));
// Data quality checks
report.push_str(&format!("\nQuality Checks:\n"));
let mut ohlcv_errors = 0;
for bar in bars {
if bar.high < bar.low
|| bar.high < bar.open
|| bar.high < bar.close
|| bar.low > bar.open
|| bar.low > bar.close
{
ohlcv_errors += 1;
}
}
report.push_str(&format!(" OHLCV errors: {}\n", ohlcv_errors));
let mut chronology_errors = 0;
for i in 1..bars.len() {
if bars[i].timestamp < bars[i - 1].timestamp {
chronology_errors += 1;
}
}
report.push_str(&format!(" Chronology errors: {}\n", chronology_errors));
report.push_str(&format!("\n=== End Report ===\n"));
report
}
// ============================================================================
// Unit Tests
// ============================================================================
#[cfg(test)]
mod tests {
use super::*;
use backtesting_service::strategy_engine::TimeFrame;
use chrono::Utc;
fn create_valid_bar() -> MarketData {
MarketData {
symbol: "TEST".to_string(),
timestamp: Utc::now(),
open: Decimal::from_f64_retain(100.0).unwrap(),
high: Decimal::from_f64_retain(105.0).unwrap(),
low: Decimal::from_f64_retain(95.0).unwrap(),
close: Decimal::from_f64_retain(102.0).unwrap(),
volume: Decimal::from_f64_retain(10000.0).unwrap(),
timeframe: TimeFrame::Daily,
}
}
#[test]
fn test_valid_ohlcv() {
let bars = vec![create_valid_bar()];
assert_valid_ohlcv(&bars);
}
#[test]
#[should_panic(expected = "High")]
fn test_invalid_ohlcv_high_low() {
let mut bar = create_valid_bar();
bar.high = Decimal::from_f64_retain(90.0).unwrap(); // High < Low
assert_valid_ohlcv(&vec![bar]);
}
#[test]
fn test_chronological() {
use chrono::Duration;
let mut bars = vec![create_valid_bar()];
let mut bar2 = create_valid_bar();
bar2.timestamp = bars[0].timestamp + Duration::minutes(1);
bars.push(bar2);
assert_chronological(&bars);
}
#[test]
#[should_panic(expected = "chronological")]
fn test_non_chronological() {
use chrono::Duration;
let mut bars = vec![create_valid_bar()];
let mut bar2 = create_valid_bar();
bar2.timestamp = bars[0].timestamp - Duration::minutes(1); // Earlier
bars.push(bar2);
assert_chronological(&bars);
}
#[test]
fn test_quality_report() {
let bars = vec![create_valid_bar()];
let report = generate_quality_report(&bars);
assert!(report.contains("Total bars: 1"));
}
}