openai_interface/embeddings/
request.rs1use serde::{Deserialize, Serialize};
2use url::Url;
3
4use crate::{
5 errors::OapiError,
6 rest::post::{Post, PostNoStream},
7};
8
9crate::request_body! {
10 #[derive(Debug, Serialize, Deserialize, Default, Clone)]
12 pub struct EmbeddingRequest {
13 pub input: EmbeddingInput,
19 pub model: String,
21 #[serde(skip_serializing_if = "Option::is_none")]
25 pub dimensions: Option<u32>,
26 #[serde(skip_serializing_if = "Option::is_none")]
32 pub encoding_format: Option<EncodingFormat>,
33 #[serde(skip_serializing_if = "Option::is_none")]
36 pub user: Option<String>,
37 }
38}
39
40#[derive(Debug, Serialize, Deserialize, Clone)]
42#[serde(untagged)]
43pub enum EmbeddingInput {
44 String(String),
46 StringArray(Vec<String>),
48 Tokens(Vec<u32>),
50 TokenArray(Vec<Vec<u32>>),
52}
53
54impl Default for EmbeddingInput {
55 fn default() -> Self {
56 Self::String(String::new())
57 }
58}
59
60#[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 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 #[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 #[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 #[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}