use arrow::record_batch::RecordBatch;
use wf_connector_api::{SourceReason, SourceResult};
use wp_connector_api::SourceBatch;
use super::payload::payload_to_bytes;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum WireFormat {
Ndjson,
ArrowStream,
ArrowFramed,
}
impl WireFormat {
pub fn from_data_format(value: Option<&str>) -> Self {
match value.unwrap_or("ndjson") {
"arrow_framed" => WireFormat::ArrowFramed,
"arrow_ipc" => WireFormat::ArrowStream,
_ => WireFormat::Ndjson,
}
}
}
pub fn decode_arrow_ipc_batches(events: &SourceBatch) -> SourceResult<Vec<RecordBatch>> {
use arrow::ipc::reader::StreamReader;
let mut batches = Vec::new();
for event in events {
let payload = payload_to_bytes(&event.payload);
let cursor = std::io::Cursor::new(payload);
let reader = StreamReader::try_new(cursor, None)
.map_err(|e| SourceReason::Decode.err_detail(format!("arrow ipc: {e}")))?;
for batch in reader {
let batch = batch
.map_err(|e| SourceReason::Decode.err_detail(format!("arrow ipc batch: {e}")))?;
batches.push(batch);
}
}
Ok(batches)
}
pub fn decode_arrow_framed_batches(events: &SourceBatch) -> SourceResult<Vec<RecordBatch>> {
use arrow::ipc::reader::StreamReader;
let mut batches = Vec::new();
for event in events {
let payload = payload_to_bytes(&event.payload);
if payload.len() < 4 {
continue;
}
let tag_len = u32::from_be_bytes([payload[0], payload[1], payload[2], payload[3]]) as usize;
let ipc_start = 4 + tag_len;
if ipc_start > payload.len() {
continue;
}
let ipc_bytes = &payload[ipc_start..];
let cursor = std::io::Cursor::new(ipc_bytes);
let reader = StreamReader::try_new(cursor, None)
.map_err(|e| SourceReason::Decode.err_detail(format!("arrow framed: {e}")))?;
for batch in reader {
let batch = batch
.map_err(|e| SourceReason::Decode.err_detail(format!("arrow framed batch: {e}")))?;
batches.push(batch);
}
}
Ok(batches)
}
#[cfg(test)]
mod tests {
use super::*;
use arrow::array::StringArray;
use arrow::datatypes::{DataType, Field, Schema};
use arrow::ipc::writer::StreamWriter;
use bytes::Bytes;
use std::sync::Arc;
use wp_connector_api::SourceEvent;
use wp_model_core::raw::RawData;
#[test]
fn wire_format_defaults_to_ndjson() {
assert_eq!(WireFormat::from_data_format(None), WireFormat::Ndjson);
}
#[test]
fn wire_format_parses_arrow_framed() {
assert_eq!(
WireFormat::from_data_format(Some("arrow_framed")),
WireFormat::ArrowFramed
);
}
#[test]
fn wire_format_parses_arrow_ipc() {
assert_eq!(
WireFormat::from_data_format(Some("arrow_ipc")),
WireFormat::ArrowStream
);
}
#[test]
fn wire_format_unknown_falls_back_to_ndjson() {
assert_eq!(
WireFormat::from_data_format(Some("nonsense")),
WireFormat::Ndjson
);
}
fn make_ipc_bytes(schema: &Arc<Schema>, values: &[&str]) -> Vec<u8> {
let mut buf = Vec::new();
let mut w = StreamWriter::try_new(&mut buf, schema).unwrap();
let batch = RecordBatch::try_new(
schema.clone(),
vec![Arc::new(StringArray::from(values.to_vec()))],
)
.unwrap();
w.write(&batch).unwrap();
w.finish().unwrap();
buf
}
fn event_from_bytes(bytes: impl Into<Bytes>) -> SourceEvent {
SourceEvent::new(
1,
"key".to_string(),
RawData::Bytes(bytes.into()),
Default::default(),
)
}
#[test]
fn decode_arrow_ipc_roundtrip() {
let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Utf8, false)]));
let ipc = make_ipc_bytes(&schema, &["hello"]);
let event = event_from_bytes(ipc);
let batches = decode_arrow_ipc_batches(&vec![event]).unwrap();
assert_eq!(batches.len(), 1);
assert_eq!(batches[0].num_rows(), 1);
}
#[test]
fn decode_arrow_ipc_multiple_batches_concat() {
let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Utf8, false)]));
let ev1 = event_from_bytes(make_ipc_bytes(&schema, &["a", "b"]));
let ev2 = event_from_bytes(make_ipc_bytes(&schema, &["c"]));
let batches = decode_arrow_ipc_batches(&vec![ev1, ev2]).unwrap();
assert_eq!(batches.len(), 2);
assert_eq!(batches[0].num_rows(), 2);
assert_eq!(batches[1].num_rows(), 1);
}
#[test]
fn decode_arrow_ipc_invalid_returns_err() {
let event = event_from_bytes(Bytes::from_static(b"ARROW?? not really"));
let result = decode_arrow_ipc_batches(&vec![event]);
assert!(result.is_err(), "expected decode error for corrupt input");
let err = result.unwrap_err();
assert!(
format!("{err:?}").to_lowercase().contains("decode"),
"error should be classified as Decode: {err:?}"
);
}
#[test]
fn decode_arrow_framed_roundtrip() {
let tag = b"my_tag";
let tag_len = (tag.len() as u32).to_be_bytes();
let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Utf8, false)]));
let ipc_buf = make_ipc_bytes(&schema, &["hello"]);
let mut frame = Vec::new();
frame.extend_from_slice(&tag_len);
frame.extend_from_slice(tag);
frame.extend_from_slice(&ipc_buf);
let event = event_from_bytes(frame);
let batches = decode_arrow_framed_batches(&vec![event]).unwrap();
assert_eq!(batches.len(), 1);
assert_eq!(batches[0].num_rows(), 1);
}
#[test]
fn decode_arrow_framed_too_short_payload_is_skipped() {
let event = event_from_bytes(Bytes::from(vec![0, 0]));
let batches = decode_arrow_framed_batches(&vec![event]).unwrap();
assert!(batches.is_empty());
}
#[test]
fn decode_arrow_framed_tag_len_exceeds_payload_is_skipped() {
let event = event_from_bytes(Bytes::from(vec![0xff, 0xff, 0xff, 0xff, 0x00]));
let batches = decode_arrow_framed_batches(&vec![event]).unwrap();
assert!(batches.is_empty());
}
#[test]
fn decode_arrow_framed_empty_tag() {
let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Utf8, false)]));
let ipc = make_ipc_bytes(&schema, &["hi"]);
let mut frame = Vec::new();
frame.extend_from_slice(&0u32.to_be_bytes()); frame.extend_from_slice(&ipc);
let event = event_from_bytes(frame);
let batches = decode_arrow_framed_batches(&vec![event]).unwrap();
assert_eq!(batches.len(), 1);
assert_eq!(batches[0].num_rows(), 1);
}
#[test]
fn decode_arrow_framed_multiple_events_concat() {
let schema = Arc::new(Schema::new(vec![Field::new("x", DataType::Utf8, false)]));
let make_framed = |tag: &str| {
let ipc = make_ipc_bytes(&schema, &["row"]);
let mut frame = Vec::new();
frame.extend_from_slice(&(tag.len() as u32).to_be_bytes());
frame.extend_from_slice(tag.as_bytes());
frame.extend_from_slice(&ipc);
event_from_bytes(frame)
};
let batches =
decode_arrow_framed_batches(&vec![make_framed("a"), make_framed("bb")]).unwrap();
assert_eq!(batches.len(), 2);
assert_eq!(batches[0].num_rows(), 1);
assert_eq!(batches[1].num_rows(), 1);
}
#[test]
fn decode_arrow_framed_invalid_ipc_returns_err() {
let mut frame = Vec::new();
frame.extend_from_slice(&1u32.to_be_bytes());
frame.extend_from_slice(b"t");
frame.extend_from_slice(b"definitely not arrow");
let event = event_from_bytes(frame);
let result = decode_arrow_framed_batches(&vec![event]);
assert!(
result.is_err(),
"expected decode error for corrupt framed body"
);
}
}