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//!
5//! Lifecycle behavior is derived from the Vercel AI SDK (Apache-2.0,
6//! Copyright 2023 Vercel, Inc.), translated to Rust and modified; see NOTICE.
7
8use std::future::IntoFuture;
9use std::sync::Arc;
10use std::time::Instant;
11
12use chrono::Utc;
13use ferrin_spec::BoxFuture;
14use ferrin_spec::JsonObject;
15use ferrin_spec::ProviderMetadata;
16use ferrin_spec::RerankingModelRef;
17use ferrin_spec::ResponseMetadata;
18use ferrin_spec::Warning;
19use ferrin_spec::error::InvalidResponseDataError;
20use ferrin_spec::error::ProviderError;
21use ferrin_spec::reranking_model::RerankDocuments;
22use ferrin_spec::reranking_model::RerankOptions;
23use serde_json::json;
24use tracing::Instrument;
25
26use crate::error::Error;
27use crate::hooks::Hooks;
28use crate::ids::default_id_generator;
29use crate::modality::ModalityOptions;
30use crate::modality::impl_modality_builder;
31use crate::modality_hooks::ModalityHooks;
32pub use crate::modality_hooks::RerankCallEndEvent;
33pub use crate::modality_hooks::RerankCallStartEvent;
34use crate::modality_hooks::impl_modality_hooks;
35use crate::registry::ProviderRegistry;
36use crate::registry::default::resolve_model;
37use crate::retry::retry;
38use crate::telemetry::ErrorEvent;
39use crate::telemetry::ErrorPhase;
40use crate::telemetry::ModelIdentity;
41use crate::telemetry::RerankEndEvent;
42use crate::telemetry::RerankStartEvent;
43use crate::telemetry::dispatcher::TelemetryDispatcher;
44use crate::telemetry::spans;
45
46/// A document to rerank: text or a JSON object.
47#[derive(Debug, Clone, PartialEq)]
48pub enum RerankDocument {
49    /// Plain text.
50    Text(String),
51    /// A structured document.
52    Object(JsonObject),
53}
54
55impl From<String> for RerankDocument {
56    fn from(text: String) -> Self {
57        Self::Text(text)
58    }
59}
60
61impl From<&str> for RerankDocument {
62    fn from(text: &str) -> Self {
63        Self::Text(text.to_owned())
64    }
65}
66
67impl From<JsonObject> for RerankDocument {
68    fn from(object: JsonObject) -> Self {
69        Self::Object(object)
70    }
71}
72
73/// One ranked document.
74#[derive(Debug, Clone, PartialEq)]
75pub struct Ranked<D> {
76    /// Index of the document in the input list.
77    pub original_index: usize,
78    /// Relevance score reported by the model.
79    pub score: f64,
80    /// The document.
81    pub document: D,
82}
83
84/// Result of [`rerank`].
85#[derive(Debug, Clone, PartialEq)]
86pub struct RerankResult<D> {
87    /// The complete original document list, in input order.
88    pub original_documents: Vec<D>,
89    /// Documents in the order returned by the model (most relevant first).
90    pub ranking: Vec<Ranked<D>>,
91    /// Adapter warnings.
92    pub warnings: Vec<Warning>,
93    /// Response metadata.
94    pub response: ResponseMetadata,
95    /// Provider-specific metadata.
96    pub provider_metadata: Option<ProviderMetadata>,
97}
98
99impl<D> RerankResult<D> {
100    /// The documents in ranked order.
101    pub fn reranked_documents(&self) -> impl Iterator<Item = &D> + '_ {
102        self.ranking.iter().map(|ranked| &ranked.document)
103    }
104}
105
106/// Reranks `documents` by relevance to `query`. Documents must be all
107/// text or all objects.
108#[must_use]
109pub fn rerank<D>(
110    model: impl Into<RerankingModelRef>,
111    query: impl Into<String>,
112    documents: Vec<D>,
113) -> Rerank<D>
114where
115    D: Into<RerankDocument> + Clone + Send + 'static,
116{
117    Rerank {
118        model: model.into(),
119        query: query.into(),
120        documents,
121        top_n: None,
122        base: ModalityOptions::default(),
123        hooks: ModalityHooks::default(),
124    }
125}
126
127/// Builder returned by [`rerank`]; `.await` runs the call.
128#[derive(Debug)]
129pub struct Rerank<D> {
130    model: RerankingModelRef,
131    query: String,
132    documents: Vec<D>,
133    top_n: Option<usize>,
134    base: ModalityOptions,
135    hooks: ModalityHooks<RerankCallStartEvent, RerankCallEndEvent>,
136}
137
138impl<D> Rerank<D> {
139    /// Returns only the `top_n` most relevant documents.
140    #[must_use]
141    pub fn top_n(mut self, top_n: usize) -> Self {
142        self.top_n = Some(top_n);
143        self
144    }
145}
146
147impl_modality_builder!(Rerank<D>);
148impl_modality_hooks!(Rerank<D>, RerankCallStartEvent, RerankCallEndEvent);
149
150impl<D> IntoFuture for Rerank<D>
151where
152    D: Into<RerankDocument> + Clone + Send + 'static,
153{
154    type Output = Result<RerankResult<D>, Error>;
155    type IntoFuture = BoxFuture<'static, Self::Output>;
156
157    fn into_future(self) -> Self::IntoFuture {
158        Box::pin(run(self))
159    }
160}
161
162/// Converts the documents to the provider representation.
163fn to_model_documents<D>(documents: &[D]) -> Result<RerankDocuments, Error>
164where
165    D: Into<RerankDocument> + Clone,
166{
167    let mut texts: Vec<String> = Vec::new();
168    let mut objects: Vec<JsonObject> = Vec::new();
169    for document in documents {
170        match document.clone().into() {
171            RerankDocument::Text(text) => texts.push(text),
172            RerankDocument::Object(object) => objects.push(object),
173        }
174    }
175    match (texts.is_empty(), objects.is_empty()) {
176        (false, true) => Ok(RerankDocuments::Text { values: texts }),
177        (true, false) => Ok(RerankDocuments::Object { values: objects }),
178        _ => Err(Error::invalid_argument(
179            "documents",
180            "documents must be all text or all objects",
181        )),
182    }
183}
184
185async fn run<D>(builder: Rerank<D>) -> Result<RerankResult<D>, Error>
186where
187    D: Into<RerankDocument> + Clone + Send + 'static,
188{
189    let model = resolve_model(&builder.model, ProviderRegistry::reranking_model)?;
190    let identity = ModelIdentity::new(model.provider().clone(), model.model_id().clone());
191    let span = spans::modality_span("rerank", &identity);
192    let base = builder.base.clone();
193    let telemetry = TelemetryDispatcher::new(base.telemetry.clone());
194    let call_id = default_id_generator().generate();
195    let Rerank {
196        query,
197        documents,
198        top_n,
199        hooks,
200        ..
201    } = builder;
202    let event_documents: Vec<RerankDocument> = documents.iter().cloned().map(Into::into).collect();
203    let start = Arc::new(RerankCallStartEvent {
204        runtime_context: Some(hooks.runtime_context.clone()),
205        call_id: call_id.clone(),
206        operation_id: "ai.rerank",
207        model: identity.clone(),
208        documents: Some(event_documents.clone()),
209        query: Some(query.clone()),
210        top_n,
211        max_retries: base.retry_policy.max_retries,
212        headers: base.headers.clone(),
213        provider_options: base.provider_options.clone(),
214    });
215
216    if documents.is_empty() {
217        tokio::join!(
218            Hooks::emit(&hooks.on_start, start.clone()),
219            telemetry.on_rerank_operation_start(&start),
220        );
221        let result = RerankResult {
222            original_documents: documents,
223            ranking: Vec::new(),
224            warnings: Vec::new(),
225            response: ResponseMetadata {
226                timestamp: Some(Utc::now()),
227                model_id: Some(identity.model_id.clone()),
228                ..ResponseMetadata::default()
229            },
230            provider_metadata: None,
231        };
232        let end = Arc::new(RerankCallEndEvent {
233            runtime_context: Some(hooks.runtime_context),
234            call_id,
235            operation_id: "ai.rerank",
236            model: identity,
237            documents: Some(event_documents),
238            query: Some(query),
239            ranking: Some(Vec::new()),
240            warnings: result.warnings.clone(),
241            provider_metadata: result.provider_metadata.clone(),
242            response: result.response.clone(),
243        });
244        tokio::join!(
245            Hooks::emit(&hooks.on_end, end.clone()),
246            telemetry.on_rerank_operation_end(&end),
247        );
248        return Ok(result);
249    }
250
251    let model_documents = to_model_documents(&documents)?;
252    base.run(|base, token| {
253        async move {
254            tokio::join!(
255                Hooks::emit(&hooks.on_start, start.clone()),
256                telemetry.on_rerank_operation_start(&start),
257            );
258            let headers = base.request_headers();
259            let outcome = retry(&base.retry_policy, &token, |_| {
260                let options = RerankOptions {
261                    query: query.clone(),
262                    documents: model_documents.clone(),
263                    top_n,
264                    provider_options: base.provider_options.clone(),
265                    headers: headers.clone(),
266                    cancellation: token.child_token(),
267                };
268                let model = &model;
269                let telemetry = &telemetry;
270                let call_id = &call_id;
271                let identity = &identity;
272                async move {
273                    let started = Instant::now();
274                    telemetry
275                        .on_rerank_start(&RerankStartEvent {
276                            call_id: call_id.clone(),
277                            model: identity.clone(),
278                            document_count: options.documents.len(),
279                            query: telemetry.record_inputs().then(|| options.query.clone()),
280                        })
281                        .await;
282                    let result = model.do_rerank(options).await.map_err(Error::from)?;
283                    telemetry
284                        .on_rerank_end(&RerankEndEvent {
285                            call_id: call_id.clone(),
286                            ranked_count: result.ranking.len(),
287                            duration: started.elapsed(),
288                        })
289                        .await;
290                    Ok(result)
291                }
292            })
293            .await;
294            let result = match outcome {
295                Ok(result) => result,
296                Err(error) => {
297                    telemetry
298                        .on_error(&ErrorEvent {
299                            call_id: &call_id,
300                            error: &error,
301                            phase: ErrorPhase::ModelCall,
302                        })
303                        .await;
304                    return Err(error);
305                }
306            };
307            spans::log_warnings(&result.warnings, &identity);
308            let mut ranking: Vec<Ranked<D>> = Vec::with_capacity(result.ranking.len());
309            for ranked in result.ranking {
310                let Some(document) = documents.get(ranked.index) else {
311                    return Err(Error::from(ProviderError::InvalidResponseData(Box::new(
312                        InvalidResponseDataError::new(
313                            format!(
314                                "ranking index {} is out of range for {} documents",
315                                ranked.index,
316                                documents.len()
317                            ),
318                            json!({ "index": ranked.index, "documents": documents.len() }),
319                        ),
320                    ))));
321                };
322                ranking.push(Ranked {
323                    original_index: ranked.index,
324                    score: ranked.relevance_score,
325                    document: document.clone(),
326                });
327            }
328            let mut response = result.response;
329            if response.timestamp.is_none() {
330                response.timestamp = Some(Utc::now());
331            }
332            if response.model_id.is_none() {
333                response.model_id = Some(identity.model_id.clone());
334            }
335            let event_ranking = ranking
336                .iter()
337                .map(|ranked| Ranked {
338                    original_index: ranked.original_index,
339                    score: ranked.score,
340                    document: ranked.document.clone().into(),
341                })
342                .collect();
343            let end = Arc::new(RerankCallEndEvent {
344                runtime_context: Some(hooks.runtime_context),
345                call_id,
346                operation_id: "ai.rerank",
347                model: identity,
348                documents: Some(event_documents),
349                query: Some(query),
350                ranking: Some(event_ranking),
351                warnings: result.warnings.clone(),
352                provider_metadata: result.provider_metadata.clone(),
353                response: response.clone(),
354            });
355            tokio::join!(
356                Hooks::emit(&hooks.on_end, end.clone()),
357                telemetry.on_rerank_operation_end(&end),
358            );
359            Ok(RerankResult {
360                original_documents: documents,
361                ranking,
362                warnings: result.warnings,
363                response,
364                provider_metadata: result.provider_metadata,
365            })
366        }
367        .instrument(span)
368    })
369    .await
370}