Skip to main content

openai_interface/embeddings/
request.rs

1use serde::Serialize;
2use url::Url;
3
4use crate::{
5    errors::OapiError,
6    rest::post::{Post, PostNoStream},
7};
8
9/// Creates an embedding vector representing the input text.
10#[derive(Debug, Serialize, Default, Clone)]
11pub struct EmbeddingRequest {
12    /// Input text to embed, encoded as a string, an array of strings, an
13    /// array of tokens, or an array of token arrays.
14    ///
15    /// The input must not exceed the max input tokens for the model (8192
16    /// tokens for all embedding models) and cannot be an empty string.
17    pub input: EmbeddingInput,
18    /// ID of the model to use, e.g. `text-embedding-v4`.
19    pub model: String,
20    /// The number of dimensions the resulting output embeddings should have.
21    ///
22    /// Only supported in `text-embedding-3` and later models.
23    #[serde(skip_serializing_if = "Option::is_none")]
24    pub dimensions: Option<u32>,
25    /// The format to return the embeddings in. Can be either `float` or
26    /// `base64`.
27    ///
28    /// Note that some OpenAI-compatible providers (e.g. Qwen) only support
29    /// `float`.
30    #[serde(skip_serializing_if = "Option::is_none")]
31    pub encoding_format: Option<EncodingFormat>,
32    /// A unique identifier representing your end-user, which can help OpenAI
33    /// to monitor and detect abuse.
34    #[serde(skip_serializing_if = "Option::is_none")]
35    pub user: Option<String>,
36    /// Add additional JSON properties to the request.
37    #[serde(flatten, skip_serializing_if = "Option::is_none")]
38    pub extra_body: Option<serde_json::Map<String, serde_json::Value>>,
39}
40
41/// The input text to embed.
42#[derive(Debug, Serialize, Clone)]
43#[serde(untagged)]
44pub enum EmbeddingInput {
45    /// A single string to embed.
46    String(String),
47    /// An array of strings to embed.
48    StringArray(Vec<String>),
49    /// A single array of tokens to embed.
50    Tokens(Vec<u32>),
51    /// An array of token arrays to embed.
52    TokenArray(Vec<Vec<u32>>),
53}
54
55impl Default for EmbeddingInput {
56    fn default() -> Self {
57        Self::String(String::new())
58    }
59}
60
61/// The format to return the embeddings in.
62#[derive(Debug, Serialize, Clone, Copy)]
63#[serde(rename_all = "lowercase")]
64pub enum EncodingFormat {
65    Float,
66    Base64,
67}
68
69impl EmbeddingRequest {
70    pub fn is_streaming(&self) -> bool {
71        false
72    }
73}
74
75impl Post for EmbeddingRequest {
76    fn is_streaming(&self) -> bool {
77        false
78    }
79
80    /// Builds the URL for the request.
81    ///
82    /// `base_url` should be like <https://api.openai.com/v1>
83    fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
84        let mut url = Url::parse(base_url.trim_end_matches('/')).map_err(OapiError::UrlError)?;
85        url.path_segments_mut()
86            .map_err(|_| OapiError::UrlCannotBeBase(base_url.to_string()))?
87            .push("embeddings");
88
89        Ok(url.to_string())
90    }
91}
92
93impl PostNoStream for EmbeddingRequest {
94    type Response = super::response::CreateEmbeddingResponse;
95}
96
97#[cfg(test)]
98mod tests {
99    use super::*;
100
101    /// Serializes a simple string input request.
102    #[test]
103    fn string_input_serialization() {
104        let request = EmbeddingRequest {
105            input: EmbeddingInput::String("Hello".to_string()),
106            model: "text-embedding-v4".to_string(),
107            dimensions: Some(1024),
108            encoding_format: Some(EncodingFormat::Float),
109            ..Default::default()
110        };
111
112        let json = serde_json::to_string(&request).unwrap();
113        assert!(json.contains(r#""input":"Hello""#), "json: {json}");
114        assert!(
115            json.contains(r#""model":"text-embedding-v4""#),
116            "json: {json}"
117        );
118        assert!(json.contains(r#""dimensions":1024"#), "json: {json}");
119        assert!(
120            json.contains(r#""encoding_format":"float""#),
121            "json: {json}"
122        );
123    }
124
125    /// Serializes an array of strings and a token array input.
126    #[test]
127    fn array_input_serialization() {
128        let request = EmbeddingRequest {
129            input: EmbeddingInput::StringArray(vec!["Hello".to_string(), "World".to_string()]),
130            model: "text-embedding-3-small".to_string(),
131            ..Default::default()
132        };
133        let json = serde_json::to_string(&request).unwrap();
134        assert!(
135            json.contains(r#""input":["Hello","World"]"#),
136            "json: {json}"
137        );
138
139        let request = EmbeddingRequest {
140            input: EmbeddingInput::TokenArray(vec![vec![1234, 5678]]),
141            model: "text-embedding-3-small".to_string(),
142            ..Default::default()
143        };
144        let json = serde_json::to_string(&request).unwrap();
145        assert!(json.contains(r#""input":[[1234,5678]]"#), "json: {json}");
146    }
147
148    /// A flat token id array is sent as `input: [..]`.
149    #[test]
150    fn flat_token_array_serialization() {
151        let request = EmbeddingRequest {
152            input: EmbeddingInput::Tokens(vec![1234, 5678]),
153            model: "text-embedding-3-small".to_string(),
154            ..Default::default()
155        };
156        let json = serde_json::to_string(&request).unwrap();
157        assert!(json.contains(r#""input":[1234,5678]"#), "json: {json}");
158    }
159
160    #[test]
161    fn test_build_url() {
162        let request = EmbeddingRequest::default();
163        let url = request.build_url("https://api.openai.com/v1/").unwrap();
164        assert_eq!(url, "https://api.openai.com/v1/embeddings");
165    }
166}