use std::sync::Arc;
use arrow_array::RecordBatch;
use arrow_array::builder::{
Float32Builder, Float64Builder, Int32Builder, 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, VariantDataRow};
impl ArrowSerializable for VariantDataRow {
fn arrow_schema(field_id_map: &FieldIdMap) -> Result<Arc<Schema>> {
let fields = vec![
field_with_id("data_stream_id", DataType::Utf8, false, field_id_map)?,
field_with_id(
"datetime",
DataType::Timestamp(TimeUnit::Microsecond, Some("UTC".into())),
false,
field_id_map,
)?,
list_field_with_id("vector", DataType::Float32, true, field_id_map)?,
field_with_id("data_type", DataType::Utf8, false, field_id_map)?,
field_with_id("specification_id", DataType::Utf8, true, field_id_map)?,
field_with_id("vector_start_bound", DataType::Float32, true, field_id_map)?,
field_with_id("vector_end_bound", DataType::Float32, true, field_id_map)?,
field_with_id(
"created_datetime",
DataType::Timestamp(TimeUnit::Microsecond, Some("UTC".into())),
true,
field_id_map,
)?,
field_with_id(
"modified_datetime",
DataType::Timestamp(TimeUnit::Microsecond, Some("UTC".into())),
true,
field_id_map,
)?,
field_with_id("variant", DataType::Utf8, true, field_id_map)?,
field_with_id("payload_int", DataType::Int32, true, field_id_map)?,
list_field_with_id("payload_int_array", DataType::Int32, true, field_id_map)?,
field_with_id("payload_float32", DataType::Float32, true, field_id_map)?,
list_field_with_id(
"payload_float32_array",
DataType::Float32,
true,
field_id_map,
)?,
field_with_id("payload_float64", DataType::Float64, true, field_id_map)?,
list_field_with_id(
"payload_float64_array",
DataType::Float64,
true,
field_id_map,
)?,
field_with_id("payload_string", DataType::Utf8, true, field_id_map)?,
field_with_id("payload_struct", DataType::Utf8, true, field_id_map)?,
];
Ok(Arc::new(Schema::new(fields)))
}
#[allow(
clippy::as_conversions,
clippy::cast_possible_truncation,
clippy::too_many_lines,
reason = "f64-to-f32 truncation is intentional for Iceberg REAL columns; \
record batch builds every payload column"
)]
fn to_record_batch(&self, field_id_map: &FieldIdMap) -> Result<RecordBatch> {
self.validate_payload()?;
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 inner_field = list_inner_field(&schema, "vector")?;
let mut builder = ListBuilder::new(Float32Builder::new()).with_field(inner_field);
match &self.vector {
Some(values) => {
let values_builder = builder.values();
for &element in values {
values_builder.append_value(element as f32);
}
builder.append(true);
}
None => builder.append(false),
}
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.map(|value| value as f32),
]);
let vector_end_bound =
arrow_array::Float32Array::from(vec![self.vector_end_bound.map(|value| value as f32)]);
let created_datetime = nullable_timestamp_array(&[Some(created)]);
let modified_datetime = nullable_timestamp_array(&[self.modified_datetime]);
let variant = arrow_array::StringArray::from(vec![self.variant.clone()]);
let payload_int = arrow_array::Int32Array::from(vec![self.payload_int]);
let payload_int_array = {
let inner_field = list_inner_field(&schema, "payload_int_array")?;
let mut builder = ListBuilder::new(Int32Builder::new()).with_field(inner_field);
match &self.payload_int_array {
Some(values) => {
let values_builder = builder.values();
for &element in values {
values_builder.append_value(element);
}
builder.append(true);
}
None => builder.append(false),
}
builder.finish()
};
let payload_float32 = arrow_array::Float32Array::from(vec![self.payload_float32]);
let payload_float32_array = {
let inner_field = list_inner_field(&schema, "payload_float32_array")?;
let mut builder = ListBuilder::new(Float32Builder::new()).with_field(inner_field);
match &self.payload_float32_array {
Some(values) => {
let values_builder = builder.values();
for &element in values {
values_builder.append_value(element);
}
builder.append(true);
}
None => builder.append(false),
}
builder.finish()
};
let payload_float64 = arrow_array::Float64Array::from(vec![self.payload_float64]);
let payload_float64_array = {
let inner_field = list_inner_field(&schema, "payload_float64_array")?;
let mut builder = ListBuilder::new(Float64Builder::new()).with_field(inner_field);
match &self.payload_float64_array {
Some(values) => {
let values_builder = builder.values();
for &element in values {
values_builder.append_value(element);
}
builder.append(true);
}
None => builder.append(false),
}
builder.finish()
};
let payload_string = arrow_array::StringArray::from(vec![self.payload_string.clone()]);
let payload_struct = arrow_array::StringArray::from(vec![
self.payload_struct.as_ref().map(ToString::to_string),
]);
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),
Arc::new(variant),
Arc::new(payload_int),
Arc::new(payload_int_array),
Arc::new(payload_float32),
Arc::new(payload_float32_array),
Arc::new(payload_float64),
Arc::new(payload_float64_array),
Arc::new(payload_string),
Arc::new(payload_struct),
],
)?)
}
}
impl ArrowSerializable for DataRow {
fn arrow_schema(field_id_map: &FieldIdMap) -> Result<Arc<Schema>> {
VariantDataRow::arrow_schema(field_id_map)
}
fn to_record_batch(&self, field_id_map: &FieldIdMap) -> Result<RecordBatch> {
VariantDataRow::from(self).to_record_batch(field_id_map)
}
}
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 field_with_id(
name: &str,
data_type: DataType,
nullable: bool,
field_id_map: &FieldIdMap,
) -> Result<Field> {
let mut field = Field::new(name, data_type, nullable);
update_field_id(&mut field, *field_id_map.get_field_id(name)?);
Ok(field)
}
fn list_field_with_id(
name: &str,
element_data_type: DataType,
nullable: bool,
field_id_map: &FieldIdMap,
) -> Result<Field> {
let element_path = format!("{name}.element");
let mut element_field = Field::new("element", element_data_type, true);
update_field_id(
&mut element_field,
*field_id_map.get_field_id(&element_path)?,
);
field_with_id(
name,
DataType::List(Arc::new(element_field)),
nullable,
field_id_map,
)
}
#[allow(
clippy::wildcard_enum_match_arm,
reason = "Arrow DataType has many variants; only List is valid for list columns"
)]
fn list_inner_field(schema: &Schema, column_name: &str) -> Result<Arc<Field>> {
let field = schema.field_with_name(column_name)?;
match field.data_type() {
DataType::List(inner) => Ok(Arc::<Field>::clone(inner)),
other => Err(IcebergError::ArrowTypeMismatch {
column_name: column_name.to_owned(),
expected: "List".to_owned(),
actual: format!("{other:?}"),
}
.into()),
}
}
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::*;
use crate::types::model::DataRowPayload;
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),
("payload_float32".to_owned(), 14),
("payload_float32_array".to_owned(), 15),
("payload_float32_array.element".to_owned(), 24),
("payload_float64".to_owned(), 16),
("payload_float64_array".to_owned(), 17),
("payload_float64_array.element".to_owned(), 25),
("payload_int".to_owned(), 12),
("payload_int_array".to_owned(), 13),
("payload_int_array.element".to_owned(), 23),
("payload_string".to_owned(), 18),
("payload_struct".to_owned(), 19),
("specification_id".to_owned(), 5),
("variant".to_owned(), 11),
("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() -> VariantDataRow {
VariantDataRow {
created_datetime: None,
data_stream_id: Uuid::nil(),
data_type: "audio".to_owned(),
datetime: Utc::now(),
modified_datetime: None,
payload_float32: None,
payload_float32_array: None,
payload_float64: None,
payload_float64_array: None,
payload_int: None,
payload_int_array: None,
payload_string: None,
payload_struct: None,
specification_id: Uuid::nil(),
variant: Some("vector".to_owned()),
vector: Some(vec![1.0, 2.0, 3.0]),
vector_end_bound: Some(100.0),
vector_start_bound: Some(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 = VariantDataRow {
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 = VariantDataRow::arrow_schema(&field_ids).unwrap();
assert_eq!(schema.fields().len(), 18);
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(), 18);
}
#[test]
fn data_row_uuids_serialized_as_strings() {
let field_ids = data_row_field_id_map();
let id = Uuid::new_v4();
let row = VariantDataRow {
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 data_row_float64_payload_serializes() {
let field_ids = data_row_field_id_map();
let row = VariantDataRow::from_payload(
Uuid::nil(),
Utc::now(),
"water_temperature".to_owned(),
Uuid::nil(),
DataRowPayload::Float64(22.5),
);
let batch = row.to_record_batch(&field_ids).unwrap();
assert_eq!(batch.num_columns(), 18);
let variant = batch
.column_by_name("variant")
.unwrap()
.as_any()
.downcast_ref::<arrow_array::StringArray>()
.unwrap();
assert_eq!(variant.value(0), "float64");
let payload = batch
.column_by_name("payload_float64")
.unwrap()
.as_any()
.downcast_ref::<arrow_array::Float64Array>()
.unwrap();
assert!((payload.value(0) - 22.5).abs() < f64::EPSILON);
}
#[test]
fn data_row_struct_payload_serializes_as_json_text() {
let field_ids = data_row_field_id_map();
let row = VariantDataRow::from_payload(
Uuid::nil(),
Utc::now(),
"cot".to_owned(),
Uuid::nil(),
DataRowPayload::Struct(serde_json::json!({"callsign": "ALPHA"})),
);
let batch = row.to_record_batch(&field_ids).unwrap();
let payload = batch
.column_by_name("payload_struct")
.unwrap()
.as_any()
.downcast_ref::<arrow_array::StringArray>()
.unwrap();
assert_eq!(payload.value(0), r#"{"callsign":"ALPHA"}"#);
}
#[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(), 18);
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 = VariantDataRow {
vector: Some(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);
}
}