use crate::Result;
use crate::error::CoreError;
use arrow_array::{ArrayRef, RecordBatch};
use arrow_avro::reader::{Decoder as ArrowAvroDecoder, ReaderBuilder};
use arrow_avro::schema::{AvroSchema as ArrowAvroSchema, SINGLE_OBJECT_MAGIC, SchemaStore};
use arrow_cast::cast;
use arrow_schema::{DataType, Field, Schema, SchemaRef};
use std::sync::Arc;
pub struct AvroBlockDecoder {
decoder: ArrowAvroDecoder,
registered: RegisteredWriterSchema,
reader_schema_json: Option<String>,
batch_size: usize,
prefix: [u8; 10],
framed: Vec<u8>,
rewrite_to: Option<SchemaRef>,
}
#[derive(Debug, Clone)]
pub struct RegisteredWriterSchema {
store: SchemaStore,
fingerprint: arrow_avro::schema::Fingerprint,
}
impl RegisteredWriterSchema {
pub fn new(writer_schema_json: &str) -> Result<Self> {
let mut store = SchemaStore::new();
let fingerprint = store
.register(ArrowAvroSchema::new(writer_schema_json.to_string()))
.map_err(|e| {
CoreError::LogBlockError(format!("Failed to register block writer schema: {e}"))
})?;
Ok(Self { store, fingerprint })
}
}
impl AvroBlockDecoder {
pub fn try_new_with_reader(
writer_schema_json: &str,
reader_schema_json: Option<&str>,
batch_size: usize,
) -> Result<Self> {
let registered = RegisteredWriterSchema::new(writer_schema_json)?;
Self::try_new_with_registered(®istered, reader_schema_json, batch_size)
}
fn build_inner(
registered: &RegisteredWriterSchema,
reader_schema_json: Option<&str>,
batch_size: usize,
) -> Result<ArrowAvroDecoder> {
let mut builder = ReaderBuilder::new()
.with_writer_schema_store(registered.store.clone())
.with_active_fingerprint(registered.fingerprint)
.with_batch_size(batch_size);
if let Some(reader_schema_json) = reader_schema_json {
builder =
builder.with_reader_schema(ArrowAvroSchema::new(reader_schema_json.to_string()));
}
builder
.build_decoder()
.map_err(|e| CoreError::LogBlockError(format!("Failed to build the Avro decoder: {e}")))
}
pub fn try_new_with_registered(
registered: &RegisteredWriterSchema,
reader_schema_json: Option<&str>,
batch_size: usize,
) -> Result<Self> {
let fingerprint = registered.fingerprint;
let arrow_avro::schema::Fingerprint::Rabin(rabin) = fingerprint else {
return Err(CoreError::LogBlockError(format!(
"Expected a Rabin fingerprint for the block writer schema, got {fingerprint:?}"
)));
};
let decoder = Self::build_inner(registered, reader_schema_json, batch_size)?;
let mut prefix = [0u8; 10];
prefix[..2].copy_from_slice(&SINGLE_OBJECT_MAGIC);
prefix[2..].copy_from_slice(&rabin.to_le_bytes());
Ok(Self {
decoder,
registered: registered.clone(),
reader_schema_json: reader_schema_json.map(str::to_string),
batch_size,
prefix,
framed: Vec::new(),
rewrite_to: None,
})
}
pub fn schema(&self) -> SchemaRef {
self.decoder.schema()
}
pub fn with_rewrite_to(mut self, schema: SchemaRef) -> Self {
self.rewrite_to = Some(schema);
self
}
pub fn decode(&mut self, body: &[u8]) -> Result<Option<RecordBatch>> {
self.framed.clear();
self.framed.reserve(self.prefix.len() + body.len());
self.framed.extend_from_slice(&self.prefix);
self.framed.extend_from_slice(body);
let consumed = self
.decoder
.decode(&self.framed)
.map_err(|e| CoreError::LogBlockError(format!("Failed to decode a log record: {e}")))?;
if consumed != self.framed.len() {
return Err(CoreError::LogBlockError(format!(
"Log record decoded partially: {consumed} of {} bytes",
self.framed.len()
)));
}
if self.decoder.batch_is_full() {
let batch = self.flush()?;
self.decoder = Self::build_inner(
&self.registered,
self.reader_schema_json.as_deref(),
self.batch_size,
)?;
return Ok(batch);
}
Ok(None)
}
pub fn flush(&mut self) -> Result<Option<RecordBatch>> {
let batch = self.decoder.flush().map_err(|e| {
CoreError::LogBlockError(format!("Failed to flush decoded records: {e}"))
})?;
let batch = batch.map(normalize_utc_timestamps).transpose()?;
match (batch, self.rewrite_to.as_ref()) {
(Some(batch), Some(target)) => {
crate::schema::batch_evolution::project_batch_to_schema(&batch, target).map(Some)
}
(batch, _) => Ok(batch),
}
}
}
fn normalize_utc_timestamps(batch: RecordBatch) -> Result<RecordBatch> {
fn is_utc_alias(tz: &str) -> bool {
matches!(tz, "+00:00" | "+0000" | "00:00" | "Z" | "z")
}
let needs_fix = batch
.schema()
.fields()
.iter()
.any(|f| matches!(f.data_type(), DataType::Timestamp(_, Some(tz)) if is_utc_alias(tz)));
if !needs_fix {
return Ok(batch);
}
let mut fields: Vec<Field> = Vec::with_capacity(batch.num_columns());
let mut columns: Vec<ArrayRef> = Vec::with_capacity(batch.num_columns());
for (field, column) in batch.schema().fields().iter().zip(batch.columns()) {
match field.data_type() {
DataType::Timestamp(unit, Some(tz)) if is_utc_alias(tz) => {
let relabeled = cast(column, &DataType::Timestamp(*unit, Some("UTC".into())))
.map_err(CoreError::ArrowError)?;
fields.push(
Field::new(
field.name(),
relabeled.data_type().clone(),
field.is_nullable(),
)
.with_metadata(field.metadata().clone()),
);
columns.push(relabeled);
}
_ => {
fields.push(field.as_ref().clone());
columns.push(column.clone());
}
}
}
RecordBatch::try_new(Arc::new(Schema::new(fields)), columns).map_err(CoreError::ArrowError)
}
#[cfg(test)]
mod multi_batch_tests {
use super::*;
use apache_avro::types::Value;
const UNION_SCHEMA: &str = r#"{"type":"record","name":"R","fields":[
{"name":"id","type":"long"},
{"name":"v","type":["null","int","string"]}]}"#;
fn null_union_records(n: usize) -> Vec<Vec<u8>> {
let schema = apache_avro::Schema::parse_str(UNION_SCHEMA).unwrap();
(0..n as i64)
.map(|i| {
let mut rec = apache_avro::types::Record::new(&schema).unwrap();
rec.put("id", i);
rec.put("v", Value::Union(0, Box::new(Value::Null)));
apache_avro::to_avro_datum(&schema, rec).unwrap()
})
.collect()
}
#[test]
fn a_block_larger_than_one_batch_decodes_with_a_union_in_the_schema() {
const BATCH: usize = 16;
const N: usize = 100;
let mut decoder =
AvroBlockDecoder::try_new_with_reader(UNION_SCHEMA, Some(UNION_SCHEMA), BATCH).unwrap();
let mut rows = 0usize;
let mut batches = 0usize;
for datum in &null_union_records(N) {
if let Some(batch) = decoder
.decode(datum)
.expect("a record after a full batch must still decode")
{
batches += 1;
rows += batch.num_rows();
}
}
if let Some(batch) = decoder.flush().expect("the final flush must succeed") {
batches += 1;
rows += batch.num_rows();
}
assert_eq!(rows, N, "every record must come back");
assert!(
batches > 1,
"the batch size must have forced more than one batch, or this proves nothing: \
{batches} batch(es) for {N} records at a batch size of {BATCH}"
);
}
#[test]
fn values_survive_the_batch_boundary_with_a_union() {
use arrow_array::cast::AsArray;
const BATCH: usize = 16;
const N: i64 = 100;
let schema = apache_avro::Schema::parse_str(UNION_SCHEMA).unwrap();
let datums: Vec<Vec<u8>> = (0..N)
.map(|i| {
let mut rec = apache_avro::types::Record::new(&schema).unwrap();
rec.put("id", i);
let v = if i % 2 == 0 {
Value::Union(1, Box::new(Value::Int(i as i32)))
} else {
Value::Union(2, Box::new(Value::String(format!("s{i}"))))
};
rec.put("v", v);
apache_avro::to_avro_datum(&schema, rec).unwrap()
})
.collect();
let mut decoder =
AvroBlockDecoder::try_new_with_reader(UNION_SCHEMA, Some(UNION_SCHEMA), BATCH).unwrap();
let mut ids: Vec<i64> = Vec::new();
let take = |batch: arrow_array::RecordBatch, ids: &mut Vec<i64>| {
let col = batch
.column_by_name("id")
.expect("the id column")
.as_primitive::<arrow_array::types::Int64Type>();
ids.extend(col.iter().flatten());
};
for datum in &datums {
if let Some(batch) = decoder.decode(datum).unwrap() {
take(batch, &mut ids);
}
}
if let Some(batch) = decoder.flush().unwrap() {
take(batch, &mut ids);
}
assert_eq!(
ids,
(0..N).collect::<Vec<_>>(),
"every row must come back once, in order, across the batch boundary"
);
}
#[test]
fn a_block_larger_than_one_batch_decodes_without_a_union() {
const SCHEMA: &str = r#"{"type":"record","name":"R","fields":[
{"name":"id","type":"long"},
{"name":"name","type":"string"}]}"#;
let schema = apache_avro::Schema::parse_str(SCHEMA).unwrap();
let datums: Vec<Vec<u8>> = (0..100i64)
.map(|i| {
let mut rec = apache_avro::types::Record::new(&schema).unwrap();
rec.put("id", i);
rec.put("name", format!("n{i}"));
apache_avro::to_avro_datum(&schema, rec).unwrap()
})
.collect();
let mut decoder = AvroBlockDecoder::try_new_with_reader(SCHEMA, Some(SCHEMA), 16).unwrap();
let mut rows = 0usize;
for d in &datums {
if let Some(b) = decoder.decode(d).unwrap() {
rows += b.num_rows();
}
}
if let Some(b) = decoder.flush().unwrap() {
rows += b.num_rows();
}
assert_eq!(rows, 100);
}
}
#[cfg(test)]
mod tests {
use super::AvroBlockDecoder;
#[test]
fn test_reader_schema_promotes_int_to_long() {
let writer = r#"{"type":"record","name":"r","fields":[{"name":"num","type":"int"}]}"#;
let reader = r#"{"type":"record","name":"r","fields":[{"name":"num","type":"long"}]}"#;
let mut decoder =
AvroBlockDecoder::try_new_with_reader(writer, Some(reader), 1024).unwrap();
decoder.decode(&[0x0E]).unwrap(); let batch = decoder.flush().unwrap().expect("a batch");
assert_eq!(
batch.schema().field(0).data_type(),
&arrow_schema::DataType::Int64
);
let col = batch
.column(0)
.as_any()
.downcast_ref::<arrow_array::Int64Array>()
.expect("promoted to i64");
assert_eq!(col.value(0), 7);
}
#[test]
fn test_reader_schema_promotes_through_a_nullable_union() {
let writer = r#"{"type":"record","name":"r","fields":[{"name":"num","type":["null","int"],"default":null}]}"#;
let reader = r#"{"type":"record","name":"r","fields":[{"name":"num","type":["null","long"],"default":null}]}"#;
let mut decoder =
AvroBlockDecoder::try_new_with_reader(writer, Some(reader), 1024).unwrap();
decoder.decode(&[0x02, 0x0E]).unwrap(); let batch = decoder.flush().unwrap().expect("a batch");
assert_eq!(
batch.schema().field(0).data_type(),
&arrow_schema::DataType::Int64
);
}
#[test]
fn test_no_reader_schema_keeps_the_writer_type() {
let writer = r#"{"type":"record","name":"r","fields":[{"name":"num","type":"int"}]}"#;
let mut decoder = AvroBlockDecoder::try_new_with_reader(writer, None, 1024).unwrap();
decoder.decode(&[0x0E]).unwrap();
let batch = decoder.flush().unwrap().expect("a batch");
assert_eq!(
batch.schema().field(0).data_type(),
&arrow_schema::DataType::Int32
);
}
}