1use 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
84const TAG_SCHEMA: u8 = 0x01;
88const TAG_RECORD_BATCH: u8 = 0x02;
89const TAG_EOS: u8 = 0x00;
90
91const CODEC_NONE: u8 = 0x00;
93const CODEC_LZ4: u8 = 0x01;
94
95const ALIGNMENT: usize = 8;
97
98#[derive(Debug, Clone, Copy, PartialEq, Eq)]
102pub enum StreamingCompression {
103 None,
105 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#[derive(Debug, Clone, Default)]
120pub struct WriterStats {
121 pub batches_written: usize,
123 pub uncompressed_bytes: u64,
125 pub compressed_bytes: u64,
128}
129
130impl WriterStats {
131 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
142pub 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 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 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 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 pub fn finish(self) -> Result<WriterStats> {
200 write_eos_message(self.writer)?;
201 Ok(self.stats)
202 }
203
204 pub fn batches_written(&self) -> usize {
206 self.stats.batches_written
207 }
208
209 pub fn stats(&self) -> &WriterStats {
211 &self.stats
212 }
213
214 pub fn schema(&self) -> &ArrowSchema {
216 &self.schema
217 }
218}
219
220pub 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 pub fn new(reader: &'a mut dyn Read) -> Result<Self> {
239 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 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 pub fn schema(&self) -> &ArrowSchema {
283 &self.schema
284 }
285
286 pub fn batches_read(&self) -> usize {
288 self.batches_read
289 }
290
291 pub fn is_finished(&self) -> bool {
293 self.finished
294 }
295
296 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
306pub 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
327pub 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
346fn 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
360fn 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
373fn 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
381fn 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
391fn 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 let rem = len % ALIGNMENT;
414 if rem != 0 {
415 let skip = ALIGNMENT - rem;
416 let mut pad_buf = vec![0u8; skip];
417 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
438fn 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
467fn serialize_schema(schema: &ArrowSchema) -> Result<Vec<u8>> {
470 let mut buf = Vec::new();
471 write_u32_le(&mut buf, schema.fields.len() as u32)?;
473
474 for field in &schema.fields {
475 write_length_prefixed_string(&mut buf, &field.name)?;
477 buf.push(dtype_tag(&field.dtype));
479 buf.push(if field.nullable { 1 } else { 0 });
481 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 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
535fn serialize_record_batch(batch: &RecordBatch) -> Result<Vec<u8>> {
538 let mut buf = Vec::new();
539 write_u64_le(&mut buf, batch.num_rows() as u64)?;
541 write_u32_le(&mut buf, batch.num_columns() as u32)?;
543
544 for col in &batch.columns {
545 buf.push(dtype_tag(&col.data_type()));
547 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 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
686fn 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
713fn 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
795fn 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#[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 assert!(r.read_next_batch().expect("after eos").is_none());
1104 assert!(r.is_finished());
1105 }
1106
1107 #[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 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 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 #[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}