use std::collections::HashMap;
use serde::{Deserialize, Serialize};
use crate::error::RerankError;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SignalWeights {
pub keyword_overlap: f32,
pub vector_similarity: f32,
pub metadata_match: f32,
pub content_quality: f32,
pub recency: f32,
#[serde(default)]
pub custom_weights: HashMap<String, f32>,
}
impl Default for SignalWeights {
fn default() -> Self {
Self {
keyword_overlap: 0.30,
vector_similarity: 0.25,
metadata_match: 0.20,
content_quality: 0.10,
recency: 0.15,
custom_weights: HashMap::new(),
}
}
}
impl SignalWeights {
pub fn validate(&self) -> Result<(), RerankError> {
let sum = self.keyword_overlap
+ self.vector_similarity
+ self.metadata_match
+ self.content_quality
+ self.recency
+ self.custom_weights.values().sum::<f32>();
if (sum - 1.0).abs() > 0.01 {
return Err(RerankError::WeightSumInvalid(sum));
}
Ok(())
}
pub fn get_weight_by_name(&self, name: &str) -> f32 {
match name {
"keyword_overlap" => self.keyword_overlap,
"vector_similarity" => self.vector_similarity,
"metadata_match" => self.metadata_match,
"content_quality" => self.content_quality,
"recency" => self.recency,
_ => self.custom_weights.get(name).copied().unwrap_or(0.0),
}
}
pub fn with_custom_weight(mut self, key: &str, weight: f32) -> Self {
self.custom_weights.insert(key.to_string(), weight);
self
}
pub fn normalize(&self) -> Self {
let sum = self.keyword_overlap
+ self.vector_similarity
+ self.metadata_match
+ self.content_quality
+ self.recency
+ self.custom_weights.values().sum::<f32>();
if sum == 0.0 {
return Self::default();
}
let custom_weights =
self.custom_weights.iter().map(|(k, v)| (k.clone(), v / sum)).collect();
Self {
keyword_overlap: self.keyword_overlap / sum,
vector_similarity: self.vector_similarity / sum,
metadata_match: self.metadata_match / sum,
content_quality: self.content_quality / sum,
recency: self.recency / sum,
custom_weights,
}
}
}