1use serde::{Deserialize, Serialize};
2use serde_json::Value;
3
4use crate::openai::common::{EmbeddingObjectType, ListObjectType, OpenAiModelId, Rest};
5
6#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, gproxy_protocol_macros::WireBuilder)]
7#[cfg_attr(not(feature = "exhaustive"), non_exhaustive)]
8pub struct CreateEmbeddingRequest {
9 pub input: EmbeddingInput,
10 pub model: OpenAiModelId,
11 #[serde(skip_serializing_if = "Option::is_none")]
12 pub dimensions: Option<u32>,
13 #[serde(skip_serializing_if = "Option::is_none")]
14 pub encoding_format: Option<EmbeddingEncodingFormat>,
15 #[serde(skip_serializing_if = "Option::is_none")]
16 pub user: Option<String>,
17 #[serde(default, flatten)]
18 pub rest: Rest,
19}
20
21#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
22#[serde(untagged)]
23pub enum EmbeddingInput {
24 Text(String),
25 TextList(Vec<String>),
26 TokenList(Vec<i64>),
27 TokenLists(Vec<Vec<i64>>),
28 Raw(Value),
29}
30
31#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
32#[serde(untagged)]
33pub enum EmbeddingEncodingFormat {
34 Known(KnownEmbeddingEncodingFormat),
35 Unknown(String),
36}
37
38#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
39#[serde(rename_all = "snake_case")]
40pub enum KnownEmbeddingEncodingFormat {
41 Float,
42 Base64,
43}
44
45#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, gproxy_protocol_macros::WireBuilder)]
46#[cfg_attr(not(feature = "exhaustive"), non_exhaustive)]
47pub struct CreateEmbeddingResponse {
48 pub data: Vec<Embedding>,
49 pub model: OpenAiModelId,
50 pub object: ListObjectType,
51 pub usage: EmbeddingUsage,
52 #[serde(default, flatten)]
53 pub rest: Rest,
54}
55
56#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, gproxy_protocol_macros::WireBuilder)]
57#[cfg_attr(not(feature = "exhaustive"), non_exhaustive)]
58pub struct Embedding {
59 pub embedding: EmbeddingVector,
60 pub index: u32,
61 pub object: EmbeddingObjectType,
62 #[serde(default, flatten)]
63 pub rest: Rest,
64}
65
66#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
67#[serde(untagged)]
68pub enum EmbeddingVector {
69 Float(Vec<f64>),
70 Base64(String),
71 Raw(Value),
72}
73
74#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, gproxy_protocol_macros::WireBuilder)]
75#[cfg_attr(not(feature = "exhaustive"), non_exhaustive)]
76pub struct EmbeddingUsage {
77 pub prompt_tokens: u64,
78 pub total_tokens: u64,
79 #[serde(default, flatten)]
80 pub rest: Rest,
81}
82
83#[cfg(test)]
84mod tests {
85 use super::*;
86 use serde_json::json;
87
88 #[test]
89 fn embedding_round_trip_preserves_unknown_fields_and_variants() {
90 let value = json!({
91 "input": {"future_input": true},
92 "model": "text-embedding-future",
93 "encoding_format": "packed",
94 "future_request": {"enabled": true}
95 });
96 let request: CreateEmbeddingRequest = serde_json::from_value(value.clone()).unwrap();
97 assert_eq!(serde_json::to_value(request).unwrap(), value);
98
99 let response = json!({
100 "data": [{
101 "embedding": "AAEC",
102 "index": 0,
103 "object": "embedding",
104 "future_embedding": 7
105 }],
106 "model": "text-embedding-future",
107 "object": "list",
108 "usage": {"prompt_tokens": 1, "total_tokens": 1, "future_usage": 2},
109 "future_response": "kept"
110 });
111 let parsed: CreateEmbeddingResponse = serde_json::from_value(response.clone()).unwrap();
112 assert_eq!(serde_json::to_value(parsed).unwrap(), response);
113 }
114}