Skip to main content

chroma_types/
base64_decode.rs

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        // handles little endian encoding
111        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        // This is valid base64 but encodes 3 bytes (not divisible by 4)
140        let invalid_length_base64 = "YWJj".to_string(); // "abc" = 3 bytes
141        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        // Valid base64 encoding 4 bytes (1 f32)
159        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}