feat(ml,trading): add gate optimizer, model registry, and P&L attribution

Gate optimizer adjusts conviction gate thresholds based on win-rate per
confidence bucket with cooldown and kill switch safety rails. Model
registry provides lifecycle management (Candidate → Staging → Production
→ Archived) with InMemoryModelRegistry for testing. P&L attribution
decomposes realized trade P&L into per-model contributions using signal
alignment.

24 new tests across 3 modules.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-02-23 13:39:22 +01:00
parent d623fb100c
commit 1b7f72a0d2
7 changed files with 1327 additions and 2 deletions

80
Cargo.lock generated
View File

@@ -2318,6 +2318,7 @@ dependencies = [
"opentelemetry-otlp",
"opentelemetry_sdk",
"prometheus",
"questdb-rs",
"rand 0.8.5",
"redis",
"rust_decimal",
@@ -3283,6 +3284,18 @@ dependencies = [
"libloading",
]
[[package]]
name = "dns-lookup"
version = "2.1.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "cf5597a4b7fe5275fc9dcf88ce26326bc8e4cb87d0130f33752d4c5f717793cf"
dependencies = [
"cfg-if",
"libc",
"socket2 0.6.0",
"windows-sys 0.60.2",
]
[[package]]
name = "doc-comment"
version = "0.3.3"
@@ -4880,6 +4893,15 @@ dependencies = [
"web-time",
]
[[package]]
name = "indoc"
version = "2.0.7"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "79cf5c93f93228cf8efb3ba362535fb11199ac548a09ce117c9b1adc3030d706"
dependencies = [
"rustversion",
]
[[package]]
name = "influxdb"
version = "0.7.2"
@@ -7288,6 +7310,36 @@ dependencies = [
"winapi",
]
[[package]]
name = "questdb-confstr"
version = "0.1.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7aceffde1cbf8e67f34cdfd70d2436396176d6ff648fa719e0231fb9856ef3e9"
[[package]]
name = "questdb-rs"
version = "4.0.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "882489d3cc6b44ff276ed90b483afd44ae1d1d3cafd3ecceff58b83a2fdc9299"
dependencies = [
"base64ct",
"dns-lookup",
"indoc",
"itoa",
"libc",
"questdb-confstr",
"ring",
"rustls 0.22.4",
"rustls-pemfile 2.2.0",
"rustls-pki-types",
"ryu",
"serde",
"serde_json",
"slugify",
"socket2 0.5.10",
"winapi",
]
[[package]]
name = "quick-error"
version = "1.2.3"
@@ -7732,7 +7784,7 @@ dependencies = [
"percent-encoding",
"pin-project-lite",
"rustls 0.21.12",
"rustls-pemfile",
"rustls-pemfile 1.0.4",
"serde",
"serde_json",
"serde_urlencoded",
@@ -8200,7 +8252,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a9aace74cb666635c918e9c12bc0d348266037aa8eb599b5cba565709a8dff00"
dependencies = [
"openssl-probe",
"rustls-pemfile",
"rustls-pemfile 1.0.4",
"schannel",
"security-framework 2.11.1",
]
@@ -8226,6 +8278,15 @@ dependencies = [
"base64 0.21.7",
]
[[package]]
name = "rustls-pemfile"
version = "2.2.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "dce314e5fee3f39953d46bb63bb8a46d40c2f8fb7cc5a3b6cab2bde9721d6e50"
dependencies = [
"rustls-pki-types",
]
[[package]]
name = "rustls-pki-types"
version = "1.12.0"
@@ -8799,6 +8860,15 @@ dependencies = [
"time",
]
[[package]]
name = "slugify"
version = "0.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d6b8cf203d2088b831d7558f8e5151bfa420c57a34240b28cee29d0ae5f2ac8b"
dependencies = [
"unidecode",
]
[[package]]
name = "smallvec"
version = "1.15.1"
@@ -10819,6 +10889,12 @@ version = "0.2.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "ebc1c04c71510c7f702b52b7c350734c9ff1295c464a03335b00bb84fc54f853"
[[package]]
name = "unidecode"
version = "0.3.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "402bb19d8e03f1d1a7450e2bd613980869438e0666331be3e073089124aa1adc"
[[package]]
name = "universal-hash"
version = "0.5.1"

View File

@@ -0,0 +1,454 @@
//! Gate threshold optimizer for conviction gates
//!
//! Adjusts conviction gate thresholds based on win-rate per confidence bucket.
//! Same safety rails pattern as the weight optimizer: bounded adjustments,
//! cooldown, freeze/unfreeze.
use super::conviction_gates::ConvictionGateConfig;
use std::time::{Duration, Instant};
use tracing::{info, warn};
/// Maximum threshold change per optimization cycle
const MAX_THRESHOLD_STEP: f64 = 0.03;
/// Minimum allowed threshold value
const MIN_THRESHOLD: f64 = 0.30;
/// Maximum allowed threshold value
const MAX_THRESHOLD: f64 = 0.90;
/// Minimum trades per bucket before adjustment
const MIN_BUCKET_TRADES: u64 = 50;
/// Win-rate metrics per confidence bucket
#[derive(Debug, Clone)]
pub struct GateBucketMetrics {
/// Confidence bucket lower bound (e.g. 0.60)
pub confidence_lower: f64,
/// Confidence bucket upper bound (e.g. 0.70)
pub confidence_upper: f64,
/// Win rate in this bucket (0.0-1.0)
pub win_rate: f64,
/// Number of trades in this bucket
pub trade_count: u64,
/// Average P&L per trade in this bucket
pub avg_pnl: f64,
}
/// Gate optimizer configuration
#[derive(Debug, Clone)]
pub struct GateOptimizerConfig {
pub max_step: f64,
pub min_threshold: f64,
pub max_threshold: f64,
pub min_bucket_trades: u64,
pub cooldown: Duration,
/// Target win rate — thresholds tighten if below, loosen if above
pub target_win_rate: f64,
}
impl Default for GateOptimizerConfig {
fn default() -> Self {
Self {
max_step: MAX_THRESHOLD_STEP,
min_threshold: MIN_THRESHOLD,
max_threshold: MAX_THRESHOLD,
min_bucket_trades: MIN_BUCKET_TRADES,
cooldown: Duration::from_secs(24 * 3600),
target_win_rate: 0.55,
}
}
}
/// Proposed threshold change
#[derive(Debug, Clone)]
pub struct ThresholdAdjustment {
pub field_name: String,
pub old_value: f64,
pub new_value: f64,
pub reason: String,
}
/// Result of a gate optimization cycle
#[derive(Debug, Clone)]
pub enum GateOptimizationResult {
/// Thresholds adjusted
Adjusted(Vec<ThresholdAdjustment>),
/// Not enough data
InsufficientData { total_trades: u64, required: u64 },
/// Still in cooldown
Cooldown { remaining: Duration },
/// Frozen by kill switch
KillSwitchActive,
}
/// Gate threshold optimizer
#[derive(Debug)]
pub struct GateOptimizer {
config: GateOptimizerConfig,
last_adjustment: Option<Instant>,
frozen: bool,
}
impl GateOptimizer {
pub fn new(config: GateOptimizerConfig) -> Self {
Self {
config,
last_adjustment: None,
frozen: false,
}
}
/// Freeze all adjustments (kill switch)
pub fn freeze(&mut self) {
self.frozen = true;
warn!("Gate optimizer frozen by kill switch");
}
/// Unfreeze adjustments
pub fn unfreeze(&mut self) {
self.frozen = false;
info!("Gate optimizer unfrozen");
}
pub fn is_frozen(&self) -> bool {
self.frozen
}
/// Run one optimization cycle
///
/// Analyzes win-rate per confidence bucket and adjusts min_confidence threshold.
/// If win rate near the current threshold is below target, threshold increases
/// (more selective). If well above target, threshold decreases (more permissive).
pub fn optimize(
&mut self,
gate_config: &ConvictionGateConfig,
buckets: &[GateBucketMetrics],
) -> GateOptimizationResult {
if self.frozen {
return GateOptimizationResult::KillSwitchActive;
}
// Check cooldown
if let Some(last) = self.last_adjustment {
let elapsed = last.elapsed();
if elapsed < self.config.cooldown {
return GateOptimizationResult::Cooldown {
remaining: self.config.cooldown - elapsed,
};
}
}
// Check minimum data
let total_trades: u64 = buckets.iter().map(|b| b.trade_count).sum();
let min_required = self.config.min_bucket_trades * 3; // At least 3 buckets worth
if total_trades < min_required {
return GateOptimizationResult::InsufficientData {
total_trades,
required: min_required,
};
}
let mut adjustments = Vec::new();
// Analyze min_confidence threshold
// Find the bucket containing the current threshold
let threshold_bucket = buckets.iter().find(|b| {
b.confidence_lower <= gate_config.min_confidence
&& gate_config.min_confidence < b.confidence_upper
&& b.trade_count >= self.config.min_bucket_trades
});
if let Some(bucket) = threshold_bucket {
let win_rate_delta = bucket.win_rate - self.config.target_win_rate;
// If win rate is too low near threshold → tighten (increase threshold)
// If win rate is high → loosen (decrease threshold)
let direction = if win_rate_delta < -0.05 {
// Win rate below target by > 5pp → tighten
1.0
} else if win_rate_delta > 0.10 {
// Win rate above target by > 10pp → loosen
-1.0
} else {
0.0 // In acceptable range
};
if direction != 0.0 {
let step = (win_rate_delta.abs() * 0.1)
.min(self.config.max_step)
.max(0.005);
let new_confidence = (gate_config.min_confidence + direction * step)
.clamp(self.config.min_threshold, self.config.max_threshold);
if (new_confidence - gate_config.min_confidence).abs() > 1e-6 {
adjustments.push(ThresholdAdjustment {
field_name: "min_confidence".into(),
old_value: gate_config.min_confidence,
new_value: new_confidence,
reason: format!(
"bucket win_rate={:.3}, target={:.3}, delta={:.3}",
bucket.win_rate, self.config.target_win_rate, win_rate_delta
),
});
}
}
}
// Analyze max_disagreement threshold
// If overall win rate on low-disagreement trades is high, can loosen
let low_disagree_buckets: Vec<&GateBucketMetrics> = buckets
.iter()
.filter(|b| b.trade_count >= self.config.min_bucket_trades)
.collect();
if !low_disagree_buckets.is_empty() {
let weighted_win_rate: f64 = low_disagree_buckets
.iter()
.map(|b| b.win_rate * b.trade_count as f64)
.sum::<f64>()
/ low_disagree_buckets
.iter()
.map(|b| b.trade_count as f64)
.sum::<f64>();
if weighted_win_rate < self.config.target_win_rate - 0.05 {
// Poor overall performance → tighten disagreement (lower max)
let new_max = (gate_config.max_disagreement - 0.01)
.clamp(0.10, 0.60);
if (new_max - gate_config.max_disagreement).abs() > 1e-6 {
adjustments.push(ThresholdAdjustment {
field_name: "max_disagreement".into(),
old_value: gate_config.max_disagreement,
new_value: new_max,
reason: format!(
"weighted win_rate={:.3} below target {:.3}",
weighted_win_rate, self.config.target_win_rate
),
});
}
}
}
if !adjustments.is_empty() {
self.last_adjustment = Some(Instant::now());
info!(
"Gate optimizer adjusted {} thresholds",
adjustments.len()
);
}
GateOptimizationResult::Adjusted(adjustments)
}
/// Apply adjustments to a ConvictionGateConfig (returns modified copy)
pub fn apply(config: &ConvictionGateConfig, adjustments: &[ThresholdAdjustment]) -> ConvictionGateConfig {
let mut new_config = config.clone();
for adj in adjustments {
match adj.field_name.as_str() {
"min_confidence" => new_config.min_confidence = adj.new_value,
"max_disagreement" => new_config.max_disagreement = adj.new_value,
"min_quorum" => new_config.min_quorum = adj.new_value,
_ => {}
}
}
new_config
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_buckets(data: &[(f64, f64, f64, u64)]) -> Vec<GateBucketMetrics> {
data.iter()
.map(|(lower, upper, win_rate, trades)| GateBucketMetrics {
confidence_lower: *lower,
confidence_upper: *upper,
win_rate: *win_rate,
trade_count: *trades,
avg_pnl: 0.0,
})
.collect()
}
#[test]
fn test_insufficient_data() {
let mut opt = GateOptimizer::new(GateOptimizerConfig::default());
let config = ConvictionGateConfig::default();
let buckets = make_buckets(&[(0.50, 0.60, 0.55, 10), (0.60, 0.70, 0.60, 10)]);
let result = opt.optimize(&config, &buckets);
assert!(matches!(
result,
GateOptimizationResult::InsufficientData { .. }
));
}
#[test]
fn test_cooldown_enforced() {
let mut opt = GateOptimizer::new(GateOptimizerConfig::default());
let config = ConvictionGateConfig::default();
let buckets = make_buckets(&[
(0.50, 0.60, 0.45, 100),
(0.60, 0.70, 0.40, 100),
(0.70, 0.80, 0.55, 100),
]);
let _ = opt.optimize(&config, &buckets);
let result2 = opt.optimize(&config, &buckets);
assert!(matches!(result2, GateOptimizationResult::Cooldown { .. }));
}
#[test]
fn test_kill_switch() {
let mut opt = GateOptimizer::new(GateOptimizerConfig::default());
opt.freeze();
let config = ConvictionGateConfig::default();
let buckets = make_buckets(&[(0.50, 0.60, 0.45, 200)]);
let result = opt.optimize(&config, &buckets);
assert!(matches!(result, GateOptimizationResult::KillSwitchActive));
}
#[test]
fn test_tightens_on_low_win_rate() {
let mut config_opt = GateOptimizerConfig::default();
config_opt.cooldown = Duration::ZERO;
let mut opt = GateOptimizer::new(config_opt);
let config = ConvictionGateConfig::default(); // min_confidence = 0.60
// Win rate at threshold bucket is poor (0.40 < 0.55 target)
let buckets = make_buckets(&[
(0.50, 0.60, 0.40, 100),
(0.60, 0.70, 0.40, 100), // Threshold bucket
(0.70, 0.80, 0.60, 100),
]);
let result = opt.optimize(&config, &buckets);
if let GateOptimizationResult::Adjusted(adjustments) = result {
let confidence_adj = adjustments
.iter()
.find(|a| a.field_name == "min_confidence");
assert!(
confidence_adj.is_some(),
"Expected min_confidence adjustment"
);
if let Some(adj) = confidence_adj {
assert!(
adj.new_value > adj.old_value,
"Expected threshold to increase (tighten) on low win rate"
);
}
}
}
#[test]
fn test_loosens_on_high_win_rate() {
let mut config_opt = GateOptimizerConfig::default();
config_opt.cooldown = Duration::ZERO;
let mut opt = GateOptimizer::new(config_opt);
let config = ConvictionGateConfig::default(); // min_confidence = 0.60
// Win rate at threshold bucket is very good (0.75 >> 0.55 target)
let buckets = make_buckets(&[
(0.50, 0.60, 0.70, 100),
(0.60, 0.70, 0.75, 100), // Threshold bucket
(0.70, 0.80, 0.80, 100),
]);
let result = opt.optimize(&config, &buckets);
if let GateOptimizationResult::Adjusted(adjustments) = result {
let confidence_adj = adjustments
.iter()
.find(|a| a.field_name == "min_confidence");
if let Some(adj) = confidence_adj {
assert!(
adj.new_value < adj.old_value,
"Expected threshold to decrease (loosen) on high win rate"
);
}
}
}
#[test]
fn test_bounds_enforced() {
let mut config_opt = GateOptimizerConfig::default();
config_opt.cooldown = Duration::ZERO;
let mut opt = GateOptimizer::new(config_opt);
// Config with threshold already near maximum
let mut config = ConvictionGateConfig::default();
config.min_confidence = 0.89;
// Very low win rate to force tightening
let buckets = make_buckets(&[
(0.80, 0.90, 0.30, 100),
(0.89, 0.95, 0.30, 100), // Threshold bucket
(0.70, 0.80, 0.40, 100),
]);
let result = opt.optimize(&config, &buckets);
if let GateOptimizationResult::Adjusted(adjustments) = result {
for adj in &adjustments {
assert!(
adj.new_value >= MIN_THRESHOLD,
"Below minimum: {}",
adj.new_value
);
assert!(
adj.new_value <= MAX_THRESHOLD,
"Above maximum: {}",
adj.new_value
);
}
}
}
#[test]
fn test_no_change_in_acceptable_range() {
let mut config_opt = GateOptimizerConfig::default();
config_opt.cooldown = Duration::ZERO;
let mut opt = GateOptimizer::new(config_opt);
let config = ConvictionGateConfig::default();
// Win rate is in acceptable range (target ± tolerance)
let buckets = make_buckets(&[
(0.50, 0.60, 0.56, 100),
(0.60, 0.70, 0.58, 100), // Just above target, within tolerance
(0.70, 0.80, 0.60, 100),
]);
let result = opt.optimize(&config, &buckets);
if let GateOptimizationResult::Adjusted(adjustments) = result {
let confidence_adj = adjustments
.iter()
.find(|a| a.field_name == "min_confidence");
assert!(
confidence_adj.is_none(),
"Expected no min_confidence adjustment when in acceptable range"
);
}
}
#[test]
fn test_apply_adjustments() {
let config = ConvictionGateConfig::default();
let adjustments = vec![
ThresholdAdjustment {
field_name: "min_confidence".into(),
old_value: 0.60,
new_value: 0.65,
reason: "test".into(),
},
ThresholdAdjustment {
field_name: "max_disagreement".into(),
old_value: 0.40,
new_value: 0.35,
reason: "test".into(),
},
];
let new_config = GateOptimizer::apply(&config, &adjustments);
assert!((new_config.min_confidence - 0.65).abs() < 1e-10);
assert!((new_config.max_disagreement - 0.35).abs() < 1e-10);
// Unchanged fields preserved
assert!((new_config.min_quorum - 0.60).abs() < 1e-10);
}
}

View File

@@ -23,6 +23,7 @@ pub mod signal;
pub mod adapters;
pub mod conviction_gates;
pub mod weight_optimizer;
pub mod gate_optimizer;
// Re-export key types that are used across ensemble modules
pub use ab_testing::{
@@ -58,6 +59,10 @@ pub use weight_optimizer::{
ModelRollingMetrics, OptimizationResult, WeightAdjustment, WeightOptimizer,
WeightOptimizerConfig,
};
pub use gate_optimizer::{
GateBucketMetrics, GateOptimizationResult, GateOptimizer, GateOptimizerConfig,
ThresholdAdjustment,
};
/// Errors that can occur in ensemble operations
#[derive(Error, Debug)]

View File

@@ -832,6 +832,7 @@ pub mod real_data_loader;
pub mod data_validation;
pub mod random_model;
pub mod model_registry;
pub mod registry; // Operational maturity: model lifecycle (Candidate → Staging → Production → Archived)
// ========== MISSING TYPES STUBS ==========

524
ml/src/registry/mod.rs Normal file
View File

@@ -0,0 +1,524 @@
//! Model Registry for lifecycle management
//!
//! Tracks training runs, model versions, and promotions through
//! Candidate → Staging → Production → Archived lifecycle.
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::time::SystemTime;
/// Model lifecycle stage
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum ModelStage {
/// Initial state after training
Candidate,
/// Passed validation, running canary
Staging,
/// Active in production ensemble
Production,
/// Replaced by newer version
Archived,
}
impl std::fmt::Display for ModelStage {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Candidate => write!(f, "candidate"),
Self::Staging => write!(f, "staging"),
Self::Production => write!(f, "production"),
Self::Archived => write!(f, "archived"),
}
}
}
/// Record of a training run
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TrainingRun {
/// Unique run identifier
pub run_id: String,
/// Experiment name (e.g. "dqn-v3-sharpe-opt")
pub experiment_name: String,
/// Model type (DQN, PPO, TFT, etc.)
pub model_type: String,
/// Hyperparameters as JSON
pub hyperparameters: serde_json::Value,
/// Git commit hash at time of training
pub git_commit: String,
/// Hash of training data for reproducibility
pub data_hash: String,
/// When training started
pub started_at: SystemTime,
/// When training finished (None if still running)
pub finished_at: Option<SystemTime>,
}
/// Metrics recorded for a model version
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelMetrics {
/// Validation Sharpe ratio
pub sharpe_ratio: f64,
/// Validation accuracy
pub accuracy: f64,
/// Validation win rate
pub win_rate: f64,
/// Maximum drawdown on validation set
pub max_drawdown: f64,
/// Any additional metrics
pub extra: HashMap<String, f64>,
}
/// A versioned model in the registry
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelVersion {
/// Unique version identifier
pub version_id: String,
/// Associated training run
pub run_id: String,
/// Model type
pub model_type: String,
/// Path to model artifact (safetensors file)
pub artifact_path: String,
/// Current lifecycle stage
pub stage: ModelStage,
/// Validation metrics
pub metrics: Option<ModelMetrics>,
/// When this version was registered
pub registered_at: SystemTime,
/// Who/what promoted this version
pub promoted_by: Option<String>,
}
/// Side-by-side run comparison
#[derive(Debug, Clone)]
pub struct RunComparison {
pub runs: Vec<(TrainingRun, Option<ModelMetrics>)>,
}
/// Model registry trait (async for database implementations)
#[async_trait::async_trait]
pub trait ModelRegistryTrait: Send + Sync {
/// Log a new training run
async fn log_run(&self, run: TrainingRun) -> Result<(), RegistryError>;
/// Log metrics for a model version
async fn log_metrics(
&self,
version_id: &str,
metrics: ModelMetrics,
) -> Result<(), RegistryError>;
/// Register a model version
async fn register_version(&self, version: ModelVersion) -> Result<(), RegistryError>;
/// Promote a model to a new stage
async fn promote(
&self,
version_id: &str,
to_stage: ModelStage,
promoted_by: &str,
) -> Result<(), RegistryError>;
/// Get the current production model for a given type
async fn get_production_model(
&self,
model_type: &str,
) -> Result<Option<ModelVersion>, RegistryError>;
/// Revert to previous production version
async fn revert(&self, model_type: &str) -> Result<ModelVersion, RegistryError>;
/// Compare metrics across runs
async fn compare_runs(&self, run_ids: &[String]) -> Result<RunComparison, RegistryError>;
}
/// Registry errors
#[derive(Debug, thiserror::Error)]
pub enum RegistryError {
#[error("Version not found: {0}")]
VersionNotFound(String),
#[error("Run not found: {0}")]
RunNotFound(String),
#[error("Invalid stage transition: {from} → {to}")]
InvalidTransition { from: String, to: String },
#[error("No previous version to revert to for model type: {0}")]
NoPreviousVersion(String),
}
/// In-memory model registry for testing
#[derive(Debug, Default)]
pub struct InMemoryModelRegistry {
runs: tokio::sync::RwLock<HashMap<String, TrainingRun>>,
versions: tokio::sync::RwLock<HashMap<String, ModelVersion>>,
metrics: tokio::sync::RwLock<HashMap<String, ModelMetrics>>,
}
impl InMemoryModelRegistry {
pub fn new() -> Self {
Self::default()
}
}
#[async_trait::async_trait]
impl ModelRegistryTrait for InMemoryModelRegistry {
async fn log_run(&self, run: TrainingRun) -> Result<(), RegistryError> {
self.runs.write().await.insert(run.run_id.clone(), run);
Ok(())
}
async fn log_metrics(
&self,
version_id: &str,
metrics: ModelMetrics,
) -> Result<(), RegistryError> {
// Also update the version's metrics
let mut versions = self.versions.write().await;
if let Some(version) = versions.get_mut(version_id) {
version.metrics = Some(metrics.clone());
}
self.metrics
.write()
.await
.insert(version_id.to_string(), metrics);
Ok(())
}
async fn register_version(&self, version: ModelVersion) -> Result<(), RegistryError> {
self.versions
.write()
.await
.insert(version.version_id.clone(), version);
Ok(())
}
async fn promote(
&self,
version_id: &str,
to_stage: ModelStage,
promoted_by: &str,
) -> Result<(), RegistryError> {
let mut versions = self.versions.write().await;
let version = versions
.get_mut(version_id)
.ok_or_else(|| RegistryError::VersionNotFound(version_id.to_string()))?;
// Validate transition
let valid = matches!(
(version.stage, to_stage),
(ModelStage::Candidate, ModelStage::Staging)
| (ModelStage::Staging, ModelStage::Production)
| (ModelStage::Production, ModelStage::Archived)
| (ModelStage::Staging, ModelStage::Archived)
| (ModelStage::Candidate, ModelStage::Archived)
);
if !valid {
return Err(RegistryError::InvalidTransition {
from: version.stage.to_string(),
to: to_stage.to_string(),
});
}
// If promoting to Production, archive the current production model of same type
if to_stage == ModelStage::Production {
let model_type = version.model_type.clone();
let current_prod: Vec<String> = versions
.iter()
.filter(|(id, v)| {
v.model_type == model_type
&& v.stage == ModelStage::Production
&& *id != version_id
})
.map(|(id, _)| id.clone())
.collect();
// Must drop the version borrow before modifying others
let version = versions.get_mut(version_id).expect("just checked");
version.stage = to_stage;
version.promoted_by = Some(promoted_by.to_string());
for old_id in current_prod {
if let Some(old_version) = versions.get_mut(&old_id) {
old_version.stage = ModelStage::Archived;
}
}
} else {
version.stage = to_stage;
version.promoted_by = Some(promoted_by.to_string());
}
Ok(())
}
async fn get_production_model(
&self,
model_type: &str,
) -> Result<Option<ModelVersion>, RegistryError> {
let versions = self.versions.read().await;
let prod = versions
.values()
.find(|v| v.model_type == model_type && v.stage == ModelStage::Production)
.cloned();
Ok(prod)
}
async fn revert(&self, model_type: &str) -> Result<ModelVersion, RegistryError> {
let mut versions = self.versions.write().await;
// Find the most recently archived version of this type
let archived: Option<String> = versions
.iter()
.filter(|(_, v)| v.model_type == model_type && v.stage == ModelStage::Archived)
.max_by_key(|(_, v)| v.registered_at)
.map(|(id, _)| id.clone());
let archived_id =
archived.ok_or_else(|| RegistryError::NoPreviousVersion(model_type.to_string()))?;
// Archive current production
let current_prod: Vec<String> = versions
.iter()
.filter(|(_, v)| v.model_type == model_type && v.stage == ModelStage::Production)
.map(|(id, _)| id.clone())
.collect();
for id in current_prod {
if let Some(v) = versions.get_mut(&id) {
v.stage = ModelStage::Archived;
}
}
// Promote archived to production
let version = versions
.get_mut(&archived_id)
.ok_or_else(|| RegistryError::VersionNotFound(archived_id.clone()))?;
version.stage = ModelStage::Production;
version.promoted_by = Some("revert".to_string());
Ok(version.clone())
}
async fn compare_runs(&self, run_ids: &[String]) -> Result<RunComparison, RegistryError> {
let runs = self.runs.read().await;
let metrics = self.metrics.read().await;
let mut comparisons = Vec::new();
for id in run_ids {
let run = runs
.get(id)
.ok_or_else(|| RegistryError::RunNotFound(id.clone()))?
.clone();
let m = metrics.get(id).cloned();
comparisons.push((run, m));
}
Ok(RunComparison {
runs: comparisons,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_run(id: &str, model_type: &str) -> TrainingRun {
TrainingRun {
run_id: id.to_string(),
experiment_name: format!("{}-exp", model_type),
model_type: model_type.to_string(),
hyperparameters: serde_json::json!({"lr": 0.001}),
git_commit: "abc123".into(),
data_hash: "sha256:deadbeef".into(),
started_at: SystemTime::now(),
finished_at: Some(SystemTime::now()),
}
}
fn make_version(id: &str, run_id: &str, model_type: &str) -> ModelVersion {
ModelVersion {
version_id: id.to_string(),
run_id: run_id.to_string(),
model_type: model_type.to_string(),
artifact_path: format!("models/{}/{}.safetensors", model_type, id),
stage: ModelStage::Candidate,
metrics: None,
registered_at: SystemTime::now(),
promoted_by: None,
}
}
#[test]
fn test_model_stage_display() {
assert_eq!(ModelStage::Candidate.to_string(), "candidate");
assert_eq!(ModelStage::Production.to_string(), "production");
}
#[tokio::test]
async fn test_log_and_register() {
let registry = InMemoryModelRegistry::new();
let run = make_run("run-1", "DQN");
registry.log_run(run).await.unwrap();
let version = make_version("v1", "run-1", "DQN");
registry.register_version(version).await.unwrap();
let prod = registry.get_production_model("DQN").await.unwrap();
assert!(prod.is_none()); // Not promoted yet
}
#[tokio::test]
async fn test_promote_lifecycle() {
let registry = InMemoryModelRegistry::new();
let run = make_run("run-1", "DQN");
registry.log_run(run).await.unwrap();
let version = make_version("v1", "run-1", "DQN");
registry.register_version(version).await.unwrap();
// Candidate → Staging
registry
.promote("v1", ModelStage::Staging, "ci")
.await
.unwrap();
// Staging → Production
registry
.promote("v1", ModelStage::Production, "ci")
.await
.unwrap();
let prod = registry.get_production_model("DQN").await.unwrap();
assert!(prod.is_some());
assert_eq!(prod.unwrap().version_id, "v1");
}
#[tokio::test]
async fn test_invalid_transition() {
let registry = InMemoryModelRegistry::new();
let version = make_version("v1", "run-1", "DQN");
registry.register_version(version).await.unwrap();
// Candidate → Production is invalid (must go through Staging)
let result = registry
.promote("v1", ModelStage::Production, "ci")
.await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_promotion_archives_old() {
let registry = InMemoryModelRegistry::new();
// Register and promote v1 to production
registry
.register_version(make_version("v1", "run-1", "DQN"))
.await
.unwrap();
registry
.promote("v1", ModelStage::Staging, "ci")
.await
.unwrap();
registry
.promote("v1", ModelStage::Production, "ci")
.await
.unwrap();
// Register and promote v2 to production
registry
.register_version(make_version("v2", "run-2", "DQN"))
.await
.unwrap();
registry
.promote("v2", ModelStage::Staging, "ci")
.await
.unwrap();
registry
.promote("v2", ModelStage::Production, "ci")
.await
.unwrap();
// v1 should be archived, v2 should be production
let prod = registry.get_production_model("DQN").await.unwrap();
assert_eq!(prod.unwrap().version_id, "v2");
}
#[tokio::test]
async fn test_revert() {
let registry = InMemoryModelRegistry::new();
// v1 → production → archived (when v2 promoted)
registry
.register_version(make_version("v1", "run-1", "DQN"))
.await
.unwrap();
registry
.promote("v1", ModelStage::Staging, "ci")
.await
.unwrap();
registry
.promote("v1", ModelStage::Production, "ci")
.await
.unwrap();
registry
.register_version(make_version("v2", "run-2", "DQN"))
.await
.unwrap();
registry
.promote("v2", ModelStage::Staging, "ci")
.await
.unwrap();
registry
.promote("v2", ModelStage::Production, "ci")
.await
.unwrap();
// Revert should bring v1 back
let reverted = registry.revert("DQN").await.unwrap();
assert_eq!(reverted.version_id, "v1");
assert_eq!(reverted.stage, ModelStage::Production);
}
#[tokio::test]
async fn test_log_metrics() {
let registry = InMemoryModelRegistry::new();
registry
.register_version(make_version("v1", "run-1", "DQN"))
.await
.unwrap();
let metrics = ModelMetrics {
sharpe_ratio: 1.5,
accuracy: 0.62,
win_rate: 0.58,
max_drawdown: -0.05,
extra: HashMap::new(),
};
registry.log_metrics("v1", metrics).await.unwrap();
// Metrics should be attached to version
let versions = registry.versions.read().await;
let v = versions.get("v1").unwrap();
assert!(v.metrics.is_some());
assert!((v.metrics.as_ref().unwrap().sharpe_ratio - 1.5).abs() < 1e-10);
}
#[tokio::test]
async fn test_compare_runs() {
let registry = InMemoryModelRegistry::new();
registry
.log_run(make_run("run-1", "DQN"))
.await
.unwrap();
registry
.log_run(make_run("run-2", "DQN"))
.await
.unwrap();
let comparison = registry
.compare_runs(&["run-1".into(), "run-2".into()])
.await
.unwrap();
assert_eq!(comparison.runs.len(), 2);
}
}

View File

@@ -0,0 +1,262 @@
//! P&L Attribution Calculator
//!
//! Decomposes realized trade P&L into per-model contributions.
//! Each model gets credit proportional to:
//! - Its ensemble weight at trade time
//! - Whether its signal aligned with the realized direction
use ml::ensemble::decision::EnsembleDecision;
/// Per-model attribution for a single trade
#[derive(Debug, Clone)]
pub struct TradeAttribution {
/// Model identifier
pub model_id: String,
/// Model's ensemble weight at trade time
pub model_weight: f64,
/// Model's raw signal (-1.0 to 1.0)
pub model_signal: f64,
/// 1.0 if signal direction matched realized direction, -1.0 otherwise
pub signal_alignment: f64,
/// Attributed P&L contribution
pub pnl_contribution: f64,
}
/// Full attribution result for a closed trade
#[derive(Debug, Clone)]
pub struct AttributionResult {
/// Per-model attributions
pub attributions: Vec<TradeAttribution>,
/// Total realized P&L (should equal sum of contributions)
pub realized_pnl: f64,
/// Residual (rounding error, should be near zero)
pub residual: f64,
}
/// Calculate per-model P&L attribution from an ensemble decision and realized P&L.
///
/// For each model that voted:
/// - `signal_alignment = 1.0` if sign(model_signal) == sign(realized_pnl), else `-1.0`
/// - `pnl_contribution = model_weight × signal_alignment × |realized_pnl|`
///
/// Models with zero signal are treated as neutral (alignment = 0.0).
pub fn attribute(decision: &EnsembleDecision, realized_pnl: f64) -> AttributionResult {
let votes = &decision.model_votes;
if votes.is_empty() || realized_pnl.abs() < 1e-12 {
return AttributionResult {
attributions: votes
.iter()
.map(|(id, v)| TradeAttribution {
model_id: id.clone(),
model_weight: v.weight,
model_signal: v.signal,
signal_alignment: 0.0,
pnl_contribution: 0.0,
})
.collect(),
realized_pnl,
residual: realized_pnl,
};
}
let realized_direction = realized_pnl.signum();
let attributions: Vec<TradeAttribution> = votes
.iter()
.map(|(id, vote)| {
let alignment = compute_alignment(vote.signal, realized_direction);
let contribution = vote.weight * alignment * realized_pnl.abs();
TradeAttribution {
model_id: id.clone(),
model_weight: vote.weight,
model_signal: vote.signal,
signal_alignment: alignment,
pnl_contribution: contribution,
}
})
.collect();
let total_attributed: f64 = attributions.iter().map(|a| a.pnl_contribution).sum();
let residual = realized_pnl - total_attributed;
AttributionResult {
attributions,
realized_pnl,
residual,
}
}
/// Compute signal alignment: does the model's signal direction match the realized direction?
///
/// - Zero signal → neutral (0.0)
/// - Same sign → aligned (1.0)
/// - Opposite sign → misaligned (-1.0)
fn compute_alignment(model_signal: f64, realized_direction: f64) -> f64 {
if model_signal.abs() < 1e-12 {
return 0.0;
}
if model_signal.signum() == realized_direction {
1.0
} else {
-1.0
}
}
#[cfg(test)]
mod tests {
use super::*;
use ml::ensemble::decision::{ModelVote, TradingAction};
use std::collections::HashMap;
fn make_decision(votes: Vec<(&str, f64, f64)>) -> EnsembleDecision {
let mut model_votes = HashMap::new();
let mut total_signal = 0.0;
for (id, signal, weight) in &votes {
total_signal += signal * weight;
model_votes.insert(
id.to_string(),
ModelVote::new(id.to_string(), *signal, signal.abs(), *weight)
.with_model_type("DQN".into()),
);
}
EnsembleDecision::new(
if total_signal > 0.0 {
TradingAction::Buy
} else if total_signal < 0.0 {
TradingAction::Sell
} else {
TradingAction::Hold
},
0.8,
total_signal,
0.0,
model_votes,
)
}
#[test]
fn test_correct_attribution_positive_pnl() {
// Two models, both bullish, trade was profitable
let decision = make_decision(vec![
("dqn", 0.8, 0.6), // 60% weight, bullish
("ppo", 0.5, 0.4), // 40% weight, bullish
]);
let result = attribute(&decision, 100.0);
assert_eq!(result.attributions.len(), 2);
// DQN: 0.6 × 1.0 × 100.0 = 60.0
let dqn = result.attributions.iter().find(|a| a.model_id == "dqn");
assert!(dqn.is_some());
if let Some(dqn) = dqn {
assert!((dqn.pnl_contribution - 60.0).abs() < 1e-10);
assert!((dqn.signal_alignment - 1.0).abs() < 1e-10);
}
// PPO: 0.4 × 1.0 × 100.0 = 40.0
let ppo = result.attributions.iter().find(|a| a.model_id == "ppo");
assert!(ppo.is_some());
if let Some(ppo) = ppo {
assert!((ppo.pnl_contribution - 40.0).abs() < 1e-10);
}
// Sum should equal realized PnL
assert!(result.residual.abs() < 1e-10);
}
#[test]
fn test_attribution_with_disagreement() {
// DQN bullish, PPO bearish, trade was profitable (bullish correct)
let decision = make_decision(vec![
("dqn", 0.8, 0.6), // bullish, correct
("ppo", -0.5, 0.4), // bearish, wrong
]);
let result = attribute(&decision, 100.0);
// DQN: 0.6 × 1.0 × 100.0 = 60.0 (aligned)
if let Some(dqn) = result.attributions.iter().find(|a| a.model_id == "dqn") {
assert!((dqn.pnl_contribution - 60.0).abs() < 1e-10);
}
// PPO: 0.4 × (-1.0) × 100.0 = -40.0 (misaligned)
if let Some(ppo) = result.attributions.iter().find(|a| a.model_id == "ppo") {
assert!((ppo.pnl_contribution - (-40.0)).abs() < 1e-10);
}
}
#[test]
fn test_zero_pnl() {
let decision = make_decision(vec![("dqn", 0.8, 0.5), ("ppo", 0.5, 0.5)]);
let result = attribute(&decision, 0.0);
// All contributions should be zero
for attr in &result.attributions {
assert!(attr.pnl_contribution.abs() < 1e-12);
assert!(attr.signal_alignment.abs() < 1e-12);
}
}
#[test]
fn test_all_models_wrong() {
// All bearish, but market went up → they were wrong
let decision = make_decision(vec![("dqn", -0.7, 0.5), ("ppo", -0.6, 0.5)]);
let result = attribute(&decision, 50.0);
// Both bearish but pnl positive → misaligned
for attr in &result.attributions {
assert!((attr.signal_alignment - (-1.0)).abs() < 1e-10);
assert!(attr.pnl_contribution < 0.0);
}
}
#[test]
fn test_all_models_correct_bearish() {
// All bearish, PnL negative → bearish direction correct
let decision = make_decision(vec![("dqn", -0.7, 0.5), ("ppo", -0.6, 0.5)]);
let result = attribute(&decision, -100.0);
// Both bearish, PnL negative → aligned
for attr in &result.attributions {
assert!((attr.signal_alignment - 1.0).abs() < 1e-10);
assert!(attr.pnl_contribution > 0.0);
}
}
#[test]
fn test_neutral_model_gets_zero() {
let decision = make_decision(vec![
("dqn", 0.8, 0.5),
("neutral", 0.0, 0.5), // Zero signal = neutral
]);
let result = attribute(&decision, 100.0);
if let Some(neutral) = result.attributions.iter().find(|a| a.model_id == "neutral") {
assert!(neutral.signal_alignment.abs() < 1e-12);
assert!(neutral.pnl_contribution.abs() < 1e-12);
}
}
#[test]
fn test_empty_votes() {
let decision = EnsembleDecision::new(
TradingAction::Hold,
0.0,
0.0,
0.0,
HashMap::new(),
);
let result = attribute(&decision, 100.0);
assert!(result.attributions.is_empty());
assert!((result.residual - 100.0).abs() < 1e-10);
}
}

View File

@@ -163,6 +163,9 @@ pub mod assets;
/// Health check endpoints for Kubernetes probes
pub mod health;
/// P&L attribution: decomposes realized trade P&L into per-model contributions
pub mod attribution;
// Re-export for tests
pub use ensemble_coordinator::EnsembleCoordinator;
pub use paper_trading_executor::PaperTradingExecutor;