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    /// An array of tokens, or an array of token arrays.
50    TokenArray(Vec<Vec<u32>>),
51}
52
53impl Default for EmbeddingInput {
54    fn default() -> Self {
55        Self::String(String::new())
56    }
57}
58
59/// The format to return the embeddings in.
60#[derive(Debug, Serialize, Clone, Copy)]
61#[serde(rename_all = "lowercase")]
62pub enum EncodingFormat {
63    Float,
64    Base64,
65}
66
67impl EmbeddingRequest {
68    pub fn is_streaming(&self) -> bool {
69        false
70    }
71}
72
73impl Post for EmbeddingRequest {
74    fn is_streaming(&self) -> bool {
75        false
76    }
77
78    /// Builds the URL for the request.
79    ///
80    /// `base_url` should be like <https://api.openai.com/v1>
81    fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
82        let mut url = Url::parse(base_url.trim_end_matches('/')).map_err(OapiError::UrlError)?;
83        url.path_segments_mut()
84            .map_err(|_| OapiError::UrlCannotBeBase(base_url.to_string()))?
85            .push("embeddings");
86
87        Ok(url.to_string())
88    }
89}
90
91impl PostNoStream for EmbeddingRequest {
92    type Response = super::response::CreateEmbeddingResponse;
93}
94
95#[cfg(test)]
96mod tests {
97    use super::*;
98
99    /// Serializes a simple string input request.
100    #[test]
101    fn string_input_serialization() {
102        let request = EmbeddingRequest {
103            input: EmbeddingInput::String("Hello".to_string()),
104            model: "text-embedding-v4".to_string(),
105            dimensions: Some(1024),
106            encoding_format: Some(EncodingFormat::Float),
107            ..Default::default()
108        };
109
110        let json = serde_json::to_string(&request).unwrap();
111        assert!(json.contains(r#""input":"Hello""#), "json: {json}");
112        assert!(
113            json.contains(r#""model":"text-embedding-v4""#),
114            "json: {json}"
115        );
116        assert!(json.contains(r#""dimensions":1024"#), "json: {json}");
117        assert!(
118            json.contains(r#""encoding_format":"float""#),
119            "json: {json}"
120        );
121    }
122
123    /// Serializes an array of strings and a token array input.
124    #[test]
125    fn array_input_serialization() {
126        let request = EmbeddingRequest {
127            input: EmbeddingInput::StringArray(vec!["Hello".to_string(), "World".to_string()]),
128            model: "text-embedding-3-small".to_string(),
129            ..Default::default()
130        };
131        let json = serde_json::to_string(&request).unwrap();
132        assert!(
133            json.contains(r#""input":["Hello","World"]"#),
134            "json: {json}"
135        );
136
137        let request = EmbeddingRequest {
138            input: EmbeddingInput::TokenArray(vec![vec![1234, 5678]]),
139            model: "text-embedding-3-small".to_string(),
140            ..Default::default()
141        };
142        let json = serde_json::to_string(&request).unwrap();
143        assert!(json.contains(r#""input":[[1234,5678]]"#), "json: {json}");
144    }
145
146    #[test]
147    fn test_build_url() {
148        let request = EmbeddingRequest::default();
149        let url = request.build_url("https://api.openai.com/v1/").unwrap();
150        assert_eq!(url, "https://api.openai.com/v1/embeddings");
151    }
152}