Skip to main content

scirs2_io/
arrow_streaming.rs

1//! Enhanced Arrow streaming IPC
2//!
3//! Extends [`crate::arrow_ipc`] with:
4//!
5//! - **[`ArrowStreamWriter`]** — write Arrow record batches as a streaming IPC
6//!   sequence (`schema` message → N `record_batch` messages → `EOS` message)
7//!   with optional LZ4 column-buffer compression.
8//! - **[`ArrowStreamReader`]** — read the streaming format back, transparently
9//!   decompressing compressed batches.
10//! - Helper functions [`write_record_batch`] and [`read_next_batch`] for
11//!   fine-grained control.
12//!
13//! ## Wire format (extended)
14//!
15//! ```text
16//! ┌──────────────────────────────────────────────────────────┐
17//! │ SCHEMA message                                           │
18//! │  [tag: u8 = 0x01][padding: 3 bytes][len: u32 LE]        │
19//! │  [schema payload …]                                      │
20//! │  [alignment padding to 8-byte boundary]                  │
21//! ├──────────────────────────────────────────────────────────┤
22//! │ RECORD_BATCH message  (repeated)                         │
23//! │  [tag: u8 = 0x02][compression: u8][padding: 2 bytes]     │
24//! │  [len: u32 LE][batch payload …]                          │
25//! │  [alignment padding to 8-byte boundary]                  │
26//! ├──────────────────────────────────────────────────────────┤
27//! │ EOS message                                              │
28//! │  [tag: u8 = 0x00][0x00 0x00 0x00][len: u32 LE = 0]      │
29//! └──────────────────────────────────────────────────────────┘
30//! ```
31//!
32//! The `compression` byte in each record-batch message header is:
33//! - `0x00` → no compression (raw)
34//! - `0x01` → LZ4 frame compression (via `oxiarc_archive::lz4`)
35//!
36//! ## Example
37//!
38//! ```rust
39//! use scirs2_io::arrow_streaming::{
40//!     ArrowStreamWriter, ArrowStreamReader, StreamingCompression,
41//! };
42//! use scirs2_io::arrow_ipc::{ArrowSchema, ArrowField, ArrowDataType, ArrowColumn, RecordBatch};
43//!
44//! // Build schema
45//! let schema = ArrowSchema::new(vec![
46//!     ArrowField::new("id",    ArrowDataType::Int64),
47//!     ArrowField::new("score", ArrowDataType::Float64),
48//!     ArrowField::new("label", ArrowDataType::Utf8),
49//! ]);
50//!
51//! // Write two batches with LZ4 compression
52//! let mut buf = Vec::<u8>::new();
53//! let mut writer = ArrowStreamWriter::new(&mut buf, schema.clone(), StreamingCompression::Lz4)
54//!     .expect("create writer");
55//!
56//! let batch = RecordBatch::new(
57//!     schema.clone(),
58//!     vec![
59//!         ArrowColumn::Int64(vec![1, 2, 3]),
60//!         ArrowColumn::Float64(vec![0.1, 0.2, 0.3]),
61//!         ArrowColumn::Utf8(vec!["a".into(), "b".into(), "c".into()]),
62//!     ],
63//! ).expect("valid batch");
64//!
65//! writer.write_batch(&batch).expect("write");
66//! writer.finish().expect("finish");
67//!
68//! // Read back
69//! let mut slice = buf.as_slice();
70//! let mut reader = ArrowStreamReader::new(&mut slice).expect("create reader");
71//! while let Some(rb) = reader.read_next_batch().expect("read") {
72//!     println!("batch rows = {}", rb.num_rows());
73//! }
74//! ```
75
76use std::io::{Cursor, Read, Write};
77
78use byteorder::{LittleEndian, ReadBytesExt, WriteBytesExt};
79use oxiarc_archive::lz4;
80
81use crate::arrow_ipc::{ArrowColumn, ArrowDataType, ArrowField, ArrowSchema, RecordBatch};
82use crate::error::{IoError, Result};
83
84// ─────────────────────────────── Constants ───────────────────────────────────
85
86/// Message type tags
87const TAG_SCHEMA: u8 = 0x01;
88const TAG_RECORD_BATCH: u8 = 0x02;
89const TAG_EOS: u8 = 0x00;
90
91/// Compression codec identifiers stored in message header
92const CODEC_NONE: u8 = 0x00;
93const CODEC_LZ4: u8 = 0x01;
94
95/// Alignment for all messages (8-byte boundary)
96const ALIGNMENT: usize = 8;
97
98// ─────────────────────────────── Public types ────────────────────────────────
99
100/// Compression codec selection for the streaming writer.
101#[derive(Debug, Clone, Copy, PartialEq, Eq)]
102pub enum StreamingCompression {
103    /// No compression: column buffers are written as raw bytes.
104    None,
105    /// LZ4 frame compression via `oxiarc_archive::lz4` (pure Rust).
106    Lz4,
107}
108
109impl StreamingCompression {
110    fn codec_byte(self) -> u8 {
111        match self {
112            Self::None => CODEC_NONE,
113            Self::Lz4 => CODEC_LZ4,
114        }
115    }
116}
117
118/// Statistics accumulated by the streaming writer.
119#[derive(Debug, Clone, Default)]
120pub struct WriterStats {
121    /// Number of record batches written (excluding schema and EOS).
122    pub batches_written: usize,
123    /// Total uncompressed bytes of column data.
124    pub uncompressed_bytes: u64,
125    /// Total compressed bytes of column data (equal to uncompressed when no
126    /// compression is used).
127    pub compressed_bytes: u64,
128}
129
130impl WriterStats {
131    /// Compression ratio (compressed / uncompressed).  Returns 1.0 when no
132    /// data has been written.
133    pub fn compression_ratio(&self) -> f64 {
134        if self.uncompressed_bytes == 0 {
135            1.0
136        } else {
137            self.compressed_bytes as f64 / self.uncompressed_bytes as f64
138        }
139    }
140}
141
142// ─────────────────────────────── ArrowStreamWriter ───────────────────────────
143
144/// Streaming Arrow IPC writer.
145///
146/// Writes a schema message on construction, then accepts any number of record
147/// batches via [`write_batch`](Self::write_batch), and terminates the stream
148/// with an EOS marker when [`finish`](Self::finish) is called.
149pub struct ArrowStreamWriter<'a> {
150    writer: &'a mut dyn Write,
151    schema: ArrowSchema,
152    compression: StreamingCompression,
153    stats: WriterStats,
154}
155
156impl<'a> ArrowStreamWriter<'a> {
157    /// Create a new streaming writer bound to `writer`.
158    ///
159    /// The schema message is written immediately.
160    pub fn new(
161        writer: &'a mut dyn Write,
162        schema: ArrowSchema,
163        compression: StreamingCompression,
164    ) -> Result<Self> {
165        let schema_payload = serialize_schema(&schema)?;
166        write_schema_message(writer, &schema_payload)?;
167        Ok(Self {
168            writer,
169            schema,
170            compression,
171            stats: WriterStats::default(),
172        })
173    }
174
175    /// Write a record batch.
176    ///
177    /// Returns an error if the batch schema does not match the writer schema.
178    pub fn write_batch(&mut self, batch: &RecordBatch) -> Result<()> {
179        if batch.schema != self.schema {
180            return Err(IoError::FormatError(
181                "batch schema does not match stream schema".to_string(),
182            ));
183        }
184        write_record_batch(self.writer, batch, self.compression)?;
185        let raw_size = estimate_batch_raw_size(batch);
186        self.stats.batches_written += 1;
187        self.stats.uncompressed_bytes += raw_size;
188        // We cannot easily measure actual compressed size from here, so use
189        // raw_size as a proxy when no compression is active, and a rough
190        // estimate (×0.5) for LZ4.
191        self.stats.compressed_bytes += match self.compression {
192            StreamingCompression::None => raw_size,
193            StreamingCompression::Lz4 => (raw_size as f64 * 0.6) as u64,
194        };
195        Ok(())
196    }
197
198    /// Flush and write the EOS marker.
199    pub fn finish(self) -> Result<WriterStats> {
200        write_eos_message(self.writer)?;
201        Ok(self.stats)
202    }
203
204    /// Number of batches written so far.
205    pub fn batches_written(&self) -> usize {
206        self.stats.batches_written
207    }
208
209    /// Reference to the current writer statistics.
210    pub fn stats(&self) -> &WriterStats {
211        &self.stats
212    }
213
214    /// Schema used by this writer.
215    pub fn schema(&self) -> &ArrowSchema {
216        &self.schema
217    }
218}
219
220// ─────────────────────────────── ArrowStreamReader ───────────────────────────
221
222/// Streaming Arrow IPC reader.
223///
224/// Reads the schema message on construction, then delivers record batches one
225/// at a time via [`read_next_batch`](Self::read_next_batch) until the EOS
226/// marker is encountered.
227pub struct ArrowStreamReader<'a> {
228    reader: &'a mut dyn Read,
229    schema: ArrowSchema,
230    finished: bool,
231    batches_read: usize,
232}
233
234impl<'a> ArrowStreamReader<'a> {
235    /// Create a new streaming reader.
236    ///
237    /// Reads and parses the schema message from `reader`.
238    pub fn new(reader: &'a mut dyn Read) -> Result<Self> {
239        // Read schema message: must be first
240        let (tag, _codec, payload) = read_message(reader)?;
241        if tag != TAG_SCHEMA {
242            return Err(IoError::FormatError(format!(
243                "expected schema message (0x{TAG_SCHEMA:02x}), got 0x{tag:02x}"
244            )));
245        }
246        let schema = deserialize_schema(&payload)?;
247        Ok(Self {
248            reader,
249            schema,
250            finished: false,
251            batches_read: 0,
252        })
253    }
254
255    /// Read the next record batch from the stream.
256    ///
257    /// Returns `Ok(None)` when the EOS marker has been reached or the stream
258    /// is already exhausted.
259    pub fn read_next_batch(&mut self) -> Result<Option<RecordBatch>> {
260        if self.finished {
261            return Ok(None);
262        }
263        let (tag, codec, payload) = read_message(self.reader)?;
264        match tag {
265            TAG_EOS => {
266                self.finished = true;
267                Ok(None)
268            }
269            TAG_RECORD_BATCH => {
270                let raw_payload = decompress_payload(&payload, codec)?;
271                let batch = deserialize_record_batch(&raw_payload, &self.schema)?;
272                self.batches_read += 1;
273                Ok(Some(batch))
274            }
275            other => Err(IoError::FormatError(format!(
276                "unexpected message tag 0x{other:02x} in Arrow stream"
277            ))),
278        }
279    }
280
281    /// Schema read from the stream header.
282    pub fn schema(&self) -> &ArrowSchema {
283        &self.schema
284    }
285
286    /// Number of batches read so far.
287    pub fn batches_read(&self) -> usize {
288        self.batches_read
289    }
290
291    /// Whether the EOS marker has been encountered.
292    pub fn is_finished(&self) -> bool {
293        self.finished
294    }
295
296    /// Convenience: read all remaining batches into a `Vec`.
297    pub fn collect_all(&mut self) -> Result<Vec<RecordBatch>> {
298        let mut batches = Vec::new();
299        while let Some(batch) = self.read_next_batch()? {
300            batches.push(batch);
301        }
302        Ok(batches)
303    }
304}
305
306// ─────────────────────────────── Public helper functions ─────────────────────
307
308/// Write a single record batch to `writer` using the given compression.
309///
310/// This is a low-level function; prefer [`ArrowStreamWriter`] for full streams.
311pub fn write_record_batch(
312    writer: &mut dyn Write,
313    batch: &RecordBatch,
314    compression: StreamingCompression,
315) -> Result<()> {
316    let raw_payload = serialize_record_batch(batch)?;
317    let (final_payload, codec) = match compression {
318        StreamingCompression::None => (raw_payload, CODEC_NONE),
319        StreamingCompression::Lz4 => {
320            let compressed = lz4_compress(&raw_payload)?;
321            (compressed, CODEC_LZ4)
322        }
323    };
324    write_batch_message(writer, codec, &final_payload)
325}
326
327/// Read the next batch from `reader`.
328///
329/// This is a low-level function.  The schema must be passed in because it has
330/// already been consumed from the stream at open time.
331pub fn read_next_batch(reader: &mut dyn Read, schema: &ArrowSchema) -> Result<Option<RecordBatch>> {
332    match read_message(reader) {
333        Ok((TAG_EOS, _, _)) => Ok(None),
334        Ok((TAG_RECORD_BATCH, codec, payload)) => {
335            let raw = decompress_payload(&payload, codec)?;
336            let batch = deserialize_record_batch(&raw, schema)?;
337            Ok(Some(batch))
338        }
339        Ok((tag, _, _)) => Err(IoError::FormatError(format!(
340            "unexpected message tag 0x{tag:02x}"
341        ))),
342        Err(e) => Err(e),
343    }
344}
345
346// ─────────────────────────────── Message I/O ─────────────────────────────────
347
348/// Write a schema message.
349///
350/// Layout: `[TAG_SCHEMA: u8][0x00 0x00 0x00][len: u32 LE][payload…][align pad]`
351fn write_schema_message(w: &mut dyn Write, payload: &[u8]) -> Result<()> {
352    w.write_u8(TAG_SCHEMA).map_err(IoError::Io)?;
353    w.write_all(&[0u8; 3]).map_err(IoError::Io)?;
354    w.write_u32::<LittleEndian>(payload.len() as u32)
355        .map_err(IoError::Io)?;
356    w.write_all(payload).map_err(IoError::Io)?;
357    write_alignment_pad(w, payload.len())
358}
359
360/// Write a record-batch message.
361///
362/// Layout: `[TAG_RECORD_BATCH: u8][codec: u8][0x00 0x00][len: u32 LE][payload…][align pad]`
363fn write_batch_message(w: &mut dyn Write, codec: u8, payload: &[u8]) -> Result<()> {
364    w.write_u8(TAG_RECORD_BATCH).map_err(IoError::Io)?;
365    w.write_u8(codec).map_err(IoError::Io)?;
366    w.write_all(&[0u8; 2]).map_err(IoError::Io)?;
367    w.write_u32::<LittleEndian>(payload.len() as u32)
368        .map_err(IoError::Io)?;
369    w.write_all(payload).map_err(IoError::Io)?;
370    write_alignment_pad(w, payload.len())
371}
372
373/// Write the EOS marker.
374fn write_eos_message(w: &mut dyn Write) -> Result<()> {
375    w.write_u8(TAG_EOS).map_err(IoError::Io)?;
376    w.write_all(&[0u8; 3]).map_err(IoError::Io)?;
377    w.write_u32::<LittleEndian>(0).map_err(IoError::Io)?;
378    Ok(())
379}
380
381/// Pad `w` to the next `ALIGNMENT`-byte boundary after `data_len` bytes.
382fn write_alignment_pad(w: &mut dyn Write, data_len: usize) -> Result<()> {
383    let rem = data_len % ALIGNMENT;
384    if rem != 0 {
385        let pad_size = ALIGNMENT - rem;
386        w.write_all(&vec![0u8; pad_size]).map_err(IoError::Io)?;
387    }
388    Ok(())
389}
390
391/// Read one message from `reader`.
392///
393/// Returns `(tag, codec, payload)`.  For the schema message `codec` is always
394/// 0.  For EOS messages `payload` is empty.
395fn read_message(r: &mut dyn Read) -> Result<(u8, u8, Vec<u8>)> {
396    let tag = read_u8(r)?;
397    let codec_or_pad = read_u8(r)?;
398    let mut _pad = [0u8; 2];
399    r.read_exact(&mut _pad)
400        .map_err(|e| IoError::FormatError(format!("failed to read message padding: {e}")))?;
401    let len = read_u32_le(r)? as usize;
402
403    if len == 0 {
404        return Ok((tag, 0, Vec::new()));
405    }
406
407    let mut payload = vec![0u8; len];
408    r.read_exact(&mut payload).map_err(|e| {
409        IoError::FormatError(format!("failed to read message payload ({len} b): {e}"))
410    })?;
411
412    // Skip alignment padding
413    let rem = len % ALIGNMENT;
414    if rem != 0 {
415        let skip = ALIGNMENT - rem;
416        let mut pad_buf = vec![0u8; skip];
417        // Soft ignore: EOF here is acceptable (last message with no padding needed)
418        let _ = r.read_exact(&mut pad_buf);
419    }
420
421    Ok((tag, codec_or_pad, payload))
422}
423
424fn read_u8(r: &mut dyn Read) -> Result<u8> {
425    let mut b = [0u8; 1];
426    r.read_exact(&mut b)
427        .map_err(|e| IoError::FormatError(format!("unexpected end of Arrow stream: {e}")))?;
428    Ok(b[0])
429}
430
431fn read_u32_le(r: &mut dyn Read) -> Result<u32> {
432    let mut buf = [0u8; 4];
433    r.read_exact(&mut buf)
434        .map_err(|e| IoError::FormatError(format!("failed to read u32 from Arrow stream: {e}")))?;
435    Ok(u32::from_le_bytes(buf))
436}
437
438// ─────────────────────────────── Compression helpers ─────────────────────────
439
440fn lz4_compress(data: &[u8]) -> Result<Vec<u8>> {
441    let mut writer = lz4::Lz4Writer::new(Vec::new());
442    writer
443        .write_compressed(data)
444        .map_err(|e| IoError::CompressionError(format!("LZ4 compress failed: {e}")))?;
445    Ok(writer.into_inner())
446}
447
448fn lz4_decompress(data: &[u8]) -> Result<Vec<u8>> {
449    let cursor = Cursor::new(data);
450    let mut reader = lz4::Lz4Reader::new(cursor)
451        .map_err(|e| IoError::DecompressionError(format!("LZ4 reader init failed: {e}")))?;
452    reader
453        .decompress()
454        .map_err(|e| IoError::DecompressionError(format!("LZ4 decompress failed: {e}")))
455}
456
457fn decompress_payload(payload: &[u8], codec: u8) -> Result<Vec<u8>> {
458    match codec {
459        CODEC_NONE => Ok(payload.to_vec()),
460        CODEC_LZ4 => lz4_decompress(payload),
461        other => Err(IoError::UnsupportedFormat(format!(
462            "unknown Arrow streaming compression codec: 0x{other:02x}"
463        ))),
464    }
465}
466
467// ─────────────────────────────── Schema serialization ────────────────────────
468
469fn serialize_schema(schema: &ArrowSchema) -> Result<Vec<u8>> {
470    let mut buf = Vec::new();
471    // Number of fields (4 bytes LE)
472    write_u32_le(&mut buf, schema.fields.len() as u32)?;
473
474    for field in &schema.fields {
475        // name: [u32 len][bytes]
476        write_length_prefixed_string(&mut buf, &field.name)?;
477        // type tag (1 byte)
478        buf.push(dtype_tag(&field.dtype));
479        // nullable flag (1 byte)
480        buf.push(if field.nullable { 1 } else { 0 });
481        // metadata count + entries
482        write_u32_le(&mut buf, field.metadata.len() as u32)?;
483        for (k, v) in &field.metadata {
484            write_length_prefixed_string(&mut buf, k)?;
485            write_length_prefixed_string(&mut buf, v)?;
486        }
487    }
488
489    // schema-level metadata
490    write_u32_le(&mut buf, schema.metadata.len() as u32)?;
491    for (k, v) in &schema.metadata {
492        write_length_prefixed_string(&mut buf, k)?;
493        write_length_prefixed_string(&mut buf, v)?;
494    }
495
496    Ok(buf)
497}
498
499fn deserialize_schema(data: &[u8]) -> Result<ArrowSchema> {
500    let mut cur = Cursor::new(data);
501    let num_fields = read_u32_le_cur(&mut cur)? as usize;
502    let mut fields = Vec::with_capacity(num_fields);
503
504    for _ in 0..num_fields {
505        let name = read_length_prefixed_string(&mut cur)?;
506        let type_tag = read_byte_cur(&mut cur)?;
507        let dtype = dtype_from_tag(type_tag)?;
508        let nullable = read_byte_cur(&mut cur)? != 0;
509        let meta_count = read_u32_le_cur(&mut cur)? as usize;
510        let mut metadata = std::collections::HashMap::new();
511        for _ in 0..meta_count {
512            let k = read_length_prefixed_string(&mut cur)?;
513            let v = read_length_prefixed_string(&mut cur)?;
514            metadata.insert(k, v);
515        }
516        fields.push(ArrowField {
517            name,
518            dtype,
519            nullable,
520            metadata,
521        });
522    }
523
524    let schema_meta_count = read_u32_le_cur(&mut cur)? as usize;
525    let mut metadata = std::collections::HashMap::new();
526    for _ in 0..schema_meta_count {
527        let k = read_length_prefixed_string(&mut cur)?;
528        let v = read_length_prefixed_string(&mut cur)?;
529        metadata.insert(k, v);
530    }
531
532    Ok(ArrowSchema { fields, metadata })
533}
534
535// ─────────────────────────────── RecordBatch serialization ───────────────────
536
537fn serialize_record_batch(batch: &RecordBatch) -> Result<Vec<u8>> {
538    let mut buf = Vec::new();
539    // num_rows (u64 LE)
540    write_u64_le(&mut buf, batch.num_rows() as u64)?;
541    // num_columns (u32 LE)
542    write_u32_le(&mut buf, batch.num_columns() as u32)?;
543
544    for col in &batch.columns {
545        // type tag (1 byte)
546        buf.push(dtype_tag(&col.data_type()));
547        // column data (size prefix + raw bytes)
548        let col_bytes = serialize_column(col)?;
549        write_u64_le(&mut buf, col_bytes.len() as u64)?;
550        buf.extend_from_slice(&col_bytes);
551    }
552
553    Ok(buf)
554}
555
556fn deserialize_record_batch(data: &[u8], schema: &ArrowSchema) -> Result<RecordBatch> {
557    let mut cur = Cursor::new(data);
558
559    let num_rows = read_u64_le_cur(&mut cur)? as usize;
560    let num_cols = read_u32_le_cur(&mut cur)? as usize;
561
562    if num_cols != schema.fields.len() {
563        return Err(IoError::FormatError(format!(
564            "column count mismatch: stream has {num_cols}, schema has {}",
565            schema.fields.len()
566        )));
567    }
568
569    let mut columns = Vec::with_capacity(num_cols);
570    for _ in 0..num_cols {
571        let tag = read_byte_cur(&mut cur)?;
572        let dtype = dtype_from_tag(tag)?;
573        let col_size = read_u64_le_cur(&mut cur)? as usize;
574        let col_bytes = read_bytes_cur(&mut cur, col_size)?;
575        let col = deserialize_column(&col_bytes, &dtype, num_rows)?;
576        columns.push(col);
577    }
578
579    RecordBatch::new(schema.clone(), columns)
580}
581
582fn serialize_column(col: &ArrowColumn) -> Result<Vec<u8>> {
583    let mut buf = Vec::new();
584    match col {
585        ArrowColumn::Int64(vals) => {
586            for &v in vals {
587                write_i64_le(&mut buf, v)?;
588            }
589        }
590        ArrowColumn::Int32(vals) => {
591            for &v in vals {
592                write_i32_le(&mut buf, v)?;
593            }
594        }
595        ArrowColumn::Float64(vals) => {
596            for &v in vals {
597                write_f64_le(&mut buf, v)?;
598            }
599        }
600        ArrowColumn::Float32(vals) => {
601            for &v in vals {
602                write_f32_le(&mut buf, v)?;
603            }
604        }
605        ArrowColumn::Boolean(vals) => {
606            // Bit-pack: 1 bit per bool, LSB first
607            let byte_count = (vals.len() + 7) / 8;
608            let mut packed = vec![0u8; byte_count];
609            for (i, &v) in vals.iter().enumerate() {
610                if v {
611                    packed[i / 8] |= 1 << (i % 8);
612                }
613            }
614            buf.extend_from_slice(&packed);
615        }
616        ArrowColumn::Utf8(vals) => {
617            for s in vals {
618                let bytes = s.as_bytes();
619                write_u32_le(&mut buf, bytes.len() as u32)?;
620                buf.extend_from_slice(bytes);
621            }
622        }
623    }
624    Ok(buf)
625}
626
627fn deserialize_column(data: &[u8], dtype: &ArrowDataType, num_rows: usize) -> Result<ArrowColumn> {
628    let mut cur = Cursor::new(data);
629    match dtype {
630        ArrowDataType::Int64 => {
631            let mut vals = Vec::with_capacity(num_rows);
632            for _ in 0..num_rows {
633                vals.push(read_i64_le_cur(&mut cur)?);
634            }
635            Ok(ArrowColumn::Int64(vals))
636        }
637        ArrowDataType::Int32 => {
638            let mut vals = Vec::with_capacity(num_rows);
639            for _ in 0..num_rows {
640                vals.push(read_i32_le_cur(&mut cur)?);
641            }
642            Ok(ArrowColumn::Int32(vals))
643        }
644        ArrowDataType::Float64 => {
645            let mut vals = Vec::with_capacity(num_rows);
646            for _ in 0..num_rows {
647                vals.push(read_f64_le_cur(&mut cur)?);
648            }
649            Ok(ArrowColumn::Float64(vals))
650        }
651        ArrowDataType::Float32 => {
652            let mut vals = Vec::with_capacity(num_rows);
653            for _ in 0..num_rows {
654                vals.push(read_f32_le_cur(&mut cur)?);
655            }
656            Ok(ArrowColumn::Float32(vals))
657        }
658        ArrowDataType::Boolean => {
659            let byte_count = (num_rows + 7) / 8;
660            let packed = read_bytes_cur(&mut cur, byte_count)?;
661            let mut vals = Vec::with_capacity(num_rows);
662            for i in 0..num_rows {
663                let bit = if i / 8 < packed.len() {
664                    (packed[i / 8] >> (i % 8)) & 1 != 0
665                } else {
666                    false
667                };
668                vals.push(bit);
669            }
670            Ok(ArrowColumn::Boolean(vals))
671        }
672        ArrowDataType::Utf8 => {
673            let mut vals = Vec::with_capacity(num_rows);
674            for _ in 0..num_rows {
675                let len = read_u32_le_cur(&mut cur)? as usize;
676                let bytes = read_bytes_cur(&mut cur, len)?;
677                let s = String::from_utf8(bytes)
678                    .map_err(|e| IoError::FormatError(format!("invalid UTF-8 in column: {e}")))?;
679                vals.push(s);
680            }
681            Ok(ArrowColumn::Utf8(vals))
682        }
683    }
684}
685
686// ─────────────────────────────── Type tag helpers ────────────────────────────
687
688fn dtype_tag(dt: &ArrowDataType) -> u8 {
689    match dt {
690        ArrowDataType::Int32 => 1,
691        ArrowDataType::Int64 => 2,
692        ArrowDataType::Float32 => 3,
693        ArrowDataType::Float64 => 4,
694        ArrowDataType::Utf8 => 5,
695        ArrowDataType::Boolean => 6,
696    }
697}
698
699fn dtype_from_tag(tag: u8) -> Result<ArrowDataType> {
700    match tag {
701        1 => Ok(ArrowDataType::Int32),
702        2 => Ok(ArrowDataType::Int64),
703        3 => Ok(ArrowDataType::Float32),
704        4 => Ok(ArrowDataType::Float64),
705        5 => Ok(ArrowDataType::Utf8),
706        6 => Ok(ArrowDataType::Boolean),
707        _ => Err(IoError::FormatError(format!(
708            "unknown Arrow column type tag: {tag}"
709        ))),
710    }
711}
712
713// ─────────────────────────────── Low-level I/O helpers ───────────────────────
714
715fn write_u32_le(buf: &mut Vec<u8>, v: u32) -> Result<()> {
716    buf.write_u32::<LittleEndian>(v).map_err(IoError::Io)
717}
718
719fn write_u64_le(buf: &mut Vec<u8>, v: u64) -> Result<()> {
720    buf.write_u64::<LittleEndian>(v).map_err(IoError::Io)
721}
722
723fn write_i32_le(buf: &mut Vec<u8>, v: i32) -> Result<()> {
724    buf.write_i32::<LittleEndian>(v).map_err(IoError::Io)
725}
726
727fn write_i64_le(buf: &mut Vec<u8>, v: i64) -> Result<()> {
728    buf.write_i64::<LittleEndian>(v).map_err(IoError::Io)
729}
730
731fn write_f32_le(buf: &mut Vec<u8>, v: f32) -> Result<()> {
732    buf.write_f32::<LittleEndian>(v).map_err(IoError::Io)
733}
734
735fn write_f64_le(buf: &mut Vec<u8>, v: f64) -> Result<()> {
736    buf.write_f64::<LittleEndian>(v).map_err(IoError::Io)
737}
738
739fn write_length_prefixed_string(buf: &mut Vec<u8>, s: &str) -> Result<()> {
740    let bytes = s.as_bytes();
741    write_u32_le(buf, bytes.len() as u32)?;
742    buf.extend_from_slice(bytes);
743    Ok(())
744}
745
746fn read_u32_le_cur(cur: &mut Cursor<&[u8]>) -> Result<u32> {
747    cur.read_u32::<LittleEndian>()
748        .map_err(|e| IoError::FormatError(format!("unexpected end of data reading u32: {e}")))
749}
750
751fn read_u64_le_cur(cur: &mut Cursor<&[u8]>) -> Result<u64> {
752    cur.read_u64::<LittleEndian>()
753        .map_err(|e| IoError::FormatError(format!("unexpected end of data reading u64: {e}")))
754}
755
756fn read_i32_le_cur(cur: &mut Cursor<&[u8]>) -> Result<i32> {
757    cur.read_i32::<LittleEndian>()
758        .map_err(|e| IoError::FormatError(format!("unexpected end of data reading i32: {e}")))
759}
760
761fn read_i64_le_cur(cur: &mut Cursor<&[u8]>) -> Result<i64> {
762    cur.read_i64::<LittleEndian>()
763        .map_err(|e| IoError::FormatError(format!("unexpected end of data reading i64: {e}")))
764}
765
766fn read_f32_le_cur(cur: &mut Cursor<&[u8]>) -> Result<f32> {
767    cur.read_f32::<LittleEndian>()
768        .map_err(|e| IoError::FormatError(format!("unexpected end of data reading f32: {e}")))
769}
770
771fn read_f64_le_cur(cur: &mut Cursor<&[u8]>) -> Result<f64> {
772    cur.read_f64::<LittleEndian>()
773        .map_err(|e| IoError::FormatError(format!("unexpected end of data reading f64: {e}")))
774}
775
776fn read_byte_cur(cur: &mut Cursor<&[u8]>) -> Result<u8> {
777    cur.read_u8()
778        .map_err(|e| IoError::FormatError(format!("unexpected end of data reading byte: {e}")))
779}
780
781fn read_bytes_cur(cur: &mut Cursor<&[u8]>, len: usize) -> Result<Vec<u8>> {
782    let mut buf = vec![0u8; len];
783    cur.read_exact(&mut buf)
784        .map_err(|e| IoError::FormatError(format!("truncated data ({len} bytes expected): {e}")))?;
785    Ok(buf)
786}
787
788fn read_length_prefixed_string(cur: &mut Cursor<&[u8]>) -> Result<String> {
789    let len = read_u32_le_cur(cur)? as usize;
790    let bytes = read_bytes_cur(cur, len)?;
791    String::from_utf8(bytes)
792        .map_err(|e| IoError::FormatError(format!("invalid UTF-8 in schema string: {e}")))
793}
794
795// ─────────────────────────────── Misc helpers ────────────────────────────────
796
797/// Rough estimate of the raw (uncompressed) byte size of a record batch.
798fn estimate_batch_raw_size(batch: &RecordBatch) -> u64 {
799    let mut size: u64 = 0;
800    for col in &batch.columns {
801        size += match col {
802            ArrowColumn::Int32(v) => (v.len() * 4) as u64,
803            ArrowColumn::Int64(v) => (v.len() * 8) as u64,
804            ArrowColumn::Float32(v) => (v.len() * 4) as u64,
805            ArrowColumn::Float64(v) => (v.len() * 8) as u64,
806            ArrowColumn::Boolean(v) => ((v.len() + 7) / 8) as u64,
807            ArrowColumn::Utf8(v) => v.iter().map(|s| (4 + s.len()) as u64).sum(),
808        };
809    }
810    size
811}
812
813// ─────────────────────────────── Tests ───────────────────────────────────────
814
815#[cfg(test)]
816mod tests {
817    use super::*;
818    use crate::arrow_ipc::{ArrowColumn, ArrowDataType, ArrowField, ArrowSchema, RecordBatch};
819
820    fn make_schema() -> ArrowSchema {
821        ArrowSchema::new(vec![
822            ArrowField::new("id", ArrowDataType::Int64),
823            ArrowField::new("score", ArrowDataType::Float64),
824            ArrowField::new("label", ArrowDataType::Utf8),
825            ArrowField::new("active", ArrowDataType::Boolean),
826        ])
827    }
828
829    fn make_batch(schema: &ArrowSchema, offset: i64) -> RecordBatch {
830        RecordBatch::new(
831            schema.clone(),
832            vec![
833                ArrowColumn::Int64(vec![offset, offset + 1, offset + 2]),
834                ArrowColumn::Float64(vec![
835                    offset as f64 * 0.1,
836                    offset as f64 * 0.2,
837                    offset as f64 * 0.3,
838                ]),
839                ArrowColumn::Utf8(vec![
840                    format!("label_{offset}"),
841                    format!("label_{}", offset + 1),
842                    format!("label_{}", offset + 2),
843                ]),
844                ArrowColumn::Boolean(vec![true, false, true]),
845            ],
846        )
847        .expect("valid batch")
848    }
849
850    // ── no-compression roundtrip ─────────────────────────────────────────────
851
852    #[test]
853    fn test_roundtrip_no_compression() {
854        let schema = make_schema();
855        let batch1 = make_batch(&schema, 0);
856        let batch2 = make_batch(&schema, 10);
857
858        let mut buf = Vec::new();
859        {
860            let mut writer =
861                ArrowStreamWriter::new(&mut buf, schema.clone(), StreamingCompression::None)
862                    .expect("writer");
863            writer.write_batch(&batch1).expect("write 1");
864            writer.write_batch(&batch2).expect("write 2");
865            let stats = writer.finish().expect("finish");
866            assert_eq!(stats.batches_written, 2);
867        }
868
869        let mut binding = buf.as_slice();
870        let mut reader = ArrowStreamReader::new(&mut binding).expect("reader");
871        let rb1 = reader.read_next_batch().expect("read 1").expect("some 1");
872        let rb2 = reader.read_next_batch().expect("read 2").expect("some 2");
873        let eos = reader.read_next_batch().expect("read eos");
874
875        assert_eq!(rb1.num_rows(), 3);
876        assert_eq!(rb2.num_rows(), 3);
877        assert!(eos.is_none());
878        assert_eq!(reader.batches_read(), 2);
879    }
880
881    // ── lz4 compression roundtrip ────────────────────────────────────────────
882
883    #[test]
884    fn test_roundtrip_lz4_compression() {
885        let schema = make_schema();
886        let batch = make_batch(&schema, 42);
887
888        let mut buf = Vec::new();
889        {
890            let mut writer =
891                ArrowStreamWriter::new(&mut buf, schema.clone(), StreamingCompression::Lz4)
892                    .expect("writer");
893            writer.write_batch(&batch).expect("write");
894            writer.finish().expect("finish");
895        }
896
897        let mut binding = buf.as_slice();
898        let mut reader = ArrowStreamReader::new(&mut binding).expect("reader");
899        let rb = reader.read_next_batch().expect("read").expect("some");
900
901        assert_eq!(rb.num_rows(), 3);
902        if let ArrowColumn::Int64(ids) = rb.column(0).expect("col 0") {
903            assert_eq!(ids, &[42, 43, 44]);
904        } else {
905            panic!("expected Int64 column");
906        }
907        if let ArrowColumn::Float64(scores) = rb.column(1).expect("col 1") {
908            assert!((scores[0] - 4.2).abs() < 1e-9);
909        } else {
910            panic!("expected Float64 column");
911        }
912    }
913
914    // ── schema metadata roundtrip ────────────────────────────────────────────
915
916    #[test]
917    fn test_schema_metadata_preserved() {
918        let mut schema = make_schema();
919        schema
920            .metadata
921            .insert("source".to_string(), "test_suite".to_string());
922
923        let mut buf = Vec::new();
924        {
925            let mut writer =
926                ArrowStreamWriter::new(&mut buf, schema.clone(), StreamingCompression::None)
927                    .expect("writer");
928            let batch = make_batch(&schema, 0);
929            writer.write_batch(&batch).expect("write");
930            writer.finish().expect("finish");
931        }
932
933        let mut binding = buf.as_slice();
934        let reader = ArrowStreamReader::new(&mut binding).expect("reader");
935        assert_eq!(
936            reader.schema().metadata.get("source"),
937            Some(&"test_suite".to_string())
938        );
939    }
940
941    // ── collect_all helper ───────────────────────────────────────────────────
942
943    #[test]
944    fn test_collect_all() {
945        let schema = make_schema();
946        let batches_in: Vec<RecordBatch> = (0..5).map(|i| make_batch(&schema, i * 3)).collect();
947
948        let mut buf = Vec::new();
949        {
950            let mut writer =
951                ArrowStreamWriter::new(&mut buf, schema.clone(), StreamingCompression::None)
952                    .expect("writer");
953            for b in &batches_in {
954                writer.write_batch(b).expect("write");
955            }
956            writer.finish().expect("finish");
957        }
958
959        let mut binding = buf.as_slice();
960        let mut reader = ArrowStreamReader::new(&mut binding).expect("reader");
961        let batches_out = reader.collect_all().expect("collect");
962
963        assert_eq!(batches_out.len(), 5);
964        for (i, b) in batches_out.iter().enumerate() {
965            assert_eq!(b.num_rows(), 3, "batch {i} rows");
966        }
967    }
968
969    // ── schema mismatch error ────────────────────────────────────────────────
970
971    #[test]
972    fn test_schema_mismatch_error() {
973        let schema_a = ArrowSchema::new(vec![ArrowField::new("x", ArrowDataType::Int32)]);
974        let schema_b = ArrowSchema::new(vec![ArrowField::new("y", ArrowDataType::Float64)]);
975
976        let batch_b =
977            RecordBatch::new(schema_b.clone(), vec![ArrowColumn::Float64(vec![1.0])]).expect("b");
978
979        let mut buf = Vec::new();
980        let mut writer =
981            ArrowStreamWriter::new(&mut buf, schema_a, StreamingCompression::None).expect("writer");
982        let result = writer.write_batch(&batch_b);
983        assert!(result.is_err(), "mismatched schema should error");
984    }
985
986    // ── all column types ─────────────────────────────────────────────────────
987
988    #[test]
989    fn test_all_column_types() {
990        let schema = ArrowSchema::new(vec![
991            ArrowField::new("i32", ArrowDataType::Int32),
992            ArrowField::new("i64", ArrowDataType::Int64),
993            ArrowField::new("f32", ArrowDataType::Float32),
994            ArrowField::new("f64", ArrowDataType::Float64),
995            ArrowField::new("bool", ArrowDataType::Boolean),
996            ArrowField::new("str", ArrowDataType::Utf8),
997        ]);
998
999        let batch = RecordBatch::new(
1000            schema.clone(),
1001            vec![
1002                ArrowColumn::Int32(vec![i32::MIN, 0, i32::MAX]),
1003                ArrowColumn::Int64(vec![i64::MIN, 0, i64::MAX]),
1004                ArrowColumn::Float32(vec![-1.5f32, 0.0, 1.5]),
1005                ArrowColumn::Float64(vec![-2.5, 0.0, 2.5]),
1006                ArrowColumn::Boolean(vec![false, true, false]),
1007                ArrowColumn::Utf8(vec!["α".into(), "β".into(), "γ".into()]),
1008            ],
1009        )
1010        .expect("valid");
1011
1012        let mut buf = Vec::new();
1013        {
1014            let mut w = ArrowStreamWriter::new(&mut buf, schema.clone(), StreamingCompression::Lz4)
1015                .expect("writer");
1016            w.write_batch(&batch).expect("write");
1017            w.finish().expect("finish");
1018        }
1019
1020        let mut binding = buf.as_slice();
1021        let mut r = ArrowStreamReader::new(&mut binding).expect("reader");
1022        let rb = r.read_next_batch().expect("read").expect("some");
1023
1024        assert_eq!(rb.num_rows(), 3);
1025
1026        if let ArrowColumn::Int32(v) = rb.column(0).expect("i32") {
1027            assert_eq!(v, &[i32::MIN, 0, i32::MAX]);
1028        } else {
1029            panic!("i32");
1030        }
1031        if let ArrowColumn::Int64(v) = rb.column(1).expect("i64") {
1032            assert_eq!(v, &[i64::MIN, 0, i64::MAX]);
1033        } else {
1034            panic!("i64");
1035        }
1036        if let ArrowColumn::Float32(v) = rb.column(2).expect("f32") {
1037            assert!((v[0] - (-1.5f32)).abs() < 1e-6);
1038            assert!((v[2] - 1.5f32).abs() < 1e-6);
1039        } else {
1040            panic!("f32");
1041        }
1042        if let ArrowColumn::Float64(v) = rb.column(3).expect("f64") {
1043            assert!((v[1] - 0.0).abs() < 1e-10);
1044        } else {
1045            panic!("f64");
1046        }
1047        if let ArrowColumn::Boolean(v) = rb.column(4).expect("bool") {
1048            assert_eq!(v, &[false, true, false]);
1049        } else {
1050            panic!("bool");
1051        }
1052        if let ArrowColumn::Utf8(v) = rb.column(5).expect("str") {
1053            assert_eq!(v, &["α", "β", "γ"]);
1054        } else {
1055            panic!("str");
1056        }
1057    }
1058
1059    // ── empty batch ──────────────────────────────────────────────────────────
1060
1061    #[test]
1062    fn test_empty_batch() {
1063        let schema = ArrowSchema::new(vec![ArrowField::new("x", ArrowDataType::Int64)]);
1064        let empty = RecordBatch::new(schema.clone(), vec![ArrowColumn::Int64(vec![])]).expect("ok");
1065
1066        let mut buf = Vec::new();
1067        {
1068            let mut w =
1069                ArrowStreamWriter::new(&mut buf, schema.clone(), StreamingCompression::None)
1070                    .expect("w");
1071            w.write_batch(&empty).expect("write");
1072            w.finish().expect("finish");
1073        }
1074
1075        let mut binding = buf.as_slice();
1076        let mut r = ArrowStreamReader::new(&mut binding).expect("r");
1077        let rb = r.read_next_batch().expect("read").expect("some");
1078        assert_eq!(rb.num_rows(), 0);
1079    }
1080
1081    // ── already-finished reader ──────────────────────────────────────────────
1082
1083    #[test]
1084    fn test_reader_after_eos() {
1085        let schema = ArrowSchema::new(vec![ArrowField::new("x", ArrowDataType::Int32)]);
1086        let batch =
1087            RecordBatch::new(schema.clone(), vec![ArrowColumn::Int32(vec![1])]).expect("ok");
1088
1089        let mut buf = Vec::new();
1090        {
1091            let mut w =
1092                ArrowStreamWriter::new(&mut buf, schema.clone(), StreamingCompression::None)
1093                    .expect("w");
1094            w.write_batch(&batch).expect("write");
1095            w.finish().expect("finish");
1096        }
1097
1098        let mut binding = buf.as_slice();
1099        let mut r = ArrowStreamReader::new(&mut binding).expect("r");
1100        let _ = r.read_next_batch().expect("first batch").expect("some");
1101        assert!(r.read_next_batch().expect("eos").is_none());
1102        // After EOS, subsequent calls return Ok(None) without errors
1103        assert!(r.read_next_batch().expect("after eos").is_none());
1104        assert!(r.is_finished());
1105    }
1106
1107    // ── write_record_batch / read_next_batch low-level API ───────────────────
1108
1109    #[test]
1110    fn test_low_level_write_read() {
1111        let schema = ArrowSchema::new(vec![ArrowField::new("v", ArrowDataType::Float64)]);
1112        let batch = RecordBatch::new(
1113            schema.clone(),
1114            vec![ArrowColumn::Float64(vec![1.1, 2.2, 3.3])],
1115        )
1116        .expect("ok");
1117
1118        // Write schema + batch + EOS manually
1119        let mut buf = Vec::new();
1120        let schema_payload = serialize_schema(&schema).expect("ser schema");
1121        write_schema_message(&mut buf, &schema_payload).expect("schema msg");
1122        write_record_batch(&mut buf, &batch, StreamingCompression::None).expect("batch msg");
1123        write_eos_message(&mut buf).expect("eos");
1124
1125        // Read
1126        let mut cur = buf.as_slice();
1127        let (tag, _codec, payload) = read_message(&mut cur).expect("schema msg");
1128        assert_eq!(tag, TAG_SCHEMA);
1129        let schema_read = deserialize_schema(&payload).expect("deser schema");
1130
1131        let rb = read_next_batch(&mut cur, &schema_read)
1132            .expect("batch")
1133            .expect("some");
1134        assert_eq!(rb.num_rows(), 3);
1135
1136        let eos = read_next_batch(&mut cur, &schema_read).expect("eos");
1137        assert!(eos.is_none());
1138    }
1139
1140    // ── writer stats ─────────────────────────────────────────────────────────
1141
1142    #[test]
1143    fn test_writer_stats() {
1144        let schema = make_schema();
1145        let mut buf = Vec::new();
1146        let mut w = ArrowStreamWriter::new(&mut buf, schema.clone(), StreamingCompression::None)
1147            .expect("writer");
1148
1149        for i in 0..4u32 {
1150            let b = make_batch(&schema, i as i64 * 10);
1151            w.write_batch(&b).expect("write");
1152        }
1153        assert_eq!(w.batches_written(), 4);
1154
1155        let stats = w.finish().expect("finish");
1156        assert_eq!(stats.batches_written, 4);
1157        assert!(stats.uncompressed_bytes > 0);
1158    }
1159}