use anyhow::{Context, Result};
use serde::{Deserialize, Serialize};
use std::collections::{HashMap, VecDeque};
use std::sync::Arc;
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use tokio::sync::RwLock;
use tracing::{debug, info, instrument, warn};
use super::query_classifier::{QueryClassification, QueryIntent};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ConformalRouterConfig {
pub risk_threshold: f32,
pub daily_budget_percent: f32,
pub confidence_level: f32,
pub min_calibration_samples: usize,
pub p95_headroom_threshold_ms: f32,
pub enabled: bool,
pub calibration_retention_hours: u64,
}
impl Default for ConformalRouterConfig {
fn default() -> Self {
Self {
risk_threshold: 0.6,
daily_budget_percent: 5.0,
confidence_level: 0.95,
min_calibration_samples: 100,
p95_headroom_threshold_ms: 10.0,
enabled: true,
calibration_retention_hours: 24,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ConformalFeatures {
pub query_length: u32,
pub word_count: u32,
pub has_special_chars: bool,
pub fuzzy_enabled: bool,
pub structural_mode: bool,
pub avg_word_length: f32,
pub query_entropy: f32,
pub identifier_density: f32,
pub semantic_complexity: f32,
pub has_file_context: bool,
pub language_detected: bool,
pub intent_confidence: f32,
pub naturalness_score: f32,
pub similar_queries_success_rate: f32,
pub user_satisfaction_history: f32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RiskAssessment {
pub risk_score: f32,
pub confidence_interval: (f32, f32),
pub nonconformity_score: f32,
pub calibrated: bool,
pub risk_factors: Vec<RiskFactor>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RiskFactor {
pub name: String,
pub contribution: f32,
pub description: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RoutingDecision {
pub should_upshift: bool,
pub upshift_type: UpshiftType,
pub budget_consumed: f32,
pub routing_reason: String,
pub expected_improvement: f32,
pub risk_assessment: RiskAssessment,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum UpshiftType {
None,
HighDimEmbeddings,
EnhancedSearch,
DiversityOptimization,
CrossEncoder,
SemanticReranking,
LSPIntegration,
ASTAnalysis,
CrossLanguageSearch,
}
impl std::fmt::Display for UpshiftType {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
UpshiftType::None => write!(f, "none"),
UpshiftType::HighDimEmbeddings => write!(f, "high_dim_embeddings"),
UpshiftType::EnhancedSearch => write!(f, "enhanced_search"),
UpshiftType::DiversityOptimization => write!(f, "diversity_optimization"),
UpshiftType::CrossEncoder => write!(f, "cross_encoder"),
UpshiftType::SemanticReranking => write!(f, "semantic_reranking"),
UpshiftType::LSPIntegration => write!(f, "lsp_integration"),
UpshiftType::ASTAnalysis => write!(f, "ast_analysis"),
UpshiftType::CrossLanguageSearch => write!(f, "cross_language_search"),
}
}
}
#[derive(Debug, Clone)]
pub struct BudgetManager {
daily_budget: f32,
current_usage: f32,
usage_history: VecDeque<(SystemTime, f32)>,
last_reset: SystemTime,
total_queries_today: u64,
}
impl BudgetManager {
pub fn new(daily_budget_percent: f32) -> Self {
Self {
daily_budget: daily_budget_percent / 100.0,
current_usage: 0.0,
usage_history: VecDeque::with_capacity(1000),
last_reset: SystemTime::now(),
total_queries_today: 0,
}
}
pub fn can_upshift(&mut self, cost: f32) -> bool {
self.maybe_reset_daily();
(self.current_usage + cost) <= self.daily_budget
}
pub fn record_upshift(&mut self, cost: f32) {
self.maybe_reset_daily();
self.current_usage += cost;
let now = SystemTime::now();
self.usage_history.push_back((now, cost));
let cutoff = now - Duration::from_secs(24 * 60 * 60);
while let Some((timestamp, _)) = self.usage_history.front() {
if *timestamp < cutoff {
self.usage_history.pop_front();
} else {
break;
}
}
}
pub fn record_query(&mut self) {
self.maybe_reset_daily();
self.total_queries_today += 1;
}
pub fn get_status(&self) -> BudgetStatus {
let usage_rate = if self.total_queries_today > 0 {
(self.current_usage * self.total_queries_today as f32) / self.total_queries_today as f32
} else {
0.0
};
BudgetStatus {
daily_budget_percent: self.daily_budget * 100.0,
current_usage_percent: (self.current_usage / self.daily_budget) * 100.0,
remaining_budget: self.daily_budget - self.current_usage,
usage_rate_percent: usage_rate * 100.0,
total_queries: self.total_queries_today,
p95_headroom_estimate: self.estimate_p95_headroom(),
}
}
fn maybe_reset_daily(&mut self) {
let now = SystemTime::now();
if let Ok(duration) = now.duration_since(self.last_reset) {
if duration >= Duration::from_secs(24 * 60 * 60) {
self.current_usage = 0.0;
self.total_queries_today = 0;
self.last_reset = now;
info!("Reset daily budget counters");
}
}
}
fn estimate_p95_headroom(&self) -> f32 {
let recent_usage: f32 = self.usage_history.iter()
.filter(|(timestamp, _)| {
SystemTime::now().duration_since(*timestamp)
.unwrap_or(Duration::ZERO) < Duration::from_secs(60 * 60)
})
.map(|(_, cost)| cost)
.sum();
if recent_usage > 0.5 {
1.0 } else if recent_usage > 0.2 {
5.0 } else {
15.0 }
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BudgetStatus {
pub daily_budget_percent: f32,
pub current_usage_percent: f32,
pub remaining_budget: f32,
pub usage_rate_percent: f32,
pub total_queries: u64,
pub p95_headroom_estimate: f32,
}
#[derive(Debug, Clone)]
pub struct CalibrationSample {
pub features: ConformalFeatures,
pub predicted_quality: f32,
pub actual_quality: f32,
pub timestamp: SystemTime,
}
pub struct ConformalPredictor {
calibration_data: VecDeque<CalibrationSample>,
nonconformity_scores: Vec<f32>,
is_calibrated: bool,
min_samples: usize,
retention_period: Duration,
}
impl ConformalPredictor {
pub fn new(min_samples: usize, retention_hours: u64) -> Self {
Self {
calibration_data: VecDeque::with_capacity(min_samples * 2),
nonconformity_scores: Vec::with_capacity(min_samples),
is_calibrated: false,
min_samples,
retention_period: Duration::from_secs(retention_hours * 60 * 60),
}
}
pub fn add_calibration_sample(&mut self, sample: CalibrationSample) {
let cutoff = SystemTime::now() - self.retention_period;
while let Some(front) = self.calibration_data.front() {
if front.timestamp < cutoff {
self.calibration_data.pop_front();
} else {
break;
}
}
self.calibration_data.push_back(sample);
if self.calibration_data.len() >= self.min_samples {
self.recalibrate();
}
}
pub fn predict_risk(
&self,
features: &ConformalFeatures,
confidence_level: f32,
) -> RiskAssessment {
let base_prediction = self.predict_base_quality(features);
if !self.is_calibrated || self.nonconformity_scores.is_empty() {
return self.heuristic_risk_assessment(features, base_prediction);
}
let alpha = 1.0 - confidence_level;
let quantile_index = ((1.0 - alpha) * (self.nonconformity_scores.len() as f32 + 1.0)) as usize;
let quantile_index = quantile_index.min(self.nonconformity_scores.len() - 1);
let nonconformity_quantile = self.nonconformity_scores[quantile_index];
let lower_bound = (base_prediction - nonconformity_quantile).max(0.0);
let upper_bound = (base_prediction + nonconformity_quantile).min(1.0);
let interval_width = upper_bound - lower_bound;
let quality_risk = 1.0 - base_prediction;
let uncertainty_risk = interval_width;
let risk_score = (quality_risk * 0.7 + uncertainty_risk * 0.3).min(1.0);
RiskAssessment {
risk_score,
confidence_interval: (lower_bound, upper_bound),
nonconformity_score: nonconformity_quantile,
calibrated: true,
risk_factors: self.identify_risk_factors(features, quality_risk, uncertainty_risk),
}
}
fn recalibrate(&mut self) {
self.nonconformity_scores.clear();
for sample in &self.calibration_data {
let predicted = self.predict_base_quality(&sample.features);
let nonconformity = (sample.actual_quality - predicted).abs();
self.nonconformity_scores.push(nonconformity);
}
self.nonconformity_scores.sort_by(|a, b| a.partial_cmp(b).unwrap());
self.is_calibrated = true;
info!(
"Conformal predictor recalibrated with {} samples, mean nonconformity: {:.3}",
self.calibration_data.len(),
self.nonconformity_scores.iter().sum::<f32>() / self.nonconformity_scores.len() as f32
);
}
fn predict_base_quality(&self, features: &ConformalFeatures) -> f32 {
let mut quality = 0.7;
if features.query_length > 100 { quality -= 0.1; }
if features.word_count > 10 { quality -= 0.05; }
if features.semantic_complexity > 0.8 { quality -= 0.15; }
if features.avg_word_length < 3.0 { quality -= 0.1; }
if features.has_file_context { quality += 0.1; }
if features.language_detected { quality += 0.05; }
if features.intent_confidence > 0.8 { quality += 0.1; }
if features.identifier_density > 0.6 { quality += 0.08; }
if features.naturalness_score > 0.7 { quality += 0.05; }
quality += features.similar_queries_success_rate * 0.1;
quality += features.user_satisfaction_history * 0.05;
quality.clamp(0.1, 0.95)
}
fn heuristic_risk_assessment(&self, features: &ConformalFeatures, base_prediction: f32) -> RiskAssessment {
let mut risk_score = 1.0 - base_prediction;
if features.semantic_complexity > 0.7 { risk_score += 0.1; }
if features.query_length > 200 { risk_score += 0.1; }
if features.naturalness_score < 0.3 { risk_score += 0.05; }
risk_score = risk_score.clamp(0.0, 1.0);
let uncertainty = 0.15;
RiskAssessment {
risk_score,
confidence_interval: (
(base_prediction - uncertainty).max(0.0),
(base_prediction + uncertainty).min(1.0)
),
nonconformity_score: uncertainty,
calibrated: false,
risk_factors: self.identify_risk_factors(features, risk_score, uncertainty),
}
}
fn identify_risk_factors(&self, features: &ConformalFeatures, quality_risk: f32, uncertainty_risk: f32) -> Vec<RiskFactor> {
let mut factors = Vec::new();
if features.semantic_complexity > 0.8 {
factors.push(RiskFactor {
name: "semantic_complexity".to_string(),
contribution: features.semantic_complexity * 0.2,
description: "High semantic complexity may reduce search accuracy".to_string(),
});
}
if features.query_length > 100 {
factors.push(RiskFactor {
name: "query_length".to_string(),
contribution: (features.query_length as f32 / 500.0).min(0.2),
description: "Long queries are harder to process accurately".to_string(),
});
}
if features.intent_confidence < 0.5 {
factors.push(RiskFactor {
name: "intent_ambiguity".to_string(),
contribution: (0.5 - features.intent_confidence) * 0.3,
description: "Ambiguous query intent increases search uncertainty".to_string(),
});
}
if !features.has_file_context {
factors.push(RiskFactor {
name: "missing_context".to_string(),
contribution: 0.1,
description: "Lack of file context limits search precision".to_string(),
});
}
factors
}
pub fn get_status(&self) -> ConformalPredictorStatus {
ConformalPredictorStatus {
is_calibrated: self.is_calibrated,
calibration_samples: self.calibration_data.len(),
min_samples_required: self.min_samples,
mean_nonconformity: if self.nonconformity_scores.is_empty() {
0.0
} else {
self.nonconformity_scores.iter().sum::<f32>() / self.nonconformity_scores.len() as f32
},
oldest_sample_age_hours: self.calibration_data.front()
.and_then(|sample| SystemTime::now().duration_since(sample.timestamp).ok())
.map(|d| d.as_secs() as f32 / 3600.0)
.unwrap_or(0.0),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ConformalPredictorStatus {
pub is_calibrated: bool,
pub calibration_samples: usize,
pub min_samples_required: usize,
pub mean_nonconformity: f32,
pub oldest_sample_age_hours: f32,
}
pub struct ConformalRouter {
config: ConformalRouterConfig,
predictor: RwLock<ConformalPredictor>,
budget_manager: RwLock<BudgetManager>,
metrics: Arc<parking_lot::RwLock<ConformalRouterMetrics>>,
upshift_costs: HashMap<UpshiftType, f32>,
}
impl ConformalRouter {
pub fn new(config: ConformalRouterConfig) -> Self {
let predictor = RwLock::new(ConformalPredictor::new(
config.min_calibration_samples,
config.calibration_retention_hours,
));
let budget_manager = RwLock::new(BudgetManager::new(config.daily_budget_percent));
let mut upshift_costs = HashMap::new();
upshift_costs.insert(UpshiftType::None, 0.0);
upshift_costs.insert(UpshiftType::HighDimEmbeddings, 0.02); upshift_costs.insert(UpshiftType::EnhancedSearch, 0.01); upshift_costs.insert(UpshiftType::DiversityOptimization, 0.015); upshift_costs.insert(UpshiftType::CrossEncoder, 0.03);
info!(
"Initialized ConformalRouter: risk_threshold={}, budget={}%, enabled={}",
config.risk_threshold, config.daily_budget_percent, config.enabled
);
Self {
config,
predictor,
budget_manager,
metrics: Arc::new(parking_lot::RwLock::new(ConformalRouterMetrics::default())),
upshift_costs,
}
}
#[instrument(skip(self, features, classification), fields(
query_len = features.query_length,
intent = %classification.intent,
confidence = classification.confidence
))]
pub async fn make_routing_decision(
&self,
features: &ConformalFeatures,
classification: &QueryClassification,
) -> Result<RoutingDecision> {
if !self.config.enabled {
return Ok(RoutingDecision {
should_upshift: false,
upshift_type: UpshiftType::None,
budget_consumed: 0.0,
routing_reason: "router_disabled".to_string(),
expected_improvement: 0.0,
risk_assessment: RiskAssessment {
risk_score: 0.0,
confidence_interval: (0.0, 1.0),
nonconformity_score: 0.0,
calibrated: false,
risk_factors: Vec::new(),
},
});
}
let start = Instant::now();
{
let mut budget = self.budget_manager.write().await;
budget.record_query();
}
let risk_assessment = {
let predictor = self.predictor.read().await;
predictor.predict_risk(features, self.config.confidence_level)
};
debug!(
"Risk assessment: score={:.3}, interval=({:.3}, {:.3}), calibrated={}",
risk_assessment.risk_score,
risk_assessment.confidence_interval.0,
risk_assessment.confidence_interval.1,
risk_assessment.calibrated
);
let should_upshift = risk_assessment.risk_score > self.config.risk_threshold;
if !should_upshift {
self.record_routing_decision(&RoutingDecision {
should_upshift: false,
upshift_type: UpshiftType::None,
budget_consumed: 0.0,
routing_reason: format!("risk_below_threshold_{:.3}", risk_assessment.risk_score),
expected_improvement: 0.0,
risk_assessment: risk_assessment.clone(),
}, start.elapsed()).await;
return Ok(RoutingDecision {
should_upshift: false,
upshift_type: UpshiftType::None,
budget_consumed: 0.0,
routing_reason: format!("risk_below_threshold_{:.3}", risk_assessment.risk_score),
expected_improvement: 0.0,
risk_assessment,
});
}
let upshift_type = self.select_upshift_type(features, &risk_assessment, classification);
let upshift_cost = *self.upshift_costs.get(&upshift_type).unwrap_or(&0.0);
let can_upshift = {
let mut budget = self.budget_manager.write().await;
let budget_status = budget.get_status();
budget.can_upshift(upshift_cost) &&
budget_status.p95_headroom_estimate >= self.config.p95_headroom_threshold_ms
};
if !can_upshift {
let reason = {
let budget = self.budget_manager.read().await;
let status = budget.get_status();
if status.remaining_budget < upshift_cost {
"budget_exhausted".to_string()
} else {
"insufficient_headroom".to_string()
}
};
let decision = RoutingDecision {
should_upshift: false,
upshift_type: UpshiftType::None,
budget_consumed: 0.0,
routing_reason: reason,
expected_improvement: 0.0,
risk_assessment,
};
self.record_routing_decision(&decision, start.elapsed()).await;
return Ok(decision);
}
{
let mut budget = self.budget_manager.write().await;
budget.record_upshift(upshift_cost);
}
let expected_improvement = self.estimate_improvement(upshift_type, &risk_assessment);
let decision = RoutingDecision {
should_upshift: true,
upshift_type,
budget_consumed: upshift_cost,
routing_reason: format!("high_risk_{:.3}", risk_assessment.risk_score),
expected_improvement,
risk_assessment,
};
self.record_routing_decision(&decision, start.elapsed()).await;
info!(
"Upshift decision: type={}, cost={:.3}, improvement={:.3}",
upshift_type, upshift_cost, expected_improvement
);
Ok(decision)
}
fn select_upshift_type(
&self,
features: &ConformalFeatures,
risk_assessment: &RiskAssessment,
classification: &QueryClassification,
) -> UpshiftType {
if features.semantic_complexity > 0.8 || features.naturalness_score > 0.7 {
return UpshiftType::HighDimEmbeddings;
}
if features.structural_mode && features.identifier_density > 0.5 {
return UpshiftType::EnhancedSearch;
}
if classification.intent == QueryIntent::NaturalLanguage && features.word_count > 5 {
return UpshiftType::DiversityOptimization;
}
let interval_width = risk_assessment.confidence_interval.1 - risk_assessment.confidence_interval.0;
if interval_width > 0.3 {
return UpshiftType::CrossEncoder;
}
UpshiftType::EnhancedSearch
}
fn estimate_improvement(&self, upshift_type: UpshiftType, risk_assessment: &RiskAssessment) -> f32 {
let base_improvements = match upshift_type {
UpshiftType::None => 0.0,
UpshiftType::HighDimEmbeddings => 0.08, UpshiftType::EnhancedSearch => 0.04, UpshiftType::DiversityOptimization => 0.06, UpshiftType::CrossEncoder => 0.12, UpshiftType::SemanticReranking => 0.10, UpshiftType::LSPIntegration => 0.15, UpshiftType::ASTAnalysis => 0.09, UpshiftType::CrossLanguageSearch => 0.11, };
let risk_multiplier = 0.5 + (risk_assessment.risk_score * 0.5);
base_improvements * risk_multiplier
}
pub async fn add_calibration_sample(
&self,
features: ConformalFeatures,
predicted_quality: f32,
actual_quality: f32,
) {
let sample = CalibrationSample {
features,
predicted_quality,
actual_quality,
timestamp: SystemTime::now(),
};
let mut predictor = self.predictor.write().await;
predictor.add_calibration_sample(sample);
}
pub async fn get_status(&self) -> ConformalRouterStatus {
let budget_status = {
let budget = self.budget_manager.read().await;
budget.get_status()
};
let predictor_status = {
let predictor = self.predictor.read().await;
predictor.get_status()
};
let metrics = self.metrics.read().clone();
ConformalRouterStatus {
enabled: self.config.enabled,
risk_threshold: self.config.risk_threshold,
budget_status,
predictor_status,
metrics,
}
}
pub async fn update_config(&mut self, new_config: ConformalRouterConfig) {
info!("Updating conformal router configuration");
if new_config.daily_budget_percent != self.config.daily_budget_percent {
let mut budget = self.budget_manager.write().await;
*budget = BudgetManager::new(new_config.daily_budget_percent);
}
if new_config.calibration_retention_hours != self.config.calibration_retention_hours ||
new_config.min_calibration_samples != self.config.min_calibration_samples {
let mut predictor = self.predictor.write().await;
*predictor = ConformalPredictor::new(
new_config.min_calibration_samples,
new_config.calibration_retention_hours,
);
}
self.config = new_config;
info!(
"Updated conformal router: risk_threshold={}, budget={}%, enabled={}",
self.config.risk_threshold, self.config.daily_budget_percent, self.config.enabled
);
}
async fn record_routing_decision(&self, decision: &RoutingDecision, latency: Duration) {
let mut metrics = self.metrics.write();
metrics.total_decisions += 1;
metrics.total_latency += latency;
if decision.should_upshift {
metrics.upshift_decisions += 1;
*metrics.upshift_type_counts.entry(decision.upshift_type).or_insert(0) += 1;
metrics.total_budget_consumed += decision.budget_consumed;
}
if decision.risk_assessment.calibrated {
metrics.calibrated_decisions += 1;
}
let risk_bucket = (decision.risk_assessment.risk_score * 10.0) as usize;
metrics.risk_score_histogram[risk_bucket.min(9)] += 1;
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ConformalRouterStatus {
pub enabled: bool,
pub risk_threshold: f32,
pub budget_status: BudgetStatus,
pub predictor_status: ConformalPredictorStatus,
pub metrics: ConformalRouterMetrics,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct ConformalRouterMetrics {
pub total_decisions: u64,
pub upshift_decisions: u64,
pub calibrated_decisions: u64,
pub total_latency: Duration,
pub total_budget_consumed: f32,
pub upshift_type_counts: HashMap<UpshiftType, u64>,
pub risk_score_histogram: [u64; 10], pub cache_hit_rate: f64, }
impl ConformalRouterMetrics {
pub fn upshift_rate(&self) -> f32 {
if self.total_decisions == 0 {
0.0
} else {
(self.upshift_decisions as f32 / self.total_decisions as f32) * 100.0
}
}
pub fn avg_latency_ms(&self) -> f64 {
if self.total_decisions == 0 {
0.0
} else {
self.total_latency.as_millis() as f64 / self.total_decisions as f64
}
}
pub fn calibration_rate(&self) -> f32 {
if self.total_decisions == 0 {
0.0
} else {
(self.calibrated_decisions as f32 / self.total_decisions as f32) * 100.0
}
}
pub fn avg_risk_score(&self) -> f64 {
let total_weight: u64 = self.risk_score_histogram.iter().sum();
if total_weight == 0 {
return 0.0;
}
let weighted_sum: f64 = self.risk_score_histogram.iter()
.enumerate()
.map(|(i, &count)| (i as f64 + 0.5) * 0.1 * count as f64)
.sum();
weighted_sum / total_weight as f64
}
}
pub fn extract_conformal_features(
query: &str,
classification: &QueryClassification,
file_context: Option<&crate::semantic::intent_router::FileContext>,
) -> ConformalFeatures {
let words: Vec<&str> = query.split_whitespace().collect();
let chars: Vec<char> = query.chars().collect();
let mut char_counts = HashMap::new();
for &ch in &chars {
*char_counts.entry(ch).or_insert(0) += 1;
}
let query_entropy = if chars.is_empty() {
0.0
} else {
char_counts.values().map(|&count| {
let p = count as f32 / chars.len() as f32;
-p * p.log2()
}).sum()
};
let identifier_pattern = regex::Regex::new(r"[a-zA-Z_][a-zA-Z0-9_]*").unwrap();
let identifiers = identifier_pattern.find_iter(query).count();
let identifier_density = if words.is_empty() {
0.0
} else {
identifiers as f32 / words.len() as f32
};
let avg_word_length = if words.is_empty() {
0.0
} else {
words.iter().map(|w| w.len()).sum::<usize>() as f32 / words.len() as f32
};
let special_chars = query.chars().any(|c| "{}[]().,;:!@#$%^&*".contains(c));
ConformalFeatures {
query_length: query.len() as u32,
word_count: words.len() as u32,
has_special_chars: special_chars,
fuzzy_enabled: false, structural_mode: matches!(classification.intent, QueryIntent::Structural),
avg_word_length,
query_entropy,
identifier_density,
semantic_complexity: classification.complexity_score,
has_file_context: file_context.is_some(),
language_detected: !classification.language_hints.is_empty(),
intent_confidence: classification.confidence,
naturalness_score: classification.naturalness_score,
similar_queries_success_rate: 0.7, user_satisfaction_history: 0.8, }
}
pub async fn initialize_conformal_router(config: &ConformalRouterConfig) -> Result<()> {
tracing::info!("Initializing conformal router module");
tracing::info!("Risk threshold: {}", config.risk_threshold);
tracing::info!("Daily budget: {}%", config.daily_budget_percent);
tracing::info!("Confidence level: {}", config.confidence_level);
tracing::info!("Min calibration samples: {}", config.min_calibration_samples);
tracing::info!("P95 headroom threshold: {}ms", config.p95_headroom_threshold_ms);
tracing::info!("Enabled: {}", config.enabled);
if config.risk_threshold < 0.0 || config.risk_threshold > 1.0 {
anyhow::bail!("Risk threshold must be in range [0.0, 1.0]");
}
if config.daily_budget_percent < 0.0 || config.daily_budget_percent > 100.0 {
anyhow::bail!("Daily budget percent must be in range [0.0, 100.0]");
}
if config.confidence_level < 0.0 || config.confidence_level > 1.0 {
anyhow::bail!("Confidence level must be in range [0.0, 1.0]");
}
if config.min_calibration_samples == 0 {
anyhow::bail!("Minimum calibration samples must be greater than 0");
}
tracing::info!("Conformal router module initialized successfully");
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::semantic::query_classifier::{QueryIntent, QueryCharacteristic, ClassifierConfig, QueryClassifier};
use smallvec::smallvec;
#[tokio::test]
async fn test_conformal_router_creation() {
let config = ConformalRouterConfig::default();
let router = ConformalRouter::new(config);
let status = router.get_status().await;
assert!(status.enabled);
assert_eq!(status.risk_threshold, 0.6);
}
#[tokio::test]
async fn test_budget_manager() {
let mut budget = BudgetManager::new(5.0);
assert!(budget.can_upshift(0.02)); budget.record_upshift(0.02);
assert!(budget.can_upshift(0.02)); budget.record_upshift(0.02);
assert!(!budget.can_upshift(0.02));
let status = budget.get_status();
assert!(status.current_usage_percent > 75.0); }
#[tokio::test]
async fn test_conformal_predictor() {
let mut predictor = ConformalPredictor::new(5, 24);
for i in 0..10 {
let features = ConformalFeatures {
query_length: 50 + i,
word_count: 5 + (i / 2),
has_special_chars: i % 2 == 0,
fuzzy_enabled: false,
structural_mode: false,
avg_word_length: 5.0,
query_entropy: 2.5,
identifier_density: 0.3,
semantic_complexity: 0.4,
has_file_context: true,
language_detected: true,
intent_confidence: 0.8,
naturalness_score: 0.6,
similar_queries_success_rate: 0.7,
user_satisfaction_history: 0.8,
};
predictor.add_calibration_sample(CalibrationSample {
features,
predicted_quality: 0.7,
actual_quality: 0.75 + (i as f32 * 0.01),
timestamp: SystemTime::now(),
});
}
let status = predictor.get_status();
assert!(status.is_calibrated);
assert!(status.calibration_samples >= 5);
let test_features = ConformalFeatures {
query_length: 60,
word_count: 6,
has_special_chars: true,
fuzzy_enabled: false,
structural_mode: false,
avg_word_length: 5.5,
query_entropy: 2.8,
identifier_density: 0.4,
semantic_complexity: 0.6,
has_file_context: true,
language_detected: true,
intent_confidence: 0.9,
naturalness_score: 0.7,
similar_queries_success_rate: 0.8,
user_satisfaction_history: 0.85,
};
let risk = predictor.predict_risk(&test_features, 0.95);
assert!(risk.risk_score >= 0.0 && risk.risk_score <= 1.0);
assert!(risk.confidence_interval.0 <= risk.confidence_interval.1);
assert!(risk.calibrated);
}
#[tokio::test]
async fn test_routing_decision() {
let config = ConformalRouterConfig::default();
let router = ConformalRouter::new(config);
let features = ConformalFeatures {
query_length: 100,
word_count: 8,
has_special_chars: true,
fuzzy_enabled: false,
structural_mode: false,
avg_word_length: 6.0,
query_entropy: 3.2,
identifier_density: 0.2,
semantic_complexity: 0.8, has_file_context: false,
language_detected: true,
intent_confidence: 0.6,
naturalness_score: 0.9, similar_queries_success_rate: 0.5,
user_satisfaction_history: 0.6,
};
let classification = crate::semantic::query_classifier::QueryClassification {
intent: QueryIntent::NaturalLanguage,
confidence: 0.8,
characteristics: vec![QueryCharacteristic::HasDescriptiveWords],
naturalness_score: 0.9,
complexity_score: 0.8,
language_hints: vec!["english".to_string()],
};
let decision = router.make_routing_decision(&features, &classification).await;
assert!(decision.is_ok());
let decision = decision.unwrap();
assert!(decision.risk_assessment.risk_score > 0.0);
}
#[test]
fn test_feature_extraction() {
let query = "how to find a function that calculates the sum of two numbers";
let classification = crate::semantic::query_classifier::QueryClassification {
intent: QueryIntent::NaturalLanguage,
confidence: 0.9,
characteristics: vec![
QueryCharacteristic::HasDescriptiveWords,
QueryCharacteristic::HasArticles
],
naturalness_score: 0.95,
complexity_score: 0.3,
language_hints: vec!["english".to_string()],
};
let features = extract_conformal_features(&query, &classification, None);
assert_eq!(features.query_length, query.len() as u32);
assert_eq!(features.word_count, 12); assert!(!features.has_special_chars);
assert!(features.naturalness_score > 0.9);
assert!(!features.has_file_context);
assert!(features.language_detected);
}
#[tokio::test]
async fn test_upshift_type_selection() {
let config = ConformalRouterConfig::default();
let router = ConformalRouter::new(config);
let features = ConformalFeatures {
semantic_complexity: 0.9,
naturalness_score: 0.8,
..Default::default()
};
let risk_assessment = RiskAssessment {
risk_score: 0.7,
confidence_interval: (0.5, 0.9),
nonconformity_score: 0.2,
calibrated: true,
risk_factors: vec![],
};
let classification = crate::semantic::query_classifier::QueryClassification {
intent: QueryIntent::NaturalLanguage,
confidence: 0.8,
characteristics: vec![],
naturalness_score: 0.8,
complexity_score: 0.9,
language_hints: vec![],
};
let upshift_type = router.select_upshift_type(&features, &risk_assessment, &classification);
assert_eq!(upshift_type, UpshiftType::HighDimEmbeddings);
}
#[tokio::test]
async fn test_configuration_validation() {
let mut config = ConformalRouterConfig::default();
config.risk_threshold = 1.5;
let result = initialize_conformal_router(&config).await;
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("Risk threshold"));
config.risk_threshold = 0.6; config.daily_budget_percent = -5.0;
let result = initialize_conformal_router(&config).await;
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("Daily budget"));
}
}
impl Default for ConformalFeatures {
fn default() -> Self {
Self {
query_length: 0,
word_count: 0,
has_special_chars: false,
fuzzy_enabled: false,
structural_mode: false,
avg_word_length: 0.0,
query_entropy: 0.0,
identifier_density: 0.0,
semantic_complexity: 0.0,
has_file_context: false,
language_detected: false,
intent_confidence: 0.0,
naturalness_score: 0.0,
similar_queries_success_rate: 0.0,
user_satisfaction_history: 0.0,
}
}
}