1use 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#[derive(Debug, Clone, PartialEq)]
48pub enum RerankDocument {
49 Text(String),
51 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#[derive(Debug, Clone, PartialEq)]
75pub struct Ranked<D> {
76 pub original_index: usize,
78 pub score: f64,
80 pub document: D,
82}
83
84#[derive(Debug, Clone, PartialEq)]
86pub struct RerankResult<D> {
87 pub original_documents: Vec<D>,
89 pub ranking: Vec<Ranked<D>>,
91 pub warnings: Vec<Warning>,
93 pub response: ResponseMetadata,
95 pub provider_metadata: Option<ProviderMetadata>,
97}
98
99impl<D> RerankResult<D> {
100 pub fn reranked_documents(&self) -> impl Iterator<Item = &D> + '_ {
102 self.ranking.iter().map(|ranked| &ranked.document)
103 }
104}
105
106#[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#[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 #[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
162fn 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}