1use 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#[derive(Debug, Clone, PartialEq)]
39pub enum RerankDocument {
40 Text(String),
42 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#[derive(Debug, Clone, PartialEq)]
66pub struct Ranked<D> {
67 pub original_index: usize,
69 pub score: f64,
71 pub document: D,
73}
74
75#[derive(Debug, Clone, PartialEq)]
77pub struct RerankResult<D> {
78 pub ranking: Vec<Ranked<D>>,
80 pub warnings: Vec<Warning>,
82 pub response: ResponseMetadata,
84 pub provider_metadata: Option<ProviderMetadata>,
86}
87
88impl<D> RerankResult<D> {
89 pub fn reranked_documents(&self) -> impl Iterator<Item = &D> + '_ {
91 self.ranking.iter().map(|ranked| &ranked.document)
92 }
93}
94
95#[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#[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 #[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
148fn 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}