use crate::core::error::{Error, Result};
use crate::ml::serving::deployment::RecordedError;
use crate::ml::serving::{DeploymentMetrics, ModelMetadata};
use serde::{Deserialize, Serialize};
use std::collections::{HashMap, VecDeque};
use std::time::{Duration, Instant};
const MAX_HTTP_REQUEST_SAMPLES: usize = 2_000;
const MAX_CONFIDENCE_SAMPLES: usize = 2_000;
const PSI_DRIFT_THRESHOLD: f64 = 0.2;
const PSI_EPSILON: f64 = 1e-4;
const DEFAULT_PSI_BINS: usize = 10;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PerformanceMetrics {
pub model_name: String,
pub model_version: String,
pub timestamp: chrono::DateTime<chrono::Utc>,
pub latency: LatencyMetrics,
pub throughput: ThroughputMetrics,
pub error_metrics: ErrorMetrics,
pub resource_utilization: ResourceUtilizationMetrics,
pub quality_metrics: QualityMetrics,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct LatencyMetrics {
pub avg_latency_ms: f64,
pub p50_latency_ms: f64,
pub p95_latency_ms: f64,
pub p99_latency_ms: f64,
pub max_latency_ms: f64,
pub min_latency_ms: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ThroughputMetrics {
pub requests_per_second: f64,
pub total_requests: u64,
pub successful_requests: u64,
pub failed_requests: u64,
pub concurrent_requests: u64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ErrorMetrics {
pub error_rate: f64,
pub error_rates_by_type: HashMap<String, f64>,
pub error_counts_by_type: HashMap<String, u64>,
pub recent_errors: Vec<ErrorEvent>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ErrorEvent {
pub error_type: String,
pub message: String,
pub timestamp: chrono::DateTime<chrono::Utc>,
pub context: Option<HashMap<String, String>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ResourceUtilizationMetrics {
pub cpu_utilization: Option<f64>,
pub memory_utilization: Option<f64>,
pub gpu_utilization: Option<f64>,
pub disk_io_utilization: Option<f64>,
pub network_io_utilization: Option<f64>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct QualityMetrics {
pub accuracy: Option<f64>,
pub confidence_scores: Option<ConfidenceMetrics>,
pub data_drift: Option<DriftMetrics>,
pub model_drift: Option<DriftMetrics>,
pub feature_importance_drift: Option<f64>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ConfidenceMetrics {
pub avg_confidence: f64,
pub min_confidence: f64,
pub max_confidence: f64,
pub low_confidence_rate: f64,
pub confidence_threshold: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DriftMetrics {
pub drift_score: f64,
pub drift_detected: bool,
pub detection_method: String,
pub threshold: f64,
pub drifting_features: Vec<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub enum AlertSeverity {
Info,
Warning,
Critical,
Emergency,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AlertConfig {
pub name: String,
pub description: String,
pub metric: String,
pub threshold: f64,
pub operator: ComparisonOperator,
pub severity: AlertSeverity,
pub evaluation_window_seconds: u64,
pub consecutive_evaluations: usize,
pub cooldown_seconds: u64,
pub enabled: bool,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum ComparisonOperator {
GreaterThan,
GreaterThanOrEqual,
LessThan,
LessThanOrEqual,
Equal,
NotEqual,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AlertEvent {
pub alert_config: AlertConfig,
pub current_value: f64,
pub threshold_value: f64,
pub message: String,
pub triggered_at: chrono::DateTime<chrono::Utc>,
pub model_name: String,
pub model_version: String,
pub context: HashMap<String, String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FeatureBaseline {
edges: Vec<f64>,
reference_frequencies: Vec<f64>,
}
impl FeatureBaseline {
pub fn from_samples(values: &[f64], n_bins: usize) -> Option<Self> {
if n_bins == 0 {
return None;
}
let mut sorted: Vec<f64> = values.iter().copied().filter(|v| v.is_finite()).collect();
if sorted.len() < n_bins {
return None;
}
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let mut edges = Vec::with_capacity(n_bins + 1);
edges.push(f64::NEG_INFINITY);
for i in 1..n_bins {
let pos = (i as f64 / n_bins as f64) * (sorted.len() - 1) as f64;
let lower = pos.floor() as usize;
let upper = pos.ceil() as usize;
let frac = pos - lower as f64;
edges.push(sorted[lower] * (1.0 - frac) + sorted[upper.min(sorted.len() - 1)] * frac);
}
edges.push(f64::INFINITY);
Some(Self {
edges,
reference_frequencies: vec![1.0 / n_bins as f64; n_bins],
})
}
fn bin_index(&self, value: f64) -> Option<usize> {
if !value.is_finite() {
return None;
}
for i in 0..self.edges.len().saturating_sub(1) {
if value >= self.edges[i] && value < self.edges[i + 1] {
return Some(i);
}
}
Some(self.reference_frequencies.len().saturating_sub(1))
}
fn n_bins(&self) -> usize {
self.reference_frequencies.len()
}
}
fn latency_metrics_from_samples(samples: &[u64]) -> LatencyMetrics {
let mut sorted = samples.to_vec();
sorted.sort_unstable();
let n = sorted.len();
let percentile = |p: f64| -> f64 {
let rank = ((p * n as f64).ceil() as usize).clamp(1, n) - 1;
sorted[rank] as f64
};
let sum: u64 = sorted.iter().sum();
LatencyMetrics {
avg_latency_ms: sum as f64 / n as f64,
p50_latency_ms: percentile(0.50),
p95_latency_ms: percentile(0.95),
p99_latency_ms: percentile(0.99),
max_latency_ms: sorted[n - 1] as f64,
min_latency_ms: sorted[0] as f64,
}
}
fn recorded_error_to_event(recorded: &RecordedError) -> ErrorEvent {
ErrorEvent {
error_type: recorded.error_type.clone(),
message: recorded.message.clone(),
timestamp: recorded.occurred_at,
context: None,
}
}
pub struct ModelMonitor {
model_metadata: ModelMetadata,
metrics_history: VecDeque<PerformanceMetrics>,
alert_configs: Vec<AlertConfig>,
alert_events: VecDeque<AlertEvent>,
alert_counters: HashMap<String, usize>,
last_alert_times: HashMap<String, Instant>,
max_history_size: usize,
collection_interval: Duration,
last_collection: Option<Instant>,
feature_baselines: HashMap<String, FeatureBaseline>,
live_feature_counts: HashMap<String, Vec<u64>>,
request_confidences: VecDeque<f64>,
http_request_latencies_ms: VecDeque<u64>,
http_success_count: u64,
http_error_count: u64,
}
impl ModelMonitor {
pub fn new(model_metadata: ModelMetadata) -> Self {
Self {
model_metadata,
metrics_history: VecDeque::new(),
alert_configs: Vec::new(),
alert_events: VecDeque::new(),
alert_counters: HashMap::new(),
last_alert_times: HashMap::new(),
max_history_size: 1440, collection_interval: Duration::from_secs(60), last_collection: None,
feature_baselines: HashMap::new(),
live_feature_counts: HashMap::new(),
request_confidences: VecDeque::new(),
http_request_latencies_ms: VecDeque::new(),
http_success_count: 0,
http_error_count: 0,
}
}
pub fn set_feature_baseline(&mut self, feature_name: &str, reference_values: &[f64]) {
self.set_feature_baseline_with_bins(feature_name, reference_values, DEFAULT_PSI_BINS);
}
pub fn set_feature_baseline_with_bins(
&mut self,
feature_name: &str,
reference_values: &[f64],
n_bins: usize,
) {
if let Some(baseline) = FeatureBaseline::from_samples(reference_values, n_bins) {
let bins = baseline.n_bins();
self.feature_baselines
.insert(feature_name.to_string(), baseline);
self.live_feature_counts
.insert(feature_name.to_string(), vec![0; bins]);
}
}
pub fn observe_feature(&mut self, feature_name: &str, value: f64) {
if let Some(baseline) = self.feature_baselines.get(feature_name) {
if let Some(bin) = baseline.bin_index(value) {
if let Some(counts) = self.live_feature_counts.get_mut(feature_name) {
if let Some(count) = counts.get_mut(bin) {
*count += 1;
}
}
}
}
}
pub fn observe_features(&mut self, data: &HashMap<String, serde_json::Value>) {
let names: Vec<String> = self.feature_baselines.keys().cloned().collect();
for name in names {
if let Some(value) = data.get(&name).and_then(|v| v.as_f64()) {
self.observe_feature(&name, value);
}
}
}
pub fn record_prediction_confidence(&mut self, confidence: f64) {
if confidence.is_finite() {
self.request_confidences.push_back(confidence);
while self.request_confidences.len() > MAX_CONFIDENCE_SAMPLES {
self.request_confidences.pop_front();
}
}
}
pub fn calculate_data_drift(&self) -> Option<DriftMetrics> {
let mut per_feature_psi: Vec<(String, f64)> = Vec::new();
for (name, baseline) in &self.feature_baselines {
let counts = match self.live_feature_counts.get(name) {
Some(c) => c,
None => continue,
};
let total: u64 = counts.iter().sum();
if total == 0 {
continue;
}
let mut psi = 0.0;
for (bin_idx, &ref_pct) in baseline.reference_frequencies.iter().enumerate() {
let live_pct = counts.get(bin_idx).copied().unwrap_or(0) as f64 / total as f64;
let ref_pct = ref_pct.max(PSI_EPSILON);
let live_pct = live_pct.max(PSI_EPSILON);
psi += (live_pct - ref_pct) * (live_pct / ref_pct).ln();
}
per_feature_psi.push((name.clone(), psi));
}
if per_feature_psi.is_empty() {
return None;
}
let drift_score = per_feature_psi
.iter()
.map(|(_, psi)| *psi)
.fold(f64::NEG_INFINITY, f64::max);
let drifting_features: Vec<String> = per_feature_psi
.into_iter()
.filter(|(_, psi)| *psi > PSI_DRIFT_THRESHOLD)
.map(|(name, _)| name)
.collect();
Some(DriftMetrics {
drift_score,
drift_detected: drift_score > PSI_DRIFT_THRESHOLD,
detection_method: "PSI".to_string(),
threshold: PSI_DRIFT_THRESHOLD,
drifting_features,
})
}
fn calculate_confidence_metrics(&self) -> Option<ConfidenceMetrics> {
if self.request_confidences.is_empty() {
return None;
}
const LOW_CONFIDENCE_THRESHOLD: f64 = 0.7;
let n = self.request_confidences.len() as f64;
let sum: f64 = self.request_confidences.iter().sum();
let min = self
.request_confidences
.iter()
.cloned()
.fold(f64::INFINITY, f64::min);
let max = self
.request_confidences
.iter()
.cloned()
.fold(f64::NEG_INFINITY, f64::max);
let low_count = self
.request_confidences
.iter()
.filter(|&&c| c < LOW_CONFIDENCE_THRESHOLD)
.count();
Some(ConfidenceMetrics {
avg_confidence: sum / n,
min_confidence: min,
max_confidence: max,
low_confidence_rate: low_count as f64 / n,
confidence_threshold: LOW_CONFIDENCE_THRESHOLD,
})
}
pub fn record_request(&mut self, success: bool, latency_ms: u64) -> Result<bool> {
self.http_request_latencies_ms.push_back(latency_ms);
while self.http_request_latencies_ms.len() > MAX_HTTP_REQUEST_SAMPLES {
self.http_request_latencies_ms.pop_front();
}
if success {
self.http_success_count += 1;
} else {
self.http_error_count += 1;
}
if let Some(last) = self.last_collection {
if last.elapsed() < self.collection_interval {
return Ok(false);
}
}
let samples: Vec<u64> = self.http_request_latencies_ms.iter().copied().collect();
let latency = latency_metrics_from_samples(&samples);
let total = self.http_success_count + self.http_error_count;
let error_rate = if total > 0 {
self.http_error_count as f64 / total as f64
} else {
0.0
};
let snapshot = PerformanceMetrics {
model_name: self.model_metadata.name.clone(),
model_version: self.model_metadata.version.clone(),
timestamp: chrono::Utc::now(),
latency,
throughput: ThroughputMetrics {
requests_per_second: 0.0,
total_requests: total,
successful_requests: self.http_success_count,
failed_requests: self.http_error_count,
concurrent_requests: 0,
},
error_metrics: ErrorMetrics {
error_rate,
error_rates_by_type: HashMap::new(),
error_counts_by_type: HashMap::new(),
recent_errors: Vec::new(),
},
resource_utilization: ResourceUtilizationMetrics {
cpu_utilization: None,
memory_utilization: None,
gpu_utilization: None,
disk_io_utilization: None,
network_io_utilization: None,
},
quality_metrics: self.calculate_quality_metrics(),
};
self.metrics_history.push_back(snapshot.clone());
while self.metrics_history.len() > self.max_history_size {
self.metrics_history.pop_front();
}
self.evaluate_alerts(&snapshot)?;
self.last_collection = Some(Instant::now());
Ok(true)
}
pub fn add_alert(&mut self, config: AlertConfig) {
self.alert_configs.push(config);
}
pub fn remove_alert(&mut self, alert_name: &str) {
self.alert_configs
.retain(|config| config.name != alert_name);
self.alert_counters.remove(alert_name);
self.last_alert_times.remove(alert_name);
}
pub fn collect_metrics(&mut self, deployment_metrics: &DeploymentMetrics) -> Result<bool> {
if let Some(last) = self.last_collection {
if last.elapsed() < self.collection_interval {
return Ok(false);
}
}
let performance_metrics = PerformanceMetrics {
model_name: self.model_metadata.name.clone(),
model_version: self.model_metadata.version.clone(),
timestamp: chrono::Utc::now(),
latency: self.calculate_latency_metrics(deployment_metrics),
throughput: self.calculate_throughput_metrics(deployment_metrics),
error_metrics: self.calculate_error_metrics(deployment_metrics),
resource_utilization: self.calculate_resource_metrics(deployment_metrics),
quality_metrics: self.calculate_quality_metrics(),
};
self.metrics_history.push_back(performance_metrics.clone());
while self.metrics_history.len() > self.max_history_size {
self.metrics_history.pop_front();
}
self.evaluate_alerts(&performance_metrics)?;
self.last_collection = Some(Instant::now());
Ok(true)
}
fn calculate_latency_metrics(&self, deployment_metrics: &DeploymentMetrics) -> LatencyMetrics {
if deployment_metrics.response_times_ms.is_empty() {
let avg_latency = deployment_metrics.avg_response_time_ms;
return LatencyMetrics {
avg_latency_ms: avg_latency,
p50_latency_ms: avg_latency,
p95_latency_ms: avg_latency,
p99_latency_ms: avg_latency,
max_latency_ms: avg_latency,
min_latency_ms: avg_latency,
};
}
latency_metrics_from_samples(&deployment_metrics.response_times_ms)
}
fn calculate_throughput_metrics(
&self,
deployment_metrics: &DeploymentMetrics,
) -> ThroughputMetrics {
ThroughputMetrics {
requests_per_second: deployment_metrics.request_rate,
total_requests: deployment_metrics.total_requests,
successful_requests: deployment_metrics.successful_requests,
failed_requests: deployment_metrics.failed_requests,
concurrent_requests: deployment_metrics.active_instances as u64,
}
}
fn calculate_error_metrics(&self, deployment_metrics: &DeploymentMetrics) -> ErrorMetrics {
let mut error_counts_by_type: HashMap<String, u64> = HashMap::new();
for recorded in &deployment_metrics.recent_errors {
*error_counts_by_type
.entry(recorded.error_type.clone())
.or_insert(0) += 1;
}
let total_recorded: u64 = error_counts_by_type.values().sum();
let error_rates_by_type: HashMap<String, f64> = if total_recorded > 0 {
error_counts_by_type
.iter()
.map(|(k, &count)| {
(
k.clone(),
(count as f64 / total_recorded as f64) * deployment_metrics.error_rate,
)
})
.collect()
} else {
HashMap::new()
};
let recent_errors = deployment_metrics
.recent_errors
.iter()
.map(|recorded| recorded_error_to_event(recorded))
.collect();
ErrorMetrics {
error_rate: deployment_metrics.error_rate,
error_rates_by_type,
error_counts_by_type,
recent_errors,
}
}
fn calculate_resource_metrics(
&self,
deployment_metrics: &DeploymentMetrics,
) -> ResourceUtilizationMetrics {
ResourceUtilizationMetrics {
cpu_utilization: deployment_metrics.cpu_utilization,
memory_utilization: deployment_metrics.memory_utilization,
gpu_utilization: None, disk_io_utilization: None,
network_io_utilization: None,
}
}
fn calculate_quality_metrics(&self) -> QualityMetrics {
QualityMetrics {
accuracy: None, confidence_scores: self.calculate_confidence_metrics(),
data_drift: self.calculate_data_drift(),
model_drift: None, feature_importance_drift: None, }
}
fn evaluate_alerts(&mut self, metrics: &PerformanceMetrics) -> Result<()> {
let alert_configs = self.alert_configs.clone();
for config in &alert_configs {
if !config.enabled {
continue;
}
if let Some(last_alert_time) = self.last_alert_times.get(&config.name) {
if last_alert_time.elapsed() < Duration::from_secs(config.cooldown_seconds) {
continue;
}
}
let current_value = match self.get_metric_value(metrics, &config.metric)? {
Some(value) => value,
None => continue,
};
let threshold_exceeded = match config.operator {
ComparisonOperator::GreaterThan => current_value > config.threshold,
ComparisonOperator::GreaterThanOrEqual => current_value >= config.threshold,
ComparisonOperator::LessThan => current_value < config.threshold,
ComparisonOperator::LessThanOrEqual => current_value <= config.threshold,
ComparisonOperator::Equal => (current_value - config.threshold).abs() < 1e-10,
ComparisonOperator::NotEqual => (current_value - config.threshold).abs() >= 1e-10,
};
if threshold_exceeded {
let should_trigger = {
let counter = self.alert_counters.entry(config.name.clone()).or_insert(0);
*counter += 1;
*counter >= config.consecutive_evaluations
};
if should_trigger {
self.trigger_alert(config, current_value)?;
self.alert_counters.insert(config.name.clone(), 0); self.last_alert_times
.insert(config.name.clone(), Instant::now());
}
} else {
self.alert_counters.insert(config.name.clone(), 0);
}
}
Ok(())
}
fn get_metric_value(
&self,
metrics: &PerformanceMetrics,
metric_name: &str,
) -> Result<Option<f64>> {
match metric_name {
"avg_latency_ms" => Ok(Some(metrics.latency.avg_latency_ms)),
"p95_latency_ms" => Ok(Some(metrics.latency.p95_latency_ms)),
"p99_latency_ms" => Ok(Some(metrics.latency.p99_latency_ms)),
"requests_per_second" => Ok(Some(metrics.throughput.requests_per_second)),
"error_rate" => Ok(Some(metrics.error_metrics.error_rate)),
"cpu_utilization" => Ok(metrics.resource_utilization.cpu_utilization),
"memory_utilization" => Ok(metrics.resource_utilization.memory_utilization),
"drift_score" => Ok(metrics
.quality_metrics
.data_drift
.as_ref()
.map(|d| d.drift_score)),
"avg_confidence" => Ok(metrics
.quality_metrics
.confidence_scores
.as_ref()
.map(|c| c.avg_confidence)),
_ => Err(Error::InvalidInput(format!(
"Unknown metric: {}",
metric_name
))),
}
}
fn trigger_alert(&mut self, config: &AlertConfig, current_value: f64) -> Result<()> {
let alert_event = AlertEvent {
alert_config: config.clone(),
current_value,
threshold_value: config.threshold,
message: format!(
"Alert '{}': {} {} {} (current: {:.4})",
config.name,
config.metric,
self.operator_to_string(&config.operator),
config.threshold,
current_value
),
triggered_at: chrono::Utc::now(),
model_name: self.model_metadata.name.clone(),
model_version: self.model_metadata.version.clone(),
context: HashMap::new(),
};
self.alert_events.push_back(alert_event.clone());
while self.alert_events.len() > 100 {
self.alert_events.pop_front();
}
match config.severity {
AlertSeverity::Info => log::info!("Alert triggered: {}", alert_event.message),
AlertSeverity::Warning => log::warn!("Alert triggered: {}", alert_event.message),
AlertSeverity::Critical | AlertSeverity::Emergency => {
log::error!("Alert triggered: {}", alert_event.message)
}
}
Ok(())
}
fn operator_to_string(&self, operator: &ComparisonOperator) -> &'static str {
match operator {
ComparisonOperator::GreaterThan => ">",
ComparisonOperator::GreaterThanOrEqual => ">=",
ComparisonOperator::LessThan => "<",
ComparisonOperator::LessThanOrEqual => "<=",
ComparisonOperator::Equal => "==",
ComparisonOperator::NotEqual => "!=",
}
}
pub fn get_recent_metrics(&self, limit: usize) -> Vec<PerformanceMetrics> {
self.metrics_history
.iter()
.rev()
.take(limit)
.cloned()
.collect()
}
pub fn get_recent_alerts(&self, limit: usize) -> Vec<AlertEvent> {
self.alert_events
.iter()
.rev()
.take(limit)
.cloned()
.collect()
}
pub fn get_alert_configs(&self) -> &[AlertConfig] {
&self.alert_configs
}
pub fn get_metrics_summary(&self, window_minutes: usize) -> Option<MetricsSummary> {
let cutoff = chrono::Utc::now() - chrono::Duration::minutes(window_minutes as i64);
let recent_metrics: Vec<_> = self
.metrics_history
.iter()
.filter(|m| m.timestamp > cutoff)
.collect();
if recent_metrics.is_empty() {
return None;
}
let avg_latency = recent_metrics
.iter()
.map(|m| m.latency.avg_latency_ms)
.sum::<f64>()
/ recent_metrics.len() as f64;
let avg_throughput = recent_metrics
.iter()
.map(|m| m.throughput.requests_per_second)
.sum::<f64>()
/ recent_metrics.len() as f64;
let avg_error_rate = recent_metrics
.iter()
.map(|m| m.error_metrics.error_rate)
.sum::<f64>()
/ recent_metrics.len() as f64;
let cpu_samples: Vec<f64> = recent_metrics
.iter()
.filter_map(|m| m.resource_utilization.cpu_utilization)
.collect();
let avg_cpu_utilization = if cpu_samples.is_empty() {
None
} else {
Some(cpu_samples.iter().sum::<f64>() / cpu_samples.len() as f64)
};
let total_requests = recent_metrics
.iter()
.map(|m| m.throughput.total_requests)
.max()
.unwrap_or(0);
Some(MetricsSummary {
window_minutes,
avg_latency_ms: avg_latency,
avg_throughput,
avg_error_rate,
avg_cpu_utilization,
total_requests,
alert_count: self
.alert_events
.iter()
.filter(|e| e.triggered_at > cutoff)
.count(),
})
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MetricsSummary {
pub window_minutes: usize,
pub avg_latency_ms: f64,
pub avg_throughput: f64,
pub avg_error_rate: f64,
pub avg_cpu_utilization: Option<f64>,
pub total_requests: u64,
pub alert_count: usize,
}
pub trait MetricsCollector: Send + Sync {
fn collect_system_metrics(&self) -> Result<SystemMetrics>;
fn collect_model_metrics(&self, model_name: &str) -> Result<ModelSpecificMetrics>;
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SystemMetrics {
pub cpu_usage: f64,
pub memory_usage: u64,
pub memory_available: u64,
pub disk_usage: f64,
pub network_bytes_sent: u64,
pub network_bytes_received: u64,
pub load_average: f64,
pub process_count: u32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelSpecificMetrics {
pub model_memory_usage: u64,
pub model_init_time_ms: u64,
pub cache_hit_rate: f64,
pub feature_processing_time_ms: u64,
pub prediction_time_ms: u64,
}
pub struct DefaultMetricsCollector;
impl MetricsCollector for DefaultMetricsCollector {
fn collect_system_metrics(&self) -> Result<SystemMetrics> {
Err(Error::NotImplemented(
"System metrics require OS-level telemetry (e.g. a sysinfo-like crate), which this \
build does not depend on; DefaultMetricsCollector cannot report real cpu/memory/\
disk/network/process data"
.to_string(),
))
}
fn collect_model_metrics(&self, _model_name: &str) -> Result<ModelSpecificMetrics> {
Err(Error::NotImplemented(
"Model-specific runtime metrics (memory usage, init time, cache hit rate) are not \
tracked anywhere in this build; DefaultMetricsCollector cannot report them"
.to_string(),
))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ml::serving::ModelMetadata;
fn create_test_metadata() -> ModelMetadata {
ModelMetadata {
name: "test_model".to_string(),
version: "1.0.0".to_string(),
model_type: "classification".to_string(),
feature_names: vec!["feature1".to_string(), "feature2".to_string()],
target_name: Some("target".to_string()),
description: "Test model".to_string(),
created_at: chrono::Utc::now(),
updated_at: chrono::Utc::now(),
metrics: HashMap::new(),
metadata: HashMap::new(),
}
}
fn create_test_deployment_metrics() -> DeploymentMetrics {
use crate::ml::serving::deployment::DeploymentStatus;
crate::ml::serving::deployment::DeploymentMetrics {
status: DeploymentStatus::Running,
active_instances: 2,
cpu_utilization: None,
memory_utilization: None,
request_rate: 50.0,
avg_response_time_ms: 120.0,
error_rate: 0.02,
total_requests: 1000,
successful_requests: 980,
failed_requests: 20,
in_flight_requests: 0,
response_times_ms: vec![100, 110, 120, 130, 140, 500],
recent_errors: vec![RecordedError {
error_type: "InvalidInput".to_string(),
message: "missing feature 'x'".to_string(),
occurred_at: chrono::Utc::now(),
}],
last_health_check: chrono::Utc::now(),
started_at: chrono::Utc::now(),
updated_at: chrono::Utc::now(),
}
}
#[test]
fn test_model_monitor_creation() {
let metadata = create_test_metadata();
let monitor = ModelMonitor::new(metadata);
assert_eq!(monitor.model_metadata.name, "test_model");
assert_eq!(monitor.alert_configs.len(), 0);
assert_eq!(monitor.metrics_history.len(), 0);
}
#[test]
fn test_alert_config() {
let config = AlertConfig {
name: "high_latency".to_string(),
description: "Alert when latency is too high".to_string(),
metric: "avg_latency_ms".to_string(),
threshold: 200.0,
operator: ComparisonOperator::GreaterThan,
severity: AlertSeverity::Warning,
evaluation_window_seconds: 300,
consecutive_evaluations: 3,
cooldown_seconds: 600,
enabled: true,
};
assert_eq!(config.name, "high_latency");
assert_eq!(config.threshold, 200.0);
assert_eq!(config.severity, AlertSeverity::Warning);
}
#[test]
fn test_metrics_collector_is_honestly_not_implemented() {
let collector = DefaultMetricsCollector;
assert!(collector.collect_system_metrics().is_err());
assert!(collector.collect_model_metrics("test_model").is_err());
}
#[test]
fn test_performance_metrics() {
let metadata = create_test_metadata();
let mut monitor = ModelMonitor::new(metadata);
let deployment_metrics = create_test_deployment_metrics();
monitor.collection_interval = Duration::from_secs(0);
let sample_taken = monitor
.collect_metrics(&deployment_metrics)
.expect("operation should succeed");
assert!(
sample_taken,
"collect_metrics must report a sample was taken"
);
assert_eq!(monitor.metrics_history.len(), 1);
let metrics = &monitor.metrics_history[0];
assert_eq!(metrics.model_name, "test_model");
assert!(metrics.latency.avg_latency_ms > 0.0);
assert!(metrics.throughput.requests_per_second > 0.0);
assert_eq!(metrics.latency.max_latency_ms, 500.0);
assert!(metrics.latency.p99_latency_ms >= 140.0);
assert_eq!(
metrics
.error_metrics
.error_counts_by_type
.get("InvalidInput"),
Some(&1)
);
assert!(!metrics
.error_metrics
.error_counts_by_type
.contains_key("prediction_error"));
}
#[test]
fn test_collect_metrics_reports_when_sample_is_skipped() {
let metadata = create_test_metadata();
let mut monitor = ModelMonitor::new(metadata);
let deployment_metrics = create_test_deployment_metrics();
let first = monitor
.collect_metrics(&deployment_metrics)
.expect("first collection succeeds");
let second = monitor
.collect_metrics(&deployment_metrics)
.expect("second call does not error");
assert!(first, "first call is not throttled");
assert!(!second, "immediate second call must report no sample taken");
}
#[test]
fn test_data_drift_is_none_until_baseline_exists() {
let metadata = create_test_metadata();
let monitor = ModelMonitor::new(metadata);
assert!(monitor.calculate_data_drift().is_none());
}
#[test]
fn test_data_drift_is_none_until_live_observations_exist() {
let metadata = create_test_metadata();
let mut monitor = ModelMonitor::new(metadata);
monitor.set_feature_baseline("x1", &(0..100).map(|i| i as f64).collect::<Vec<_>>());
assert!(monitor.calculate_data_drift().is_none());
}
#[test]
fn test_real_psi_drift_detects_a_shifted_distribution() {
let metadata = create_test_metadata();
let mut monitor = ModelMonitor::new(metadata);
let reference: Vec<f64> = (0..1000).map(|i| (i % 100) as f64).collect();
monitor.set_feature_baseline("x1", &reference);
for _ in 0..200 {
monitor.observe_feature("x1", 99.0);
}
let drift = monitor
.calculate_data_drift()
.expect("baseline + live observations exist");
assert_eq!(drift.detection_method, "PSI");
assert!(
drift.drift_score > PSI_DRIFT_THRESHOLD,
"a fully-shifted distribution must exceed the drift threshold, got {}",
drift.drift_score
);
assert!(drift.drift_detected);
assert!(drift.drifting_features.contains(&"x1".to_string()));
}
#[test]
fn test_real_psi_drift_is_low_for_matching_distribution() {
let metadata = create_test_metadata();
let mut monitor = ModelMonitor::new(metadata);
let reference: Vec<f64> = (0..1000).map(|i| (i % 100) as f64).collect();
monitor.set_feature_baseline("x1", &reference);
for i in 0..1000 {
monitor.observe_feature("x1", (i % 100) as f64);
}
let drift = monitor
.calculate_data_drift()
.expect("baseline + live observations exist");
assert!(
drift.drift_score < PSI_DRIFT_THRESHOLD,
"a matching distribution must not exceed the drift threshold, got {}",
drift.drift_score
);
assert!(!drift.drift_detected);
}
#[test]
fn test_confidence_metrics_none_until_recorded() {
let metadata = create_test_metadata();
let mut monitor = ModelMonitor::new(metadata);
assert!(monitor.calculate_confidence_metrics().is_none());
monitor.record_prediction_confidence(0.9);
monitor.record_prediction_confidence(0.5);
let confidence = monitor
.calculate_confidence_metrics()
.expect("samples were recorded");
assert!((confidence.avg_confidence - 0.7).abs() < 1e-9);
assert_eq!(confidence.min_confidence, 0.5);
assert_eq!(confidence.max_confidence, 0.9);
}
#[test]
fn test_alert_severity_maps_to_log_level_without_panicking() {
let metadata = create_test_metadata();
let mut monitor = ModelMonitor::new(metadata);
for severity in [
AlertSeverity::Info,
AlertSeverity::Warning,
AlertSeverity::Critical,
AlertSeverity::Emergency,
] {
let config = AlertConfig {
name: format!("{:?}", severity),
description: "test".to_string(),
metric: "avg_latency_ms".to_string(),
threshold: 0.0,
operator: ComparisonOperator::GreaterThanOrEqual,
severity,
evaluation_window_seconds: 60,
consecutive_evaluations: 1,
cooldown_seconds: 0,
enabled: true,
};
monitor
.trigger_alert(&config, 1.0)
.expect("trigger_alert does not error");
}
assert_eq!(monitor.get_recent_alerts(10).len(), 4);
}
#[test]
fn test_latency_percentiles_are_real_order_statistics_not_fixed_multiples_of_the_mean() {
let mut samples: Vec<u64> = vec![1; 90];
samples.extend(std::iter::repeat(1000).take(10));
let metrics = latency_metrics_from_samples(&samples);
assert!((metrics.avg_latency_ms - 100.9).abs() < 1e-9);
assert_eq!(
metrics.p50_latency_ms, 1.0,
"p50 must reflect the real bulk of fast requests, not a multiple of the mean"
);
assert!(metrics.p50_latency_ms < 10.0);
assert_eq!(metrics.max_latency_ms, 1000.0);
assert_eq!(metrics.min_latency_ms, 1.0);
assert_eq!(metrics.p95_latency_ms, 1000.0);
assert_eq!(metrics.p99_latency_ms, 1000.0);
}
#[test]
fn test_get_metric_value_returns_none_for_unavailable_data_not_an_error() {
let metadata = create_test_metadata();
let monitor = ModelMonitor::new(metadata.clone());
let deployment_metrics = create_test_deployment_metrics();
let snapshot = PerformanceMetrics {
model_name: metadata.name.clone(),
model_version: metadata.version.clone(),
timestamp: chrono::Utc::now(),
latency: monitor.calculate_latency_metrics(&deployment_metrics),
throughput: monitor.calculate_throughput_metrics(&deployment_metrics),
error_metrics: monitor.calculate_error_metrics(&deployment_metrics),
resource_utilization: monitor.calculate_resource_metrics(&deployment_metrics),
quality_metrics: monitor.calculate_quality_metrics(),
};
assert_eq!(
monitor
.get_metric_value(&snapshot, "cpu_utilization")
.unwrap(),
None
);
assert_eq!(
monitor.get_metric_value(&snapshot, "drift_score").unwrap(),
None
);
assert_eq!(
monitor
.get_metric_value(&snapshot, "avg_confidence")
.unwrap(),
None
);
assert!(monitor.get_metric_value(&snapshot, "not_a_metric").is_err());
}
}