1use base64::{engine::general_purpose, Engine as _};
2use chroma_error::{ChromaError, ErrorCodes};
3use serde::{Deserialize, Serialize};
4use thiserror::Error;
5
6#[derive(Error, Debug)]
7pub enum Base64DecodeError {
8 #[error("Invalid base64 string: {0}")]
9 InvalidBase64(#[from] base64::DecodeError),
10 #[error("Invalid byte length: {byte_length} bytes cannot be converted to f32 values (must be divisible by 4)")]
11 InvalidByteLength { byte_length: usize },
12 #[error("Failed to convert embedding {embedding_index} to byte array")]
13 EmbeddingConversionFailed { embedding_index: usize },
14 #[error("Non-finite float value at embedding index {embedding_index}: {value}")]
15 NonFiniteFloatValue { embedding_index: usize, value: f32 },
16}
17
18impl ChromaError for Base64DecodeError {
19 fn code(&self) -> ErrorCodes {
20 ErrorCodes::InvalidArgument
21 }
22}
23
24#[derive(Serialize, Deserialize, Debug, Clone)]
25#[cfg_attr(feature = "utoipa", derive(utoipa::ToSchema))]
26#[serde(untagged)]
27pub enum EmbeddingsPayload {
28 JsonArrays(Vec<Vec<f32>>),
29 Base64Binary(Vec<String>),
30}
31
32#[derive(Serialize, Deserialize, Debug, Clone)]
33#[cfg_attr(feature = "utoipa", derive(utoipa::ToSchema))]
34#[serde(untagged)]
35pub enum UpdateEmbeddingsPayload {
36 JsonArrays(Vec<Option<Vec<f32>>>),
37 Base64Binary(Vec<Option<String>>),
38}
39
40pub fn decode_embeddings(
41 embeddings: EmbeddingsPayload,
42) -> Result<Vec<Vec<f32>>, Base64DecodeError> {
43 match embeddings {
44 EmbeddingsPayload::Base64Binary(base64_strings) => {
45 Ok(decode_base64_embeddings(&base64_strings)?)
46 }
47 EmbeddingsPayload::JsonArrays(arrays) => Ok(arrays),
48 }
49}
50
51pub fn maybe_decode_update_embeddings(
52 embeddings: Option<UpdateEmbeddingsPayload>,
53) -> Result<Option<Vec<Option<Vec<f32>>>>, Base64DecodeError> {
54 match embeddings {
55 Some(UpdateEmbeddingsPayload::Base64Binary(base64_data)) => {
56 Ok(Some(decode_base64_update_embeddings(&base64_data)?))
57 }
58 Some(UpdateEmbeddingsPayload::JsonArrays(arrays)) => Ok(Some(arrays)),
59 None => Ok(None),
60 }
61}
62
63pub fn decode_base64_embeddings(
64 base64_strings: &Vec<String>,
65) -> Result<Vec<Vec<f32>>, Base64DecodeError> {
66 let mut result = Vec::with_capacity(base64_strings.len());
67
68 for base64_str in base64_strings {
69 let floats = decode_base64_embedding(base64_str)?;
70
71 result.push(floats);
72 }
73
74 Ok(result)
75}
76
77pub fn decode_base64_update_embeddings(
78 base64_data: &Vec<Option<String>>,
79) -> Result<Vec<Option<Vec<f32>>>, Base64DecodeError> {
80 let mut result = Vec::with_capacity(base64_data.len());
81
82 for base64_str in base64_data {
83 if let Some(base64_str) = base64_str {
84 let floats = decode_base64_embedding(base64_str)?;
85
86 result.push(Some(floats));
87 } else {
88 result.push(None);
89 }
90 }
91
92 Ok(result)
93}
94
95pub fn decode_base64_embedding(base64_str: &String) -> Result<Vec<f32>, Base64DecodeError> {
96 let bytes = general_purpose::STANDARD.decode(base64_str)?;
97
98 let float_count = bytes.len() / 4;
99 if bytes.len() % 4 != 0 {
100 return Err(Base64DecodeError::InvalidByteLength {
101 byte_length: bytes.len(),
102 });
103 }
104
105 let mut floats = Vec::with_capacity(float_count);
106 for (embedding_index, chunk) in bytes.chunks_exact(4).enumerate() {
107 let float_bytes: [u8; 4] = chunk
108 .try_into()
109 .map_err(|_| Base64DecodeError::EmbeddingConversionFailed { embedding_index })?;
110 let f = f32::from_le_bytes(float_bytes);
112 if !f.is_finite() {
113 return Err(Base64DecodeError::NonFiniteFloatValue {
114 embedding_index,
115 value: f,
116 });
117 }
118 floats.push(f);
119 }
120
121 Ok(floats)
122}
123
124#[cfg(test)]
125mod tests {
126 use super::*;
127 #[cfg(feature = "testing")]
128 use proptest::prelude::*;
129
130 #[test]
131 fn test_invalid_base64_returns_error() {
132 let invalid_base64 = "invalid!@#$".to_string();
133 let result = decode_base64_embedding(&invalid_base64);
134 assert!(matches!(result, Err(Base64DecodeError::InvalidBase64(_))));
135 }
136
137 #[test]
138 fn test_invalid_byte_length_returns_error() {
139 let invalid_length_base64 = "YWJj".to_string(); let result = decode_base64_embedding(&invalid_length_base64);
142 assert!(matches!(
143 result,
144 Err(Base64DecodeError::InvalidByteLength { byte_length: 3 })
145 ));
146 }
147
148 #[test]
149 fn test_get_embeddings_propagates_error() {
150 let invalid_embeddings = EmbeddingsPayload::Base64Binary(vec!["invalid!@#$".to_string()]);
151 let result = decode_embeddings(invalid_embeddings);
152
153 assert!(matches!(result, Err(Base64DecodeError::InvalidBase64(_))));
154 }
155
156 #[test]
157 fn test_valid_base64_decoding() {
158 let valid_base64 = base64::Engine::encode(
160 &base64::engine::general_purpose::STANDARD,
161 1.0f32.to_le_bytes(),
162 );
163 let result = decode_base64_embedding(&valid_base64);
164 assert!(result.is_ok());
165 assert_eq!(result.unwrap(), vec![1.0f32]);
166 }
167
168 #[test]
169 fn test_multiple_embeddings_with_one_invalid() {
170 let valid_base64 = base64::Engine::encode(
171 &base64::engine::general_purpose::STANDARD,
172 1.0f32.to_le_bytes(),
173 );
174 let embeddings =
175 EmbeddingsPayload::Base64Binary(vec![valid_base64, "invalid!@#$".to_string()]);
176
177 let result = decode_embeddings(embeddings);
178 assert!(matches!(result, Err(Base64DecodeError::InvalidBase64(_))));
179 }
180
181 #[test]
182 fn test_decode_base64_embedding() {
183 let valid_base64 = base64::Engine::encode(
184 &base64::engine::general_purpose::STANDARD,
185 1.0f32.to_le_bytes(),
186 );
187 let result = decode_base64_embedding(&valid_base64);
188 assert!(result.is_ok());
189 assert_eq!(result.unwrap(), vec![1.0f32]);
190 }
191
192 #[test]
193 fn test_decode_base64_update_embeddings() {
194 let valid_base64s: Vec<Option<String>> = vec![
195 Some(base64::Engine::encode(
196 &base64::engine::general_purpose::STANDARD,
197 1.0f32.to_le_bytes(),
198 )),
199 Some(base64::Engine::encode(
200 &base64::engine::general_purpose::STANDARD,
201 2.0f32.to_le_bytes(),
202 )),
203 None,
204 Some(base64::Engine::encode(
205 &base64::engine::general_purpose::STANDARD,
206 3.0f32.to_le_bytes(),
207 )),
208 None,
209 ];
210 let result = decode_base64_update_embeddings(&valid_base64s);
211 assert!(result.is_ok());
212 assert_eq!(
213 result.unwrap(),
214 vec![
215 Some(vec![1.0f32]),
216 Some(vec![2.0f32]),
217 None,
218 Some(vec![3.0f32]),
219 None,
220 ]
221 );
222 }
223
224 #[test]
225 fn test_decode_base64_embeddings() {
226 let valid_base64s = vec![
227 base64::Engine::encode(
228 &base64::engine::general_purpose::STANDARD,
229 1.0f32.to_le_bytes(),
230 ),
231 base64::Engine::encode(
232 &base64::engine::general_purpose::STANDARD,
233 2.0f32.to_le_bytes(),
234 ),
235 base64::Engine::encode(
236 &base64::engine::general_purpose::STANDARD,
237 3.0f32.to_le_bytes(),
238 ),
239 ];
240 let result = decode_base64_embeddings(&valid_base64s);
241 assert!(result.is_ok());
242 assert_eq!(
243 result.unwrap(),
244 vec![vec![1.0f32], vec![2.0f32], vec![3.0f32]]
245 );
246 }
247
248 #[cfg(feature = "testing")]
249 fn encode_floats_to_base64(floats: &[f32]) -> String {
250 let mut bytes = Vec::with_capacity(floats.len() * 4);
251 for &f in floats {
252 bytes.extend_from_slice(&f.to_le_bytes());
253 }
254 general_purpose::STANDARD.encode(&bytes)
255 }
256
257 #[test]
258 fn test_nan_base64_rejected() {
259 let nan_base64 = base64::Engine::encode(
260 &base64::engine::general_purpose::STANDARD,
261 f32::NAN.to_le_bytes(),
262 );
263 let result = decode_base64_embedding(&nan_base64);
264 assert!(matches!(
265 result,
266 Err(Base64DecodeError::NonFiniteFloatValue { .. })
267 ));
268 }
269
270 #[test]
271 fn test_infinity_base64_rejected() {
272 let inf_base64 = base64::Engine::encode(
273 &base64::engine::general_purpose::STANDARD,
274 f32::INFINITY.to_le_bytes(),
275 );
276 let result = decode_base64_embedding(&inf_base64);
277 assert!(matches!(
278 result,
279 Err(Base64DecodeError::NonFiniteFloatValue { .. })
280 ));
281 }
282
283 #[test]
284 fn test_neg_infinity_base64_rejected() {
285 let neg_inf_base64 = base64::Engine::encode(
286 &base64::engine::general_purpose::STANDARD,
287 f32::NEG_INFINITY.to_le_bytes(),
288 );
289 let result = decode_base64_embedding(&neg_inf_base64);
290 assert!(matches!(
291 result,
292 Err(Base64DecodeError::NonFiniteFloatValue { .. })
293 ));
294 }
295
296 #[test]
297 fn test_nan_in_multi_embedding_rejected() {
298 let mut bytes = Vec::new();
299 bytes.extend_from_slice(&1.0f32.to_le_bytes());
300 bytes.extend_from_slice(&f32::NAN.to_le_bytes());
301 bytes.extend_from_slice(&3.0f32.to_le_bytes());
302 let base64_str = base64::Engine::encode(&base64::engine::general_purpose::STANDARD, &bytes);
303 let result = decode_base64_embedding(&base64_str);
304 assert!(matches!(
305 result,
306 Err(Base64DecodeError::NonFiniteFloatValue {
307 embedding_index: 1,
308 ..
309 })
310 ));
311 }
312
313 #[cfg(feature = "testing")]
314 fn embeddings_strategy() -> impl Strategy<Value = Vec<Vec<f32>>> {
315 proptest::collection::vec(
316 proptest::collection::vec(
317 proptest::num::f32::NORMAL
318 | proptest::num::f32::SUBNORMAL
319 | proptest::num::f32::ZERO,
320 0..10,
321 ),
322 0..5,
323 )
324 }
325
326 #[cfg(feature = "testing")]
327 proptest! {
328 #[test]
329 fn test_decode_base64_embeddings_prop(embeddings in embeddings_strategy()) {
330 let base64_strings = embeddings.iter().map(|e| encode_floats_to_base64(e)).collect();
331 let result = decode_base64_embeddings(&base64_strings).unwrap();
332 for (original, decoded) in embeddings.iter().zip(result.iter()) {
333 prop_assert_eq!(original, decoded);
334 }
335 }
336 }
337}