Skip to main content

ferrin_core/
rerank.rs

1//! Reranking: [`rerank`] orders documents by relevance to a query.
2//!
3//! Design: `docs/01-architecture/11-other-modalities.md` ยง5.
4
5use std::future::IntoFuture;
6use std::time::Instant;
7
8use chrono::Utc;
9use ferrin_spec::BoxFuture;
10use ferrin_spec::JsonObject;
11use ferrin_spec::ProviderMetadata;
12use ferrin_spec::RerankingModelRef;
13use ferrin_spec::ResponseMetadata;
14use ferrin_spec::Warning;
15use ferrin_spec::error::InvalidResponseDataError;
16use ferrin_spec::error::ProviderError;
17use ferrin_spec::reranking_model::RerankDocuments;
18use ferrin_spec::reranking_model::RerankOptions;
19use serde_json::json;
20use tracing::Instrument;
21
22use crate::error::Error;
23use crate::ids::default_id_generator;
24use crate::modality::ModalityOptions;
25use crate::modality::impl_modality_builder;
26use crate::registry::ProviderRegistry;
27use crate::registry::default::resolve_model;
28use crate::retry::retry;
29use crate::telemetry::ErrorEvent;
30use crate::telemetry::ErrorPhase;
31use crate::telemetry::ModelIdentity;
32use crate::telemetry::RerankEndEvent;
33use crate::telemetry::RerankStartEvent;
34use crate::telemetry::dispatcher::TelemetryDispatcher;
35use crate::telemetry::spans;
36
37/// A document to rerank: text or a JSON object.
38#[derive(Debug, Clone, PartialEq)]
39pub enum RerankDocument {
40    /// Plain text.
41    Text(String),
42    /// A structured document.
43    Object(JsonObject),
44}
45
46impl From<String> for RerankDocument {
47    fn from(text: String) -> Self {
48        Self::Text(text)
49    }
50}
51
52impl From<&str> for RerankDocument {
53    fn from(text: &str) -> Self {
54        Self::Text(text.to_owned())
55    }
56}
57
58impl From<JsonObject> for RerankDocument {
59    fn from(object: JsonObject) -> Self {
60        Self::Object(object)
61    }
62}
63
64/// One ranked document.
65#[derive(Debug, Clone, PartialEq)]
66pub struct Ranked<D> {
67    /// Index of the document in the input list.
68    pub original_index: usize,
69    /// Relevance score reported by the model.
70    pub score: f64,
71    /// The document.
72    pub document: D,
73}
74
75/// Result of [`rerank`].
76#[derive(Debug, Clone, PartialEq)]
77pub struct RerankResult<D> {
78    /// Documents in the order returned by the model (most relevant first).
79    pub ranking: Vec<Ranked<D>>,
80    /// Adapter warnings.
81    pub warnings: Vec<Warning>,
82    /// Response metadata.
83    pub response: ResponseMetadata,
84    /// Provider-specific metadata.
85    pub provider_metadata: Option<ProviderMetadata>,
86}
87
88impl<D> RerankResult<D> {
89    /// The documents in ranked order.
90    pub fn reranked_documents(&self) -> impl Iterator<Item = &D> + '_ {
91        self.ranking.iter().map(|ranked| &ranked.document)
92    }
93}
94
95/// Reranks `documents` by relevance to `query`. Documents must be all
96/// text or all objects.
97#[must_use]
98pub fn rerank<D>(
99    model: impl Into<RerankingModelRef>,
100    query: impl Into<String>,
101    documents: Vec<D>,
102) -> Rerank<D>
103where
104    D: Into<RerankDocument> + Clone + Send + 'static,
105{
106    Rerank {
107        model: model.into(),
108        query: query.into(),
109        documents,
110        top_n: None,
111        base: ModalityOptions::default(),
112    }
113}
114
115/// Builder returned by [`rerank`]; `.await` runs the call.
116#[derive(Debug)]
117pub struct Rerank<D> {
118    model: RerankingModelRef,
119    query: String,
120    documents: Vec<D>,
121    top_n: Option<usize>,
122    base: ModalityOptions,
123}
124
125impl<D> Rerank<D> {
126    /// Returns only the `top_n` most relevant documents.
127    #[must_use]
128    pub fn top_n(mut self, top_n: usize) -> Self {
129        self.top_n = Some(top_n);
130        self
131    }
132}
133
134impl_modality_builder!(Rerank<D>);
135
136impl<D> IntoFuture for Rerank<D>
137where
138    D: Into<RerankDocument> + Clone + Send + 'static,
139{
140    type Output = Result<RerankResult<D>, Error>;
141    type IntoFuture = BoxFuture<'static, Self::Output>;
142
143    fn into_future(self) -> Self::IntoFuture {
144        Box::pin(run(self))
145    }
146}
147
148/// Converts the documents to the provider representation.
149fn to_model_documents<D>(documents: &[D]) -> Result<RerankDocuments, Error>
150where
151    D: Into<RerankDocument> + Clone,
152{
153    let mut texts: Vec<String> = Vec::new();
154    let mut objects: Vec<JsonObject> = Vec::new();
155    for document in documents {
156        match document.clone().into() {
157            RerankDocument::Text(text) => texts.push(text),
158            RerankDocument::Object(object) => objects.push(object),
159        }
160    }
161    match (texts.is_empty(), objects.is_empty()) {
162        (false, true) => Ok(RerankDocuments::Text { values: texts }),
163        (true, false) => Ok(RerankDocuments::Object { values: objects }),
164        _ => Err(Error::invalid_argument(
165            "documents",
166            "documents must be all text or all objects",
167        )),
168    }
169}
170
171async fn run<D>(builder: Rerank<D>) -> Result<RerankResult<D>, Error>
172where
173    D: Into<RerankDocument> + Clone + Send + 'static,
174{
175    let model = resolve_model(&builder.model, ProviderRegistry::reranking_model)?;
176    let identity = ModelIdentity::new(model.provider().clone(), model.model_id().clone());
177    let span = spans::modality_span("rerank", &identity);
178    let base = builder.base.clone();
179    let telemetry = TelemetryDispatcher::new(base.telemetry.clone());
180    let call_id = default_id_generator().generate();
181    let Rerank {
182        query,
183        documents,
184        top_n,
185        ..
186    } = builder;
187
188    if documents.is_empty() {
189        telemetry.on_rerank_start(&RerankStartEvent {
190            call_id: call_id.clone(),
191            model: identity.clone(),
192            document_count: 0,
193            query: telemetry.record_inputs().then(|| query.clone()),
194        });
195        telemetry.on_rerank_end(&RerankEndEvent {
196            call_id,
197            ranked_count: 0,
198            duration: std::time::Duration::ZERO,
199        });
200        return Ok(RerankResult {
201            ranking: Vec::new(),
202            warnings: Vec::new(),
203            response: ResponseMetadata {
204                timestamp: Some(Utc::now()),
205                model_id: Some(identity.model_id.clone()),
206                ..ResponseMetadata::default()
207            },
208            provider_metadata: None,
209        });
210    }
211
212    let model_documents = to_model_documents(&documents)?;
213    base.run(|base, token| {
214        async move {
215            let headers = base.request_headers();
216            let outcome = retry(&base.retry_policy, &token, |_| {
217                let options = RerankOptions {
218                    query: query.clone(),
219                    documents: model_documents.clone(),
220                    top_n,
221                    provider_options: base.provider_options.clone(),
222                    headers: headers.clone(),
223                    cancellation: token.child_token(),
224                };
225                let model = &model;
226                let telemetry = &telemetry;
227                let call_id = &call_id;
228                let identity = &identity;
229                async move {
230                    let started = Instant::now();
231                    telemetry.on_rerank_start(&RerankStartEvent {
232                        call_id: call_id.clone(),
233                        model: identity.clone(),
234                        document_count: options.documents.len(),
235                        query: telemetry.record_inputs().then(|| options.query.clone()),
236                    });
237                    let result = model.do_rerank(options).await.map_err(Error::from)?;
238                    telemetry.on_rerank_end(&RerankEndEvent {
239                        call_id: call_id.clone(),
240                        ranked_count: result.ranking.len(),
241                        duration: started.elapsed(),
242                    });
243                    Ok(result)
244                }
245            })
246            .await;
247            let result = match outcome {
248                Ok(result) => result,
249                Err(error) => {
250                    telemetry.on_error(&ErrorEvent {
251                        call_id: &call_id,
252                        error: &error,
253                        phase: ErrorPhase::ModelCall,
254                    });
255                    return Err(error);
256                }
257            };
258            spans::log_warnings(&result.warnings, &identity);
259            let mut ranking: Vec<Ranked<D>> = Vec::with_capacity(result.ranking.len());
260            for ranked in result.ranking {
261                let Some(document) = documents.get(ranked.index) else {
262                    return Err(Error::from(ProviderError::InvalidResponseData(Box::new(
263                        InvalidResponseDataError::new(
264                            format!(
265                                "ranking index {} is out of range for {} documents",
266                                ranked.index,
267                                documents.len()
268                            ),
269                            json!({ "index": ranked.index, "documents": documents.len() }),
270                        ),
271                    ))));
272                };
273                ranking.push(Ranked {
274                    original_index: ranked.index,
275                    score: ranked.relevance_score,
276                    document: document.clone(),
277                });
278            }
279            let mut response = result.response;
280            if response.timestamp.is_none() {
281                response.timestamp = Some(Utc::now());
282            }
283            if response.model_id.is_none() {
284                response.model_id = Some(identity.model_id.clone());
285            }
286            Ok(RerankResult {
287                ranking,
288                warnings: result.warnings,
289                response,
290                provider_metadata: result.provider_metadata,
291            })
292        }
293        .instrument(span)
294    })
295    .await
296}