Merge branch 'feature/operational-maturity'

# Conflicts:
#	ml/src/lib.rs
This commit is contained in:
jgrusewski
2026-02-23 14:20:23 +01:00
18 changed files with 3768 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

@@ -222,6 +222,7 @@ dashmap = { version = "6.0", features = ["serde"] }
bytes = "1.5"
smallvec = { version = "1.11", features = ["serde", "const_generics"] }
prometheus = "0.14"
questdb-rs = { version = "4", default-features = false }
# MINIMAL numerical libraries only - ALL ML/GPU frameworks REMOVED from workspace
nalgebra = { version = "0.33", features = ["serde", "rand"] }

View File

@@ -60,6 +60,7 @@ config = { path = "../config" }
uuid = { workspace = true, features = ["v4", "serde"] }
once_cell.workspace = true
rand.workspace = true
questdb-rs = { workspace = true, optional = true }
# Metrics (Wave 5)
prometheus.workspace = true
@@ -76,6 +77,7 @@ tempfile = "3.8"
[features]
default = ["database"]
database = ["sqlx"]
questdb = ["questdb-rs"]
[[bench]]
name = "ml_strategy_bench"

View File

@@ -39,6 +39,8 @@ pub mod thresholds;
pub mod tls;
pub mod traits;
pub mod types;
#[cfg(feature = "questdb")]
pub mod questdb;
// Re-export database types for external use
pub use database::{DatabaseConfig, DatabaseError, DatabasePool};

302
common/src/questdb.rs Normal file
View File

@@ -0,0 +1,302 @@
//! QuestDB client with ring buffer for graceful degradation
//!
//! Non-critical path: if QuestDB is unavailable, metrics buffer locally
//! and flush when connection is restored. Trading never depends on QuestDB.
use std::collections::VecDeque;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::Mutex;
use tracing::{debug, info, warn};
/// Maximum entries in the ring buffer before oldest are dropped
const DEFAULT_BUFFER_CAPACITY: usize = 10_000;
/// Health check interval
const HEALTH_CHECK_INTERVAL: Duration = Duration::from_secs(30);
/// A metric entry waiting to be flushed to QuestDB
#[derive(Debug, Clone)]
pub struct MetricEntry {
pub table: String,
pub columns: Vec<(String, MetricValue)>,
pub timestamp_ns: i64,
pub created_at: Instant,
}
/// Supported metric value types for ILP
#[derive(Debug, Clone)]
pub enum MetricValue {
Symbol(String),
Float(f64),
Int(i64),
Bool(bool),
}
/// QuestDB connection configuration
#[derive(Debug, Clone)]
pub struct QuestDBConfig {
/// ILP host for writes (default: localhost)
pub ilp_host: String,
/// ILP port for writes (default: 9009)
pub ilp_port: u16,
/// PostgreSQL host for queries (default: localhost)
pub pg_host: String,
/// PostgreSQL port for queries (default: 8812)
pub pg_port: u16,
/// Ring buffer capacity (default: 10,000)
pub buffer_capacity: usize,
}
impl Default for QuestDBConfig {
fn default() -> Self {
Self {
ilp_host: "localhost".into(),
ilp_port: 9009,
pg_host: "localhost".into(),
pg_port: 8812,
buffer_capacity: DEFAULT_BUFFER_CAPACITY,
}
}
}
/// QuestDB client with graceful degradation via ring buffer
pub struct QuestDBClient {
config: QuestDBConfig,
buffer: Arc<Mutex<VecDeque<MetricEntry>>>,
connected: Arc<AtomicBool>,
last_health_check: Arc<Mutex<Instant>>,
}
impl QuestDBClient {
/// Create a new QuestDB client (does not connect immediately)
pub fn new(config: QuestDBConfig) -> Self {
Self {
buffer: Arc::new(Mutex::new(VecDeque::with_capacity(config.buffer_capacity))),
connected: Arc::new(AtomicBool::new(false)),
last_health_check: Arc::new(Mutex::new(Instant::now())),
config,
}
}
/// Check if QuestDB is currently connected
pub fn is_connected(&self) -> bool {
self.connected.load(Ordering::Relaxed)
}
/// Get current buffer size (metrics waiting to be flushed)
pub async fn buffer_size(&self) -> usize {
self.buffer.lock().await.len()
}
/// Get age of oldest buffered entry
pub async fn buffer_age(&self) -> Duration {
let buf = self.buffer.lock().await;
buf.front()
.map(|e| e.created_at.elapsed())
.unwrap_or(Duration::ZERO)
}
/// Write a metric entry (buffers if QuestDB is unavailable)
pub async fn write(&self, entry: MetricEntry) {
let mut buf = self.buffer.lock().await;
// If buffer is full, drop oldest entry
if buf.len() >= self.config.buffer_capacity {
buf.pop_front();
debug!("QuestDB buffer full, dropped oldest entry");
}
buf.push_back(entry);
// Try to flush if connected
if self.connected.load(Ordering::Relaxed) {
drop(buf); // Release lock before flush
if let Err(e) = self.try_flush().await {
warn!("QuestDB flush failed: {}", e);
self.connected.store(false, Ordering::Relaxed);
}
}
}
/// Attempt to flush buffered entries to QuestDB via ILP
async fn try_flush(&self) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
let mut buf = self.buffer.lock().await;
if buf.is_empty() {
return Ok(());
}
let mut sender = questdb::ingress::SenderBuilder::new(
questdb::ingress::Protocol::Tcp,
&self.config.ilp_host,
self.config.ilp_port,
)
.build()?;
let mut ilp_buf = questdb::ingress::Buffer::new();
let entries_to_flush: Vec<MetricEntry> = buf.drain(..).collect();
let count = entries_to_flush.len();
for entry in &entries_to_flush {
ilp_buf.table(entry.table.as_str())?;
for (name, value) in &entry.columns {
match value {
MetricValue::Symbol(s) => {
ilp_buf.symbol(name.as_str(), s.as_str())?;
}
MetricValue::Float(f) => {
ilp_buf.column_f64(name.as_str(), *f)?;
}
MetricValue::Int(i) => {
ilp_buf.column_i64(name.as_str(), *i)?;
}
MetricValue::Bool(b) => {
ilp_buf.column_bool(name.as_str(), *b)?;
}
}
}
ilp_buf.at(questdb::ingress::TimestampNanos::new(entry.timestamp_ns))?;
}
sender.flush(&mut ilp_buf)?;
info!("Flushed {} entries to QuestDB", count);
Ok(())
}
/// Periodic health check — call from a background task
pub async fn health_check(&self) -> bool {
let mut last = self.last_health_check.lock().await;
if last.elapsed() < HEALTH_CHECK_INTERVAL {
return self.connected.load(Ordering::Relaxed);
}
*last = Instant::now();
drop(last);
// Try connecting via ILP
let connected = match questdb::ingress::SenderBuilder::new(
questdb::ingress::Protocol::Tcp,
&self.config.ilp_host,
self.config.ilp_port,
)
.build()
{
Ok(_sender) => {
if !self.connected.load(Ordering::Relaxed) {
info!("QuestDB connection restored");
}
true
}
Err(e) => {
if self.connected.load(Ordering::Relaxed) {
warn!("QuestDB connection lost: {}", e);
}
false
}
};
self.connected.store(connected, Ordering::Relaxed);
// If reconnected, try flushing buffer
if connected {
if let Err(e) = self.try_flush().await {
warn!("QuestDB reconnect flush failed: {}", e);
}
}
connected
}
/// Get ILP connection string for the configured host
pub fn ilp_address(&self) -> String {
format!("{}:{}", self.config.ilp_host, self.config.ilp_port)
}
/// Get PostgreSQL connection string for queries
pub fn pg_connection_string(&self) -> String {
format!(
"postgresql://admin:quest@{}:{}/qdb",
self.config.pg_host, self.config.pg_port
)
}
}
impl std::fmt::Debug for QuestDBClient {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("QuestDBClient")
.field("connected", &self.connected.load(Ordering::Relaxed))
.field("config", &self.config)
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_client_creation() {
let client = QuestDBClient::new(QuestDBConfig::default());
assert!(!client.is_connected());
assert_eq!(client.buffer_size().await, 0);
}
#[tokio::test]
async fn test_buffer_stores_entries() {
let client = QuestDBClient::new(QuestDBConfig::default());
let entry = MetricEntry {
table: "test".into(),
columns: vec![("value".into(), MetricValue::Float(1.0))],
timestamp_ns: 1_234_567_890,
created_at: Instant::now(),
};
client.write(entry).await;
assert_eq!(client.buffer_size().await, 1);
}
#[tokio::test]
async fn test_buffer_capacity_limit() {
let config = QuestDBConfig {
buffer_capacity: 3,
..Default::default()
};
let client = QuestDBClient::new(config);
for i in 0..5 {
let entry = MetricEntry {
table: "test".into(),
columns: vec![("i".into(), MetricValue::Int(i))],
timestamp_ns: i,
created_at: Instant::now(),
};
client.write(entry).await;
}
// Buffer capacity is 3, so oldest 2 should be dropped
assert_eq!(client.buffer_size().await, 3);
}
#[tokio::test]
async fn test_buffer_age_empty() {
let client = QuestDBClient::new(QuestDBConfig::default());
assert_eq!(client.buffer_age().await, Duration::ZERO);
}
#[tokio::test]
async fn test_default_config() {
let config = QuestDBConfig::default();
assert_eq!(config.ilp_host, "localhost");
assert_eq!(config.ilp_port, 9009);
assert_eq!(config.pg_host, "localhost");
assert_eq!(config.pg_port, 8812);
assert_eq!(config.buffer_capacity, 10_000);
}
#[test]
fn test_pg_connection_string() {
let client = QuestDBClient::new(QuestDBConfig::default());
assert_eq!(
client.pg_connection_string(),
"postgresql://admin:quest@localhost:8812/qdb"
);
}
}

View File

@@ -41,6 +41,27 @@ services:
networks:
- foxhunt-network
# QuestDB - High-performance time-series DB for ML metrics and feedback loop
questdb:
image: questdb/questdb:8.2.3
container_name: foxhunt-questdb
ports:
- "9009:9009" # ILP ingestion (Influx Line Protocol)
- "8812:8812" # PostgreSQL wire protocol (SQL queries)
- "9003:9003" # HTTP REST API + Web Console
volumes:
- questdb_data:/var/lib/questdb
environment:
- QDB_PG_ENABLED=true
- QDB_LINE_TCP_NET_BIND_TO=0.0.0.0:9009
healthcheck:
test: ["CMD-SHELL", "wget --no-verbose --tries=1 -O /dev/null http://localhost:9003/exec?query=SELECT%201 || exit 1"]
interval: 10s
timeout: 5s
retries: 5
networks:
- foxhunt-network
# InfluxDB - Time-series data for HFT metrics
influxdb:
image: influxdb:2.7-alpine
@@ -190,6 +211,9 @@ services:
# mTLS Validation Options
- MTLS_ENABLE_REVOCATION_CHECK=${MTLS_ENABLE_REVOCATION_CHECK:-false}
- MTLS_CRL_URL=${MTLS_CRL_URL:-}
# QuestDB for ML feedback loop metrics
- QUESTDB_ILP_HOST=questdb:9009
- QUESTDB_PG_URL=postgresql://admin:quest@questdb:8812/qdb
- KILL_SWITCH_SOCKET_PATH=/tmp/kill_switch.sock
- GRPC_PORT=50051
- RUST_LOG=info
@@ -477,6 +501,7 @@ services:
volumes:
postgres_data:
redis_data:
questdb_data:
influxdb_data:
vault_data:
prometheus_data:

View File

@@ -0,0 +1,566 @@
//! 7-Gate Conviction System
//!
//! Evaluates ensemble predictions through a series of quality gates before
//! allowing trade execution. Gates are evaluated in order; any failure
//! results in HOLD.
use serde::{Deserialize, Serialize};
/// Trading session windows (Eastern Time)
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum TradingSession {
PreMarket,
Regular,
AfterHours,
}
/// Configuration for the 7-gate conviction system
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ConvictionGateConfig {
pub model_health_threshold: f64,
pub allowed_sessions: Vec<TradingSession>,
pub min_confidence: f64,
pub max_disagreement: f64,
pub min_quorum: f64,
pub regime_tightening_factor: f64,
pub high_vol_threshold: f64,
pub conviction_scaling_enabled: bool,
}
impl Default for ConvictionGateConfig {
fn default() -> Self {
Self {
model_health_threshold: 0.70,
allowed_sessions: vec![TradingSession::Regular],
min_confidence: 0.60,
max_disagreement: 0.40,
min_quorum: 0.60,
regime_tightening_factor: 0.80,
high_vol_threshold: 0.03,
conviction_scaling_enabled: true,
}
}
}
/// Which gate rejected the decision (if any)
#[derive(Debug, Clone, PartialEq)]
pub enum GateRejection {
ModelHealth {
healthy_ratio: f64,
threshold: f64,
},
TimeOfDay {
current_session: TradingSession,
},
Confidence {
confidence: f64,
threshold: f64,
},
Agreement {
disagreement: f64,
threshold: f64,
},
Quorum {
quorum_ratio: f64,
threshold: f64,
},
Regime {
adjusted_confidence: f64,
adjusted_threshold: f64,
},
}
/// Result of passing all conviction gates
#[derive(Debug, Clone)]
pub struct GatePassResult {
pub conviction_score: f64,
pub gate_details: Vec<GateEvaluation>,
}
/// Individual gate evaluation result
#[derive(Debug, Clone)]
pub struct GateEvaluation {
pub gate_name: String,
pub value: f64,
pub threshold: f64,
pub passed: bool,
}
/// Outcome of evaluating all gates
#[derive(Debug, Clone)]
pub enum ConvictionGateOutcome {
Passed(GatePassResult),
Rejected(GateRejection),
}
/// Input data for gate evaluation
#[derive(Debug, Clone)]
pub struct GateInput {
pub confidence: f64,
pub disagreement_rate: f64,
pub quorum_ratio: f64,
pub healthy_models: usize,
pub total_models: usize,
pub current_session: TradingSession,
pub regime_volatility: f64,
}
/// Evaluates ensemble decisions through 7 conviction gates
#[derive(Debug, Clone)]
pub struct ConvictionGateEvaluator {
config: ConvictionGateConfig,
}
impl ConvictionGateEvaluator {
pub fn new(config: ConvictionGateConfig) -> Self {
Self { config }
}
pub fn config(&self) -> &ConvictionGateConfig {
&self.config
}
pub fn config_mut(&mut self) -> &mut ConvictionGateConfig {
&mut self.config
}
/// Evaluate all 7 gates in order. Returns on first rejection.
pub fn evaluate(&self, input: &GateInput) -> ConvictionGateOutcome {
let mut gate_details = Vec::with_capacity(7);
// Gate 1: Model Health - healthy_ratio >= threshold
let healthy_ratio = if input.total_models == 0 {
0.0
} else {
input.healthy_models as f64 / input.total_models as f64
};
let gate1_passed = healthy_ratio >= self.config.model_health_threshold;
gate_details.push(GateEvaluation {
gate_name: "model_health".to_string(),
value: healthy_ratio,
threshold: self.config.model_health_threshold,
passed: gate1_passed,
});
if !gate1_passed {
return ConvictionGateOutcome::Rejected(GateRejection::ModelHealth {
healthy_ratio,
threshold: self.config.model_health_threshold,
});
}
// Gate 2: Time-of-Day - current_session in allowed_sessions
let gate2_passed = self.config.allowed_sessions.contains(&input.current_session);
gate_details.push(GateEvaluation {
gate_name: "time_of_day".to_string(),
value: if gate2_passed { 1.0 } else { 0.0 },
threshold: 1.0,
passed: gate2_passed,
});
if !gate2_passed {
return ConvictionGateOutcome::Rejected(GateRejection::TimeOfDay {
current_session: input.current_session,
});
}
// Gate 6 prep: If regime_volatility > high_vol_threshold, tighten thresholds
let high_vol = input.regime_volatility > self.config.high_vol_threshold;
let (adj_min_confidence, adj_max_disagreement, adj_min_quorum) = if high_vol {
let factor = self.config.regime_tightening_factor;
let adj_conf = (self.config.min_confidence / factor).min(0.95);
let adj_disagree = (self.config.max_disagreement * factor).max(0.05);
let adj_quorum = (self.config.min_quorum / factor).min(0.95);
(adj_conf, adj_disagree, adj_quorum)
} else {
(
self.config.min_confidence,
self.config.max_disagreement,
self.config.min_quorum,
)
};
// Gate 3: Confidence >= adj_min_confidence
let gate3_passed = input.confidence >= adj_min_confidence;
gate_details.push(GateEvaluation {
gate_name: "confidence".to_string(),
value: input.confidence,
threshold: adj_min_confidence,
passed: gate3_passed,
});
if !gate3_passed {
return ConvictionGateOutcome::Rejected(GateRejection::Confidence {
confidence: input.confidence,
threshold: adj_min_confidence,
});
}
// Gate 4: Disagreement <= adj_max_disagreement
let gate4_passed = input.disagreement_rate <= adj_max_disagreement;
gate_details.push(GateEvaluation {
gate_name: "agreement".to_string(),
value: input.disagreement_rate,
threshold: adj_max_disagreement,
passed: gate4_passed,
});
if !gate4_passed {
return ConvictionGateOutcome::Rejected(GateRejection::Agreement {
disagreement: input.disagreement_rate,
threshold: adj_max_disagreement,
});
}
// Gate 5: Quorum >= adj_min_quorum
let gate5_passed = input.quorum_ratio >= adj_min_quorum;
gate_details.push(GateEvaluation {
gate_name: "quorum".to_string(),
value: input.quorum_ratio,
threshold: adj_min_quorum,
passed: gate5_passed,
});
if !gate5_passed {
return ConvictionGateOutcome::Rejected(GateRejection::Quorum {
quorum_ratio: input.quorum_ratio,
threshold: adj_min_quorum,
});
}
// Gate 6: Regime (already applied via threshold adjustments)
// Record the regime gate evaluation
let regime_value = if high_vol {
input.confidence
} else {
input.confidence
};
let regime_threshold = adj_min_confidence;
gate_details.push(GateEvaluation {
gate_name: "regime".to_string(),
value: regime_value,
threshold: regime_threshold,
passed: true, // Already enforced via adjusted thresholds in gates 3-5
});
// Gate 7: Conviction Sizing
let conviction_score = if self.config.conviction_scaling_enabled {
input.confidence
* (1.0 - input.disagreement_rate)
* input.quorum_ratio
* healthy_ratio
} else {
1.0
};
gate_details.push(GateEvaluation {
gate_name: "conviction_sizing".to_string(),
value: conviction_score,
threshold: 0.0, // No minimum threshold for conviction sizing
passed: true,
});
ConvictionGateOutcome::Passed(GatePassResult {
conviction_score,
gate_details,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
fn default_passing_input() -> GateInput {
GateInput {
confidence: 0.80,
disagreement_rate: 0.10,
quorum_ratio: 0.80,
healthy_models: 8,
total_models: 10,
current_session: TradingSession::Regular,
regime_volatility: 0.01,
}
}
#[test]
fn test_default_config_values() {
let config = ConvictionGateConfig::default();
assert!((config.model_health_threshold - 0.70).abs() < f64::EPSILON);
assert_eq!(config.allowed_sessions, vec![TradingSession::Regular]);
assert!((config.min_confidence - 0.60).abs() < f64::EPSILON);
assert!((config.max_disagreement - 0.40).abs() < f64::EPSILON);
assert!((config.min_quorum - 0.60).abs() < f64::EPSILON);
assert!((config.regime_tightening_factor - 0.80).abs() < f64::EPSILON);
assert!((config.high_vol_threshold - 0.03).abs() < f64::EPSILON);
assert!(config.conviction_scaling_enabled);
}
#[test]
fn test_gate_rejection_variants() {
let r1 = GateRejection::ModelHealth {
healthy_ratio: 0.5,
threshold: 0.7,
};
let r2 = GateRejection::ModelHealth {
healthy_ratio: 0.5,
threshold: 0.7,
};
assert_eq!(r1, r2);
let r3 = GateRejection::TimeOfDay {
current_session: TradingSession::PreMarket,
};
assert_ne!(r1, r3);
}
#[test]
fn test_gate1_model_health_passes() {
let evaluator = ConvictionGateEvaluator::new(ConvictionGateConfig::default());
let input = GateInput {
healthy_models: 8,
total_models: 10,
..default_passing_input()
};
let outcome = evaluator.evaluate(&input);
assert!(matches!(outcome, ConvictionGateOutcome::Passed(_)));
}
#[test]
fn test_gate1_model_health_rejects() {
let evaluator = ConvictionGateEvaluator::new(ConvictionGateConfig::default());
let input = GateInput {
healthy_models: 5,
total_models: 10,
..default_passing_input()
};
let outcome = evaluator.evaluate(&input);
if let ConvictionGateOutcome::Rejected(GateRejection::ModelHealth {
healthy_ratio,
threshold,
}) = outcome
{
assert!((healthy_ratio - 0.5).abs() < f64::EPSILON);
assert!((threshold - 0.70).abs() < f64::EPSILON);
} else {
panic!("Expected ModelHealth rejection, got {:?}", outcome);
}
}
#[test]
fn test_gate2_time_of_day_rejects_premarket() {
let evaluator = ConvictionGateEvaluator::new(ConvictionGateConfig::default());
let input = GateInput {
current_session: TradingSession::PreMarket,
..default_passing_input()
};
let outcome = evaluator.evaluate(&input);
if let ConvictionGateOutcome::Rejected(GateRejection::TimeOfDay { current_session }) =
outcome
{
assert_eq!(current_session, TradingSession::PreMarket);
} else {
panic!("Expected TimeOfDay rejection, got {:?}", outcome);
}
}
#[test]
fn test_gate3_confidence_rejects_low() {
let evaluator = ConvictionGateEvaluator::new(ConvictionGateConfig::default());
let input = GateInput {
confidence: 0.45,
..default_passing_input()
};
let outcome = evaluator.evaluate(&input);
if let ConvictionGateOutcome::Rejected(GateRejection::Confidence {
confidence,
threshold,
}) = outcome
{
assert!((confidence - 0.45).abs() < f64::EPSILON);
assert!((threshold - 0.60).abs() < f64::EPSILON);
} else {
panic!("Expected Confidence rejection, got {:?}", outcome);
}
}
#[test]
fn test_gate4_agreement_rejects_high_disagreement() {
let evaluator = ConvictionGateEvaluator::new(ConvictionGateConfig::default());
let input = GateInput {
disagreement_rate: 0.55,
..default_passing_input()
};
let outcome = evaluator.evaluate(&input);
if let ConvictionGateOutcome::Rejected(GateRejection::Agreement {
disagreement,
threshold,
}) = outcome
{
assert!((disagreement - 0.55).abs() < f64::EPSILON);
assert!((threshold - 0.40).abs() < f64::EPSILON);
} else {
panic!("Expected Agreement rejection, got {:?}", outcome);
}
}
#[test]
fn test_gate5_quorum_rejects_low() {
let evaluator = ConvictionGateEvaluator::new(ConvictionGateConfig::default());
let input = GateInput {
quorum_ratio: 0.40,
..default_passing_input()
};
let outcome = evaluator.evaluate(&input);
if let ConvictionGateOutcome::Rejected(GateRejection::Quorum {
quorum_ratio,
threshold,
}) = outcome
{
assert!((quorum_ratio - 0.40).abs() < f64::EPSILON);
assert!((threshold - 0.60).abs() < f64::EPSILON);
} else {
panic!("Expected Quorum rejection, got {:?}", outcome);
}
}
#[test]
fn test_conviction_score_calculation() {
let evaluator = ConvictionGateEvaluator::new(ConvictionGateConfig::default());
let input = GateInput {
confidence: 0.80,
disagreement_rate: 0.10,
quorum_ratio: 0.80,
healthy_models: 9,
total_models: 10,
current_session: TradingSession::Regular,
regime_volatility: 0.01,
};
let outcome = evaluator.evaluate(&input);
if let ConvictionGateOutcome::Passed(result) = outcome {
// score = 0.80 * (1.0 - 0.10) * 0.80 * (9/10) = 0.80 * 0.90 * 0.80 * 0.90 = 0.5184
let expected = 0.80 * 0.90 * 0.80 * 0.90;
assert!(
(result.conviction_score - expected).abs() < 1e-10,
"Expected {}, got {}",
expected,
result.conviction_score
);
} else {
panic!("Expected Passed outcome, got {:?}", outcome);
}
}
#[test]
fn test_gate6_regime_tightens_in_high_vol() {
let evaluator = ConvictionGateEvaluator::new(ConvictionGateConfig::default());
// With default config: min_confidence=0.60, tightening=0.80
// High vol: adj_min_confidence = 0.60 / 0.80 = 0.75
// 0.70 confidence < 0.75 threshold => rejected
let input = GateInput {
confidence: 0.70,
regime_volatility: 0.05,
..default_passing_input()
};
let outcome = evaluator.evaluate(&input);
if let ConvictionGateOutcome::Rejected(GateRejection::Confidence {
confidence,
threshold,
}) = outcome
{
assert!((confidence - 0.70).abs() < f64::EPSILON);
assert!((threshold - 0.75).abs() < f64::EPSILON);
} else {
panic!(
"Expected Confidence rejection due to regime tightening, got {:?}",
outcome
);
}
}
#[test]
fn test_gate6_regime_no_effect_in_low_vol() {
let evaluator = ConvictionGateEvaluator::new(ConvictionGateConfig::default());
// Low vol: thresholds unchanged, min_confidence=0.60
// 0.65 confidence >= 0.60 => passes
let input = GateInput {
confidence: 0.65,
regime_volatility: 0.01,
..default_passing_input()
};
let outcome = evaluator.evaluate(&input);
assert!(
matches!(outcome, ConvictionGateOutcome::Passed(_)),
"Expected Passed in low vol, got {:?}",
outcome
);
}
#[test]
fn test_zero_models_rejects() {
let evaluator = ConvictionGateEvaluator::new(ConvictionGateConfig::default());
let input = GateInput {
healthy_models: 0,
total_models: 0,
..default_passing_input()
};
let outcome = evaluator.evaluate(&input);
if let ConvictionGateOutcome::Rejected(GateRejection::ModelHealth {
healthy_ratio,
threshold,
}) = outcome
{
assert!((healthy_ratio - 0.0).abs() < f64::EPSILON);
assert!((threshold - 0.70).abs() < f64::EPSILON);
} else {
panic!("Expected ModelHealth rejection for 0/0 models, got {:?}", outcome);
}
}
#[test]
fn test_conviction_scaling_disabled() {
let mut config = ConvictionGateConfig::default();
config.conviction_scaling_enabled = false;
let evaluator = ConvictionGateEvaluator::new(config);
let input = default_passing_input();
let outcome = evaluator.evaluate(&input);
if let ConvictionGateOutcome::Passed(result) = outcome {
assert!(
(result.conviction_score - 1.0).abs() < f64::EPSILON,
"Expected 1.0 when scaling disabled, got {}",
result.conviction_score
);
} else {
panic!("Expected Passed outcome, got {:?}", outcome);
}
}
#[test]
fn test_gate_details_count() {
let evaluator = ConvictionGateEvaluator::new(ConvictionGateConfig::default());
let input = default_passing_input();
let outcome = evaluator.evaluate(&input);
if let ConvictionGateOutcome::Passed(result) = outcome {
assert_eq!(
result.gate_details.len(),
7,
"Expected 7 gate evaluations, got {}",
result.gate_details.len()
);
} else {
panic!("Expected Passed outcome, got {:?}", outcome);
}
}
#[test]
fn test_after_hours_allowed_when_configured() {
let mut config = ConvictionGateConfig::default();
config.allowed_sessions.push(TradingSession::AfterHours);
let evaluator = ConvictionGateEvaluator::new(config);
let input = GateInput {
current_session: TradingSession::AfterHours,
..default_passing_input()
};
let outcome = evaluator.evaluate(&input);
assert!(
matches!(outcome, ConvictionGateOutcome::Passed(_)),
"Expected Passed for AfterHours when configured, got {:?}",
outcome
);
}
}

View File

@@ -4,9 +4,13 @@
//! from multiple ML models (DQN, PPO, TFT, MAMBA-2, Liquid, TLOB) for production trading decisions.
//! Supports dynamic weighting based on performance and model diversity metrics.
use crate::ensemble::conviction_gates::{
ConvictionGateConfig, ConvictionGateEvaluator, ConvictionGateOutcome, GateInput, TradingSession,
};
use crate::ensemble::inference_adapter::{FeatureVector, ModelInferenceAdapter};
use crate::ensemble::{EnsembleDecision, ModelVote, ModelWeight, TradingAction};
use crate::{Features, MLError, MLResult, ModelPrediction};
use chrono::{Timelike, Utc};
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
@@ -34,6 +38,9 @@ pub struct EnsembleCoordinator {
/// Real model inference adapters for production predictions
adapters: Vec<Box<dyn ModelInferenceAdapter>>,
/// Conviction gate evaluator (optional — None means no gating)
conviction_gates: Option<ConvictionGateEvaluator>,
}
impl std::fmt::Debug for EnsembleCoordinator {
@@ -77,9 +84,28 @@ impl EnsembleCoordinator {
model_weights: Arc::new(RwLock::new(HashMap::new())),
config,
adapters: Vec::new(),
conviction_gates: Some(ConvictionGateEvaluator::new(ConvictionGateConfig {
// Allow all sessions by default — trading service narrows to Regular
allowed_sessions: vec![
TradingSession::PreMarket,
TradingSession::Regular,
TradingSession::AfterHours,
],
..ConvictionGateConfig::default()
})),
}
}
/// Replace conviction gate configuration
pub fn set_conviction_gates(&mut self, config: ConvictionGateConfig) {
self.conviction_gates = Some(ConvictionGateEvaluator::new(config));
}
/// Get mutable reference to conviction gate evaluator
pub fn conviction_gates_mut(&mut self) -> Option<&mut ConvictionGateEvaluator> {
self.conviction_gates.as_mut()
}
/// Add a real model inference adapter to the ensemble.
/// When an adapter's model_name matches a registered model, its predictions
/// are used for ensemble inference.
@@ -129,6 +155,50 @@ impl EnsembleCoordinator {
.aggregate(predictions, &*self.model_weights.read().await)
.await?;
// Apply conviction gates if configured
if let Some(ref gates) = self.conviction_gates {
let gate_input = GateInput {
confidence: decision.confidence,
disagreement_rate: decision.disagreement_rate,
quorum_ratio: Self::calculate_quorum_ratio(&decision),
healthy_models: self.adapters.iter().filter(|a| a.is_ready()).count(),
total_models: self.adapters.len(),
current_session: Self::current_trading_session(),
regime_volatility: Self::extract_regime_volatility(features),
};
match gates.evaluate(&gate_input) {
ConvictionGateOutcome::Passed(pass_result) => {
info!(
"Conviction gates passed: score={:.3}, {} gates evaluated",
pass_result.conviction_score,
pass_result.gate_details.len()
);
let mut gated_decision = decision;
gated_decision.metadata.insert(
"conviction_score".into(),
serde_json::Value::from(pass_result.conviction_score),
);
return Ok(gated_decision);
}
ConvictionGateOutcome::Rejected(rejection) => {
info!("Conviction gate rejected: {:?} — forcing HOLD", rejection);
let mut hold_decision = EnsembleDecision::new(
TradingAction::Hold,
decision.confidence,
0.0,
decision.disagreement_rate,
decision.model_votes,
);
hold_decision.metadata.insert(
"gate_rejection".into(),
serde_json::Value::String(format!("{:?}", rejection)),
);
return Ok(hold_decision);
}
}
}
info!(
"Ensemble decision: {:?}, confidence: {:.3}, disagreement: {:.3}",
decision.action, decision.confidence, decision.disagreement_rate
@@ -292,6 +362,46 @@ impl EnsembleCoordinator {
Ok(())
}
/// Calculate quorum ratio (fraction of models agreeing on majority direction)
fn calculate_quorum_ratio(decision: &EnsembleDecision) -> f64 {
if decision.model_votes.is_empty() {
return 0.0;
}
let majority_action = &decision.action;
let agreeing = decision
.model_votes
.values()
.filter(|v| TradingAction::from_signal(v.signal, 0.3) == *majority_action)
.count();
agreeing as f64 / decision.model_votes.len() as f64
}
/// Determine current trading session (Eastern Time, simplified UTC-5)
fn current_trading_session() -> TradingSession {
let now = Utc::now();
let et_hour = (now.hour() + 24 - 5) % 24;
let et_minute = now.minute();
let et_time = et_hour * 60 + et_minute;
match et_time {
t if t < 240 => TradingSession::AfterHours, // 00:00-04:00 ET
t if t < 570 => TradingSession::PreMarket, // 04:00-09:30 ET
t if t < 960 => TradingSession::Regular, // 09:30-16:00 ET
_ => TradingSession::AfterHours, // 16:00-24:00 ET
}
}
/// Extract regime volatility from feature vector
/// Regime features are at indices 48-50 in the 51-dim standard feature vector
fn extract_regime_volatility(features: &Features) -> f64 {
features
.values
.get(48)
.copied()
.unwrap_or(0.01)
.abs()
}
}
impl Default for EnsembleCoordinator {

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

@@ -21,6 +21,9 @@ pub mod inference_adapter;
pub mod inference_ensemble;
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::{
@@ -48,6 +51,18 @@ pub use metrics::{
};
pub use training_integration::EnsembleTrainingIntegration;
pub use inference_adapter::{EnsemblePrediction, FeatureVector, ModelInferenceAdapter, PredictionMeta};
pub use conviction_gates::{
ConvictionGateConfig, ConvictionGateEvaluator, ConvictionGateOutcome, GateEvaluation,
GateInput, GatePassResult, GateRejection, TradingSession,
};
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

@@ -0,0 +1,426 @@
//! Autonomous weight optimizer for ensemble models
//!
//! Adjusts model weights based on rolling Sharpe ratios using
//! exponentially-weighted moving average. Bounded by safety rails.
use std::collections::HashMap;
use std::time::{Duration, Instant};
use tracing::{info, warn};
/// Minimum weight for any model (prevents zeroing out)
const MIN_WEIGHT: f64 = 0.05;
/// Maximum weight for any model (prevents dominance)
const MAX_WEIGHT: f64 = 0.40;
/// Maximum weight change per optimization cycle
const MAX_STEP: f64 = 0.03;
/// Minimum number of trades before first adjustment
const MIN_OBSERVATIONS: u64 = 100;
/// Cooldown between adjustments
const COOLDOWN: Duration = Duration::from_secs(24 * 3600); // 24 hours
/// Grace period for newly deployed models
const GRACE_PERIOD: Duration = Duration::from_secs(7 * 24 * 3600); // 7 days
/// Per-model rolling performance metrics (fed from QuestDB queries)
#[derive(Debug, Clone)]
pub struct ModelRollingMetrics {
pub model_id: String,
pub sharpe_30d: f64,
pub win_rate_30d: f64,
pub prediction_accuracy: f64,
pub trade_count: u64,
pub deployed_at: Instant,
}
/// Weight optimizer configuration
#[derive(Debug, Clone)]
pub struct WeightOptimizerConfig {
pub ema_alpha: f64,
pub min_weight: f64,
pub max_weight: f64,
pub max_step: f64,
pub min_observations: u64,
pub cooldown: Duration,
pub grace_period: Duration,
}
impl Default for WeightOptimizerConfig {
fn default() -> Self {
Self {
ema_alpha: 0.1,
min_weight: MIN_WEIGHT,
max_weight: MAX_WEIGHT,
max_step: MAX_STEP,
min_observations: MIN_OBSERVATIONS,
cooldown: COOLDOWN,
grace_period: GRACE_PERIOD,
}
}
}
/// Proposed weight changes from the optimizer
#[derive(Debug, Clone)]
pub struct WeightAdjustment {
pub model_id: String,
pub old_weight: f64,
pub new_weight: f64,
pub reason: String,
}
/// Result of an optimization cycle
#[derive(Debug, Clone)]
pub enum OptimizationResult {
/// Weights adjusted successfully
Adjusted(Vec<WeightAdjustment>),
/// Not enough observations yet
InsufficientData { total_trades: u64, required: u64 },
/// Still in cooldown from last adjustment
Cooldown { remaining: Duration },
/// Kill switch is active, no adjustments allowed
KillSwitchActive,
}
/// Autonomous weight optimizer
#[derive(Debug)]
pub struct WeightOptimizer {
config: WeightOptimizerConfig,
last_adjustment: Option<Instant>,
ema_sharpe: HashMap<String, f64>,
frozen: bool,
}
impl WeightOptimizer {
pub fn new(config: WeightOptimizerConfig) -> Self {
Self {
config,
last_adjustment: None,
ema_sharpe: HashMap::new(),
frozen: false,
}
}
/// Freeze all adjustments (kill switch)
pub fn freeze(&mut self) {
self.frozen = true;
warn!("Weight optimizer frozen by kill switch");
}
/// Unfreeze adjustments (human acknowledgment)
pub fn unfreeze(&mut self) {
self.frozen = false;
info!("Weight optimizer unfrozen");
}
pub fn is_frozen(&self) -> bool {
self.frozen
}
/// Run one optimization cycle
///
/// Takes current weights and rolling metrics, returns proposed adjustments.
pub fn optimize(
&mut self,
current_weights: &HashMap<String, f64>,
metrics: &[ModelRollingMetrics],
) -> OptimizationResult {
if self.frozen {
return OptimizationResult::KillSwitchActive;
}
// Check cooldown
if let Some(last) = self.last_adjustment {
let elapsed = last.elapsed();
if elapsed < self.config.cooldown {
return OptimizationResult::Cooldown {
remaining: self.config.cooldown - elapsed,
};
}
}
// Check minimum observations
let total_trades: u64 = metrics.iter().map(|m| m.trade_count).sum();
if total_trades < self.config.min_observations {
return OptimizationResult::InsufficientData {
total_trades,
required: self.config.min_observations,
};
}
// Filter out models in grace period
let eligible: Vec<&ModelRollingMetrics> = metrics
.iter()
.filter(|m| m.deployed_at.elapsed() >= self.config.grace_period)
.collect();
if eligible.is_empty() {
return OptimizationResult::InsufficientData {
total_trades,
required: self.config.min_observations,
};
}
// Update EMA of Sharpe ratios
for m in &eligible {
let prev = self
.ema_sharpe
.get(&m.model_id)
.copied()
.unwrap_or(m.sharpe_30d);
let new_ema =
self.config.ema_alpha * m.sharpe_30d + (1.0 - self.config.ema_alpha) * prev;
self.ema_sharpe.insert(m.model_id.clone(), new_ema);
}
// Calculate raw weights from EMA Sharpe (shift to positive range)
let min_ema = self
.ema_sharpe
.values()
.cloned()
.fold(f64::INFINITY, f64::min);
let shift = if min_ema < 0.0 {
min_ema.abs() + 0.1
} else {
0.0
};
let raw_weights: HashMap<String, f64> = self
.ema_sharpe
.iter()
.map(|(id, &ema)| (id.clone(), (ema + shift).max(0.01)))
.collect();
let total_raw: f64 = raw_weights.values().sum();
if total_raw <= 0.0 {
return OptimizationResult::InsufficientData {
total_trades,
required: self.config.min_observations,
};
}
// Normalize and clamp
let mut target_weights: HashMap<String, f64> = raw_weights
.iter()
.map(|(id, &raw)| {
let normalized = raw / total_raw;
let clamped = normalized.clamp(self.config.min_weight, self.config.max_weight);
(id.clone(), clamped)
})
.collect();
// Re-normalize after clamping
let total_clamped: f64 = target_weights.values().sum();
if total_clamped > 0.0 {
for w in target_weights.values_mut() {
*w /= total_clamped;
}
}
// Apply max step limit
let mut adjustments = Vec::new();
for (model_id, &target) in &target_weights {
let current = current_weights.get(model_id).copied().unwrap_or(target);
let delta = (target - current).clamp(-self.config.max_step, self.config.max_step);
let new_weight =
(current + delta).clamp(self.config.min_weight, self.config.max_weight);
if (new_weight - current).abs() > 1e-6 {
adjustments.push(WeightAdjustment {
model_id: model_id.clone(),
old_weight: current,
new_weight,
reason: format!(
"EMA Sharpe: {:.3}, target: {:.3}",
self.ema_sharpe.get(model_id).unwrap_or(&0.0),
target
),
});
}
}
if !adjustments.is_empty() {
self.last_adjustment = Some(Instant::now());
info!("Weight optimizer adjusted {} models", adjustments.len());
}
OptimizationResult::Adjusted(adjustments)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_metrics(models: &[(&str, f64, u64)]) -> Vec<ModelRollingMetrics> {
models
.iter()
.map(|(id, sharpe, trades)| ModelRollingMetrics {
model_id: id.to_string(),
sharpe_30d: *sharpe,
win_rate_30d: 0.55,
prediction_accuracy: 0.60,
trade_count: *trades,
deployed_at: Instant::now() - Duration::from_secs(30 * 24 * 3600),
})
.collect()
}
fn equal_weights(models: &[&str]) -> HashMap<String, f64> {
let w = 1.0 / models.len() as f64;
models.iter().map(|id| (id.to_string(), w)).collect()
}
#[test]
fn test_insufficient_data() {
let mut opt = WeightOptimizer::new(WeightOptimizerConfig::default());
let weights = equal_weights(&["dqn", "ppo"]);
let metrics = make_metrics(&[("dqn", 1.0, 10), ("ppo", 0.5, 10)]);
let result = opt.optimize(&weights, &metrics);
assert!(matches!(
result,
OptimizationResult::InsufficientData { .. }
));
}
#[test]
fn test_cooldown_enforced() {
let mut opt = WeightOptimizer::new(WeightOptimizerConfig::default());
let weights = equal_weights(&["dqn", "ppo"]);
let metrics = make_metrics(&[("dqn", 1.5, 200), ("ppo", 0.5, 200)]);
// First optimization succeeds
let result = opt.optimize(&weights, &metrics);
assert!(matches!(result, OptimizationResult::Adjusted(_)));
// Second immediately after should be cooldown
let result2 = opt.optimize(&weights, &metrics);
assert!(matches!(result2, OptimizationResult::Cooldown { .. }));
}
#[test]
fn test_kill_switch_freezes() {
let mut opt = WeightOptimizer::new(WeightOptimizerConfig::default());
opt.freeze();
let weights = equal_weights(&["dqn"]);
let metrics = make_metrics(&[("dqn", 1.0, 200)]);
let result = opt.optimize(&weights, &metrics);
assert!(matches!(result, OptimizationResult::KillSwitchActive));
}
#[test]
fn test_unfreeze_allows_optimization() {
let mut opt = WeightOptimizer::new(WeightOptimizerConfig::default());
opt.freeze();
assert!(opt.is_frozen());
opt.unfreeze();
assert!(!opt.is_frozen());
}
#[test]
fn test_weights_bounded() {
let mut config = WeightOptimizerConfig::default();
config.cooldown = Duration::ZERO;
let mut opt = WeightOptimizer::new(config);
let weights = equal_weights(&["dqn", "ppo", "tft"]);
let metrics = make_metrics(&[("dqn", 5.0, 200), ("ppo", -2.0, 200), ("tft", 0.5, 200)]);
let result = opt.optimize(&weights, &metrics);
if let OptimizationResult::Adjusted(adjustments) = result {
for adj in &adjustments {
assert!(
adj.new_weight >= MIN_WEIGHT,
"Weight below minimum: {}",
adj.new_weight
);
assert!(
adj.new_weight <= MAX_WEIGHT,
"Weight above maximum: {}",
adj.new_weight
);
}
}
}
#[test]
fn test_max_step_enforced() {
let mut config = WeightOptimizerConfig::default();
config.cooldown = Duration::ZERO;
let mut opt = WeightOptimizer::new(config);
// Both weights within [MIN_WEIGHT, MAX_WEIGHT] range
let mut weights = HashMap::new();
weights.insert("dqn".into(), 0.30);
weights.insert("ppo".into(), 0.30);
weights.insert("tft".into(), 0.40);
let metrics = make_metrics(&[("dqn", 5.0, 200), ("ppo", -1.0, 200), ("tft", 0.5, 200)]);
let result = opt.optimize(&weights, &metrics);
if let OptimizationResult::Adjusted(adjustments) = result {
for adj in &adjustments {
let delta = (adj.new_weight - adj.old_weight).abs();
// Step should be bounded by MAX_STEP, unless weight bounds force a correction
assert!(
delta <= MAX_STEP + 1e-6 || adj.new_weight == MIN_WEIGHT || adj.new_weight == MAX_WEIGHT,
"Step too large: {} for {} (old={}, new={})",
delta,
adj.model_id,
adj.old_weight,
adj.new_weight
);
}
}
}
#[test]
fn test_grace_period_excludes_new_models() {
let mut config = WeightOptimizerConfig::default();
config.cooldown = Duration::ZERO;
let mut opt = WeightOptimizer::new(config);
let weights = equal_weights(&["dqn", "new_model"]);
let mut metrics = make_metrics(&[("dqn", 1.0, 200)]);
// New model deployed 1 day ago (within 7-day grace period)
metrics.push(ModelRollingMetrics {
model_id: "new_model".into(),
sharpe_30d: 2.0,
win_rate_30d: 0.70,
prediction_accuracy: 0.80,
trade_count: 50,
deployed_at: Instant::now() - Duration::from_secs(24 * 3600),
});
let result = opt.optimize(&weights, &metrics);
if let OptimizationResult::Adjusted(adjustments) = result {
// new_model should not be in adjustments (grace period)
assert!(adjustments.iter().all(|a| a.model_id != "new_model"));
}
}
#[test]
fn test_weights_sum_approximately_one() {
let mut config = WeightOptimizerConfig::default();
config.cooldown = Duration::ZERO;
let mut opt = WeightOptimizer::new(config);
let weights = equal_weights(&["a", "b", "c", "d"]);
let metrics = make_metrics(&[
("a", 1.0, 200),
("b", 1.5, 200),
("c", 0.5, 200),
("d", 0.8, 200),
]);
let result = opt.optimize(&weights, &metrics);
if let OptimizationResult::Adjusted(adjustments) = result {
let mut new_weights = weights.clone();
for adj in &adjustments {
new_weights.insert(adj.model_id.clone(), adj.new_weight);
}
let sum: f64 = new_weights.values().sum();
// Due to step limits, sum may not be exactly 1.0 but should be close
assert!(
(sum - 1.0).abs() < 0.2,
"Weights sum to {} (expected ~1.0)",
sum
);
}
}
}

View File

@@ -834,6 +834,7 @@ pub mod random_model;
pub mod model_registry;
pub mod data_pipeline;
pub mod asset_selection;
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

@@ -0,0 +1,529 @@
//! Autonomous Feedback Loop Orchestrator
//!
//! Runs as a background tokio task. Periodically:
//! 1. Queries rolling performance metrics (from QuestDB or in-memory)
//! 2. Runs weight optimizer → applies weight adjustments
//! 3. Runs gate optimizer → applies threshold adjustments
//! 4. Checks retraining triggers (Sharpe < threshold, accuracy < threshold)
//! 5. Monitors ensemble Sharpe for kill switch activation
use ml::ensemble::conviction_gates::ConvictionGateConfig;
use ml::ensemble::gate_optimizer::{
GateBucketMetrics, GateOptimizationResult, GateOptimizer, GateOptimizerConfig,
};
use ml::ensemble::weight_optimizer::{
ModelRollingMetrics, OptimizationResult, WeightOptimizer, WeightOptimizerConfig,
};
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::Mutex;
use tracing::{error, info, warn};
/// Retraining trigger reasons
#[derive(Debug, Clone)]
pub enum RetrainingTrigger {
/// Sharpe ratio dropped below threshold
LowSharpe { model_id: String, sharpe: f64 },
/// Accuracy dropped below threshold
LowAccuracy { model_id: String, accuracy: f64 },
/// Models are too correlated
HighCorrelation { pair: (String, String), correlation: f64 },
}
/// Configuration for the feedback loop
#[derive(Debug, Clone)]
pub struct FeedbackLoopConfig {
/// How often to run the optimization cycle
pub cycle_interval: Duration,
/// Kill switch: if ensemble Sharpe drops below this, freeze all optimizers
pub kill_switch_threshold: f64,
/// Sharpe ratio threshold for retraining trigger
pub retrain_sharpe_threshold: f64,
/// Accuracy threshold for retraining trigger
pub retrain_accuracy_threshold: f64,
/// Correlation threshold for retraining trigger
pub retrain_correlation_threshold: f64,
}
impl Default for FeedbackLoopConfig {
fn default() -> Self {
Self {
cycle_interval: Duration::from_secs(24 * 3600), // 24h
kill_switch_threshold: -1.0,
retrain_sharpe_threshold: -0.5,
retrain_accuracy_threshold: 0.45,
retrain_correlation_threshold: 0.90,
}
}
}
/// Result of a single feedback loop cycle
#[derive(Debug, Clone)]
pub struct CycleResult {
/// Weight adjustments applied (if any)
pub weight_result: OptimizationResult,
/// Gate adjustments applied (if any)
pub gate_result: GateOptimizationResult,
/// Retraining triggers detected
pub retraining_triggers: Vec<RetrainingTrigger>,
/// Whether kill switch was activated this cycle
pub kill_switch_activated: bool,
}
/// Trait for providing rolling metrics to the feedback loop.
/// In production, this queries QuestDB. In tests, it's mocked.
pub trait MetricsProvider: Send + Sync {
/// Get rolling model-level metrics for weight optimization
fn get_model_metrics(&self) -> Vec<ModelRollingMetrics>;
/// Get confidence-bucketed metrics for gate optimization
fn get_gate_buckets(&self) -> Vec<GateBucketMetrics>;
/// Get current model weights
fn get_current_weights(&self) -> HashMap<String, f64>;
/// Get current conviction gate config
fn get_gate_config(&self) -> ConvictionGateConfig;
/// Get 7-day ensemble Sharpe ratio
fn get_ensemble_sharpe_7d(&self) -> f64;
/// Get model pairwise correlations (for retraining triggers)
fn get_model_correlations(&self) -> Vec<((String, String), f64)>;
}
/// Autonomous feedback loop orchestrator
pub struct FeedbackLoop {
weight_optimizer: Arc<Mutex<WeightOptimizer>>,
gate_optimizer: Arc<Mutex<GateOptimizer>>,
config: FeedbackLoopConfig,
}
impl FeedbackLoop {
pub fn new(config: FeedbackLoopConfig) -> Self {
Self {
weight_optimizer: Arc::new(Mutex::new(WeightOptimizer::new(
WeightOptimizerConfig::default(),
))),
gate_optimizer: Arc::new(Mutex::new(GateOptimizer::new(
GateOptimizerConfig::default(),
))),
config,
}
}
/// Create with custom optimizer configs
pub fn with_optimizers(
config: FeedbackLoopConfig,
weight_config: WeightOptimizerConfig,
gate_config: GateOptimizerConfig,
) -> Self {
Self {
weight_optimizer: Arc::new(Mutex::new(WeightOptimizer::new(weight_config))),
gate_optimizer: Arc::new(Mutex::new(GateOptimizer::new(gate_config))),
config,
}
}
/// Run one optimization cycle. Returns the results for observability.
pub async fn run_cycle(&self, metrics: &dyn MetricsProvider) -> CycleResult {
// 1. Check kill switch
let ensemble_sharpe = metrics.get_ensemble_sharpe_7d();
let kill_switch_activated = ensemble_sharpe < self.config.kill_switch_threshold;
if kill_switch_activated {
warn!(
ensemble_sharpe = ensemble_sharpe,
threshold = self.config.kill_switch_threshold,
"Kill switch activated — freezing all optimizers"
);
self.weight_optimizer.lock().await.freeze();
self.gate_optimizer.lock().await.freeze();
return CycleResult {
weight_result: OptimizationResult::KillSwitchActive,
gate_result: GateOptimizationResult::KillSwitchActive,
retraining_triggers: vec![],
kill_switch_activated: true,
};
}
// 2. Run weight optimization
let model_metrics = metrics.get_model_metrics();
let current_weights = metrics.get_current_weights();
let weight_result = self
.weight_optimizer
.lock()
.await
.optimize(&current_weights, &model_metrics);
if let OptimizationResult::Adjusted(ref adjustments) = weight_result {
info!(
count = adjustments.len(),
"Feedback loop: weight adjustments proposed"
);
}
// 3. Run gate optimization
let gate_config = metrics.get_gate_config();
let buckets = metrics.get_gate_buckets();
let gate_result = self
.gate_optimizer
.lock()
.await
.optimize(&gate_config, &buckets);
if let GateOptimizationResult::Adjusted(ref adjustments) = gate_result {
info!(
count = adjustments.len(),
"Feedback loop: gate adjustments proposed"
);
}
// 4. Check retraining triggers
let retraining_triggers = self.check_retraining_triggers(metrics);
if !retraining_triggers.is_empty() {
warn!(
count = retraining_triggers.len(),
"Feedback loop: retraining triggers detected"
);
}
CycleResult {
weight_result,
gate_result,
retraining_triggers,
kill_switch_activated: false,
}
}
/// Spawn the feedback loop as a background tokio task.
///
/// Returns a handle that can be used to abort the loop.
pub fn spawn(
self: Arc<Self>,
metrics: Arc<dyn MetricsProvider>,
) -> tokio::task::JoinHandle<()> {
let interval = self.config.cycle_interval;
tokio::spawn(async move {
let mut ticker = tokio::time::interval(interval);
// Skip the first immediate tick
ticker.tick().await;
loop {
ticker.tick().await;
info!("Feedback loop: starting optimization cycle");
let result = self.run_cycle(metrics.as_ref()).await;
if result.kill_switch_activated {
error!("Feedback loop: kill switch active, pausing until manual intervention");
// Sleep longer when kill switch is active to avoid log spam
tokio::time::sleep(Duration::from_secs(3600)).await;
}
}
})
}
/// Freeze both optimizers (manual kill switch)
pub async fn freeze(&self) {
self.weight_optimizer.lock().await.freeze();
self.gate_optimizer.lock().await.freeze();
}
/// Unfreeze both optimizers (requires human acknowledgment)
pub async fn unfreeze(&self) {
self.weight_optimizer.lock().await.unfreeze();
self.gate_optimizer.lock().await.unfreeze();
}
/// Check if either optimizer is frozen
pub async fn is_frozen(&self) -> bool {
self.weight_optimizer.lock().await.is_frozen()
|| self.gate_optimizer.lock().await.is_frozen()
}
fn check_retraining_triggers(&self, metrics: &dyn MetricsProvider) -> Vec<RetrainingTrigger> {
let mut triggers = Vec::new();
// Check per-model metrics
for m in &metrics.get_model_metrics() {
if m.sharpe_30d < self.config.retrain_sharpe_threshold {
triggers.push(RetrainingTrigger::LowSharpe {
model_id: m.model_id.clone(),
sharpe: m.sharpe_30d,
});
}
if m.prediction_accuracy < self.config.retrain_accuracy_threshold {
triggers.push(RetrainingTrigger::LowAccuracy {
model_id: m.model_id.clone(),
accuracy: m.prediction_accuracy,
});
}
}
// Check correlations
for ((a, b), corr) in &metrics.get_model_correlations() {
if *corr > self.config.retrain_correlation_threshold {
triggers.push(RetrainingTrigger::HighCorrelation {
pair: (a.clone(), b.clone()),
correlation: *corr,
});
}
}
triggers
}
}
impl std::fmt::Debug for FeedbackLoop {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("FeedbackLoop")
.field("config", &self.config)
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Instant;
/// Mock metrics provider for testing
struct MockMetrics {
model_metrics: Vec<ModelRollingMetrics>,
gate_buckets: Vec<GateBucketMetrics>,
weights: HashMap<String, f64>,
gate_config: ConvictionGateConfig,
ensemble_sharpe: f64,
correlations: Vec<((String, String), f64)>,
}
impl MockMetrics {
fn healthy() -> Self {
Self {
model_metrics: vec![
ModelRollingMetrics {
model_id: "dqn".into(),
sharpe_30d: 1.5,
win_rate_30d: 0.58,
prediction_accuracy: 0.62,
trade_count: 200,
deployed_at: Instant::now() - Duration::from_secs(30 * 24 * 3600),
},
ModelRollingMetrics {
model_id: "ppo".into(),
sharpe_30d: 0.8,
win_rate_30d: 0.55,
prediction_accuracy: 0.58,
trade_count: 200,
deployed_at: Instant::now() - Duration::from_secs(30 * 24 * 3600),
},
],
gate_buckets: vec![
GateBucketMetrics {
confidence_lower: 0.50,
confidence_upper: 0.60,
win_rate: 0.56,
trade_count: 100,
avg_pnl: 0.002,
},
GateBucketMetrics {
confidence_lower: 0.60,
confidence_upper: 0.70,
win_rate: 0.58,
trade_count: 100,
avg_pnl: 0.003,
},
GateBucketMetrics {
confidence_lower: 0.70,
confidence_upper: 0.80,
win_rate: 0.62,
trade_count: 100,
avg_pnl: 0.005,
},
],
weights: {
let mut w = HashMap::new();
w.insert("dqn".into(), 0.5);
w.insert("ppo".into(), 0.5);
w
},
gate_config: ConvictionGateConfig::default(),
ensemble_sharpe: 1.2,
correlations: vec![],
}
}
}
impl MetricsProvider for MockMetrics {
fn get_model_metrics(&self) -> Vec<ModelRollingMetrics> {
self.model_metrics.clone()
}
fn get_gate_buckets(&self) -> Vec<GateBucketMetrics> {
self.gate_buckets.clone()
}
fn get_current_weights(&self) -> HashMap<String, f64> {
self.weights.clone()
}
fn get_gate_config(&self) -> ConvictionGateConfig {
self.gate_config.clone()
}
fn get_ensemble_sharpe_7d(&self) -> f64 {
self.ensemble_sharpe
}
fn get_model_correlations(&self) -> Vec<((String, String), f64)> {
self.correlations.clone()
}
}
#[tokio::test]
async fn test_healthy_cycle_adjusts_weights() {
let mut config = FeedbackLoopConfig::default();
config.cycle_interval = Duration::from_millis(100);
let feedback = FeedbackLoop::with_optimizers(
config,
WeightOptimizerConfig {
cooldown: Duration::ZERO,
..Default::default()
},
GateOptimizerConfig {
cooldown: Duration::ZERO,
..Default::default()
},
);
let metrics = MockMetrics::healthy();
let result = feedback.run_cycle(&metrics).await;
// Should produce weight adjustments (models have different Sharpe)
assert!(matches!(result.weight_result, OptimizationResult::Adjusted(_)));
assert!(!result.kill_switch_activated);
assert!(result.retraining_triggers.is_empty());
}
#[tokio::test]
async fn test_kill_switch_freezes_optimizers() {
let feedback = FeedbackLoop::new(FeedbackLoopConfig {
kill_switch_threshold: -1.0,
..Default::default()
});
let mut metrics = MockMetrics::healthy();
metrics.ensemble_sharpe = -2.0; // Below kill switch
let result = feedback.run_cycle(&metrics).await;
assert!(result.kill_switch_activated);
assert!(matches!(
result.weight_result,
OptimizationResult::KillSwitchActive
));
assert!(matches!(
result.gate_result,
GateOptimizationResult::KillSwitchActive
));
assert!(feedback.is_frozen().await);
}
#[tokio::test]
async fn test_manual_freeze_unfreeze() {
let feedback = FeedbackLoop::new(FeedbackLoopConfig::default());
assert!(!feedback.is_frozen().await);
feedback.freeze().await;
assert!(feedback.is_frozen().await);
let metrics = MockMetrics::healthy();
let result = feedback.run_cycle(&metrics).await;
assert!(matches!(
result.weight_result,
OptimizationResult::KillSwitchActive
));
feedback.unfreeze().await;
assert!(!feedback.is_frozen().await);
}
#[tokio::test]
async fn test_retraining_trigger_low_sharpe() {
let config = FeedbackLoopConfig {
retrain_sharpe_threshold: -0.5,
..Default::default()
};
let feedback = FeedbackLoop::new(config);
let mut metrics = MockMetrics::healthy();
metrics.model_metrics[1].sharpe_30d = -1.0; // PPO tanking
let result = feedback.run_cycle(&metrics).await;
assert_eq!(result.retraining_triggers.len(), 1);
if let RetrainingTrigger::LowSharpe { ref model_id, .. } =
result.retraining_triggers[0]
{
assert_eq!(model_id, "ppo");
} else {
panic!("Expected LowSharpe trigger");
}
}
#[tokio::test]
async fn test_retraining_trigger_low_accuracy() {
let config = FeedbackLoopConfig {
retrain_accuracy_threshold: 0.45,
..Default::default()
};
let feedback = FeedbackLoop::new(config);
let mut metrics = MockMetrics::healthy();
metrics.model_metrics[0].prediction_accuracy = 0.40; // DQN accuracy dropped
let result = feedback.run_cycle(&metrics).await;
let accuracy_triggers: Vec<_> = result
.retraining_triggers
.iter()
.filter(|t| matches!(t, RetrainingTrigger::LowAccuracy { .. }))
.collect();
assert_eq!(accuracy_triggers.len(), 1);
}
#[tokio::test]
async fn test_retraining_trigger_high_correlation() {
let config = FeedbackLoopConfig {
retrain_correlation_threshold: 0.90,
..Default::default()
};
let feedback = FeedbackLoop::new(config);
let mut metrics = MockMetrics::healthy();
metrics.correlations = vec![
(("dqn".into(), "ppo".into()), 0.95), // Too correlated
];
let result = feedback.run_cycle(&metrics).await;
let corr_triggers: Vec<_> = result
.retraining_triggers
.iter()
.filter(|t| matches!(t, RetrainingTrigger::HighCorrelation { .. }))
.collect();
assert_eq!(corr_triggers.len(), 1);
}
#[tokio::test]
async fn test_no_triggers_healthy_system() {
let feedback = FeedbackLoop::new(FeedbackLoopConfig::default());
let metrics = MockMetrics::healthy();
let result = feedback.run_cycle(&metrics).await;
assert!(result.retraining_triggers.is_empty());
assert!(!result.kill_switch_activated);
}
}

View File

@@ -163,6 +163,15 @@ 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;
/// Autonomous feedback loop: weight/gate optimization with kill switch
pub mod feedback_loop;
/// QuestDB-backed metrics provider for the feedback loop
pub mod questdb_metrics;
// Re-export for tests
pub use ensemble_coordinator::EnsembleCoordinator;
pub use paper_trading_executor::PaperTradingExecutor;

View File

@@ -37,6 +37,8 @@ use trading_service::services::ml_performance_monitor::MLPerformanceMonitor;
use trading_service::services::monitoring::MonitoringServiceImpl;
use trading_service::services::risk::RiskServiceImpl;
use trading_service::services::trading::TradingServiceImpl;
use trading_service::feedback_loop::{FeedbackLoop, FeedbackLoopConfig};
use trading_service::questdb_metrics::QuestDBMetricsProvider;
use trading_service::state::TradingServiceState;
/// Default configuration values
@@ -528,6 +530,44 @@ async fn main() -> Result<()> {
warn!("⚠️ Ensemble coordinator not available - prediction generation loop disabled");
}
// Initialize autonomous feedback loop (weight + gate optimization with kill switch)
{
let feedback_config = FeedbackLoopConfig {
cycle_interval: Duration::from_secs(
std::env::var("FEEDBACK_LOOP_INTERVAL_SECS")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or(24 * 3600),
),
kill_switch_threshold: std::env::var("FEEDBACK_KILL_SWITCH_SHARPE")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or(-1.0),
..Default::default()
};
let feedback_loop = Arc::new(FeedbackLoop::new(feedback_config));
// Connect to QuestDB for real metrics (pg wire protocol on port 8812)
let questdb_url = std::env::var("QUESTDB_PG_URL")
.unwrap_or_else(|_| "postgresql://admin:quest@localhost:8812/qdb".to_string());
match QuestDBMetricsProvider::try_new(&questdb_url).await {
Some(provider) => {
if let Err(e) = provider.ensure_tables().await {
warn!("Failed to create QuestDB tables: {} — feedback loop will retry", e);
}
let metrics_provider: Arc<dyn trading_service::feedback_loop::MetricsProvider> =
Arc::new(provider);
let _feedback_handle = Arc::clone(&feedback_loop).spawn(metrics_provider);
info!("Autonomous feedback loop spawned with QuestDB metrics provider");
}
None => {
warn!("QuestDB unavailable — feedback loop disabled until QUESTDB_PG_URL is reachable");
}
}
}
// Subscribe to ML performance alerts
let monitor_clone = Arc::clone(&ml_performance_monitor);
tokio::spawn(async move {

View File

@@ -0,0 +1,422 @@
//! QuestDB-backed MetricsProvider for the feedback loop
//!
//! Queries QuestDB via PostgreSQL wire protocol (port 8812) to provide
//! rolling model metrics, confidence buckets, and ensemble Sharpe for
//! the autonomous weight/gate optimizers.
use crate::feedback_loop::MetricsProvider;
use ml::ensemble::conviction_gates::ConvictionGateConfig;
use ml::ensemble::gate_optimizer::GateBucketMetrics;
use ml::ensemble::weight_optimizer::ModelRollingMetrics;
use sqlx::postgres::PgPool;
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Instant;
use tokio::sync::RwLock;
use tracing::{debug, warn};
/// QuestDB metrics provider that queries real time-series data.
///
/// Tables expected in QuestDB (created on first write via ILP):
/// - `model_predictions`: model_id, signal, confidence, timestamp
/// - `trade_outcomes`: model_id, realized_pnl, signal_alignment, timestamp
/// - `ensemble_metrics`: sharpe_7d, sharpe_30d, timestamp
pub struct QuestDBMetricsProvider {
pool: PgPool,
/// Cached gate config (updated from external source)
gate_config: Arc<RwLock<ConvictionGateConfig>>,
/// Cached weights (updated by the feedback loop itself)
current_weights: Arc<RwLock<HashMap<String, f64>>>,
}
impl QuestDBMetricsProvider {
pub async fn new(questdb_pg_url: &str) -> Result<Self, sqlx::Error> {
let pool = PgPool::connect(questdb_pg_url).await?;
Ok(Self {
pool,
gate_config: Arc::new(RwLock::new(ConvictionGateConfig::default())),
current_weights: Arc::new(RwLock::new(HashMap::new())),
})
}
/// Try to connect, returning None if QuestDB is unavailable
pub async fn try_new(questdb_pg_url: &str) -> Option<Self> {
match Self::new(questdb_pg_url).await {
Ok(provider) => Some(provider),
Err(e) => {
warn!("QuestDB unavailable at {}: {} — feedback loop will use defaults", questdb_pg_url, e);
None
}
}
}
/// Update the cached gate config (called when config changes)
pub async fn set_gate_config(&self, config: ConvictionGateConfig) {
*self.gate_config.write().await = config;
}
/// Update the cached weights (called after weight adjustments)
pub async fn set_weights(&self, weights: HashMap<String, f64>) {
*self.current_weights.write().await = weights;
}
/// Create the required tables if they don't exist.
/// QuestDB auto-creates tables on ILP write, but we create them
/// explicitly for querying so tests can verify structure.
pub async fn ensure_tables(&self) -> Result<(), sqlx::Error> {
// QuestDB uses CREATE TABLE IF NOT EXISTS with designated timestamp
sqlx::query(
"CREATE TABLE IF NOT EXISTS model_predictions (
model_id SYMBOL,
signal DOUBLE,
confidence DOUBLE,
prediction_accuracy DOUBLE,
win_rate DOUBLE,
timestamp TIMESTAMP
) TIMESTAMP(timestamp) PARTITION BY DAY;"
)
.execute(&self.pool)
.await?;
sqlx::query(
"CREATE TABLE IF NOT EXISTS trade_outcomes (
model_id SYMBOL,
realized_pnl DOUBLE,
signal_alignment DOUBLE,
confidence_bucket_lower DOUBLE,
confidence_bucket_upper DOUBLE,
timestamp TIMESTAMP
) TIMESTAMP(timestamp) PARTITION BY DAY;"
)
.execute(&self.pool)
.await?;
sqlx::query(
"CREATE TABLE IF NOT EXISTS ensemble_metrics (
sharpe_7d DOUBLE,
sharpe_30d DOUBLE,
total_pnl DOUBLE,
timestamp TIMESTAMP
) TIMESTAMP(timestamp) PARTITION BY DAY;"
)
.execute(&self.pool)
.await?;
Ok(())
}
/// Query 30-day rolling Sharpe per model
async fn query_model_sharpe(&self) -> HashMap<String, f64> {
let result: Result<Vec<(String, f64)>, _> = sqlx::query_as(
"SELECT model_id, coalesce(avg(signal) / (stddev(signal) + 1e-10), 0.0) as sharpe
FROM model_predictions
WHERE timestamp > dateadd('d', -30, now())
GROUP BY model_id"
)
.fetch_all(&self.pool)
.await;
match result {
Ok(rows) => rows.into_iter().collect(),
Err(e) => {
debug!("QuestDB model Sharpe query failed: {}", e);
HashMap::new()
}
}
}
/// Query 30-day win rate per model
async fn query_model_win_rates(&self) -> HashMap<String, (f64, u64)> {
let result: Result<Vec<(String, f64, i64)>, _> = sqlx::query_as(
"SELECT model_id,
sum(CASE WHEN signal_alignment > 0 THEN 1.0 ELSE 0.0 END) / count(1) as win_rate,
count(*) as trade_count
FROM trade_outcomes
WHERE timestamp > dateadd('d', -30, now())
GROUP BY model_id"
)
.fetch_all(&self.pool)
.await;
match result {
Ok(rows) => rows
.into_iter()
.map(|(id, wr, cnt)| (id, (wr, cnt as u64)))
.collect(),
Err(e) => {
debug!("QuestDB win rate query failed: {}", e);
HashMap::new()
}
}
}
/// Query confidence-bucketed win rates for gate optimization
async fn query_confidence_buckets(&self) -> Vec<GateBucketMetrics> {
// 5 buckets: [0.5-0.6), [0.6-0.7), [0.7-0.8), [0.8-0.9), [0.9-1.0]
let result: Result<Vec<(f64, f64, f64, i64, f64)>, _> = sqlx::query_as(
"SELECT confidence_bucket_lower,
confidence_bucket_upper,
sum(CASE WHEN signal_alignment > 0 THEN 1.0 ELSE 0.0 END) / count(1) as win_rate,
count(*) as trade_count,
avg(realized_pnl) as avg_pnl
FROM trade_outcomes
WHERE timestamp > dateadd('d', -30, now())
AND confidence_bucket_lower IS NOT NULL
GROUP BY confidence_bucket_lower, confidence_bucket_upper
ORDER BY confidence_bucket_lower"
)
.fetch_all(&self.pool)
.await;
match result {
Ok(rows) => rows
.into_iter()
.map(|(lower, upper, win_rate, count, avg_pnl)| GateBucketMetrics {
confidence_lower: lower,
confidence_upper: upper,
win_rate,
trade_count: count as u64,
avg_pnl,
})
.collect(),
Err(e) => {
debug!("QuestDB confidence bucket query failed: {}", e);
Vec::new()
}
}
}
/// Query latest ensemble Sharpe (7-day)
async fn query_ensemble_sharpe(&self) -> f64 {
let result: Result<Option<(f64,)>, _> = sqlx::query_as(
"SELECT sharpe_7d FROM ensemble_metrics ORDER BY timestamp DESC LIMIT 1"
)
.fetch_optional(&self.pool)
.await;
match result {
Ok(Some((sharpe,))) => sharpe,
Ok(None) => 0.0,
Err(e) => {
debug!("QuestDB ensemble Sharpe query failed: {}", e);
0.0
}
}
}
}
impl MetricsProvider for QuestDBMetricsProvider {
fn get_model_metrics(&self) -> Vec<ModelRollingMetrics> {
// block_in_place allows blocking inside a multi-threaded tokio runtime
// by moving other tasks off this thread. Safe because the feedback loop
// runs periodically (24h cycle) — not on the hot path.
let handle = tokio::runtime::Handle::current();
tokio::task::block_in_place(|| {
let sharpes = handle.block_on(self.query_model_sharpe());
let win_rates = handle.block_on(self.query_model_win_rates());
let mut metrics = Vec::new();
for (model_id, sharpe) in &sharpes {
let (win_rate, trade_count) = win_rates
.get(model_id)
.copied()
.unwrap_or((0.5, 0));
metrics.push(ModelRollingMetrics {
model_id: model_id.clone(),
sharpe_30d: *sharpe,
win_rate_30d: win_rate,
prediction_accuracy: win_rate, // Using win_rate as proxy
trade_count,
deployed_at: Instant::now() - std::time::Duration::from_secs(30 * 24 * 3600),
});
}
metrics
})
}
fn get_gate_buckets(&self) -> Vec<GateBucketMetrics> {
let handle = tokio::runtime::Handle::current();
tokio::task::block_in_place(|| handle.block_on(self.query_confidence_buckets()))
}
fn get_current_weights(&self) -> HashMap<String, f64> {
let handle = tokio::runtime::Handle::current();
tokio::task::block_in_place(|| {
handle.block_on(async { self.current_weights.read().await.clone() })
})
}
fn get_gate_config(&self) -> ConvictionGateConfig {
let handle = tokio::runtime::Handle::current();
tokio::task::block_in_place(|| {
handle.block_on(async { self.gate_config.read().await.clone() })
})
}
fn get_ensemble_sharpe_7d(&self) -> f64 {
let handle = tokio::runtime::Handle::current();
tokio::task::block_in_place(|| handle.block_on(self.query_ensemble_sharpe()))
}
fn get_model_correlations(&self) -> Vec<((String, String), f64)> {
// Correlation requires cross-model signal comparison
// For now, query is deferred until we have enough data
Vec::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
/// Drop and recreate all test tables for clean isolation between runs.
/// QuestDB TRUNCATE via PG wire can be flaky; DROP + CREATE is reliable.
async fn reset_tables(pool: &PgPool) {
for table in &["model_predictions", "trade_outcomes", "ensemble_metrics"] {
let drop_q = format!("DROP TABLE IF EXISTS {table};");
let _ = sqlx::query(&drop_q).execute(pool).await;
}
}
/// Poll QuestDB until a table has at least `min_rows` rows (up to 5s).
async fn wait_for_rows(pool: &PgPool, table: &str, min_rows: i64) {
for _ in 0..50 {
let query = format!("SELECT count(*) FROM {table}");
let row: Option<(i64,)> = sqlx::query_as(&query)
.fetch_optional(pool)
.await
.ok()
.flatten();
if row.map(|(c,)| c).unwrap_or(0) >= min_rows {
return;
}
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
}
}
/// Comprehensive integration test: connect, write, read, verify metrics.
/// Single test avoids parallel interference on shared QuestDB tables.
///
/// Run with: docker-compose up -d questdb
/// Then: SQLX_OFFLINE=true cargo test -p trading_service --lib -- questdb_metrics --include-ignored
#[tokio::test(flavor = "multi_thread")]
#[ignore = "Requires QuestDB on localhost:8812"]
async fn test_questdb_provider_full_lifecycle() {
let provider = QuestDBMetricsProvider::new("postgresql://admin:quest@localhost:8812/qdb")
.await
.expect("QuestDB should be running");
// Drop + recreate for clean slate (prior test runs may have left data)
reset_tables(&provider.pool).await;
provider.ensure_tables().await.expect("Tables should be created");
// Phase 1: Empty tables — metrics should return defaults
let metrics = provider.get_model_metrics();
assert!(metrics.is_empty(), "Empty tables should yield no metrics");
let sharpe = provider.get_ensemble_sharpe_7d();
assert!((sharpe - 0.0).abs() < 1e-10, "Empty ensemble Sharpe should be 0.0");
// Phase 2: Write test data (multiple rows per model for valid stddev)
sqlx::query(
"INSERT INTO model_predictions(model_id, signal, confidence, prediction_accuracy, win_rate, timestamp)
VALUES ('dqn', 0.8, 0.7, 0.62, 0.58, systimestamp()),
('dqn', 0.6, 0.8, 0.65, 0.60, systimestamp()),
('ppo', 0.3, 0.6, 0.55, 0.52, systimestamp()),
('ppo', 0.5, 0.7, 0.58, 0.54, systimestamp())"
)
.execute(&provider.pool)
.await
.expect("Insert predictions should succeed");
sqlx::query(
"INSERT INTO trade_outcomes(model_id, realized_pnl, signal_alignment, confidence_bucket_lower, confidence_bucket_upper, timestamp)
VALUES ('dqn', 100.0, 1.0, 0.60, 0.70, systimestamp()),
('dqn', -50.0, -1.0, 0.60, 0.70, systimestamp()),
('ppo', 75.0, 1.0, 0.50, 0.60, systimestamp())"
)
.execute(&provider.pool)
.await
.expect("Insert outcomes should succeed");
sqlx::query(
"INSERT INTO ensemble_metrics(sharpe_7d, sharpe_30d, total_pnl, timestamp)
VALUES (1.5, 1.2, 5000.0, systimestamp())"
)
.execute(&provider.pool)
.await
.expect("Insert ensemble metrics should succeed");
// Wait for QuestDB WAL to commit all tables
wait_for_rows(&provider.pool, "model_predictions", 4).await;
wait_for_rows(&provider.pool, "trade_outcomes", 3).await;
wait_for_rows(&provider.pool, "ensemble_metrics", 1).await;
// Phase 3: Read back and verify
let metrics = provider.get_model_metrics();
assert!(!metrics.is_empty(), "Should have model metrics after insert");
assert!(metrics.len() >= 2, "Should have at least DQN and PPO");
let sharpe = provider.get_ensemble_sharpe_7d();
assert!((sharpe - 1.5).abs() < 0.1, "Sharpe should be ~1.5, got {sharpe}");
let buckets = provider.get_gate_buckets();
assert!(!buckets.is_empty(), "Should have confidence buckets");
}
#[tokio::test(flavor = "multi_thread")]
#[ignore = "Requires QuestDB on localhost:8812"]
async fn test_questdb_provider_try_new_succeeds() {
let provider = QuestDBMetricsProvider::try_new("postgresql://admin:quest@localhost:8812/qdb").await;
assert!(provider.is_some(), "Should connect to QuestDB");
}
#[tokio::test(flavor = "multi_thread")]
async fn test_questdb_provider_try_new_fails_gracefully() {
// Connect to a non-existent QuestDB — should return None, not panic
let provider = QuestDBMetricsProvider::try_new("postgresql://admin:quest@localhost:19999/qdb").await;
assert!(provider.is_none(), "Should fail gracefully");
}
#[tokio::test(flavor = "multi_thread")]
#[ignore = "Requires QuestDB on localhost:8812"]
async fn test_feedback_loop_with_questdb() {
use crate::feedback_loop::{FeedbackLoop, FeedbackLoopConfig};
use ml::ensemble::gate_optimizer::GateOptimizerConfig;
use ml::ensemble::weight_optimizer::WeightOptimizerConfig;
use std::time::Duration;
let provider = QuestDBMetricsProvider::new("postgresql://admin:quest@localhost:8812/qdb")
.await
.expect("QuestDB should be running");
provider.ensure_tables().await.expect("Tables should be created");
// Set up some weights
let mut weights = HashMap::new();
weights.insert("dqn".to_string(), 0.5);
weights.insert("ppo".to_string(), 0.5);
provider.set_weights(weights).await;
let feedback = FeedbackLoop::with_optimizers(
FeedbackLoopConfig {
cycle_interval: Duration::from_millis(100),
..Default::default()
},
WeightOptimizerConfig {
cooldown: Duration::ZERO,
..Default::default()
},
GateOptimizerConfig {
cooldown: Duration::ZERO,
..Default::default()
},
);
// Run a cycle — with empty QuestDB should report InsufficientData
let result = feedback.run_cycle(&provider).await;
assert!(!result.kill_switch_activated);
}
}