openai_protocol/
rerank.rs1use std::collections::HashMap;
2
3use serde::{Deserialize, Serialize};
4use serde_json::Value;
5use validator::Validate;
6
7use super::common::{default_true, GenerationRequest, StringOrArray, UsageInfo};
8
9fn default_rerank_object() -> String {
10 "rerank".to_string()
11}
12
13fn current_timestamp() -> i64 {
14 std::time::SystemTime::now()
15 .duration_since(std::time::UNIX_EPOCH)
16 .unwrap_or_else(|_| std::time::Duration::from_secs(0))
17 .as_secs() as i64
18}
19
20#[derive(Debug, Clone, Deserialize, Serialize, Validate, schemars::JsonSchema)]
25#[validate(schema(function = "validate_rerank_request"))]
26pub struct RerankRequest {
27 #[validate(custom(function = "validate_query"))]
29 pub query: String,
30
31 #[validate(custom(function = "validate_documents"))]
33 pub documents: Vec<String>,
34
35 pub model: String,
37
38 #[serde(skip_serializing_if = "Option::is_none")]
40 #[validate(range(min = 1))]
41 pub top_k: Option<usize>,
42
43 #[serde(default = "default_true")]
45 pub return_documents: bool,
46
47 pub rid: Option<StringOrArray>,
50
51 pub user: Option<String>,
53}
54
55impl GenerationRequest for RerankRequest {
56 fn get_model(&self) -> Option<&str> {
57 Some(&self.model)
58 }
59
60 fn is_stream(&self) -> bool {
61 false }
63
64 fn extract_text_for_routing(&self) -> String {
65 self.query.clone()
66 }
67}
68
69impl super::validated::Normalizable for RerankRequest {
70 }
72
73fn validate_query(query: &str) -> Result<(), validator::ValidationError> {
79 if query.trim().is_empty() {
80 return Err(validator::ValidationError::new("query cannot be empty"));
81 }
82 Ok(())
83}
84
85fn validate_documents(documents: &[String]) -> Result<(), validator::ValidationError> {
87 if documents.is_empty() {
88 return Err(validator::ValidationError::new(
89 "documents list cannot be empty",
90 ));
91 }
92 Ok(())
93}
94
95#[expect(
97 clippy::unnecessary_wraps,
98 reason = "validator crate requires Result return type"
99)]
100fn validate_rerank_request(req: &RerankRequest) -> Result<(), validator::ValidationError> {
101 if let Some(k) = req.top_k {
103 if k > req.documents.len() {
104 tracing::warn!(
105 "top_k ({}) is greater than number of documents ({})",
106 k,
107 req.documents.len()
108 );
109 }
110 }
111 Ok(())
112}
113
114impl RerankRequest {
115 pub fn effective_top_k(&self) -> usize {
117 self.top_k.unwrap_or(self.documents.len())
118 }
119}
120
121#[serde_with::skip_serializing_none]
123#[derive(Debug, Clone, Serialize, Deserialize, schemars::JsonSchema)]
124pub struct RerankResult {
125 pub score: f32,
127
128 pub document: Option<String>,
130
131 pub index: usize,
133
134 pub meta_info: Option<HashMap<String, Value>>,
136}
137
138#[derive(Debug, Clone, Serialize, Deserialize, schemars::JsonSchema)]
140pub struct RerankResponse {
141 pub results: Vec<RerankResult>,
143
144 pub model: String,
146
147 pub usage: Option<UsageInfo>,
149
150 #[serde(default = "default_rerank_object")]
152 pub object: String,
153
154 pub id: Option<StringOrArray>,
156
157 pub created: i64,
159}
160
161impl RerankResponse {
162 pub fn new(
164 results: Vec<RerankResult>,
165 model: String,
166 request_id: Option<StringOrArray>,
167 ) -> Self {
168 RerankResponse {
169 results,
170 model,
171 usage: None,
172 object: default_rerank_object(),
173 id: request_id,
174 created: current_timestamp(),
175 }
176 }
177
178 pub fn apply_top_k(&mut self, k: usize) {
180 self.results.truncate(k);
181 }
182
183 pub fn drop_documents(&mut self) {
185 for result in &mut self.results {
186 result.document = None;
187 }
188 }
189}
190
191#[derive(Debug, Clone, Serialize, Deserialize, schemars::JsonSchema)]
194pub struct V1RerankReqInput {
195 pub query: String,
196 pub documents: Vec<String>,
197}
198
199impl From<V1RerankReqInput> for RerankRequest {
201 fn from(v1: V1RerankReqInput) -> Self {
202 RerankRequest {
203 query: v1.query,
204 documents: v1.documents,
205 model: super::UNKNOWN_MODEL_ID.to_string(),
206 top_k: None,
207 return_documents: true,
208 rid: None,
209 user: None,
210 }
211 }
212}