openai_interface/embeddings/
request.rs1use serde::Serialize;
2use url::Url;
3
4use crate::{
5 errors::OapiError,
6 rest::post::{Post, PostNoStream},
7};
8
9#[derive(Debug, Serialize, Default, Clone)]
11pub struct EmbeddingRequest {
12 pub input: EmbeddingInput,
18 pub model: String,
20 #[serde(skip_serializing_if = "Option::is_none")]
24 pub dimensions: Option<u32>,
25 #[serde(skip_serializing_if = "Option::is_none")]
31 pub encoding_format: Option<EncodingFormat>,
32 #[serde(skip_serializing_if = "Option::is_none")]
35 pub user: Option<String>,
36 #[serde(flatten, skip_serializing_if = "Option::is_none")]
38 pub extra_body: Option<serde_json::Map<String, serde_json::Value>>,
39}
40
41#[derive(Debug, Serialize, Clone)]
43#[serde(untagged)]
44pub enum EmbeddingInput {
45 String(String),
47 StringArray(Vec<String>),
49 TokenArray(Vec<Vec<u32>>),
51}
52
53impl Default for EmbeddingInput {
54 fn default() -> Self {
55 Self::String(String::new())
56 }
57}
58
59#[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 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 #[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 #[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}