xz-rerank 0.1.1

检索结果重排序 — 本地多信号融合 + 远程 Rerank API 适配
Documentation
use async_trait::async_trait;

use crate::error::RerankError;
use crate::traits::SignalPlugin;
use crate::types::RerankCandidate;

/// 向量相似度信号
///
/// 需要 query embedding 由调用方在 rerank 前计算。
#[derive(Debug)]
pub struct VectorSimilaritySignal {
    /// query embedding(由调用方在 rerank 前计算)
    query_embedding: Option<Vec<f32>>,
}

impl VectorSimilaritySignal {
    /// 创建新的向量相似度信号(不含 query embedding)
    ///
    /// 调用方需要在 rerank 前通过 `with_query_embedding` 设置 query embedding,
    /// 否则 scoring 时会返回 [`RerankError::MissingQueryEmbedding`]。
    pub fn new() -> Self {
        Self { query_embedding: None }
    }

    /// 创建带有 query embedding 的向量相似度信号
    pub fn with_query_embedding(embedding: Vec<f32>) -> Self {
        Self { query_embedding: Some(embedding) }
    }
}

impl Default for VectorSimilaritySignal {
    /// 返回默认的向量相似度信号(不含 query embedding)
    fn default() -> Self {
        Self::new()
    }
}

#[async_trait]
impl SignalPlugin for VectorSimilaritySignal {
    fn name(&self) -> &str {
        "vector_similarity"
    }

    fn weight_key(&self) -> &'static str {
        "vector_similarity"
    }

    async fn score(&self, _query: &str, candidate: &RerankCandidate) -> Result<f32, RerankError> {
        let q = match self.query_embedding.as_ref() {
            Some(q) => q,
            None => return Err(RerankError::MissingQueryEmbedding),
        };
        let empty = vec![];
        let c = candidate.embedding.as_ref().unwrap_or(&empty);

        if q.is_empty() || c.is_empty() {
            return Ok(0.0);
        }

        let dot: f32 = q.iter().zip(c).map(|(a, b)| a * b).sum();
        let q_norm: f32 = q.iter().map(|x| x * x).sum::<f32>().sqrt();
        let c_norm: f32 = c.iter().map(|x| x * x).sum::<f32>().sqrt();

        if q_norm == 0.0 || c_norm == 0.0 {
            return Ok(0.0);
        }

        let similarity = dot / (q_norm * c_norm);
        Ok((similarity + 1.0) / 2.0) // [-1, 1] → [0, 1]
    }
}

#[async_trait]
impl SignalPlugin for &VectorSimilaritySignal {
    fn name(&self) -> &str {
        "vector_similarity"
    }

    fn weight_key(&self) -> &'static str {
        "vector_similarity"
    }

    async fn score(&self, _query: &str, candidate: &RerankCandidate) -> Result<f32, RerankError> {
        VectorSimilaritySignal::score(self, _query, candidate).await
    }
}