gproxy_protocol/openai/
rerank.rs1use 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}