Skip to main content

gproxy_protocol/openai/
rerank.rs

1//! OpenAI-shaped rerank compatibility wire.
2//!
3//! OpenAI has no standalone rerank endpoint in the local documentation
4//! snapshot. This shape is retained from v2 for compatible providers.
5
6use serde::{Deserialize, Serialize};
7use serde_json::Value;
8
9use crate::openai::common::{OpenAiModelId, Rest};
10
11#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, gproxy_protocol_macros::WireBuilder)]
12#[cfg_attr(not(feature = "exhaustive"), non_exhaustive)]
13pub struct RerankRequest {
14    pub model: OpenAiModelId,
15    pub query: String,
16    pub documents: Vec<RerankDocument>,
17    #[serde(skip_serializing_if = "Option::is_none")]
18    pub top_n: Option<u32>,
19    #[serde(skip_serializing_if = "Option::is_none")]
20    pub instruct: Option<String>,
21    #[serde(skip_serializing_if = "Option::is_none")]
22    pub return_documents: Option<bool>,
23    #[serde(default, flatten)]
24    pub rest: Rest,
25}
26
27#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
28#[serde(untagged)]
29pub enum RerankDocument {
30    Text(String),
31    Structured(RerankDocumentContent),
32    Raw(Value),
33}
34
35#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, gproxy_protocol_macros::WireBuilder)]
36#[cfg_attr(not(feature = "exhaustive"), non_exhaustive)]
37pub struct RerankDocumentContent {
38    #[serde(skip_serializing_if = "Option::is_none")]
39    pub text: Option<String>,
40    #[serde(skip_serializing_if = "Option::is_none")]
41    pub image: Option<String>,
42    #[serde(default, flatten)]
43    pub rest: Rest,
44}
45
46#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, gproxy_protocol_macros::WireBuilder)]
47#[cfg_attr(not(feature = "exhaustive"), non_exhaustive)]
48pub struct RerankResponse {
49    pub model: OpenAiModelId,
50    pub results: Vec<RerankResult>,
51    #[serde(skip_serializing_if = "Option::is_none")]
52    pub id: Option<String>,
53    #[serde(skip_serializing_if = "Option::is_none")]
54    pub object: Option<String>,
55    #[serde(skip_serializing_if = "Option::is_none")]
56    pub provider: Option<String>,
57    #[serde(skip_serializing_if = "Option::is_none")]
58    pub usage: Option<RerankUsage>,
59    #[serde(default, flatten)]
60    pub rest: Rest,
61}
62
63#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, gproxy_protocol_macros::WireBuilder)]
64#[cfg_attr(not(feature = "exhaustive"), non_exhaustive)]
65pub struct RerankResult {
66    pub index: u32,
67    pub relevance_score: f64,
68    #[serde(skip_serializing_if = "Option::is_none")]
69    pub document: Option<RerankDocumentContent>,
70    #[serde(default, flatten)]
71    pub rest: Rest,
72}
73
74#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, gproxy_protocol_macros::WireBuilder)]
75#[cfg_attr(not(feature = "exhaustive"), non_exhaustive)]
76pub struct RerankUsage {
77    #[serde(skip_serializing_if = "Option::is_none")]
78    pub total_tokens: Option<u64>,
79    #[serde(skip_serializing_if = "Option::is_none")]
80    pub search_units: Option<u64>,
81    #[serde(skip_serializing_if = "Option::is_none")]
82    pub cost: Option<f64>,
83    #[serde(default, flatten)]
84    pub rest: Rest,
85}
86
87#[cfg(test)]
88mod tests {
89    use super::*;
90    use serde_json::json;
91
92    #[test]
93    fn rerank_round_trip_preserves_unknown_fields_and_documents() {
94        let value = json!({
95            "model":"rerank-future",
96            "query":"q",
97            "documents":["text", {"future_document":true}],
98            "future_request":{"x":1}
99        });
100        let parsed: RerankRequest = serde_json::from_value(value.clone()).unwrap();
101        assert_eq!(serde_json::to_value(parsed).unwrap(), value);
102    }
103}