Skip to main content

openrouter_rs/api/
rerank.rs

1use derive_builder::Builder;
2use reqwest::Client as HttpClient;
3use serde::{Deserialize, Serialize};
4
5use crate::{
6    error::OpenRouterError,
7    transport::{request as transport_request, response as transport_response},
8    types::ProviderPreferences,
9};
10
11/// One document input accepted by rerank requests.
12#[derive(Serialize, Deserialize, Debug, Clone, PartialEq)]
13#[serde(untagged)]
14#[non_exhaustive]
15pub enum RerankDocumentInput {
16    Text(String),
17    Multimodal {
18        #[serde(skip_serializing_if = "Option::is_none")]
19        text: Option<String>,
20        #[serde(skip_serializing_if = "Option::is_none")]
21        image: Option<String>,
22    },
23}
24
25impl RerankDocumentInput {
26    pub fn text(value: impl Into<String>) -> Self {
27        Self::Text(value.into())
28    }
29
30    pub fn multimodal<T, U>(text: Option<T>, image: Option<U>) -> Self
31    where
32        T: Into<String>,
33        U: Into<String>,
34    {
35        Self::Multimodal {
36            text: text.map(Into::into),
37            image: image.map(Into::into),
38        }
39    }
40}
41
42impl From<String> for RerankDocumentInput {
43    fn from(value: String) -> Self {
44        Self::Text(value)
45    }
46}
47
48impl From<&str> for RerankDocumentInput {
49    fn from(value: &str) -> Self {
50        Self::Text(value.to_string())
51    }
52}
53
54/// Request payload for `POST /rerank`.
55#[derive(Serialize, Deserialize, Debug, Clone, Builder)]
56#[builder(build_fn(error = "OpenRouterError"))]
57#[non_exhaustive]
58pub struct RerankRequest {
59    #[builder(setter(into))]
60    pub model: String,
61    #[builder(setter(into))]
62    pub query: String,
63    #[builder(setter(custom))]
64    pub documents: Vec<RerankDocumentInput>,
65    #[builder(setter(strip_option), default)]
66    #[serde(skip_serializing_if = "Option::is_none")]
67    pub top_n: Option<u32>,
68    #[builder(setter(strip_option), default)]
69    #[serde(skip_serializing_if = "Option::is_none")]
70    pub provider: Option<ProviderPreferences>,
71}
72
73impl RerankRequest {
74    pub fn builder() -> RerankRequestBuilder {
75        RerankRequestBuilder::default()
76    }
77}
78
79impl RerankRequestBuilder {
80    pub fn documents<T, S>(&mut self, items: T) -> &mut Self
81    where
82        T: IntoIterator<Item = S>,
83        S: Into<RerankDocumentInput>,
84    {
85        self.documents = Some(items.into_iter().map(Into::into).collect());
86        self
87    }
88}
89
90/// The original document returned in a rerank result.
91#[derive(Serialize, Deserialize, Debug, Clone)]
92#[non_exhaustive]
93pub struct RerankDocument {
94    #[serde(skip_serializing_if = "Option::is_none")]
95    pub text: Option<String>,
96    #[serde(skip_serializing_if = "Option::is_none")]
97    pub image: Option<String>,
98}
99
100/// One scored rerank result.
101#[derive(Serialize, Deserialize, Debug, Clone)]
102#[non_exhaustive]
103pub struct RerankResult {
104    pub index: u64,
105    pub relevance_score: f64,
106    pub document: RerankDocument,
107}
108
109/// Usage statistics returned by rerank providers.
110#[derive(Serialize, Deserialize, Debug, Clone)]
111#[non_exhaustive]
112pub struct RerankUsage {
113    #[serde(skip_serializing_if = "Option::is_none")]
114    pub cost: Option<f64>,
115    #[serde(skip_serializing_if = "Option::is_none")]
116    pub search_units: Option<u64>,
117    #[serde(skip_serializing_if = "Option::is_none")]
118    pub total_tokens: Option<u64>,
119}
120
121/// Response payload for `POST /rerank`.
122#[derive(Serialize, Deserialize, Debug, Clone)]
123#[non_exhaustive]
124pub struct RerankResponse {
125    #[serde(skip_serializing_if = "Option::is_none")]
126    pub id: Option<String>,
127    pub model: String,
128    #[serde(skip_serializing_if = "Option::is_none")]
129    pub provider: Option<String>,
130    pub results: Vec<RerankResult>,
131    #[serde(skip_serializing_if = "Option::is_none")]
132    pub usage: Option<RerankUsage>,
133}
134
135/// Submit a rerank request.
136pub async fn create_rerank(
137    base_url: &str,
138    api_key: &str,
139    x_title: &Option<String>,
140    http_referer: &Option<String>,
141    app_categories: &Option<Vec<String>>,
142    request: &RerankRequest,
143) -> Result<RerankResponse, OpenRouterError> {
144    let http_client = crate::transport::new_client()?;
145    create_rerank_with_client(
146        &http_client,
147        base_url,
148        api_key,
149        x_title,
150        http_referer,
151        app_categories,
152        request,
153    )
154    .await
155}
156
157pub(crate) async fn create_rerank_with_client(
158    http_client: &HttpClient,
159    base_url: &str,
160    api_key: &str,
161    x_title: &Option<String>,
162    http_referer: &Option<String>,
163    app_categories: &Option<Vec<String>>,
164    request: &RerankRequest,
165) -> Result<RerankResponse, OpenRouterError> {
166    let url = format!("{base_url}/rerank");
167    let response = transport_request::with_client_request_headers(
168        transport_request::post(http_client, &url),
169        api_key,
170        x_title,
171        http_referer,
172        app_categories,
173    )?
174    .json(request)
175    .send()
176    .await?;
177
178    if response.status().is_success() {
179        transport_response::parse_json_response(response, "rerank").await
180    } else {
181        transport_response::handle_error(response).await?;
182        unreachable!()
183    }
184}