1pub const INLINE_THRESHOLD_BYTES: usize = 10 * 1024 * 1024; use crate::cache::CacheKey;
15use crate::error::{Result, SomaError};
16use crate::value::Value;
17use serde::{Deserialize, Serialize};
18
19#[derive(Debug, Clone, Serialize, Deserialize)]
21pub struct StoreMeta {
22 pub total_rows: usize,
24 pub shape_tail: Vec<usize>,
26 pub dtype: String,
28}
29
30impl StoreMeta {
31 pub fn from_value(value: &Value) -> Self {
33 match value {
34 Value::Tensor { shape, .. } => Self {
35 total_rows: shape.first().copied().unwrap_or(0),
36 shape_tail: shape.get(1..).unwrap_or_default().to_vec(),
37 dtype: "tensor".into(),
38 },
39 Value::Text(_) => Self {
40 total_rows: 1,
41 shape_tail: vec![],
42 dtype: "text".into(),
43 },
44 Value::Json(_) => Self {
45 total_rows: 1,
46 shape_tail: vec![],
47 dtype: "json".into(),
48 },
49 Value::Bytes(b) | Value::Object(b) => Self {
50 total_rows: b.len(),
51 shape_tail: vec![],
52 dtype: "bytes".into(),
53 },
54 Value::Empty => Self {
55 total_rows: 0,
56 shape_tail: vec![],
57 dtype: "empty".into(),
58 },
59 }
60 }
61}
62
63pub fn slice_tensor_rows(value: &Value, start: usize, len: usize) -> Result<Value> {
65 match value {
66 Value::Tensor { values, shape } => {
67 if shape.is_empty() {
68 return Err(SomaError::DataStore("cannot slice scalar tensor".into()));
69 }
70 let cols: usize = shape[1..].iter().product::<usize>().max(1);
71 let row_start = start * cols;
72 let row_end = (start + len) * cols;
73 if row_end > values.len() {
74 return Err(SomaError::DataStore(format!(
75 "row range {start}..{} out of bounds (total rows: {})",
76 start + len,
77 shape[0]
78 )));
79 }
80 let mut new_shape = shape.clone();
81 new_shape[0] = len;
82 Ok(Value::tensor(
83 values[row_start..row_end].to_vec(),
84 new_shape,
85 ))
86 }
87 _ => Err(SomaError::DataStore(
88 "get_rows only works on Tensor values".into(),
89 )),
90 }
91}
92
93#[derive(Debug, Clone, Serialize, Deserialize)]
96#[serde(tag = "type")]
97#[non_exhaustive]
98pub enum DataRef {
99 Local {
101 path: String,
103 },
104 S3 {
106 bucket: String,
108 key: String,
110 region: Option<String>,
112 },
113 Cached {
115 cache_key: CacheKey,
117 },
118 Stream {
120 endpoint: String,
122 format: StreamFormat,
124 },
125 Inline {
127 value: Value,
129 },
130 Zarr {
132 bucket: String,
134 array_path: String,
136 region: Option<String>,
138 },
139}
140
141#[derive(Debug, Clone, Serialize, Deserialize, Default)]
143#[serde(rename_all = "snake_case")]
144#[non_exhaustive]
145pub enum StreamFormat {
146 #[default]
148 JsonLines,
149 Csv,
151 Arrow,
153 Protobuf,
155}
156
157#[derive(Debug, Clone, Serialize, Deserialize)]
159#[serde(tag = "type")]
160#[non_exhaustive]
161pub enum StorageConfig {
162 #[serde(rename = "local")]
164 Local {
165 base_path: String,
167 },
168 #[serde(rename = "s3")]
170 S3 {
171 bucket: String,
173 prefix: String,
175 region: Option<String>,
177 endpoint: Option<String>,
179 },
180 #[serde(rename = "zarr")]
182 Zarr {
183 bucket: String,
185 prefix: String,
187 region: Option<String>,
189 endpoint: Option<String>,
191 chunk_rows: usize,
193 },
194}
195
196impl Default for StorageConfig {
197 fn default() -> Self {
198 Self::Local {
199 base_path: "/tmp/soma-data".to_string(),
200 }
201 }
202}
203
204pub trait DataStore: Send + Sync {
209 fn put(&self, key: &CacheKey, data: &Value) -> Result<DataRef>;
211
212 fn get(&self, data_ref: &DataRef) -> Result<Value>;
214
215 fn exists(&self, data_ref: &DataRef) -> Result<bool>;
217
218 fn remove(&self, data_ref: &DataRef) -> Result<()>;
220
221 fn config(&self) -> &StorageConfig;
223
224 fn get_rows(&self, data_ref: &DataRef, start: usize, len: usize) -> Result<Value> {
228 let value = self.get(data_ref)?;
229 slice_tensor_rows(&value, start, len)
230 }
231
232 fn meta(&self, data_ref: &DataRef) -> Result<StoreMeta> {
235 let value = self.get(data_ref)?;
236 Ok(StoreMeta::from_value(&value))
237 }
238}
239
240pub struct LocalDataStore {
242 config: StorageConfig,
243 base_path: std::path::PathBuf,
244}
245
246impl LocalDataStore {
247 pub fn new(base_path: impl Into<std::path::PathBuf>) -> Self {
251 let base = base_path.into();
252 std::fs::create_dir_all(&base).ok();
253 Self {
254 config: StorageConfig::Local {
255 base_path: base.to_string_lossy().to_string(),
256 },
257 base_path: base,
258 }
259 }
260}
261
262impl DataStore for LocalDataStore {
263 fn put(&self, key: &CacheKey, data: &Value) -> Result<DataRef> {
264 let path = self.base_path.join(key.to_hex());
265 let bytes = serde_json::to_vec(data)
266 .map_err(|e| crate::error::SomaError::DataStore(e.to_string()))?;
267 std::fs::write(&path, &bytes)
268 .map_err(|e| crate::error::SomaError::DataStore(e.to_string()))?;
269 Ok(DataRef::Local {
270 path: path.to_string_lossy().to_string(),
271 })
272 }
273
274 fn get(&self, data_ref: &DataRef) -> Result<Value> {
275 match data_ref {
276 DataRef::Local { path } => {
277 let bytes = std::fs::read(path)
278 .map_err(|e| crate::error::SomaError::DataStore(e.to_string()))?;
279 serde_json::from_slice(&bytes)
280 .map_err(|e| crate::error::SomaError::DataStore(e.to_string()))
281 }
282 DataRef::Cached { cache_key } => {
283 let path = self.base_path.join(cache_key.to_hex());
284 let bytes = std::fs::read(&path)
285 .map_err(|e| crate::error::SomaError::DataStore(e.to_string()))?;
286 serde_json::from_slice(&bytes)
287 .map_err(|e| crate::error::SomaError::DataStore(e.to_string()))
288 }
289 DataRef::Inline { value } => Ok(value.clone()),
290 _ => Err(crate::error::SomaError::DataStore(
291 "Cannot get non-local DataRef from LocalDataStore".into(),
292 )),
293 }
294 }
295
296 fn exists(&self, data_ref: &DataRef) -> Result<bool> {
297 match data_ref {
298 DataRef::Local { path } => Ok(std::path::Path::new(path).exists()),
299 DataRef::Cached { cache_key } => Ok(self.base_path.join(cache_key.to_hex()).exists()),
300 DataRef::Inline { .. } => Ok(true),
301 _ => Ok(false),
302 }
303 }
304
305 fn remove(&self, data_ref: &DataRef) -> Result<()> {
306 if let DataRef::Local { path } = data_ref {
307 std::fs::remove_file(path).ok();
308 }
309 Ok(())
310 }
311
312 fn config(&self) -> &StorageConfig {
313 &self.config
314 }
315}
316
317#[cfg(test)]
318mod tests {
319 use super::*;
320
321 #[test]
322 fn local_data_store_roundtrip() {
323 let dir = std::env::temp_dir().join("soma-ds-test");
324 let store = LocalDataStore::new(&dir);
325
326 let key = CacheKey::hash_data(b"test_data");
327 let value = Value::tensor(vec![1.0, 2.0, 3.0], vec![3]);
328
329 let data_ref = store.put(&key, &value).unwrap();
330 assert!(store.exists(&data_ref).unwrap());
331
332 let retrieved = store.get(&data_ref).unwrap();
333 let (data, _) = retrieved.as_tensor().unwrap();
334 assert_eq!(data, &[1.0, 2.0, 3.0]);
335
336 store.remove(&data_ref).unwrap();
337 assert!(!store.exists(&data_ref).unwrap());
338
339 let _ = std::fs::remove_dir_all(&dir);
340 }
341
342 #[test]
343 fn inline_data_ref() {
344 let dir = std::env::temp_dir().join("soma-ds-test-inline");
345 let store = LocalDataStore::new(&dir);
346
347 let data_ref = DataRef::Inline {
348 value: Value::tensor(vec![42.0], vec![1]),
349 };
350
351 assert!(store.exists(&data_ref).unwrap());
352 let v = store.get(&data_ref).unwrap();
353 let (data, _) = v.as_tensor().unwrap();
354 assert_eq!(data, &[42.0]);
355
356 let _ = std::fs::remove_dir_all(&dir);
357 }
358
359 #[test]
360 fn storage_config_serde() {
361 let s3 = StorageConfig::S3 {
362 bucket: "my-lab".into(),
363 prefix: "experiments/".into(),
364 region: Some("eu-west-1".into()),
365 endpoint: None,
366 };
367 let json = serde_json::to_string(&s3).unwrap();
368 assert!(json.contains("my-lab"));
369
370 let local = StorageConfig::Local {
371 base_path: "/data".into(),
372 };
373 let json = serde_json::to_string(&local).unwrap();
374 assert!(json.contains("/data"));
375 }
376
377 #[test]
378 fn data_ref_serde() {
379 let refs = vec![
380 DataRef::Local {
381 path: "/tmp/x".into(),
382 },
383 DataRef::S3 {
384 bucket: "b".into(),
385 key: "k".into(),
386 region: None,
387 },
388 DataRef::Cached {
389 cache_key: CacheKey::hash_data(b"x"),
390 },
391 DataRef::Inline {
392 value: Value::Empty,
393 },
394 DataRef::Zarr {
395 bucket: "b".into(),
396 array_path: "data/abc".into(),
397 region: None,
398 },
399 ];
400 for r in &refs {
401 let json = serde_json::to_string(r).unwrap();
402 let _: DataRef = serde_json::from_str(&json).unwrap();
403 }
404 }
405
406 #[test]
407 fn slice_tensor_rows_basic() {
408 let v = Value::tensor(
410 vec![
411 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0,
412 ],
413 vec![4, 3],
414 );
415 let sliced = slice_tensor_rows(&v, 1, 2).unwrap();
417 let (data, shape) = sliced.as_tensor().unwrap();
418 assert_eq!(shape, &[2, 3]);
419 assert_eq!(data, &[4.0, 5.0, 6.0, 7.0, 8.0, 9.0]);
420 }
421
422 #[test]
423 fn slice_tensor_rows_single() {
424 let v = Value::tensor(vec![10.0, 20.0, 30.0], vec![3]);
425 let sliced = slice_tensor_rows(&v, 1, 1).unwrap();
426 let (data, shape) = sliced.as_tensor().unwrap();
427 assert_eq!(shape, &[1]);
428 assert_eq!(data, &[20.0]);
429 }
430
431 #[test]
432 fn slice_tensor_rows_out_of_bounds() {
433 let v = Value::tensor(vec![1.0, 2.0, 3.0], vec![3]);
434 assert!(slice_tensor_rows(&v, 2, 5).is_err());
435 }
436
437 #[test]
438 fn store_meta_from_tensor() {
439 let v = Value::tensor(vec![0.0; 12], vec![4, 3]);
440 let meta = StoreMeta::from_value(&v);
441 assert_eq!(meta.total_rows, 4);
442 assert_eq!(meta.shape_tail, vec![3]);
443 assert_eq!(meta.dtype, "tensor");
444 }
445
446 #[test]
447 fn store_meta_from_json() {
448 let v = Value::json(serde_json::json!({"a": 1}));
449 let meta = StoreMeta::from_value(&v);
450 assert_eq!(meta.dtype, "json");
451 assert_eq!(meta.total_rows, 1);
452 }
453
454 #[test]
455 fn default_get_rows_on_local_store() {
456 let dir = std::env::temp_dir().join("soma-ds-test-getrows");
457 let store = LocalDataStore::new(&dir);
458
459 let key = CacheKey::hash_data(b"rows_test");
460 let value = Value::tensor(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![3, 2]);
461 let data_ref = store.put(&key, &value).unwrap();
462
463 let sliced = store.get_rows(&data_ref, 1, 2).unwrap();
465 let (data, shape) = sliced.as_tensor().unwrap();
466 assert_eq!(shape, &[2, 2]);
467 assert_eq!(data, &[3.0, 4.0, 5.0, 6.0]);
468
469 let _ = std::fs::remove_dir_all(&dir);
470 }
471
472 #[test]
473 fn default_meta_on_local_store() {
474 let dir = std::env::temp_dir().join("soma-ds-test-meta");
475 let store = LocalDataStore::new(&dir);
476
477 let key = CacheKey::hash_data(b"meta_test");
478 let value = Value::tensor(vec![0.0; 20], vec![5, 4]);
479 let data_ref = store.put(&key, &value).unwrap();
480
481 let meta = store.meta(&data_ref).unwrap();
482 assert_eq!(meta.total_rows, 5);
483 assert_eq!(meta.shape_tail, vec![4]);
484 assert_eq!(meta.dtype, "tensor");
485
486 let _ = std::fs::remove_dir_all(&dir);
487 }
488
489 #[test]
490 fn zarr_storage_config_serde() {
491 let zarr = StorageConfig::Zarr {
492 bucket: "soma-research".into(),
493 prefix: "data/".into(),
494 region: None,
495 endpoint: Some("s3.eu-central-003.backblazeb2.com".into()),
496 chunk_rows: 1024,
497 };
498 let json = serde_json::to_string(&zarr).unwrap();
499 assert!(json.contains("soma-research"));
500 assert!(json.contains("1024"));
501 let _: StorageConfig = serde_json::from_str(&json).unwrap();
502 }
503}