use std::future::Future;
use serde::Deserialize;
use serde::Serialize;
use tokio_util::sync::CancellationToken;
use crate::error::ProviderError;
use crate::json::JsonObject;
use crate::language_model::ResponseMetadata;
use crate::shared::Headers;
use crate::shared::ModelId;
use crate::shared::ProviderId;
use crate::shared::ProviderMetadata;
use crate::shared::ProviderOptions;
use crate::shared::Warning;
pub trait RerankingModel: Send + Sync + 'static {
fn provider(&self) -> &ProviderId;
fn model_id(&self) -> &ModelId;
fn do_rerank(
&self,
options: RerankOptions,
) -> impl Future<Output = Result<RerankResult, ProviderError>> + Send;
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "lowercase")]
#[non_exhaustive]
pub enum RerankDocuments {
Text {
values: Vec<String>,
},
Object {
values: Vec<JsonObject>,
},
}
impl RerankDocuments {
#[must_use]
pub fn len(&self) -> usize {
match self {
Self::Text { values } => values.len(),
Self::Object { values } => values.len(),
}
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
#[derive(Debug, Clone)]
pub struct RerankOptions {
pub query: String,
pub documents: RerankDocuments,
pub top_n: Option<usize>,
pub provider_options: ProviderOptions,
pub headers: Headers,
pub cancellation: CancellationToken,
}
impl RerankOptions {
#[must_use]
pub fn new(query: impl Into<String>, documents: RerankDocuments) -> Self {
Self {
query: query.into(),
documents,
top_n: None,
provider_options: ProviderOptions::new(),
headers: Headers::new(),
cancellation: CancellationToken::new(),
}
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct RankedDocument {
pub index: usize,
pub relevance_score: f64,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct RerankResult {
pub ranking: Vec<RankedDocument>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_metadata: Option<ProviderMetadata>,
#[serde(default)]
pub warnings: Vec<Warning>,
#[serde(default)]
pub response: ResponseMetadata,
}