use async_trait::async_trait;
use std::collections::HashMap;
use std::time::Instant;
use crate::error::RerankError;
use crate::signals::{
ContentQualitySignal, KeywordOverlapSignal, MetadataMatchSignal, RecencySignal,
VectorSimilaritySignal,
};
use crate::traits::{Reranker, RerankerBackendType, RerankerInfo, RerankerPricing, SignalPlugin};
use crate::types::{
ChannelRecencyRule, RecencyMode, RerankCandidate, RerankConfig, RerankHit, RerankResult,
RerankStats, ScoreBreakdown, SignalScore, SignalWeights,
};
pub struct LocalSignalReranker {
weights: SignalWeights,
signals: Vec<Box<dyn SignalPlugin>>,
default_recency_mode: RecencyMode,
channel_recency: Vec<ChannelRecencyRule>,
info: RerankerInfo,
}
impl std::fmt::Debug for LocalSignalReranker {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("LocalSignalReranker")
.field("weights", &self.weights)
.field(
"signal_names",
&self.signals.iter().map(|s| s.name().to_string()).collect::<Vec<_>>(),
)
.field("default_recency_mode", &self.default_recency_mode)
.finish()
}
}
impl LocalSignalReranker {
pub fn new(weights: SignalWeights) -> Self {
Self {
weights,
signals: Vec::new(),
default_recency_mode: RecencyMode::ExponentialDecay { decay_rate: 0.01 },
channel_recency: Vec::new(),
info: RerankerInfo {
name: "local-signal".into(),
display_name: "Local Signal Reranker".into(),
backend_type: RerankerBackendType::Local,
supports_batch: true,
max_candidates: None,
pricing: Some(RerankerPricing { cost_per_search: 0.0 }),
},
}
}
pub fn with_signal(mut self, signal: Box<dyn SignalPlugin>) -> Self {
self.signals.push(signal);
self
}
pub fn with_channel_recency(mut self, rules: Vec<ChannelRecencyRule>) -> Self {
self.channel_recency = rules;
self
}
pub fn with_recency_mode(mut self, mode: RecencyMode) -> Self {
self.default_recency_mode = mode;
self
}
fn get_recency_mode_for(&self, channel: Option<&String>) -> RecencyMode {
if let Some(ch) = channel {
for rule in &self.channel_recency {
if rule.channel == *ch {
return rule.mode.clone();
}
}
}
self.default_recency_mode.clone()
}
async fn compute_scores(
&self,
query: &str,
candidates: &[RerankCandidate],
query_embedding_override: Option<&Vec<f32>>,
) -> Result<(Vec<Vec<f32>>, HashMap<String, u64>), RerankError> {
use futures::future::join_all;
let futures: Vec<_> = self
.signals
.iter()
.map(|signal| {
let signal_name = signal.name().to_string();
let q = query.to_string();
let c = candidates.to_vec();
let qe: Option<Vec<f32>> = query_embedding_override.cloned();
let is_vs = signal_name == "vector_similarity";
async move {
let start = Instant::now();
let result = if is_vs {
if let Some(ref emb) = qe {
let emb: Vec<f32> = emb.clone();
let vs = VectorSimilaritySignal::with_query_embedding(emb);
vs.score_batch(&q, &c).await
} else {
signal.score_batch(&q, &c).await
}
} else {
signal.score_batch(&q, &c).await
};
let elapsed = start.elapsed().as_micros() as u64;
(signal_name, result, elapsed)
}
})
.collect();
let results = join_all(futures).await;
let mut signal_scores = Vec::with_capacity(results.len());
let mut timings = HashMap::new();
for (name, result, elapsed) in results {
signal_scores.push(result?);
timings.insert(name, elapsed);
}
Ok((signal_scores, timings))
}
fn weighted_sum(&self, signal_scores: &[Vec<f32>], weights: &SignalWeights) -> Vec<f32> {
let n = signal_scores.first().map(|s| s.len()).unwrap_or(0);
let mut final_scores = vec![0.0f32; n];
for (signal_idx, signal_score_vec) in signal_scores.iter().enumerate() {
let weight_key = self.signals.get(signal_idx).map(|s| s.weight_key()).unwrap_or("");
let w = weights.get_weight_by_name(weight_key);
for (i, &s) in signal_score_vec.iter().enumerate() {
final_scores[i] += s * w;
}
}
final_scores
}
}
impl Default for LocalSignalReranker {
fn default() -> Self {
let weights = SignalWeights::default();
Self {
weights: weights.clone(),
signals: vec![
Box::new(KeywordOverlapSignal),
Box::new(VectorSimilaritySignal::new()),
Box::new(MetadataMatchSignal::default()),
Box::new(ContentQualitySignal),
Box::new(RecencySignal::new(RecencyMode::ExponentialDecay { decay_rate: 0.01 })),
],
default_recency_mode: RecencyMode::ExponentialDecay { decay_rate: 0.01 },
channel_recency: Vec::new(),
info: RerankerInfo {
name: "local-signal".into(),
display_name: "Local Signal Reranker".into(),
backend_type: RerankerBackendType::Local,
supports_batch: true,
max_candidates: None,
pricing: Some(RerankerPricing { cost_per_search: 0.0 }),
},
}
}
}
#[async_trait]
impl Reranker for LocalSignalReranker {
async fn rerank(
&self,
query: &str,
candidates: Vec<RerankCandidate>,
config: &RerankConfig,
) -> Result<RerankResult, RerankError> {
let start = Instant::now();
if candidates.is_empty() {
return Err(RerankError::EmptyCandidates);
}
let total_candidates = candidates.len();
let (mut signal_scores, mut signal_timings) =
self.compute_scores(query, &candidates, config.query_embedding.as_ref()).await?;
if let Some(recency_idx) = self.signals.iter().position(|s| s.name() == "recency") {
if let Some(ref mode) = config.recency_mode {
let recency_signal = RecencySignal::new(mode.clone());
let recency_start = Instant::now();
signal_scores[recency_idx] = recency_signal.score_batch(query, &candidates).await?;
signal_timings
.insert("recency".to_string(), recency_start.elapsed().as_micros() as u64);
} else {
let now_ms = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_millis() as u64;
let recency_start = Instant::now();
let mut recency_scores = Vec::with_capacity(candidates.len());
for candidate in &candidates {
let mode = self.get_recency_mode_for(candidate.channel.as_ref());
let signal = RecencySignal::new(mode).with_now(now_ms);
recency_scores.push(signal.score(query, candidate).await?);
}
signal_scores[recency_idx] = recency_scores;
signal_timings
.insert("recency".to_string(), recency_start.elapsed().as_micros() as u64);
}
}
let final_scores = self.weighted_sum(&signal_scores, &self.weights);
let mut scored: Vec<(usize, f32)> =
final_scores.iter().enumerate().map(|(i, &s)| (i, s)).collect();
scored.sort_by(|a, b| b.1.total_cmp(&a.1));
let min_score = config.min_score.unwrap_or(0.0);
let mut filtered_out = 0usize;
let mut hits = Vec::new();
for (idx, score) in scored {
if score < min_score {
filtered_out += 1;
continue;
}
if hits.len() >= config.top_k {
continue;
}
let candidate = &candidates[idx];
let score_breakdown = if config.include_score_breakdown {
let signals: Vec<SignalScore> = signal_scores
.iter()
.enumerate()
.map(|(si, scores)| {
let raw_score = scores[idx];
let weight = self
.signals
.get(si)
.map(|s| self.weights.get_weight_by_name(s.weight_key()))
.unwrap_or(0.0);
SignalScore {
name: self.get_signal_name(si),
raw_score,
weight,
contribution: raw_score * weight,
}
})
.collect();
Some(ScoreBreakdown { signals, final_score: score })
} else {
None
};
hits.push(RerankHit {
candidate_id: candidate.id.clone(),
score,
score_breakdown,
candidate: candidate.clone(),
});
}
let final_scores_vec: Vec<f32> = hits.iter().map(|h| h.score).collect();
let n = final_scores_vec.len();
let max_score = final_scores_vec.first().copied().unwrap_or(0.0);
let min_score_final = final_scores_vec.last().copied().unwrap_or(0.0);
let avg_score = if n > 0 { final_scores_vec.iter().sum::<f32>() / n as f32 } else { 0.0 };
let mut sorted = final_scores_vec.clone();
sorted.sort_by(|a, b| a.total_cmp(b));
let median_score = if n > 0 { sorted[n / 2] } else { 0.0 };
Ok(RerankResult {
hits,
stats: RerankStats {
total_candidates,
filtered_out,
max_score,
min_score: min_score_final,
avg_score,
median_score,
signal_timings,
},
reranker: self.info.name.clone(),
latency_ms: start.elapsed().as_millis() as u64,
})
}
fn reranker_info(&self) -> &RerankerInfo {
&self.info
}
}
impl LocalSignalReranker {
fn get_signal_name(&self, idx: usize) -> String {
self.signals
.get(idx)
.map(|s| s.name().to_string())
.unwrap_or_else(|| format!("signal_{idx}"))
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashMap;
#[test]
fn test_sort_with_nan_scores_descending() {
let mut scored: Vec<(usize, f32)> =
vec![(0, f32::NAN), (1, 0.5), (2, f32::NAN), (3, -1.0), (4, f32::NAN)];
scored.sort_by(|a, b| b.1.total_cmp(&a.1));
assert_eq!(scored.len(), 5);
assert!(scored[0].1.is_nan());
}
#[test]
fn test_sort_with_nan_scores_ascending() {
let mut sorted: Vec<f32> = vec![f32::NAN, 0.5, f32::NAN, -1.0, f32::NAN];
sorted.sort_by(|a, b| a.total_cmp(b));
assert_eq!(sorted.len(), 5);
}
#[tokio::test]
async fn test_recency_no_decay_override_equal_scores() {
let weights = SignalWeights {
keyword_overlap: 0.0,
vector_similarity: 0.0,
metadata_match: 0.0,
content_quality: 0.0,
recency: 1.0,
custom_weights: std::collections::HashMap::new(),
};
let reranker = LocalSignalReranker::new(weights).with_signal(Box::new(RecencySignal::new(
RecencyMode::ExponentialDecay { decay_rate: 0.01 },
)));
let now_ms = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.ok()
.map(|d| d.as_millis() as u64)
.unwrap_or(0);
let recent = RerankCandidate {
id: "recent".into(),
content: "doc".into(),
metadata: HashMap::new(),
retrieval_score: None,
channel: None,
created_at: Some(now_ms - 60_000), embedding: None,
};
let old = RerankCandidate {
id: "old".into(),
content: "doc".into(),
metadata: HashMap::new(),
retrieval_score: None,
channel: None,
created_at: Some(now_ms - 86_400_000), embedding: None,
};
let config = RerankConfig {
top_k: 10,
min_score: None,
include_score_breakdown: false,
recency_mode: Some(RecencyMode::NoDecay),
query_embedding: None,
};
let result = reranker.rerank("test", vec![recent, old], &config).await;
let result = match result {
Ok(r) => r,
Err(e) => panic!("rerank failed: {e}"),
};
assert_eq!(result.hits.len(), 2, "should return both candidates");
let scores: Vec<f32> = result.hits.iter().map(|h| h.score).collect();
let diff = (scores[0] - scores[1]).abs();
assert!(
diff < 1e-6,
"with NoDecay override, recent and old docs must have equal scores; got diff={diff}, scores={scores:?}"
);
}
}