use async_trait::async_trait;
use crate::error::RerankError;
use crate::traits::Reranker;
use crate::types::{RerankCandidate, RerankConfig, RerankResult};
pub struct MultiStageReranker<S1: Reranker, S2: Reranker> {
stage1: S1,
stage2: S2,
top_k_after_stage1: usize,
}
impl<S1: Reranker, S2: Reranker> MultiStageReranker<S1, S2> {
pub fn new(stage1: S1, stage2: S2, top_k_after_stage1: usize) -> Self {
Self { stage1, stage2, top_k_after_stage1 }
}
}
impl<S1: Reranker, S2: Reranker> std::fmt::Debug for MultiStageReranker<S1, S2> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MultiStageReranker")
.field("stage1", &self.stage1.reranker_info().name)
.field("stage2", &self.stage2.reranker_info().name)
.field("top_k_after_stage1", &self.top_k_after_stage1)
.finish()
}
}
#[async_trait]
impl<S1: Reranker, S2: Reranker> Reranker for MultiStageReranker<S1, S2> {
async fn rerank(
&self,
query: &str,
candidates: Vec<RerankCandidate>,
config: &RerankConfig,
) -> Result<RerankResult, RerankError> {
let stage1_config = RerankConfig { top_k: self.top_k_after_stage1, ..config.clone() };
let stage1_result = self.stage1.rerank(query, candidates, &stage1_config).await?;
let top_candidates: Vec<RerankCandidate> = stage1_result
.hits
.into_iter()
.take(self.top_k_after_stage1)
.map(|h| h.candidate)
.collect();
self.stage2.rerank(query, top_candidates, config).await
}
fn reranker_info(&self) -> &crate::traits::RerankerInfo {
self.stage2.reranker_info()
}
}