use std::sync::Arc;
use arrow_array::RecordBatch;
use arrow_array::builder::{Float32Builder, ListBuilder, TimestampMicrosecondBuilder};
use arrow_ipc::writer::StreamWriter;
use arrow_schema::{DataType, Field, Schema, TimeUnit};
use arrow_select::concat::concat_batches as arrow_concat_batches;
use chrono::{DateTime, Utc};
use crate::types::arrow::{ArrowSerializable, FieldIdMap, update_field_id};
use crate::types::error::{IcebergError, Result};
use crate::types::model::{DataRow, MetadataRow};
impl ArrowSerializable for DataRow {
fn arrow_schema(field_id_map: &FieldIdMap) -> Result<Arc<Schema>> {
let mut element_field = Field::new("element", DataType::Float32, true);
update_field_id(
&mut element_field,
*field_id_map.get_field_id("vector.element")?,
);
let mut data_stream_id_field = Field::new("data_stream_id", DataType::Utf8, false);
update_field_id(
&mut data_stream_id_field,
*field_id_map.get_field_id("data_stream_id")?,
);
let mut datetime_field = Field::new(
"datetime",
DataType::Timestamp(TimeUnit::Microsecond, Some("UTC".into())),
false,
);
update_field_id(&mut datetime_field, *field_id_map.get_field_id("datetime")?);
let mut vector_field = Field::new("vector", DataType::List(Arc::new(element_field)), false);
update_field_id(&mut vector_field, *field_id_map.get_field_id("vector")?);
let mut data_type_field = Field::new("data_type", DataType::Utf8, false);
update_field_id(
&mut data_type_field,
*field_id_map.get_field_id("data_type")?,
);
let mut specification_id_field = Field::new("specification_id", DataType::Utf8, true);
update_field_id(
&mut specification_id_field,
*field_id_map.get_field_id("specification_id")?,
);
let mut vector_start_bound_field =
Field::new("vector_start_bound", DataType::Float32, false);
update_field_id(
&mut vector_start_bound_field,
*field_id_map.get_field_id("vector_start_bound")?,
);
let mut vector_end_bound_field = Field::new("vector_end_bound", DataType::Float32, false);
update_field_id(
&mut vector_end_bound_field,
*field_id_map.get_field_id("vector_end_bound")?,
);
let mut created_datetime_field = Field::new(
"created_datetime",
DataType::Timestamp(TimeUnit::Microsecond, Some("UTC".into())),
true,
);
update_field_id(
&mut created_datetime_field,
*field_id_map.get_field_id("created_datetime")?,
);
let mut modified_datetime_field = Field::new(
"modified_datetime",
DataType::Timestamp(TimeUnit::Microsecond, Some("UTC".into())),
true,
);
update_field_id(
&mut modified_datetime_field,
*field_id_map.get_field_id("modified_datetime")?,
);
let fields = vec![
data_stream_id_field,
datetime_field,
vector_field,
data_type_field,
specification_id_field,
vector_start_bound_field,
vector_end_bound_field,
created_datetime_field,
modified_datetime_field,
];
Ok(Arc::new(Schema::new(fields)))
}
#[allow(
clippy::as_conversions,
clippy::cast_possible_truncation,
clippy::wildcard_enum_match_arm,
reason = "f64-to-f32 truncation is intentional: Iceberg stores vectors as Float32; \
wildcard match on DataType has 30+ variants where only List is valid"
)]
fn to_record_batch(&self, field_id_map: &FieldIdMap) -> Result<RecordBatch> {
let schema = Self::arrow_schema(field_id_map)?;
let created = self.created_datetime.unwrap_or_else(Utc::now);
let data_stream_id = arrow_array::StringArray::from(vec![self.data_stream_id.to_string()]);
let datetime = timestamp_array(&[self.datetime]);
let vector = {
let element_field = schema.field_with_name("vector")?;
let inner_field = match element_field.data_type() {
DataType::List(f) => Arc::<Field>::clone(f),
other => {
return Err(IcebergError::ArrowTypeMismatch {
column_name: "vector".to_owned(),
expected: "List".to_owned(),
actual: format!("{other:?}"),
}
.into());
}
};
let mut builder = ListBuilder::new(Float32Builder::new()).with_field(inner_field);
let values = builder.values();
for &element in &self.vector {
values.append_value(element as f32);
}
builder.append(true);
builder.finish()
};
let data_type = arrow_array::StringArray::from(vec![self.data_type.clone()]);
let specification_id =
arrow_array::StringArray::from(vec![self.specification_id.to_string()]);
let vector_start_bound =
arrow_array::Float32Array::from(vec![self.vector_start_bound as f32]);
let vector_end_bound = arrow_array::Float32Array::from(vec![self.vector_end_bound as f32]);
let created_datetime = nullable_timestamp_array(&[Some(created)]);
let modified_datetime = nullable_timestamp_array(&[self.modified_datetime]);
Ok(RecordBatch::try_new(
schema,
vec![
Arc::new(data_stream_id),
Arc::new(datetime),
Arc::new(vector),
Arc::new(data_type),
Arc::new(specification_id),
Arc::new(vector_start_bound),
Arc::new(vector_end_bound),
Arc::new(created_datetime),
Arc::new(modified_datetime),
],
)?)
}
}
impl ArrowSerializable for MetadataRow {
fn arrow_schema(field_id_map: &FieldIdMap) -> Result<Arc<Schema>> {
let mut data_stream_id_field = Field::new("data_stream_id", DataType::Utf8, false);
update_field_id(
&mut data_stream_id_field,
*field_id_map.get_field_id("data_stream_id")?,
);
let mut datetime_field = Field::new(
"datetime",
DataType::Timestamp(TimeUnit::Microsecond, Some("UTC".into())),
false,
);
update_field_id(&mut datetime_field, *field_id_map.get_field_id("datetime")?);
let mut latitude_field = Field::new("latitude", DataType::Float32, true);
update_field_id(&mut latitude_field, *field_id_map.get_field_id("latitude")?);
let mut longitude_field = Field::new("longitude", DataType::Float32, true);
update_field_id(
&mut longitude_field,
*field_id_map.get_field_id("longitude")?,
);
let mut altitude_field = Field::new("altitude", DataType::Float32, true);
update_field_id(&mut altitude_field, *field_id_map.get_field_id("altitude")?);
let mut speed_field = Field::new("speed", DataType::Float32, true);
update_field_id(&mut speed_field, *field_id_map.get_field_id("speed")?);
let mut heading_field = Field::new("heading", DataType::Float32, true);
update_field_id(&mut heading_field, *field_id_map.get_field_id("heading")?);
let mut pitch_field = Field::new("pitch", DataType::Float32, true);
update_field_id(&mut pitch_field, *field_id_map.get_field_id("pitch")?);
let mut roll_field = Field::new("roll", DataType::Float32, true);
update_field_id(&mut roll_field, *field_id_map.get_field_id("roll")?);
let mut speed_over_ground_field = Field::new("speed_over_ground", DataType::Float32, true);
update_field_id(
&mut speed_over_ground_field,
*field_id_map.get_field_id("speed_over_ground")?,
);
let mut created_datetime_field = Field::new(
"created_datetime",
DataType::Timestamp(TimeUnit::Microsecond, Some("UTC".into())),
true,
);
update_field_id(
&mut created_datetime_field,
*field_id_map.get_field_id("created_datetime")?,
);
let mut modified_datetime_field = Field::new(
"modified_datetime",
DataType::Timestamp(TimeUnit::Microsecond, Some("UTC".into())),
true,
);
update_field_id(
&mut modified_datetime_field,
*field_id_map.get_field_id("modified_datetime")?,
);
let fields = vec![
data_stream_id_field,
datetime_field,
latitude_field,
longitude_field,
altitude_field,
speed_field,
heading_field,
pitch_field,
roll_field,
speed_over_ground_field,
created_datetime_field,
modified_datetime_field,
];
Ok(Arc::new(Schema::new(fields)))
}
#[allow(
clippy::cast_possible_truncation,
clippy::as_conversions,
reason = "f64-to-f32 truncation is intentional: Iceberg stores metadata as Float32"
)]
fn to_record_batch(&self, field_id_map: &FieldIdMap) -> Result<RecordBatch> {
let schema = Self::arrow_schema(field_id_map)?;
let created = self.created_datetime.unwrap_or_else(Utc::now);
let data_stream_id = arrow_array::StringArray::from(vec![self.data_stream_id.to_string()]);
let datetime = timestamp_array(&[self.datetime]);
let latitude = nullable_f32_array(&[self.latitude.map(|v| v as f32)]);
let longitude = nullable_f32_array(&[self.longitude.map(|v| v as f32)]);
let altitude = nullable_f32_array(&[self.altitude.map(|v| v as f32)]);
let speed = nullable_f32_array(&[self.speed.map(|v| v as f32)]);
let heading = nullable_f32_array(&[self.heading.map(|v| v as f32)]);
let pitch = nullable_f32_array(&[self.pitch.map(|v| v as f32)]);
let roll = nullable_f32_array(&[self.roll.map(|v| v as f32)]);
let speed_over_ground = nullable_f32_array(&[self.speed_over_ground.map(|v| v as f32)]);
let created_datetime = nullable_timestamp_array(&[Some(created)]);
let modified_datetime = nullable_timestamp_array(&[self.modified_datetime]);
Ok(RecordBatch::try_new(
schema,
vec![
Arc::new(data_stream_id),
Arc::new(datetime),
Arc::new(latitude),
Arc::new(longitude),
Arc::new(altitude),
Arc::new(speed),
Arc::new(heading),
Arc::new(pitch),
Arc::new(roll),
Arc::new(speed_over_ground),
Arc::new(created_datetime),
Arc::new(modified_datetime),
],
)?)
}
}
pub fn concat_batches(batches: &[RecordBatch]) -> Result<RecordBatch> {
let first = batches
.first()
.ok_or_else(|| IcebergError::EmptyBatchList {
operation: "concatenate".to_owned(),
})?;
Ok(arrow_concat_batches(&first.schema(), batches)?)
}
fn nullable_f32_array(values: &[Option<f32>]) -> arrow_array::Float32Array {
let mut builder = Float32Builder::new();
for v in values {
match v {
Some(val) => builder.append_value(*val),
None => builder.append_null(),
}
}
builder.finish()
}
fn nullable_timestamp_array(
values: &[Option<DateTime<Utc>>],
) -> arrow_array::TimestampMicrosecondArray {
let mut builder = TimestampMicrosecondBuilder::new().with_timezone("UTC");
for v in values {
match v {
Some(dt) => builder.append_value(dt.timestamp_micros()),
None => builder.append_null(),
}
}
builder.finish()
}
pub fn serialize_to_arrow_ipc(batches: &[RecordBatch]) -> Result<Vec<u8>> {
let first = batches
.first()
.ok_or_else(|| IcebergError::EmptyBatchList {
operation: "serialize".to_owned(),
})?;
let schema = first.schema();
let mut buf = Vec::new();
let mut writer = StreamWriter::try_new(&mut buf, &schema)?;
for batch in batches {
writer.write(batch)?;
}
writer.finish()?;
drop(writer);
Ok(buf)
}
fn timestamp_array(values: &[DateTime<Utc>]) -> arrow_array::TimestampMicrosecondArray {
let mut builder = TimestampMicrosecondBuilder::new().with_timezone("UTC");
for v in values {
builder.append_value(v.timestamp_micros());
}
builder.finish()
}
#[cfg(test)]
#[allow(
clippy::default_numeric_fallback,
clippy::indexing_slicing,
clippy::panic,
clippy::unwrap_used,
reason = "test code uses literals, indexing, panic, and unwrap for clarity"
)]
mod tests {
use std::collections::HashMap;
use std::io::Cursor;
use std::result::Result as StdResult;
use arrow_array::Array as _;
use arrow_ipc::reader::StreamReader;
use uuid::Uuid;
use super::*;
fn data_row_field_id_map() -> FieldIdMap {
let mapping = HashMap::from([
("created_datetime".to_owned(), 8),
("data_stream_id".to_owned(), 1),
("data_type".to_owned(), 4),
("datetime".to_owned(), 2),
("modified_datetime".to_owned(), 9),
("specification_id".to_owned(), 5),
("vector".to_owned(), 3),
("vector.element".to_owned(), 10),
("vector_end_bound".to_owned(), 7),
("vector_start_bound".to_owned(), 6),
]);
FieldIdMap::new("horizon_public.data_row".to_owned(), mapping)
}
fn metadata_row_field_id_map() -> FieldIdMap {
let mapping = HashMap::from([
("altitude".to_owned(), 5),
("created_datetime".to_owned(), 11),
("data_stream_id".to_owned(), 1),
("datetime".to_owned(), 2),
("heading".to_owned(), 7),
("latitude".to_owned(), 3),
("longitude".to_owned(), 4),
("modified_datetime".to_owned(), 12),
("pitch".to_owned(), 8),
("roll".to_owned(), 9),
("speed".to_owned(), 6),
("speed_over_ground".to_owned(), 10),
]);
FieldIdMap::new("horizon_public.metadata_row".to_owned(), mapping)
}
fn sample_data_row() -> DataRow {
DataRow {
created_datetime: None,
data_stream_id: Uuid::nil(),
data_type: "audio".to_owned(),
datetime: Utc::now(),
modified_datetime: None,
specification_id: Uuid::nil(),
vector: vec![1.0, 2.0, 3.0],
vector_end_bound: 100.0,
vector_start_bound: 0.0,
}
}
fn sample_metadata_row() -> MetadataRow {
MetadataRow {
altitude: Some(10.0),
created_datetime: None,
data_stream_id: Uuid::nil(),
datetime: Utc::now(),
heading: Some(90.0),
latitude: Some(37.7749),
longitude: Some(-122.4194),
modified_datetime: None,
pitch: None,
roll: None,
speed: Some(5.5),
speed_over_ground: None,
}
}
#[test]
fn concat_batches_merges_rows() {
let field_ids = data_row_field_id_map();
let batch1 = sample_data_row().to_record_batch(&field_ids).unwrap();
let batch2 = sample_data_row().to_record_batch(&field_ids).unwrap();
let combined = concat_batches(&[batch1, batch2]).unwrap();
assert_eq!(combined.num_rows(), 2);
}
#[test]
fn created_datetime_defaults_to_now() {
let field_ids = data_row_field_id_map();
let row = DataRow {
created_datetime: None,
..sample_data_row()
};
let batch = row.to_record_batch(&field_ids).unwrap();
let ts_col = batch
.column(7)
.as_any()
.downcast_ref::<arrow_array::TimestampMicrosecondArray>()
.unwrap();
assert!(!ts_col.is_null(0));
}
#[test]
fn data_row_arrow_schema_has_field_ids() {
let field_ids = data_row_field_id_map();
let schema = DataRow::arrow_schema(&field_ids).unwrap();
assert_eq!(schema.fields().len(), 9);
let ds_field = schema.field_with_name("data_stream_id").unwrap();
assert_eq!(
ds_field.metadata().get("PARQUET:field_id"),
Some(&"1".to_owned())
);
let vector_field = schema.field_with_name("vector").unwrap();
assert_eq!(
vector_field.metadata().get("PARQUET:field_id"),
Some(&"3".to_owned())
);
if let DataType::List(element) = vector_field.data_type() {
assert_eq!(
element.metadata().get("PARQUET:field_id"),
Some(&"10".to_owned())
);
} else {
panic!("vector should be List type");
}
}
#[test]
fn data_row_to_record_batch_round_trip() {
let field_ids = data_row_field_id_map();
let row = sample_data_row();
let batch = row.to_record_batch(&field_ids).unwrap();
assert_eq!(batch.num_rows(), 1);
assert_eq!(batch.num_columns(), 9);
}
#[test]
fn data_row_uuids_serialized_as_strings() {
let field_ids = data_row_field_id_map();
let id = Uuid::new_v4();
let row = DataRow {
data_stream_id: id,
..sample_data_row()
};
let batch = row.to_record_batch(&field_ids).unwrap();
let str_col = batch
.column(0)
.as_any()
.downcast_ref::<arrow_array::StringArray>()
.unwrap();
assert_eq!(str_col.value(0), id.to_string());
}
#[test]
fn metadata_row_arrow_schema_has_field_ids() {
let field_ids = metadata_row_field_id_map();
let schema = MetadataRow::arrow_schema(&field_ids).unwrap();
assert_eq!(schema.fields().len(), 12);
let lat_field = schema.field_with_name("latitude").unwrap();
assert_eq!(
lat_field.metadata().get("PARQUET:field_id"),
Some(&"3".to_owned())
);
assert!(lat_field.is_nullable());
}
#[test]
fn metadata_row_to_record_batch_round_trip() {
let field_ids = metadata_row_field_id_map();
let row = sample_metadata_row();
let batch = row.to_record_batch(&field_ids).unwrap();
assert_eq!(batch.num_rows(), 1);
assert_eq!(batch.num_columns(), 12);
}
#[test]
fn serialize_and_decode_data_row_arrow_ipc() {
let field_ids = data_row_field_id_map();
let row = sample_data_row();
let batch = row.to_record_batch(&field_ids).unwrap();
let bytes = serialize_to_arrow_ipc(&[batch]).unwrap();
assert!(!bytes.is_empty());
let reader = StreamReader::try_new(Cursor::new(&bytes), None).unwrap();
let decoded: Vec<RecordBatch> = reader.collect::<StdResult<Vec<_>, _>>().unwrap();
assert_eq!(decoded.len(), 1);
assert_eq!(decoded[0].num_rows(), 1);
assert_eq!(decoded[0].num_columns(), 9);
let schema = decoded[0].schema();
let ds_field = schema.field_with_name("data_stream_id").unwrap();
assert_eq!(
ds_field.metadata().get("PARQUET:field_id"),
Some(&"1".to_owned())
);
}
#[test]
fn serialize_and_decode_metadata_row_arrow_ipc() {
let field_ids = metadata_row_field_id_map();
let row = sample_metadata_row();
let batch = row.to_record_batch(&field_ids).unwrap();
let bytes = serialize_to_arrow_ipc(&[batch]).unwrap();
let reader = StreamReader::try_new(Cursor::new(&bytes), None).unwrap();
let decoded: Vec<RecordBatch> = reader.collect::<StdResult<Vec<_>, _>>().unwrap();
assert_eq!(decoded.len(), 1);
assert_eq!(decoded[0].num_rows(), 1);
}
#[test]
fn serialize_multiple_batches() {
let field_ids = data_row_field_id_map();
let row1 = sample_data_row();
let row2 = DataRow {
vector: vec![4.0, 5.0],
..sample_data_row()
};
let batch1 = row1.to_record_batch(&field_ids).unwrap();
let batch2 = row2.to_record_batch(&field_ids).unwrap();
let bytes = serialize_to_arrow_ipc(&[batch1, batch2]).unwrap();
let reader = StreamReader::try_new(Cursor::new(&bytes), None).unwrap();
let decoded: Vec<RecordBatch> = reader.collect::<StdResult<Vec<_>, _>>().unwrap();
assert_eq!(decoded.len(), 2);
assert_eq!(decoded[0].num_rows(), 1);
assert_eq!(decoded[1].num_rows(), 1);
}
}