Skip to main content

lens_core/semantic/
calibration.rs

1//! # Calibration Preservation System
2//!
3//! Ensures semantic features don't cause calibration shock:
4//! - Cap dense features in log-odds space to prevent extreme predictions
5//! - Maintain Expected Calibration Error (ECE) ≤ 0.005 drift from baseline
6//! - Temperature scaling and isotonic calibration as safety mechanisms
7//! - Monitor calibration across query types and languages
8
9use anyhow::{Context, Result};
10use serde::{Deserialize, Serialize};
11use std::collections::HashMap;
12use std::sync::Arc;
13use tokio::sync::RwLock;
14use tracing::{debug, info, warn};
15
16/// Calibration preservation system for semantic search
17pub struct CalibrationSystem {
18    config: CalibrationConfig,
19    /// Baseline calibration measurements (pre-semantic)
20    baseline_calibration: Arc<RwLock<Option<CalibrationMeasurement>>>,
21    /// Current calibration state
22    current_calibration: Arc<RwLock<CalibrationMeasurement>>,
23    /// Temperature scaling models per query type/language
24    temperature_models: Arc<RwLock<HashMap<String, TemperatureModel>>>,
25    /// Feature scaling parameters
26    feature_scalers: Arc<RwLock<FeatureScalerSet>>,
27    /// Calibration history for drift detection
28    calibration_history: Arc<RwLock<Vec<CalibrationSnapshot>>>,
29}
30
31#[derive(Debug, Clone, Serialize, Deserialize)]
32pub struct CalibrationConfig {
33    /// Maximum ECE drift allowed from baseline
34    pub max_ece_drift: f32,
35    /// Cap for dense features in log-odds space
36    pub log_odds_cap: f32,
37    /// Default temperature scaling factor
38    pub temperature: f32,
39    /// Minimum samples for calibration measurement
40    pub min_samples_for_calibration: usize,
41    /// Calibration measurement window size
42    pub measurement_window_size: usize,
43    /// Enable automatic temperature adjustment
44    pub auto_temperature_adjustment: bool,
45}
46
47/// Calibration measurement with confidence intervals
48#[derive(Debug, Clone, Serialize, Deserialize)]
49pub struct CalibrationMeasurement {
50    /// Expected Calibration Error
51    pub ece: f32,
52    /// Calibration error confidence interval
53    pub ece_ci_lower: f32,
54    pub ece_ci_upper: f32,
55    /// Maximum Calibration Error
56    pub mce: f32,
57    /// Reliability diagram data points
58    pub reliability_data: Vec<ReliabilityPoint>,
59    /// Sample count for measurement
60    pub sample_count: usize,
61    /// Measurement timestamp
62    pub timestamp: std::time::SystemTime,
63    /// Query type breakdown
64    pub by_query_type: HashMap<String, f32>,
65    /// Language breakdown
66    pub by_language: HashMap<String, f32>,
67}
68
69/// Point in reliability diagram
70#[derive(Debug, Clone, Serialize, Deserialize)]
71pub struct ReliabilityPoint {
72    /// Predicted probability bin
73    pub predicted_prob: f32,
74    /// Actual positive rate in this bin
75    pub actual_positive_rate: f32,
76    /// Sample count in bin
77    pub bin_count: usize,
78    /// Confidence interval for actual rate
79    pub confidence_interval: (f32, f32),
80}
81
82/// Temperature scaling model for post-hoc calibration
83#[derive(Debug, Clone, Serialize, Deserialize)]
84pub struct TemperatureModel {
85    /// Learned temperature parameter
86    pub temperature: f32,
87    /// Bias term
88    pub bias: f32,
89    /// Model performance metrics
90    pub ece_before: f32,
91    pub ece_after: f32,
92    /// Training sample count
93    pub training_samples: usize,
94    /// Model scope (query_type:language)
95    pub scope: String,
96}
97
98/// Feature scaling to prevent calibration shock
99#[derive(Debug, Clone, Serialize, Deserialize)]
100pub struct FeatureScalerSet {
101    /// Log-odds caps for semantic features
102    pub semantic_feature_caps: HashMap<String, f32>,
103    /// Linear scaling factors
104    pub linear_scalers: HashMap<String, LinearScaler>,
105    /// Quantile-based robust scalers
106    pub robust_scalers: HashMap<String, RobustScaler>,
107}
108
109#[derive(Debug, Clone, Serialize, Deserialize)]
110pub struct LinearScaler {
111    pub mean: f32,
112    pub std: f32,
113}
114
115#[derive(Debug, Clone, Serialize, Deserialize)]
116pub struct RobustScaler {
117    pub median: f32,
118    pub iqr: f32, // Interquartile range
119}
120
121/// Snapshot of calibration at a point in time
122#[derive(Debug, Clone, Serialize, Deserialize)]
123pub struct CalibrationSnapshot {
124    pub timestamp: std::time::SystemTime,
125    pub ece: f32,
126    pub mce: f32,
127    pub sample_count: usize,
128    pub semantic_features_active: bool,
129    pub temperature: f32,
130}
131
132/// Prediction with confidence for calibration analysis
133#[derive(Debug, Clone)]
134pub struct CalibratedPrediction {
135    pub raw_score: f32,
136    pub calibrated_score: f32,
137    pub confidence: f32,
138    pub query_type: String,
139    pub language: Option<String>,
140    pub features_used: Vec<String>,
141}
142
143/// Training sample for calibration
144#[derive(Debug, Clone)]
145pub struct CalibrationSample {
146    pub prediction: f32,
147    pub actual_relevance: f32, // 0.0 or 1.0 for binary relevance
148    pub query_type: String,
149    pub language: Option<String>,
150}
151
152impl CalibrationSystem {
153    /// Create new calibration system
154    pub async fn new(config: CalibrationConfig) -> Result<Self> {
155        info!("Creating calibration preservation system");
156        info!("ECE drift limit: {:.4}, log-odds cap: {:.2}, temperature: {:.2}", 
157              config.max_ece_drift, config.log_odds_cap, config.temperature);
158        
159        Ok(Self {
160            config,
161            baseline_calibration: Arc::new(RwLock::new(None)),
162            current_calibration: Arc::new(RwLock::new(CalibrationMeasurement::default())),
163            temperature_models: Arc::new(RwLock::new(HashMap::new())),
164            feature_scalers: Arc::new(RwLock::new(FeatureScalerSet::default())),
165            calibration_history: Arc::new(RwLock::new(Vec::new())),
166        })
167    }
168    
169    /// Establish baseline calibration before semantic features
170    pub async fn establish_baseline(&self, samples: &[CalibrationSample]) -> Result<()> {
171        info!("Establishing baseline calibration with {} samples", samples.len());
172        
173        if samples.len() < self.config.min_samples_for_calibration {
174            anyhow::bail!("Insufficient samples for baseline: {} < {}", 
175                         samples.len(), self.config.min_samples_for_calibration);
176        }
177        
178        // Calculate baseline calibration
179        let measurement = self.calculate_calibration_measurement(samples).await?;
180        
181        info!("Baseline calibration established: ECE = {:.4}, MCE = {:.4}", 
182              measurement.ece, measurement.mce);
183        
184        *self.baseline_calibration.write().await = Some(measurement.clone());
185        // Add to history before moving measurement
186        self.add_calibration_snapshot(measurement.ece, measurement.mce, false, self.config.temperature).await;
187        
188        // Store current measurement
189        *self.current_calibration.write().await = measurement;
190        
191        Ok(())
192    }
193    
194    /// Apply calibration-aware scaling to features
195    pub async fn scale_features(&self, features: &HashMap<String, f32>) -> Result<HashMap<String, f32>> {
196        let scalers = self.feature_scalers.read().await;
197        let mut scaled_features = HashMap::new();
198        
199        for (name, value) in features {
200            let scaled_value = if name.contains("semantic") {
201                // Apply log-odds capping for semantic features
202                let capped_value = if let Some(&cap) = scalers.semantic_feature_caps.get(name) {
203                    value.clamp(-cap, cap)
204                } else {
205                    value.clamp(-self.config.log_odds_cap, self.config.log_odds_cap)
206                };
207                
208                // Apply robust scaling if available
209                if let Some(robust_scaler) = scalers.robust_scalers.get(name) {
210                    (capped_value - robust_scaler.median) / (robust_scaler.iqr + 1e-8)
211                } else {
212                    capped_value
213                }
214            } else {
215                // Standard scaling for non-semantic features
216                if let Some(linear_scaler) = scalers.linear_scalers.get(name) {
217                    (value - linear_scaler.mean) / (linear_scaler.std + 1e-8)
218                } else {
219                    *value
220                }
221            };
222            
223            scaled_features.insert(name.clone(), scaled_value);
224        }
225        
226        debug!("Scaled {} features with calibration preservation", scaled_features.len());
227        Ok(scaled_features)
228    }
229    
230    /// Apply temperature scaling to prediction
231    pub async fn apply_temperature_scaling(&self, prediction: f32, query_type: &str, language: Option<&str>) -> Result<f32> {
232        let models = self.temperature_models.read().await;
233        
234        // Try to find specific model for query type + language
235        let scope_key = match language {
236            Some(lang) => format!("{}:{}", query_type, lang),
237            None => query_type.to_string(),
238        };
239        
240        let temperature = if let Some(model) = models.get(&scope_key) {
241            model.temperature
242        } else if let Some(model) = models.get(query_type) {
243            model.temperature
244        } else {
245            self.config.temperature
246        };
247        
248        // Apply temperature scaling: p_calibrated = sigmoid(logit(p) / T)
249        let logit_p = (prediction / (1.0 - prediction + 1e-8)).ln();
250        let scaled_logit = logit_p / temperature;
251        let calibrated_p = 1.0 / (1.0 + (-scaled_logit).exp());
252        
253        Ok(calibrated_p.clamp(0.001, 0.999)) // Avoid extreme probabilities
254    }
255    
256    /// Update calibration measurements with new samples
257    pub async fn update_calibration(&self, samples: &[CalibrationSample]) -> Result<CalibrationStatus> {
258        if samples.len() < self.config.min_samples_for_calibration / 4 {
259            return Ok(CalibrationStatus::InsufficientData);
260        }
261        
262        // Calculate current calibration
263        let current_measurement = self.calculate_calibration_measurement(samples).await?;
264        
265        // Check for drift from baseline
266        let baseline_guard = self.baseline_calibration.read().await;
267        let drift_status = if let Some(baseline) = baseline_guard.as_ref() {
268            let ece_drift = (current_measurement.ece - baseline.ece).abs();
269            
270            if ece_drift > self.config.max_ece_drift {
271                warn!("Calibration drift detected: ECE drift {:.4} > limit {:.4}", 
272                      ece_drift, self.config.max_ece_drift);
273                CalibrationStatus::DriftDetected { 
274                    ece_drift,
275                    current_ece: current_measurement.ece,
276                    baseline_ece: baseline.ece,
277                }
278            } else {
279                CalibrationStatus::WithinLimits {
280                    ece_drift,
281                    current_ece: current_measurement.ece,
282                }
283            }
284        } else {
285            CalibrationStatus::NoBaseline
286        };
287        
288        // Update current measurement
289        *self.current_calibration.write().await = current_measurement.clone();
290        
291        // Add to history
292        self.add_calibration_snapshot(current_measurement.ece, current_measurement.mce, true, self.config.temperature).await;
293        
294        // Trigger automatic temperature adjustment if needed
295        if matches!(drift_status, CalibrationStatus::DriftDetected { .. }) && self.config.auto_temperature_adjustment {
296            self.adjust_temperature_models(samples).await?;
297        }
298        
299        info!("Calibration updated: ECE = {:.4}, status = {:?}", 
300              current_measurement.ece, drift_status);
301        
302        Ok(drift_status)
303    }
304    
305    /// Train temperature scaling models
306    pub async fn train_temperature_models(&self, samples: &[CalibrationSample]) -> Result<()> {
307        info!("Training temperature scaling models with {} samples", samples.len());
308        
309        // Group samples by query type and language
310        let mut grouped_samples: HashMap<String, Vec<&CalibrationSample>> = HashMap::new();
311        
312        for sample in samples {
313            let key = match &sample.language {
314                Some(lang) => format!("{}:{}", sample.query_type, lang),
315                None => sample.query_type.clone(),
316            };
317            grouped_samples.entry(key.clone()).or_default().push(sample);
318            
319            // Also add to query type only group
320            grouped_samples.entry(sample.query_type.clone()).or_default().push(sample);
321        }
322        
323        let mut models = self.temperature_models.write().await;
324        
325        for (scope, group_samples) in grouped_samples {
326            if group_samples.len() >= self.config.min_samples_for_calibration / 4 {
327                let model = self.fit_temperature_model(group_samples, &scope).await?;
328                models.insert(scope, model);
329            }
330        }
331        
332        info!("Trained {} temperature models", models.len());
333        Ok(())
334    }
335    
336    /// Get current calibration metrics
337    pub async fn get_calibration_metrics(&self) -> CalibrationMetrics {
338        let current = self.current_calibration.read().await;
339        let baseline = self.baseline_calibration.read().await;
340        let history = self.calibration_history.read().await;
341        
342        let ece_drift = if let Some(baseline_measurement) = baseline.as_ref() {
343            (current.ece - baseline_measurement.ece).abs()
344        } else {
345            0.0
346        };
347        
348        let recent_trend = if history.len() >= 5 {
349            let recent_eces: Vec<f32> = history.iter().rev().take(5).map(|s| s.ece).collect();
350            let first = recent_eces[4];
351            let last = recent_eces[0];
352            (last - first) / 5.0 // Average change per measurement
353        } else {
354            0.0
355        };
356        
357        CalibrationMetrics {
358            current_ece: current.ece,
359            baseline_ece: baseline.as_ref().map(|b| b.ece).unwrap_or(0.0),
360            ece_drift,
361            max_allowed_drift: self.config.max_ece_drift,
362            within_limits: ece_drift <= self.config.max_ece_drift,
363            recent_trend,
364            measurement_count: history.len(),
365            sample_count: current.sample_count,
366        }
367    }
368    
369    // Private implementation methods
370    
371    fn calculate_calibration_measurement<'a>(&'a self, samples: &'a [CalibrationSample]) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<CalibrationMeasurement>> + Send + 'a>> {
372        Box::pin(async move {
373        let num_bins = 10;
374        let mut bins = vec![Vec::new(); num_bins];
375        
376        // Assign samples to bins based on predicted probability
377        for sample in samples {
378            let bin_idx = ((sample.prediction * num_bins as f32) as usize).min(num_bins - 1);
379            bins[bin_idx].push(sample);
380        }
381        
382        // Calculate reliability points
383        let mut reliability_data = Vec::new();
384        let mut total_ece = 0.0;
385        let mut total_samples = 0;
386        let mut max_calibration_error = 0.0;
387        
388        for (i, bin) in bins.iter().enumerate() {
389            if bin.is_empty() {
390                continue;
391            }
392            
393            let bin_center = (i as f32 + 0.5) / num_bins as f32;
394            let actual_positive_rate = bin.iter()
395                .map(|s| s.actual_relevance)
396                .sum::<f32>() / bin.len() as f32;
397            
398            let bin_error = (bin_center - actual_positive_rate).abs();
399            let bin_weight = bin.len() as f32 / samples.len() as f32;
400            
401            total_ece += bin_error * bin_weight;
402            max_calibration_error = f32::max(max_calibration_error, bin_error);
403            total_samples += bin.len();
404            
405            // Calculate confidence interval (Wilson score interval)
406            let (ci_lower, ci_upper) = self.calculate_wilson_confidence_interval(
407                actual_positive_rate, bin.len()
408            );
409            
410            reliability_data.push(ReliabilityPoint {
411                predicted_prob: bin_center,
412                actual_positive_rate,
413                bin_count: bin.len(),
414                confidence_interval: (ci_lower, ci_upper),
415            });
416        }
417        
418        // Calculate ECE confidence interval
419        let ece_std_error = (total_ece * (1.0 - total_ece) / samples.len() as f32).sqrt();
420        let ece_ci_lower = (total_ece - 1.96 * ece_std_error).max(0.0);
421        let ece_ci_upper = (total_ece + 1.96 * ece_std_error).min(1.0);
422        
423        // Calculate breakdown by query type and language
424        let by_query_type = self.calculate_ece_by_category(samples, |s| &s.query_type).await?;
425        let by_language = self.calculate_ece_by_language(samples).await?;
426        
427        Ok(CalibrationMeasurement {
428            ece: total_ece,
429            ece_ci_lower,
430            ece_ci_upper,
431            mce: max_calibration_error,
432            reliability_data,
433            sample_count: samples.len(),
434            timestamp: std::time::SystemTime::now(),
435            by_query_type,
436            by_language,
437        })
438        })
439    }
440    
441    fn calculate_wilson_confidence_interval(&self, p: f32, n: usize) -> (f32, f32) {
442        if n == 0 {
443            return (0.0, 1.0);
444        }
445        
446        let z = 1.96; // 95% confidence
447        let n_f = n as f32;
448        
449        let center = p + z * z / (2.0 * n_f);
450        let half_width = z * ((p * (1.0 - p) / n_f) + (z * z / (4.0 * n_f * n_f))).sqrt();
451        let denominator = 1.0 + z * z / n_f;
452        
453        let lower = ((center - half_width) / denominator).max(0.0);
454        let upper = ((center + half_width) / denominator).min(1.0);
455        
456        (lower, upper)
457    }
458    
459    async fn calculate_ece_by_category<F>(&self, samples: &[CalibrationSample], category_fn: F) -> Result<HashMap<String, f32>>
460    where
461        F: Fn(&CalibrationSample) -> &String,
462    {
463        let mut category_groups: HashMap<String, Vec<&CalibrationSample>> = HashMap::new();
464        
465        for sample in samples {
466            let category = category_fn(sample);
467            category_groups.entry(category.clone()).or_default().push(sample);
468        }
469        
470        let mut category_eces = HashMap::new();
471        
472        for (category, group) in category_groups {
473            if group.len() >= 10 { // Minimum samples for reliable ECE
474                let group_samples: Vec<CalibrationSample> = group.into_iter().cloned().collect();
475                let measurement = self.calculate_calibration_measurement(&group_samples).await?;
476                category_eces.insert(category, measurement.ece);
477            }
478        }
479        
480        Ok(category_eces)
481    }
482    
483    async fn calculate_ece_by_language(&self, samples: &[CalibrationSample]) -> Result<HashMap<String, f32>> {
484        let mut language_groups: HashMap<String, Vec<&CalibrationSample>> = HashMap::new();
485        
486        for sample in samples {
487            let language = sample.language.as_deref().unwrap_or("unknown");
488            language_groups.entry(language.to_string()).or_default().push(sample);
489        }
490        
491        let mut language_eces = HashMap::new();
492        
493        for (language, group) in language_groups {
494            if group.len() >= 10 {
495                let group_samples: Vec<CalibrationSample> = group.into_iter().cloned().collect();
496                let measurement = self.calculate_calibration_measurement(&group_samples).await?;
497                language_eces.insert(language, measurement.ece);
498            }
499        }
500        
501        Ok(language_eces)
502    }
503    
504    async fn fit_temperature_model(&self, samples: Vec<&CalibrationSample>, scope: &str) -> Result<TemperatureModel> {
505        // Calculate ECE before temperature scaling
506        let samples_owned: Vec<CalibrationSample> = samples.into_iter().cloned().collect();
507        let before_measurement = self.calculate_calibration_measurement(&samples_owned).await?;
508        
509        // Use simple grid search to find optimal temperature
510        let mut best_temperature = 1.0;
511        let mut best_ece = before_measurement.ece;
512        
513        for temp_candidate in (5..=200).step_by(5) {
514            let temperature = temp_candidate as f32 / 100.0; // 0.05 to 2.0
515            
516            // Apply temperature scaling and calculate ECE
517            let calibrated_samples: Vec<CalibrationSample> = samples_owned.iter()
518                .map(|s| {
519                    let logit_p = (s.prediction / (1.0 - s.prediction + 1e-8)).ln();
520                    let scaled_logit = logit_p / temperature;
521                    let calibrated_p = 1.0 / (1.0 + (-scaled_logit).exp());
522                    
523                    CalibrationSample {
524                        prediction: calibrated_p.clamp(0.001, 0.999),
525                        actual_relevance: s.actual_relevance,
526                        query_type: s.query_type.clone(),
527                        language: s.language.clone(),
528                    }
529                })
530                .collect();
531            
532            let temp_measurement = self.calculate_calibration_measurement(&calibrated_samples).await?;
533            
534            if temp_measurement.ece < best_ece {
535                best_ece = temp_measurement.ece;
536                best_temperature = temperature;
537            }
538        }
539        
540        Ok(TemperatureModel {
541            temperature: best_temperature,
542            bias: 0.0, // Not used in current implementation
543            ece_before: before_measurement.ece,
544            ece_after: best_ece,
545            training_samples: samples_owned.len(),
546            scope: scope.to_string(),
547        })
548    }
549    
550    async fn adjust_temperature_models(&self, samples: &[CalibrationSample]) -> Result<()> {
551        info!("Adjusting temperature models due to calibration drift");
552        
553        // Retrain temperature models with current data
554        self.train_temperature_models(samples).await?;
555        
556        Ok(())
557    }
558    
559    async fn add_calibration_snapshot(&self, ece: f32, mce: f32, semantic_active: bool, temperature: f32) {
560        let snapshot = CalibrationSnapshot {
561            timestamp: std::time::SystemTime::now(),
562            ece,
563            mce,
564            sample_count: 0, // Would be filled in real implementation
565            semantic_features_active: semantic_active,
566            temperature,
567        };
568        
569        let mut history = self.calibration_history.write().await;
570        history.push(snapshot);
571        
572        // Keep only recent history
573        let history_len = history.len();
574        if history_len > self.config.measurement_window_size {
575            history.drain(0..history_len - self.config.measurement_window_size);
576        }
577    }
578}
579
580// Default implementations
581
582impl Default for CalibrationMeasurement {
583    fn default() -> Self {
584        Self {
585            ece: 0.0,
586            ece_ci_lower: 0.0,
587            ece_ci_upper: 0.0,
588            mce: 0.0,
589            reliability_data: Vec::new(),
590            sample_count: 0,
591            timestamp: std::time::SystemTime::now(),
592            by_query_type: HashMap::new(),
593            by_language: HashMap::new(),
594        }
595    }
596}
597
598impl Default for FeatureScalerSet {
599    fn default() -> Self {
600        Self {
601            semantic_feature_caps: HashMap::new(),
602            linear_scalers: HashMap::new(),
603            robust_scalers: HashMap::new(),
604        }
605    }
606}
607
608/// Calibration drift status
609#[derive(Debug, Clone)]
610pub enum CalibrationStatus {
611    WithinLimits { ece_drift: f32, current_ece: f32 },
612    DriftDetected { ece_drift: f32, current_ece: f32, baseline_ece: f32 },
613    NoBaseline,
614    InsufficientData,
615}
616
617/// Calibration system metrics
618#[derive(Debug, Clone, Serialize, Deserialize)]
619pub struct CalibrationMetrics {
620    pub current_ece: f32,
621    pub baseline_ece: f32,
622    pub ece_drift: f32,
623    pub max_allowed_drift: f32,
624    pub within_limits: bool,
625    pub recent_trend: f32,
626    pub measurement_count: usize,
627    pub sample_count: usize,
628}
629
630/// Initialize calibration system
631pub async fn initialize_calibration(config: &CalibrationConfig) -> Result<()> {
632    info!("Initializing calibration preservation system");
633    info!("ECE drift limit: {:.4}, log-odds cap: {:.2}", 
634          config.max_ece_drift, config.log_odds_cap);
635    
636    if config.max_ece_drift > 0.01 {
637        warn!("ECE drift limit {:.4} may be too permissive", config.max_ece_drift);
638    }
639    
640    if config.log_odds_cap < 3.0 {
641        warn!("Log-odds cap {:.2} may be too restrictive", config.log_odds_cap);
642    }
643    
644    info!("Calibration system initialized");
645    Ok(())
646}
647
648#[cfg(test)]
649mod tests {
650    use super::*;
651
652    #[tokio::test]
653    async fn test_calibration_system_creation() {
654        let config = CalibrationConfig {
655            max_ece_drift: 0.005,
656            log_odds_cap: 5.0,
657            temperature: 1.0,
658            min_samples_for_calibration: 100,
659            measurement_window_size: 50,
660            auto_temperature_adjustment: true,
661        };
662        
663        let system = CalibrationSystem::new(config).await.unwrap();
664        let metrics = system.get_calibration_metrics().await;
665        assert_eq!(metrics.measurement_count, 0);
666    }
667
668    #[tokio::test]
669    async fn test_feature_scaling() {
670        let config = CalibrationConfig {
671            max_ece_drift: 0.005,
672            log_odds_cap: 3.0,
673            temperature: 1.0,
674            min_samples_for_calibration: 100,
675            measurement_window_size: 50,
676            auto_temperature_adjustment: false,
677        };
678        
679        let system = CalibrationSystem::new(config).await.unwrap();
680        
681        let mut features = HashMap::new();
682        features.insert("semantic_similarity".to_string(), 8.0); // High value
683        features.insert("lexical_score".to_string(), 0.7);
684        
685        let scaled = system.scale_features(&features).await.unwrap();
686        
687        // Semantic feature should be capped
688        assert!(scaled["semantic_similarity"].abs() <= 3.0);
689        // Non-semantic feature should pass through (no scaler configured)
690        assert_eq!(scaled["lexical_score"], 0.7);
691    }
692
693    #[tokio::test]
694    async fn test_temperature_scaling() {
695        let config = CalibrationConfig {
696            max_ece_drift: 0.005,
697            log_odds_cap: 5.0,
698            temperature: 2.0, // Higher temperature for smoothing
699            min_samples_for_calibration: 100,
700            measurement_window_size: 50,
701            auto_temperature_adjustment: false,
702        };
703        
704        let system = CalibrationSystem::new(config).await.unwrap();
705        
706        let extreme_prediction = 0.95;
707        let calibrated = system.apply_temperature_scaling(extreme_prediction, "natural_language", None).await.unwrap();
708        
709        // Temperature scaling should reduce extreme predictions
710        assert!(calibrated < extreme_prediction);
711        assert!(calibrated > 0.5); // Should still be above 0.5 for high confidence
712    }
713
714    #[test]
715    fn test_wilson_confidence_interval() {
716        let config = CalibrationConfig {
717            max_ece_drift: 0.005,
718            log_odds_cap: 5.0,
719            temperature: 1.0,
720            min_samples_for_calibration: 100,
721            measurement_window_size: 50,
722            auto_temperature_adjustment: false,
723        };
724        
725        let system = CalibrationSystem {
726            config,
727            baseline_calibration: Arc::new(RwLock::new(None)),
728            current_calibration: Arc::new(RwLock::new(CalibrationMeasurement::default())),
729            temperature_models: Arc::new(RwLock::new(HashMap::new())),
730            feature_scalers: Arc::new(RwLock::new(FeatureScalerSet::default())),
731            calibration_history: Arc::new(RwLock::new(Vec::new())),
732        };
733        
734        let (lower, upper) = system.calculate_wilson_confidence_interval(0.5, 100);
735        
736        // Should have reasonable confidence interval
737        assert!(lower < 0.5);
738        assert!(upper > 0.5);
739        assert!(upper - lower < 0.2); // Not too wide for n=100
740    }
741
742    #[tokio::test]
743    async fn test_calibration_measurement() {
744        let config = CalibrationConfig {
745            max_ece_drift: 0.005,
746            log_odds_cap: 5.0,
747            temperature: 1.0,
748            min_samples_for_calibration: 10,
749            measurement_window_size: 50,
750            auto_temperature_adjustment: false,
751        };
752        
753        let system = CalibrationSystem::new(config).await.unwrap();
754        
755        // Create perfectly calibrated samples
756        let samples = vec![
757            CalibrationSample { prediction: 0.1, actual_relevance: 0.0, query_type: "test".to_string(), language: None },
758            CalibrationSample { prediction: 0.3, actual_relevance: 0.0, query_type: "test".to_string(), language: None },
759            CalibrationSample { prediction: 0.5, actual_relevance: 1.0, query_type: "test".to_string(), language: None },
760            CalibrationSample { prediction: 0.7, actual_relevance: 1.0, query_type: "test".to_string(), language: None },
761            CalibrationSample { prediction: 0.9, actual_relevance: 1.0, query_type: "test".to_string(), language: None },
762        ];
763        
764        let measurement = system.calculate_calibration_measurement(&samples).await.unwrap();
765        
766        assert!(measurement.ece >= 0.0);
767        assert!(measurement.ece <= 1.0);
768        assert!(measurement.mce >= 0.0);
769        assert_eq!(measurement.sample_count, 5);
770    }
771}