1use std::sync::Arc;
2
3use crate::{
4 LadduDataError, LadduDataResult,
5 data::{EventBatch, OwnedEvent},
6 io::{
7 DataFragment, EventSink, EventSource, FragmentedSource, ReadPlan, SourceCapabilities,
8 WritePlan, fragmented_batches,
9 },
10 schema::Schema,
11};
12
13#[derive(Clone, Debug)]
15pub struct MemorySource {
16 schema: Arc<Schema>,
17 batches: Arc<[EventBatch]>,
18}
19
20#[derive(Clone, Copy, Debug)]
22pub struct MemoryFragmentKey {
23 batch_index: usize,
24}
25
26impl MemorySource {
27 pub fn new(batch: EventBatch) -> Self {
29 Self {
30 schema: Arc::clone(batch.schema()),
31 batches: Arc::from([batch]),
32 }
33 }
34
35 pub fn from_batches(batches: Vec<EventBatch>) -> LadduDataResult<Self> {
42 if batches.is_empty() {
43 return Err(LadduDataError::InvalidArgument(
44 "memory source requires at least one batch",
45 ));
46 }
47
48 let schema = Arc::clone(batches[0].schema());
49
50 for batch in &batches {
51 if batch.schema().as_ref() != schema.as_ref() {
52 return Err(LadduDataError::Schema(
53 "memory source batches have different schemas".into(),
54 ));
55 }
56 }
57
58 Ok(Self {
59 schema,
60 batches: batches.into(),
61 })
62 }
63
64 pub fn from_events<I>(schema: Arc<Schema>, events: I) -> LadduDataResult<Self>
71 where
72 I: IntoIterator<Item = OwnedEvent>,
73 {
74 let batch = EventBatch::from_events(schema, events)?;
75 Ok(Self::new(batch))
76 }
77
78 pub fn schema_arc(&self) -> &Arc<Schema> {
80 &self.schema
81 }
82
83 pub fn batches_slice(&self) -> &[EventBatch] {
85 &self.batches
86 }
87
88 pub fn into_batches(self) -> Arc<[EventBatch]> {
90 self.batches
91 }
92
93 pub fn into_batch(self) -> LadduDataResult<EventBatch> {
100 EventBatch::concat(&self.batches)
101 }
102}
103
104impl EventSource for MemorySource {
105 fn schema(&self) -> LadduDataResult<Arc<Schema>> {
106 Ok(Arc::clone(&self.schema))
107 }
108
109 fn capabilities(&self) -> SourceCapabilities {
110 SourceCapabilities {
111 exact_len: true,
112 exact_weighted_total: true,
113 random_access: true,
114 deterministic_partitioning: true,
115 predicate_pushdown: false,
116 projection_pushdown: false,
117 streaming: false,
118 }
119 }
120
121 fn num_events(&self) -> LadduDataResult<Option<u64>> {
122 Ok(Some(self.batches.iter().map(|b| b.len() as u64).sum()))
123 }
124
125 fn weighted_total(&self) -> LadduDataResult<Option<f64>> {
126 let total = self
127 .batches
128 .iter()
129 .map(|batch| (0..batch.len()).map(|i| batch.weights_at(i)).sum::<f64>())
130 .sum();
131
132 Ok(Some(total))
133 }
134
135 fn batches(
136 &self,
137 plan: ReadPlan,
138 ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
139 fragmented_batches(Arc::new(self.clone()), plan)
140 }
141}
142
143impl FragmentedSource for MemorySource {
144 type Key = MemoryFragmentKey;
145
146 fn fragments(&self) -> LadduDataResult<Vec<DataFragment<Self::Key>>> {
147 let mut fragments = Vec::with_capacity(self.batches.len());
148 let mut global_start = 0_u64;
149
150 for (batch_index, batch) in self.batches.iter().enumerate() {
151 let rows = batch.len() as u64;
152
153 fragments.push(DataFragment {
154 key: MemoryFragmentKey { batch_index },
155 global_start,
156 rows,
157 });
158
159 global_start += rows;
160 }
161
162 Ok(fragments)
163 }
164
165 fn read_fragment_range(
166 &self,
167 key: &Self::Key,
168 local_start: usize,
169 local_len: usize,
170 chunk_size: Option<usize>,
171 ) -> LadduDataResult<Box<dyn Iterator<Item = LadduDataResult<EventBatch>> + Send>> {
172 if matches!(chunk_size, Some(0)) {
173 return Err(LadduDataError::InvalidArgument(
174 "chunk_size must be nonzero",
175 ));
176 }
177
178 let batch = self
179 .batches
180 .get(key.batch_index)
181 .ok_or_else(|| LadduDataError::Source("invalid memory batch index".into()))?
182 .clone();
183
184 let end = local_start
185 .checked_add(local_len)
186 .ok_or(LadduDataError::InvalidArgument(
187 "slice range overflows usize",
188 ))?;
189
190 if end > batch.len() {
191 return Err(LadduDataError::InvalidArgument(
192 "memory fragment range exceeds batch length",
193 ));
194 }
195
196 Ok(Box::new(MemoryRangeIter {
197 batch,
198 pos: local_start,
199 end,
200 chunk_size: chunk_size.unwrap_or(local_len.max(1)),
201 }))
202 }
203}
204
205struct MemoryRangeIter {
206 batch: EventBatch,
207 pos: usize,
208 end: usize,
209 chunk_size: usize,
210}
211
212impl Iterator for MemoryRangeIter {
213 type Item = LadduDataResult<EventBatch>;
214
215 fn next(&mut self) -> Option<Self::Item> {
216 if self.pos >= self.end {
217 return None;
218 }
219
220 let next = (self.pos + self.chunk_size).min(self.end);
221 let batch = self.batch.slice(self.pos, next);
222 self.pos = next;
223
224 Some(Ok(batch))
225 }
226}
227
228#[derive(Clone, Debug, Default)]
230pub struct MemorySink {
231 schema: Option<Arc<Schema>>,
232 batches: Vec<EventBatch>,
233 finished: bool,
234}
235
236impl MemorySink {
237 pub fn new() -> Self {
239 Self::default()
240 }
241
242 pub fn schema(&self) -> Option<&Arc<Schema>> {
244 self.schema.as_ref()
245 }
246
247 pub fn batches(&self) -> &[EventBatch] {
249 &self.batches
250 }
251
252 pub fn into_batches(self) -> Vec<EventBatch> {
254 self.batches
255 }
256
257 pub fn into_source(self) -> LadduDataResult<MemorySource> {
264 MemorySource::from_batches(self.batches)
265 }
266
267 pub fn into_batch(self) -> LadduDataResult<EventBatch> {
274 EventBatch::concat(&self.batches)
275 }
276
277 pub fn clear(&mut self) {
279 self.schema = None;
280 self.batches.clear();
281 self.finished = false;
282 }
283}
284
285impl EventSink for MemorySink {
286 fn retains_batches(&self) -> bool {
287 true
288 }
289
290 fn begin(&mut self, schema: Arc<Schema>, _plan: WritePlan) -> LadduDataResult<()> {
291 self.schema = Some(schema);
292 self.batches.clear();
293 self.finished = false;
294 Ok(())
295 }
296
297 fn write_batch(&mut self, batch: &EventBatch) -> LadduDataResult<()> {
298 let schema = self
299 .schema
300 .as_ref()
301 .ok_or_else(|| LadduDataError::Sink("memory sink not initialized".into()))?;
302
303 if schema.as_ref() != batch.schema().as_ref() {
304 return Err(LadduDataError::Sink(
305 "batch schema does not match memory sink schema".into(),
306 ));
307 }
308
309 self.batches.push(batch.clone());
310 Ok(())
311 }
312
313 fn finish(&mut self) -> LadduDataResult<()> {
314 self.finished = true;
315 Ok(())
316 }
317}
318
319#[cfg(test)]
320mod tests {
321 use laddu_physics::vectors::RealVec4;
322
323 use super::*;
324 use crate::data::EventBatchBuilder;
325
326 fn v(x: f64) -> RealVec4 {
327 RealVec4 {
328 e: x,
329 px: x,
330 py: x,
331 pz: x,
332 }
333 }
334
335 fn schema() -> Arc<Schema> {
336 Arc::new(Schema::new(["p"], ["id"], true).unwrap())
337 }
338
339 fn batch(start: usize, len: usize) -> EventBatch {
340 let schema = schema();
341 let mut builder = EventBatchBuilder::with_capacity(schema, len);
342
343 for i in start..start + len {
344 builder
345 .push_weighted([v(i as f64)], [i as f64], 1.0 + i as f64)
346 .unwrap();
347 }
348
349 builder.finish().unwrap()
350 }
351
352 fn concat_scalars(batches: Vec<EventBatch>) -> Vec<f64> {
353 EventBatch::concat(&batches)
354 .unwrap()
355 .scalar_column(0)
356 .to_vec()
357 }
358
359 #[test]
360 fn memory_source_reports_exact_capabilities_counts_and_weighted_total() {
361 let source = MemorySource::from_batches(vec![batch(0, 2), batch(2, 3)]).unwrap();
362
363 let caps = source.capabilities();
364
365 assert!(caps.exact_len);
366 assert!(caps.exact_weighted_total);
367 assert!(caps.random_access);
368 assert!(caps.deterministic_partitioning);
369 assert!(!caps.streaming);
370
371 assert_eq!(source.num_events().unwrap(), Some(5));
372 assert_eq!(
373 source.weighted_total().unwrap(),
374 Some(1.0 + 2.0 + 3.0 + 4.0 + 5.0)
375 );
376 }
377
378 #[test]
379 fn memory_source_coalesces_by_default_but_respects_chunked_read_plan() {
380 let source = MemorySource::from_batches(vec![batch(0, 2), batch(2, 3)]).unwrap();
381
382 let default_batches: Vec<EventBatch> = source
383 .batches(ReadPlan::default())
384 .unwrap()
385 .map(Result::unwrap)
386 .collect();
387
388 assert_eq!(default_batches.len(), 1);
389 assert_eq!(
390 default_batches[0].scalar_column(0),
391 &[0.0, 1.0, 2.0, 3.0, 4.0]
392 );
393
394 let chunked_batches: Vec<EventBatch> = source
395 .batches(ReadPlan {
396 chunk_size: Some(2),
397 ..ReadPlan::default()
398 })
399 .unwrap()
400 .map(Result::unwrap)
401 .collect();
402
403 assert_eq!(
404 chunked_batches
405 .iter()
406 .map(EventBatch::len)
407 .collect::<Vec<_>>(),
408 vec![2, 2, 1]
409 );
410 assert_eq!(
411 concat_scalars(chunked_batches),
412 vec![0.0, 1.0, 2.0, 3.0, 4.0]
413 );
414 }
415
416 #[test]
417 fn memory_source_rejects_empty_or_schema_mismatched_batches() {
418 assert!(matches!(
419 MemorySource::from_batches(vec![]),
420 Err(LadduDataError::InvalidArgument(_))
421 ));
422
423 let first = batch(0, 1);
424
425 let other_schema = Arc::new(Schema::new(["q"], ["id"], true).unwrap());
426 let mut builder = EventBatchBuilder::new(other_schema);
427 builder.push_weighted([v(10.0)], [10.0], 1.0).unwrap();
428 let second = builder.finish().unwrap();
429
430 assert!(matches!(
431 MemorySource::from_batches(vec![first, second]),
432 Err(LadduDataError::Schema(_))
433 ));
434 }
435
436 #[test]
437 fn memory_sink_validates_lifecycle_schema_and_can_be_reused() {
438 let mut sink = MemorySink::new();
439 let first = batch(0, 2);
440
441 assert!(matches!(
442 sink.write_batch(&first),
443 Err(LadduDataError::Sink(_))
444 ));
445
446 sink.begin(Arc::clone(first.schema()), WritePlan::default())
447 .unwrap();
448 sink.write_batch(&first).unwrap();
449 sink.finish().unwrap();
450
451 assert_eq!(sink.batches().len(), 1);
452 assert_eq!(sink.batches()[0].scalar_column(0), &[0.0, 1.0]);
453
454 let mismatched_schema = Arc::new(Schema::new(["other"], ["id"], true).unwrap());
455 let mut builder = EventBatchBuilder::new(mismatched_schema);
456 builder.push_weighted([v(9.0)], [9.0], 9.0).unwrap();
457 let mismatched = builder.finish().unwrap();
458
459 assert!(matches!(
460 sink.write_batch(&mismatched),
461 Err(LadduDataError::Sink(_))
462 ));
463
464 sink.clear();
465 assert!(sink.schema().is_none());
466 assert!(sink.batches().is_empty());
467 }
468
469 #[test]
470 fn memory_sink_into_source_roundtrips_captured_batches() {
471 let mut sink = MemorySink::new();
472 let first = batch(0, 2);
473 let second = batch(2, 2);
474
475 sink.begin(Arc::clone(first.schema()), WritePlan::default())
476 .unwrap();
477 sink.write_batch(&first).unwrap();
478 sink.write_batch(&second).unwrap();
479 sink.finish().unwrap();
480
481 let source = sink.into_source().unwrap();
482 let merged = source.into_batch().unwrap();
483
484 assert_eq!(merged.scalar_column(0), &[0.0, 1.0, 2.0, 3.0]);
485 assert_eq!(merged.weights_column().unwrap(), &[1.0, 2.0, 3.0, 4.0]);
486 }
487}