Skip to main content

gproxy_protocol/openai/
embeddings.rs

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}