Skip to main content

openai_interface/embeddings/
request.rs

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