use std::path::{Component, PathBuf};
use rust_decimal::Decimal;
use serde::{Deserialize, Serialize};
use crate::data::DataSourceSpec;
use crate::domain::ExecutionMode;
use crate::error::BacktestError;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct FeeSchedule {
pub per_contract_cents: u64,
pub per_order_cents: u64,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "model", rename_all = "snake_case")]
pub enum SlippageModel {
None,
FixedCents {
cents: u64,
},
SpreadFraction {
fraction: Decimal,
},
SizeProportional {
cents_per_contract: u64,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum TouchSize {
QuotedSize,
Flat {
contracts: u32,
},
}
impl Default for TouchSize {
fn default() -> Self {
Self::QuotedSize
}
}
const fn default_depth_levels() -> u32 {
5
}
fn default_decay() -> Decimal {
Decimal::new(5, 1)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct LiquidityProfile {
#[serde(default)]
pub touch_size: TouchSize,
#[serde(default = "default_depth_levels")]
pub depth_levels: u32,
#[serde(default = "default_decay")]
pub decay: Decimal,
}
impl Default for LiquidityProfile {
fn default() -> Self {
Self {
touch_size: TouchSize::default(),
depth_levels: default_depth_levels(),
decay: default_decay(),
}
}
}
impl LiquidityProfile {
pub const MAX_DEPTH_LEVELS: u32 = 64;
#[must_use = "an unvalidated liquidity profile must not reach the seeder"]
pub fn validate(&self) -> Result<(), BacktestError> {
if self.depth_levels > Self::MAX_DEPTH_LEVELS {
return Err(BacktestError::Config(format!(
"liquidity_profile.depth_levels = {} exceeds hard cap {}",
self.depth_levels,
Self::MAX_DEPTH_LEVELS
)));
}
if self.decay <= Decimal::ZERO || self.decay > Decimal::ONE {
return Err(BacktestError::Config(format!(
"liquidity_profile.decay must be in (0, 1], got {}",
self.decay
)));
}
if let TouchSize::Flat { contracts } = self.touch_size
&& contracts == 0
{
return Err(BacktestError::Config(
"liquidity_profile.touch_size flat contracts must be > 0, got 0".to_string(),
));
}
Ok(())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ResourceLimits {
pub max_steps: u64,
pub max_contracts_per_snapshot: u32,
pub max_total_bytes: u64,
pub max_file_bytes: u64,
pub max_decompressed_bytes: u64,
pub max_manifest_bytes: u64,
pub max_string_len: u32,
pub max_rows_per_table: u64,
}
impl ResourceLimits {
pub const CAP_MAX_STEPS: u64 = 10_000_000;
pub const CAP_MAX_CONTRACTS_PER_SNAPSHOT: u32 = 5_000_000;
pub const CAP_MAX_TOTAL_BYTES: u64 = 32 * (1 << 30);
pub const CAP_MAX_FILE_BYTES: u64 = 64 * (1 << 30);
pub const CAP_MAX_DECOMPRESSED_BYTES: u64 = 64 * (1 << 30);
pub const CAP_MAX_MANIFEST_BYTES: u64 = 256 * (1 << 20);
pub const CAP_MAX_STRING_LEN: u32 = 16 * (1 << 20);
pub const CAP_MAX_ROWS_PER_TABLE: u64 = 1 << 40;
#[must_use = "the limits are not enforced unless the result is checked"]
pub fn validate(&self) -> Result<(), BacktestError> {
fn over(field: &str, value: u64, cap: u64) -> BacktestError {
BacktestError::Config(format!("limits.{field} = {value} exceeds hard cap {cap}"))
}
if self.max_steps > Self::CAP_MAX_STEPS {
return Err(over("max_steps", self.max_steps, Self::CAP_MAX_STEPS));
}
if self.max_contracts_per_snapshot > Self::CAP_MAX_CONTRACTS_PER_SNAPSHOT {
return Err(over(
"max_contracts_per_snapshot",
u64::from(self.max_contracts_per_snapshot),
u64::from(Self::CAP_MAX_CONTRACTS_PER_SNAPSHOT),
));
}
if self.max_total_bytes > Self::CAP_MAX_TOTAL_BYTES {
return Err(over(
"max_total_bytes",
self.max_total_bytes,
Self::CAP_MAX_TOTAL_BYTES,
));
}
if self.max_file_bytes > Self::CAP_MAX_FILE_BYTES {
return Err(over(
"max_file_bytes",
self.max_file_bytes,
Self::CAP_MAX_FILE_BYTES,
));
}
if self.max_decompressed_bytes > Self::CAP_MAX_DECOMPRESSED_BYTES {
return Err(over(
"max_decompressed_bytes",
self.max_decompressed_bytes,
Self::CAP_MAX_DECOMPRESSED_BYTES,
));
}
if self.max_manifest_bytes > Self::CAP_MAX_MANIFEST_BYTES {
return Err(over(
"max_manifest_bytes",
self.max_manifest_bytes,
Self::CAP_MAX_MANIFEST_BYTES,
));
}
if self.max_string_len > Self::CAP_MAX_STRING_LEN {
return Err(over(
"max_string_len",
u64::from(self.max_string_len),
u64::from(Self::CAP_MAX_STRING_LEN),
));
}
if self.max_rows_per_table > Self::CAP_MAX_ROWS_PER_TABLE {
return Err(over(
"max_rows_per_table",
self.max_rows_per_table,
Self::CAP_MAX_ROWS_PER_TABLE,
));
}
Ok(())
}
}
impl Default for ResourceLimits {
fn default() -> Self {
Self {
max_steps: 100_000,
max_contracts_per_snapshot: 50_000,
max_total_bytes: 2 * (1 << 30), max_file_bytes: 4 * (1 << 30), max_decompressed_bytes: 8 * (1 << 30), max_manifest_bytes: 16 * (1 << 20), max_string_len: 64 * (1 << 10), max_rows_per_table: 100_000_000,
}
}
}
const fn default_marketable_cap_ticks() -> u32 {
10
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct BacktestConfig {
pub data_source: DataSourceSpec,
pub mode: ExecutionMode,
pub seed: u64,
pub initial_capital: u64,
pub fees: FeeSchedule,
pub slippage: SlippageModel,
#[serde(default = "default_marketable_cap_ticks")]
pub marketable_cap_ticks: u32,
#[serde(default)]
pub liquidity_profile: LiquidityProfile,
#[serde(default)]
pub limits: ResourceLimits,
pub output_dir: PathBuf,
#[serde(default)]
pub overwrite: bool,
}
impl BacktestConfig {
#[must_use = "an unvalidated config must not reach the engine"]
pub fn validate(&self) -> Result<(), BacktestError> {
if self.initial_capital == 0 {
return Err(BacktestError::Config(
"initial capital must be positive, got 0".to_string(),
));
}
if self.marketable_cap_ticks == 0 {
return Err(BacktestError::Config(
"marketable_cap_ticks must be positive, got 0".to_string(),
));
}
self.limits.validate()?;
self.liquidity_profile.validate()?;
if let SlippageModel::SpreadFraction { fraction } = &self.slippage
&& fraction.is_sign_negative()
{
return Err(BacktestError::Config(format!(
"slippage spread fraction must be non-negative, got {fraction}"
)));
}
if self.output_dir.as_os_str().is_empty() {
return Err(BacktestError::Config(
"output dir must not be empty".to_string(),
));
}
if self
.output_dir
.components()
.any(|c| matches!(c, Component::ParentDir))
{
return Err(BacktestError::Config(format!(
"output dir must not contain '..' components, got {}",
self.output_dir.display()
)));
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use rust_decimal_macros::dec;
use super::{
BacktestConfig, FeeSchedule, LiquidityProfile, ResourceLimits, SlippageModel, TouchSize,
};
use crate::data::DataSourceSpec;
use crate::domain::ExecutionMode;
use crate::error::BacktestError;
fn valid_config() -> BacktestConfig {
BacktestConfig {
data_source: DataSourceSpec::Parquet {
path: "chains/spx.parquet".to_string(),
sha256: String::new(),
},
mode: ExecutionMode::Naive,
seed: 42,
initial_capital: 10_000_000,
fees: FeeSchedule {
per_contract_cents: 65,
per_order_cents: 100,
},
slippage: SlippageModel::SpreadFraction {
fraction: dec!(0.5),
},
marketable_cap_ticks: 10,
liquidity_profile: LiquidityProfile::default(),
limits: ResourceLimits::default(),
output_dir: "runs/out".into(),
overwrite: false,
}
}
fn assert_config_error(result: Result<(), BacktestError>, needle: &str) {
match result {
Err(BacktestError::Config(msg)) => {
assert!(msg.contains(needle), "message {msg:?} misses {needle:?}");
}
other => panic!("expected BacktestError::Config, got {other:?}"),
}
}
#[test]
fn test_config_accepts_valid_config_ok() {
assert!(valid_config().validate().is_ok());
}
#[test]
fn test_config_rejects_zero_capital_config_error() {
let mut cfg = valid_config();
cfg.initial_capital = 0;
assert_config_error(cfg.validate(), "initial capital");
}
type LimitBump = fn(&mut ResourceLimits);
#[test]
fn test_config_rejects_each_limits_field_over_cap_config_error() {
let cases: [(&str, LimitBump); 8] = [
("max_steps", |l| {
l.max_steps = ResourceLimits::CAP_MAX_STEPS + 1;
}),
("max_contracts_per_snapshot", |l| {
l.max_contracts_per_snapshot = ResourceLimits::CAP_MAX_CONTRACTS_PER_SNAPSHOT + 1;
}),
("max_total_bytes", |l| {
l.max_total_bytes = ResourceLimits::CAP_MAX_TOTAL_BYTES + 1;
}),
("max_file_bytes", |l| {
l.max_file_bytes = ResourceLimits::CAP_MAX_FILE_BYTES + 1;
}),
("max_decompressed_bytes", |l| {
l.max_decompressed_bytes = ResourceLimits::CAP_MAX_DECOMPRESSED_BYTES + 1;
}),
("max_manifest_bytes", |l| {
l.max_manifest_bytes = ResourceLimits::CAP_MAX_MANIFEST_BYTES + 1;
}),
("max_string_len", |l| {
l.max_string_len = ResourceLimits::CAP_MAX_STRING_LEN + 1;
}),
("max_rows_per_table", |l| {
l.max_rows_per_table = ResourceLimits::CAP_MAX_ROWS_PER_TABLE + 1;
}),
];
for (field, bump) in cases {
let mut cfg = valid_config();
bump(&mut cfg.limits);
assert_config_error(cfg.validate(), field);
}
}
#[test]
fn test_config_rejects_parent_dir_output_path_config_error() {
let mut cfg = valid_config();
cfg.output_dir = "runs/../../etc".into();
assert_config_error(cfg.validate(), "..");
}
#[test]
fn test_config_rejects_empty_output_path_config_error() {
let mut cfg = valid_config();
cfg.output_dir = std::path::PathBuf::new();
assert_config_error(cfg.validate(), "output dir");
}
#[test]
fn test_config_rejects_zero_marketable_cap_ticks_config_error() {
let mut cfg = valid_config();
cfg.marketable_cap_ticks = 0;
assert_config_error(cfg.validate(), "marketable_cap_ticks");
}
#[test]
fn test_config_marketable_cap_ticks_defaults_to_ten_when_absent() {
let json = r#"{
"data_source": {"kind": "parquet", "path": "p.parquet", "sha256": ""},
"mode": "naive",
"seed": 1,
"initial_capital": 1000,
"fees": {"per_contract_cents": 0, "per_order_cents": 0},
"slippage": {"model": "none"},
"output_dir": "out"
}"#;
let parsed: Result<BacktestConfig, _> = serde_json::from_str(json);
assert!(matches!(parsed, Ok(ref c) if c.marketable_cap_ticks == 10));
}
#[test]
fn test_liquidity_profile_defaults_match_documented_values() {
let profile = LiquidityProfile::default();
assert_eq!(profile.touch_size, TouchSize::QuotedSize);
assert_eq!(profile.depth_levels, 5);
assert_eq!(profile.decay, dec!(0.5));
}
#[test]
fn test_config_rejects_depth_levels_over_cap_config_error() {
let mut cfg = valid_config();
cfg.liquidity_profile.depth_levels = LiquidityProfile::MAX_DEPTH_LEVELS + 1;
assert_config_error(cfg.validate(), "depth_levels");
}
#[test]
fn test_config_rejects_decay_out_of_range_config_error() {
for bad in [dec!(0), dec!(-0.1), dec!(1.5)] {
let mut cfg = valid_config();
cfg.liquidity_profile.decay = bad;
assert_config_error(cfg.validate(), "decay");
}
}
#[test]
fn test_config_accepts_decay_of_one_uniform_ladder_ok() {
let mut cfg = valid_config();
cfg.liquidity_profile.decay = dec!(1);
assert!(cfg.validate().is_ok());
}
#[test]
fn test_config_rejects_zero_flat_touch_size_config_error() {
let mut cfg = valid_config();
cfg.liquidity_profile.touch_size = TouchSize::Flat { contracts: 0 };
assert_config_error(cfg.validate(), "flat contracts");
}
#[test]
fn test_config_liquidity_profile_defaults_when_absent() {
let json = r#"{
"data_source": {"kind": "parquet", "path": "p.parquet", "sha256": ""},
"mode": "realistic",
"seed": 1,
"initial_capital": 1000,
"fees": {"per_contract_cents": 0, "per_order_cents": 0},
"slippage": {"model": "none"},
"output_dir": "out"
}"#;
let parsed: Result<BacktestConfig, _> = serde_json::from_str(json);
assert!(matches!(
parsed,
Ok(ref c)
if c.liquidity_profile == LiquidityProfile::default()
&& c.liquidity_profile.decay == dec!(0.5)
));
}
#[test]
fn test_config_liquidity_profile_flat_touch_round_trips() {
let mut cfg = valid_config();
cfg.liquidity_profile = LiquidityProfile {
touch_size: TouchSize::Flat { contracts: 25 },
depth_levels: 3,
decay: dec!(0.7),
};
let json = serde_json::to_string(&cfg).unwrap_or_default();
let back: Result<BacktestConfig, _> = serde_json::from_str(&json);
assert!(matches!(back, Ok(ref c) if *c == cfg));
}
#[test]
fn test_config_rejects_negative_slippage_fraction_config_error() {
let mut cfg = valid_config();
cfg.slippage = SlippageModel::SpreadFraction {
fraction: dec!(-0.1),
};
assert_config_error(cfg.validate(), "non-negative");
}
#[test]
fn test_limits_defaults_match_documented_values() {
let limits = ResourceLimits::default();
assert_eq!(limits.max_steps, 100_000);
assert_eq!(limits.max_contracts_per_snapshot, 50_000);
assert_eq!(limits.max_total_bytes, 2 * (1 << 30));
assert_eq!(limits.max_file_bytes, 4 * (1 << 30));
assert_eq!(limits.max_decompressed_bytes, 8 * (1 << 30));
assert_eq!(limits.max_manifest_bytes, 16 * (1 << 20));
assert_eq!(limits.max_string_len, 64 * (1 << 10));
assert_eq!(limits.max_rows_per_table, 100_000_000);
}
#[test]
fn test_config_unknown_field_fails_deserialisation() {
let json = r#"{
"data_source": {"kind": "parquet", "path": "p.parquet", "sha256": ""},
"mode": "naive",
"seed": 1,
"initial_capital": 1000,
"fees": {"per_contract_cents": 0, "per_order_cents": 0},
"slippage": {"model": "none"},
"output_dir": "out",
"not_a_real_field": true
}"#;
let parsed: Result<BacktestConfig, _> = serde_json::from_str(json);
assert!(parsed.is_err(), "unknown key must fail deserialisation");
}
#[test]
fn test_config_rejects_unknown_mode_value_fails_deserialisation() {
for bad in ["\"bogus\"", "\"realistic \"", "\"Naive\"", "\"\""] {
let json = format!(
r#"{{
"data_source": {{"kind": "parquet", "path": "p.parquet", "sha256": ""}},
"mode": {bad},
"seed": 1,
"initial_capital": 1000,
"fees": {{"per_contract_cents": 0, "per_order_cents": 0}},
"slippage": {{"model": "none"}},
"output_dir": "out"
}}"#
);
let parsed: Result<BacktestConfig, _> = serde_json::from_str(&json);
assert!(
parsed.is_err(),
"unknown mode value {bad} must fail deserialisation"
);
}
}
#[test]
fn test_config_accepts_both_known_mode_values() {
for (mode, expected) in [
("naive", ExecutionMode::Naive),
("realistic", ExecutionMode::Realistic),
] {
let json = format!(
r#"{{
"data_source": {{"kind": "parquet", "path": "p.parquet", "sha256": ""}},
"mode": "{mode}",
"seed": 1,
"initial_capital": 1000,
"fees": {{"per_contract_cents": 0, "per_order_cents": 0}},
"slippage": {{"model": "none"}},
"output_dir": "out"
}}"#
);
let parsed: Result<BacktestConfig, _> = serde_json::from_str(&json);
assert!(
matches!(parsed, Ok(ref c) if c.mode == expected && c.validate().is_ok()),
"mode {mode:?} must deserialise to {expected:?} and validate"
);
}
}
#[test]
fn test_limits_unknown_field_fails_deserialisation() {
let json = r#"{
"max_steps": 1, "max_contracts_per_snapshot": 1, "max_total_bytes": 1,
"max_file_bytes": 1, "max_decompressed_bytes": 1, "max_manifest_bytes": 1,
"max_string_len": 1, "max_rows_per_table": 1, "bogus": 0
}"#;
let parsed: Result<ResourceLimits, _> = serde_json::from_str(json);
assert!(parsed.is_err(), "unknown key must fail deserialisation");
}
#[test]
fn test_config_serde_round_trip_preserves_fields() {
let cfg = valid_config();
let json = serde_json::to_string(&cfg).unwrap_or_default();
let back: Result<BacktestConfig, _> = serde_json::from_str(&json);
assert!(matches!(back, Ok(ref c) if *c == cfg));
}
#[test]
fn test_config_serialized_field_set_is_pinned() {
let value = match serde_json::to_value(valid_config()) {
Ok(value) => value,
Err(err) => panic!("BacktestConfig must serialize: {err}"),
};
let object = match value.as_object() {
Some(object) => object,
None => panic!("BacktestConfig must serialize to a JSON object"),
};
let mut keys: Vec<&str> = object.keys().map(String::as_str).collect();
keys.sort_unstable();
let expected = [
"data_source",
"fees",
"initial_capital",
"limits",
"liquidity_profile",
"marketable_cap_ticks",
"mode",
"output_dir",
"overwrite",
"seed",
"slippage",
];
assert_eq!(
keys, expected,
"BacktestConfig serialized field set drifted from the frozen v1.0 \
config surface; update the pinned list here and record the SemVer \
event (docs/SEMVER.md §\"v1.0 commitments\")"
);
}
}