diff --git a/Cargo.lock b/Cargo.lock index 7fb5a2af4..c31fc9b4d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2497,6 +2497,7 @@ dependencies = [ "insta", "parquet", "prost", + "rmp-serde", "serde", "tokio", "tokio-stream", @@ -6313,6 +6314,25 @@ dependencies = [ "windows-sys 0.52.0", ] +[[package]] +name = "rmp" +version = "0.8.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4ba8be72d372b2c9b35542551678538b562e7cf86c3315773cae48dfbfe7790c" +dependencies = [ + "num-traits", +] + +[[package]] +name = "rmp-serde" +version = "1.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72f81bee8c8ef9b577d1681a70ebbc962c232461e397b22c208c43c04b67a155" +dependencies = [ + "rmp", + "serde", +] + [[package]] name = "roaring" version = "0.11.5" diff --git a/iceberg/Cargo.toml b/iceberg/Cargo.toml index 5d09b28eb..110221a2a 100644 --- a/iceberg/Cargo.toml +++ b/iceberg/Cargo.toml @@ -20,6 +20,7 @@ typetag = "0.2" datafusion-distributed = { path = "..", features = ["sql"] } delegate = "0.13" prost = "0.14.1" +rmp-serde = "1" tokio = { version = "1.48", features = ["full"] } tokio-stream = "0.1" diff --git a/iceberg/src/codec.rs b/iceberg/src/codec.rs index 2f5104a0f..f597ff52b 100644 --- a/iceberg/src/codec.rs +++ b/iceberg/src/codec.rs @@ -1,11 +1,7 @@ -use crate::work_unit_feed::FileScanTaskMessage; -use bytes::{Buf, BufMut}; use datafusion::common::Result; use datafusion::execution::TaskContext; use datafusion::physical_plan::ExecutionPlan; use datafusion_proto::physical_plan::PhysicalExtensionCodec; -use prost::encoding::{DecodeContext, WireType}; -use prost::{DecodeError, Message}; use std::sync::Arc; #[derive(Debug)] @@ -26,40 +22,3 @@ impl PhysicalExtensionCodec for IcebergCodec { unimplemented!() } } - -// TODO: Implement serde for FileScanTaskMessage -// This message is an individual WorkUnit, but it cannot be serialized yet. During distributed -// execution, this will be streamed over the wire from coordinator to workers, but for that to -// happen, it will need to be represented as a prost::Message. -// -// WARNING: for the ones who end up implementing this. I have no idea if implementing Message here -// is really the best option for serialization, it might not be. -impl Message for FileScanTaskMessage { - fn encode_raw(&self, _buf: &mut impl BufMut) - where - Self: Sized, - { - unimplemented!() - } - - fn merge_field( - &mut self, - _tag: u32, - _wire_type: WireType, - _buf: &mut impl Buf, - _ctx: DecodeContext, - ) -> std::result::Result<(), DecodeError> - where - Self: Sized, - { - unimplemented!() - } - - fn encoded_len(&self) -> usize { - unimplemented!() - } - - fn clear(&mut self) { - unimplemented!() - } -} diff --git a/iceberg/src/data_source.rs b/iceberg/src/data_source.rs index 6a3666822..4d27fba1d 100644 --- a/iceberg/src/data_source.rs +++ b/iceberg/src/data_source.rs @@ -2,7 +2,7 @@ use std::sync::Arc; use datafusion::arrow::datatypes::SchemaRef; use datafusion::common::stats::Precision; -use datafusion::common::{ColumnStatistics, Statistics, exec_datafusion_err}; +use datafusion::common::{ColumnStatistics, Statistics}; use datafusion::config::ConfigOptions; use datafusion::datasource::source::DataSource; use datafusion::error::Result; @@ -22,6 +22,7 @@ use iceberg::arrow::ArrowReaderBuilder; use iceberg::spec::SnapshotRef; use crate::common::{convert_filters_to_predicate, df_err, iceberg_err}; +use crate::work_unit_wire::FileScanTaskDecoder; use crate::{IcebergConfig, IcebergWorkUnitFeed}; /// Snapshot summary keys defined by the Iceberg table spec: @@ -209,14 +210,12 @@ impl DataSource for IcebergDataSource { .with_row_selection_enabled(config.row_selection_enabled) .build(); + let mut decoder = FileScanTaskDecoder::default(); let feed = self .feed .feed(partition, context)? - .map(|msg_or_err| match msg_or_err { - Ok(msg) => match msg.inner { - Some(msg) => Ok(msg), - None => Err(iceberg_err(exec_datafusion_err!("Missing inner"))), - }, + .map(move |msg_or_err| match msg_or_err { + Ok(msg) => decoder.decode(msg).map_err(iceberg_err), Err(err) => Err(iceberg_err(err)), }) .boxed(); diff --git a/iceberg/src/lib.rs b/iceberg/src/lib.rs index cb90c5e8a..e104550f1 100644 --- a/iceberg/src/lib.rs +++ b/iceberg/src/lib.rs @@ -13,6 +13,7 @@ mod distributed_desired_task_count_handler; mod iceberg_ext; mod table_provider; mod work_unit_feed; +mod work_unit_wire; mod codec; #[doc(hidden)] diff --git a/iceberg/src/test_utils/harness.rs b/iceberg/src/test_utils/harness.rs index b837b2ca1..46e32e83d 100644 --- a/iceberg/src/test_utils/harness.rs +++ b/iceberg/src/test_utils/harness.rs @@ -62,6 +62,11 @@ impl IcebergTestHarness { pub async fn physical_plan(&self, sql: &str) -> Result> { self.ctx.sql(sql).await?.create_physical_plan().await } + + #[cfg(test)] + pub(crate) fn task_context(&self) -> Arc { + self.ctx.task_ctx() + } } #[derive(Debug, Clone, Serialize, Deserialize)] diff --git a/iceberg/src/work_unit_feed.rs b/iceberg/src/work_unit_feed.rs index 6e73eb3d1..af9cea052 100644 --- a/iceberg/src/work_unit_feed.rs +++ b/iceberg/src/work_unit_feed.rs @@ -6,14 +6,14 @@ use datafusion::error::DataFusionError; use datafusion::execution::TaskContext; use datafusion::physical_expr::Partitioning; use datafusion_distributed::{DistributedWorkUnitFeedContext, WorkUnitFeedProvider}; +use futures::StreamExt; use futures::stream::BoxStream; -use futures::{StreamExt, TryStreamExt}; use iceberg::expr::Predicate; -use iceberg::scan::FileScanTask; use tokio::sync::mpsc::UnboundedReceiver; use tokio_stream::wrappers::UnboundedReceiverStream; use crate::common::df_err; +use crate::work_unit_wire::{FileScanTaskEncoder, FileScanTaskMessage}; /// Work unit feed implementation that yields [FileScanTask] messages at execution time. /// @@ -116,17 +116,6 @@ pub(crate) struct SyncManager { feeds: TakeableVec>>, } -#[derive(Debug, Clone, Default)] -pub struct FileScanTaskMessage { - pub(crate) inner: Option, -} - -impl FileScanTaskMessage { - fn new(inner: FileScanTask) -> Self { - Self { inner: Some(inner) } - } -} - impl WorkUnitFeedProvider for IcebergWorkUnitFeed { type WorkUnit = FileScanTaskMessage; @@ -173,7 +162,7 @@ impl WorkUnitFeedProvider for IcebergWorkUnitFeed { // the return streams of the `feed()` method. let task = SpawnedTask::spawn(async move { let mut stream = match table_scan.plan_files().await { - Ok(stream) => stream.map_ok(FileScanTaskMessage::new), + Ok(stream) => stream, Err(err) => { let _ = txs[0].send(Err(df_err(err))); return; @@ -183,9 +172,16 @@ impl WorkUnitFeedProvider for IcebergWorkUnitFeed { // Round robing across output partitions. // TODO: this is fine for Partitioning::UnknownPartitioning, but any other // partitioning will require smarter routing across output channels. + let mut encoders = (0..txs.len()) + .map(|_| FileScanTaskEncoder::default()) + .collect::>(); let mut i = 0; while let Some(scan_task_or_err) = stream.next().await { - let _ = txs[i % txs.len()].send(scan_task_or_err.map_err(df_err)); + let partition = i % txs.len(); + let message = scan_task_or_err + .map_err(df_err) + .and_then(|task| encoders[partition].encode(task)); + let _ = txs[partition].send(message); i += 1; } }); diff --git a/iceberg/src/work_unit_wire.rs b/iceberg/src/work_unit_wire.rs new file mode 100644 index 000000000..923d0c1ea --- /dev/null +++ b/iceberg/src/work_unit_wire.rs @@ -0,0 +1,371 @@ +use std::collections::HashMap; +use std::sync::Arc; + +use datafusion::common::{Result, exec_datafusion_err, exec_err}; +use iceberg::expr::BoundPredicate; +use iceberg::scan::{FileScanTask, FileScanTaskDeleteFile}; +use iceberg::spec::{ + DataFileFormat, Datum, Literal, NameMapping, PartitionSpec, PrimitiveLiteral, PrimitiveType, + Schema, SchemaRef, Struct, Type, +}; +use prost::Message; +use serde::{Deserialize, Serialize}; + +/// Wire representation of one Iceberg file scan task. +#[derive(Clone, PartialEq, Message)] +pub struct FileScanTaskMessage { + #[prost(bytes = "vec", tag = "1")] + payload: Vec, +} + +impl FileScanTaskMessage { + fn from_wire(wire: &FileScanTaskWire) -> Result { + let payload = rmp_serde::to_vec_named(wire).map_err(|error| { + exec_datafusion_err!("failed to serialize Iceberg file scan task: {error}") + })?; + Ok(Self { payload }) + } + + fn into_wire(self) -> Result { + rmp_serde::from_slice(&self.payload).map_err(|error| { + exec_datafusion_err!("failed to deserialize Iceberg file scan task: {error}") + }) + } +} + +#[derive(Serialize, Deserialize)] +struct FileScanTaskWire { + context_id: u32, + context: Option, + body: FileScanTaskBody, +} + +#[derive(PartialEq, Serialize, Deserialize)] +struct SharedTaskContext { + schema: SchemaRef, + project_field_ids: Vec, + predicate: Option, + partition_spec: Option>, + name_mapping: Option>, + case_sensitive: bool, +} + +impl SharedTaskContext { + fn from_task(task: &FileScanTask) -> Self { + Self { + schema: Arc::clone(&task.schema), + project_field_ids: task.project_field_ids.clone(), + predicate: task.predicate.clone(), + partition_spec: task.partition_spec.clone(), + name_mapping: task.name_mapping.clone(), + case_sensitive: task.case_sensitive, + } + } + + fn matches(&self, task: &FileScanTask) -> bool { + self.schema == task.schema + && self.project_field_ids == task.project_field_ids + && self.predicate == task.predicate + && self.partition_spec == task.partition_spec + && self.name_mapping == task.name_mapping + && self.case_sensitive == task.case_sensitive + } + + fn validate(mut self) -> Result { + for field_id in &self.project_field_ids { + if self.schema.field_by_id(*field_id).is_none() { + return exec_err!("Iceberg work unit projects unknown field id {field_id}"); + } + } + self.partition_spec = self + .partition_spec + .take() + .map(|spec| { + Arc::unwrap_or_clone(spec) + .into_unbound() + .bind(Arc::clone(&self.schema)) + .map(Arc::new) + .map_err(|error| { + exec_datafusion_err!("invalid Iceberg work unit partition spec: {error}") + }) + }) + .transpose()?; + Ok(self) + } +} + +#[derive(Serialize, Deserialize)] +struct FileScanTaskBody { + file_size_in_bytes: u64, + start: u64, + length: u64, + record_count: Option, + data_file_path: String, + data_file_format: DataFileFormat, + deletes: Vec, + partition: Option>>>, +} + +impl FileScanTaskBody { + fn try_from_task(mut task: FileScanTask) -> Result { + validate_file_range(task.file_size_in_bytes, task.start, task.length)?; + let partition = partition_to_wire( + task.partition.take(), + task.partition_spec.as_deref(), + &task.schema, + )?; + Ok(Self { + file_size_in_bytes: task.file_size_in_bytes, + start: task.start, + length: task.length, + record_count: task.record_count, + data_file_path: task.data_file_path, + data_file_format: task.data_file_format, + deletes: task.deletes, + partition, + }) + } + + fn into_task(self, context: &SharedTaskContext) -> Result { + validate_file_range(self.file_size_in_bytes, self.start, self.length)?; + let partition = partition_from_wire( + self.partition, + context.partition_spec.as_deref(), + &context.schema, + )?; + Ok(FileScanTask { + file_size_in_bytes: self.file_size_in_bytes, + start: self.start, + length: self.length, + record_count: self.record_count, + data_file_path: self.data_file_path, + data_file_format: self.data_file_format, + schema: Arc::clone(&context.schema), + project_field_ids: context.project_field_ids.clone(), + predicate: context.predicate.clone(), + deletes: self.deletes, + partition, + partition_spec: context.partition_spec.clone(), + name_mapping: context.name_mapping.clone(), + case_sensitive: context.case_sensitive, + }) + } +} + +/// Encodes tasks while defining shared serde context only once per feed. +#[derive(Default)] +pub(crate) struct FileScanTaskEncoder { + contexts: Vec, +} + +impl FileScanTaskEncoder { + pub(crate) fn encode(&mut self, task: FileScanTask) -> Result { + let existing = self + .contexts + .iter() + .position(|context| context.matches(&task)); + let (context_id, context) = match existing { + Some(index) => (context_id(index)?, None), + None => ( + context_id(self.contexts.len())?, + Some(SharedTaskContext::from_task(&task)), + ), + }; + let mut wire = FileScanTaskWire { + context_id, + context, + body: FileScanTaskBody::try_from_task(task)?, + }; + let message = FileScanTaskMessage::from_wire(&wire)?; + if let Some(context) = wire.context.take() { + self.contexts.push(context); + } + Ok(message) + } +} + +/// Decodes a single feed, retaining shared serde contexts referenced by later tasks. +#[derive(Default)] +pub(crate) struct FileScanTaskDecoder { + contexts: HashMap, +} + +impl FileScanTaskDecoder { + pub(crate) fn decode(&mut self, message: FileScanTaskMessage) -> Result { + let wire = message.into_wire()?; + if wire.context_id == 0 { + return exec_err!("Iceberg work unit context id must be non-zero"); + } + let definition = wire.context.map(SharedTaskContext::validate).transpose()?; + if let (Some(existing), Some(definition)) = + (self.contexts.get(&wire.context_id), definition.as_ref()) + && existing != definition + { + return exec_err!( + "Iceberg work unit context id {} was redefined", + wire.context_id + ); + } + let Some(context) = definition + .as_ref() + .or_else(|| self.contexts.get(&wire.context_id)) + else { + return exec_err!( + "Iceberg work unit references undefined context id {}", + wire.context_id + ); + }; + + let task = wire.body.into_task(context)?; + if let Some(definition) = definition { + self.contexts.entry(wire.context_id).or_insert(definition); + } + Ok(task) + } +} + +fn context_id(index: usize) -> Result { + u32::try_from(index) + .ok() + .and_then(|index| index.checked_add(1)) + .ok_or_else(|| exec_datafusion_err!("too many Iceberg task contexts in one feed")) +} + +fn validate_file_range(file_size: u64, start: u64, length: u64) -> Result<()> { + let Some(end) = start.checked_add(length) else { + return exec_err!("Iceberg work unit file range overflows u64"); + }; + if end > file_size { + return exec_err!( + "Iceberg work unit file range {start}..{end} exceeds file size {file_size}" + ); + } + Ok(()) +} + +fn partition_to_wire( + partition: Option, + spec: Option<&PartitionSpec>, + schema: &Schema, +) -> Result>>>> { + // iceberg-rust 0.10 does not propagate the manifest partition spec into planned tasks. + // Omit values that cannot be interpreted without their matching type metadata. + let Some(spec) = spec else { + return Ok(None); + }; + let Some(partition) = partition else { + return exec_err!("Iceberg work unit has a partition spec without partition values"); + }; + let partition_type = spec.partition_type(schema).map_err(|error| { + exec_datafusion_err!("invalid Iceberg work unit partition spec: {error}") + })?; + if partition.iter().len() != partition_type.fields().len() { + return exec_err!( + "Iceberg work unit has {} partition values for {} partition fields", + partition.iter().len(), + partition_type.fields().len() + ); + } + + partition + .into_iter() + .zip(partition_type.fields()) + .enumerate() + .map(|(index, (value, field))| { + let Some(value) = value else { + return Ok(None); + }; + let Literal::Primitive(value) = value else { + return exec_err!("Iceberg partition value {index} is not primitive"); + }; + let Type::Primitive(data_type) = field.field_type.as_ref() else { + return exec_err!("Iceberg partition field {index} is not primitive"); + }; + if !data_type.compatible(&value) { + return exec_err!( + "Iceberg partition value {index} is incompatible with type {data_type}" + ); + } + primitive_literal_bytes(value).map(Some) + }) + .collect::>>() + .map(Some) +} + +fn partition_from_wire( + partition: Option>>>, + spec: Option<&PartitionSpec>, + schema: &Schema, +) -> Result> { + let (partition, spec) = match (partition, spec) { + (None, None) => return Ok(None), + (Some(partition), Some(spec)) => (partition, spec), + (Some(_), None) => { + return exec_err!("Iceberg work unit has partition values without a partition spec"); + } + (None, Some(_)) => { + return exec_err!("Iceberg work unit has a partition spec without partition values"); + } + }; + let partition_type = spec.partition_type(schema).map_err(|error| { + exec_datafusion_err!("invalid Iceberg work unit partition spec: {error}") + })?; + if partition.len() != partition_type.fields().len() { + return exec_err!( + "Iceberg work unit has {} partition values for {} partition fields", + partition.len(), + partition_type.fields().len() + ); + } + + partition + .into_iter() + .zip(partition_type.fields()) + .enumerate() + .map(|(index, (value, field))| { + let Some(value) = value else { + return Ok(None); + }; + let Type::Primitive(data_type) = field.field_type.as_ref() else { + return exec_err!("Iceberg partition field {index} is not primitive"); + }; + if let PrimitiveType::Fixed(length) = data_type + && usize::try_from(*length).ok() != Some(value.len()) + { + return exec_err!( + "Iceberg partition value {index} has {} bytes for fixed[{length}]", + value.len() + ); + } + let datum = Datum::try_from_bytes(&value, data_type.clone()).map_err(|error| { + exec_datafusion_err!("invalid Iceberg partition value {index}: {error}") + })?; + let literal = datum.literal().clone(); + if primitive_literal_bytes(literal.clone())? != value { + return exec_err!("Iceberg partition value {index} has non-canonical bytes"); + } + Ok(Some(Literal::Primitive(literal))) + }) + .collect::>() + .map(Some) +} + +fn primitive_literal_bytes(literal: PrimitiveLiteral) -> Result> { + Ok(match literal { + PrimitiveLiteral::Boolean(value) => vec![u8::from(value)], + PrimitiveLiteral::Int(value) => value.to_le_bytes().to_vec(), + PrimitiveLiteral::Long(value) => value.to_le_bytes().to_vec(), + PrimitiveLiteral::Float(value) => value.to_le_bytes().to_vec(), + PrimitiveLiteral::Double(value) => value.to_le_bytes().to_vec(), + PrimitiveLiteral::String(value) => value.into_bytes(), + PrimitiveLiteral::Binary(value) => value, + PrimitiveLiteral::Int128(value) => value.to_be_bytes().to_vec(), + PrimitiveLiteral::UInt128(value) => value.to_be_bytes().to_vec(), + PrimitiveLiteral::AboveMax | PrimitiveLiteral::BelowMin => { + return exec_err!("Iceberg partition values cannot be range sentinels"); + } + }) +} + +#[cfg(test)] +include!("work_unit_wire_tests.rs"); diff --git a/iceberg/src/work_unit_wire_tests.rs b/iceberg/src/work_unit_wire_tests.rs new file mode 100644 index 000000000..843107cd9 --- /dev/null +++ b/iceberg/src/work_unit_wire_tests.rs @@ -0,0 +1,331 @@ +mod tests { + use super::*; + use crate::IcebergDataSource; + use crate::test_utils::IcebergTestHarness; + use datafusion::datasource::source::DataSourceExec; + use datafusion::physical_plan::ExecutionPlan; + use futures::TryStreamExt; + use iceberg::expr::{Bind, Reference}; + use iceberg::spec::{Datum, NestedField, PrimitiveType, Transform}; + + #[test] + fn roundtrips_complete_partitioned_task() { + let task = sample_task("file.parquet", 10); + + assert_eq!(roundtrip(task.clone()).unwrap(), task); + } + + #[test] + fn defines_reuses_and_replaces_shared_context() { + let mut encoder = FileScanTaskEncoder::default(); + let first = encoder.encode(sample_task("first.parquet", 10)).unwrap(); + let second = encoder.encode(sample_task("second.parquet", 20)).unwrap(); + let mut changed_task = sample_task("third.parquet", 30); + changed_task.case_sensitive = false; + let changed = encoder.encode(changed_task).unwrap(); + + assert_eq!(context_state(&first), (1, true)); + assert_eq!(context_state(&second), (1, false)); + assert_eq!(context_state(&changed), (2, true)); + + let mut decoder = FileScanTaskDecoder::default(); + let first = decode(&mut decoder, first).unwrap(); + let second = decode(&mut decoder, second).unwrap(); + assert_eq!(first.data_file_path, "first.parquet"); + assert_eq!(second.data_file_path, "second.parquet"); + assert!(Arc::ptr_eq(&first.schema, &second.schema)); + assert!(Arc::ptr_eq( + first.partition_spec.as_ref().unwrap(), + second.partition_spec.as_ref().unwrap() + )); + assert!(!decode(&mut decoder, changed).unwrap().case_sensitive); + } + + #[tokio::test] + async fn taxi_tasks_benefit_from_shared_context() { + let (tasks, shared_size) = taxi_file_scan_tasks().await.unwrap(); + assert_eq!(tasks.len(), 7); + + let repeated_size = tasks + .into_iter() + .map(|task| { + FileScanTaskEncoder::default() + .encode(task) + .unwrap() + .encoded_len() + }) + .sum::(); + + assert!( + shared_size * 4 < repeated_size * 3, + "expected shared taxi context to reduce encoded size by at least 25%: shared={shared_size}, repeated={repeated_size}" + ); + } + + #[test] + fn omits_partition_values_when_dependency_omits_their_spec() { + let mut task = sample_task("file.parquet", 10); + task.partition_spec = None; + + let decoded = roundtrip(task).unwrap(); + + assert!(decoded.partition.is_none()); + assert!(decoded.partition_spec.is_none()); + } + + #[test] + fn roundtrips_iceberg_primitive_value_bytes() { + let mut datums = vec![ + Datum::bool(true), + Datum::int(-12), + Datum::long(i64::MIN + 1), + Datum::float(f32::from_bits(0x8000_0000)), + Datum::double(f64::from_bits(0x7ff8_0000_0000_0001)), + Datum::string("iceberg"), + Datum::binary([0, 1, 255]), + Datum::fixed([0, 1, 255]), + Datum::date(1), + Datum::timestamp_micros(1), + ]; + datums.push( + Datum::try_from_bytes(&i128::MIN.to_be_bytes(), PrimitiveType::Decimal { + precision: 38, + scale: 0, + }) + .unwrap(), + ); + datums.push(Datum::try_from_bytes(&u128::MAX.to_be_bytes(), PrimitiveType::Uuid).unwrap()); + + for datum in datums { + let bytes = primitive_literal_bytes(datum.literal().clone()).unwrap(); + let decoded = Datum::try_from_bytes(&bytes, datum.data_type().clone()).unwrap(); + assert_eq!(decoded, datum); + } + for sentinel in [PrimitiveLiteral::AboveMax, PrimitiveLiteral::BelowMin] { + assert!(primitive_literal_bytes(sentinel).is_err()); + } + } + + #[test] + fn rejects_non_primitive_partition_values() { + let mut task = sample_task("file.parquet", 10); + task.partition = Some( + [Some(Literal::List(vec![Some(Literal::int(10))]))] + .into_iter() + .collect(), + ); + + let error = FileScanTaskEncoder::default().encode(task).unwrap_err(); + + assert!(error.to_string().contains("not primitive")); + } + + #[test] + fn rejects_invalid_decoded_tasks() { + assert_rejected( + |wire| { + wire.body.start = u64::MAX; + wire.body.length = 2; + wire.body.file_size_in_bytes = u64::MAX; + }, + "overflows", + ); + assert_rejected( + |wire| wire.body.partition.as_mut().unwrap().clear(), + "partition values for 1 partition fields", + ); + assert_rejected( + |wire| wire.body.partition = Some(vec![Some(vec![0; 3])]), + "invalid Iceberg partition value 0", + ); + assert_rejected( + |wire| wire.context.as_mut().unwrap().project_field_ids.push(99), + "projects unknown field id 99", + ); + } + + #[test] + fn failed_task_does_not_cache_its_context() { + let valid = encoded_task(); + let context_id = wire(&valid).context_id; + let invalid = edit(valid.clone(), |wire| { + wire.body.start = u64::MAX; + wire.body.length = 2; + wire.body.file_size_in_bytes = u64::MAX; + }); + let reference = edit(valid, |wire| wire.context = None); + let mut decoder = FileScanTaskDecoder::default(); + + assert!(decoder.decode(invalid).is_err()); + let error = decoder.decode(reference).unwrap_err(); + + assert!(error + .to_string() + .contains(&format!("undefined context id {context_id}"))); + } + + #[test] + fn rebinds_partition_spec_even_when_partition_value_is_null() { + let message = edit(encoded_task(), |wire| { + let context = wire.context.as_mut().unwrap(); + context.schema = mismatched_schema(); + context.project_field_ids = vec![2]; + context.predicate = None; + wire.body.partition = Some(vec![None]); + }); + + assert_decode_error(message, "Cannot find partition source field with id `1`"); + } + + async fn taxi_file_scan_tasks() -> Result<(Vec, usize)> { + let harness = IcebergTestHarness::new().await?; + let plan = harness.physical_plan("SELECT * FROM taxi").await?; + let exec = find_iceberg_exec(&plan).expect("plan contains an Iceberg DataSourceExec"); + let data_source = exec + .data_source() + .downcast_ref::() + .expect("DataSourceExec contains an IcebergDataSource"); + let partition_count = exec.properties().partitioning.partition_count(); + let context = harness.task_context(); + let streams = (0..partition_count) + .map(|partition| data_source.feed().feed(partition, Arc::clone(&context))) + .collect::>>()?; + + let mut tasks = vec![]; + let mut encoded_size = 0; + for stream in streams { + let messages = stream.try_collect::>().await?; + let mut decoder = FileScanTaskDecoder::default(); + for message in messages { + encoded_size += message.encoded_len(); + tasks.push(decoder.decode(message)?); + } + } + Ok((tasks, encoded_size)) + } + + fn find_iceberg_exec(plan: &Arc) -> Option> { + if let Some(exec) = plan.downcast_ref::() + && exec + .data_source() + .downcast_ref::() + .is_some() + { + return Some(Arc::new(exec.clone())); + } + plan.children().into_iter().find_map(find_iceberg_exec) + } + + fn roundtrip(task: FileScanTask) -> Result { + let message = FileScanTaskEncoder::default().encode(task)?; + decode(&mut FileScanTaskDecoder::default(), message) + } + + fn decode( + decoder: &mut FileScanTaskDecoder, + message: FileScanTaskMessage, + ) -> Result { + decoder.decode(prost_roundtrip(message)) + } + + fn context_state(message: &FileScanTaskMessage) -> (u32, bool) { + let wire = wire(message); + (wire.context_id, wire.context.is_some()) + } + + fn encoded_task() -> FileScanTaskMessage { + FileScanTaskEncoder::default() + .encode(sample_task("file.parquet", 10)) + .unwrap() + } + + fn edit( + message: FileScanTaskMessage, + edit: impl FnOnce(&mut FileScanTaskWire), + ) -> FileScanTaskMessage { + let mut wire = message.into_wire().unwrap(); + edit(&mut wire); + FileScanTaskMessage::from_wire(&wire).unwrap() + } + + fn wire(message: &FileScanTaskMessage) -> FileScanTaskWire { + message.clone().into_wire().unwrap() + } + + fn prost_roundtrip(message: FileScanTaskMessage) -> FileScanTaskMessage { + FileScanTaskMessage::decode(message.encode_to_vec().as_slice()).unwrap() + } + + #[track_caller] + fn assert_rejected(edit_wire: impl FnOnce(&mut FileScanTaskWire), expected: &str) { + assert_decode_error(edit(encoded_task(), edit_wire), expected); + } + + #[track_caller] + fn assert_decode_error(message: FileScanTaskMessage, expected: &str) { + let error = FileScanTaskDecoder::default() + .decode(prost_roundtrip(message)) + .unwrap_err(); + assert!(error.to_string().contains(expected), "{error}"); + } + + fn sample_task(path: &str, partition_value: i32) -> FileScanTask { + let schema = test_schema(); + let partition_spec = Arc::new( + PartitionSpec::builder(Arc::clone(&schema)) + .with_spec_id(7) + .add_partition_field("id", "id", Transform::Identity) + .unwrap() + .build() + .unwrap(), + ); + let predicate = Reference::new("id") + .greater_than(Datum::int(5)) + .bind(Arc::clone(&schema), true) + .unwrap(); + FileScanTask { + file_size_in_bytes: 100, + start: 10, + length: 80, + record_count: Some(12), + data_file_path: path.to_owned(), + data_file_format: DataFileFormat::Parquet, + schema, + project_field_ids: vec![1], + predicate: Some(predicate), + deletes: vec![], + partition: Some([Some(Literal::int(partition_value))].into_iter().collect()), + partition_spec: Some(partition_spec), + name_mapping: None, + case_sensitive: true, + } + } + + fn test_schema() -> SchemaRef { + Arc::new( + Schema::builder() + .with_schema_id(3) + .with_fields([Arc::new(NestedField::required( + 1, + "id", + Type::Primitive(PrimitiveType::Int), + ))]) + .with_identifier_field_ids([1]) + .build() + .unwrap(), + ) + } + + fn mismatched_schema() -> SchemaRef { + Arc::new( + Schema::builder() + .with_fields([Arc::new(NestedField::optional( + 2, + "other", + Type::Primitive(PrimitiveType::Int), + ))]) + .build() + .unwrap(), + ) + } +} diff --git a/src/test_utils/test_work_unit_feed.rs b/src/test_utils/test_work_unit_feed.rs index e05f1cf82..7fd27bd72 100644 --- a/src/test_utils/test_work_unit_feed.rs +++ b/src/test_utils/test_work_unit_feed.rs @@ -10,7 +10,7 @@ use datafusion::arrow::record_batch::RecordBatch; use datafusion::catalog::{Session, TableFunctionImpl}; use datafusion::common::stats::Precision; use datafusion::common::tree_node::{Transformed, TreeNode}; -use datafusion::common::{Result, ScalarValue, Statistics, internal_err, plan_err}; +use datafusion::common::{Result, ScalarValue, Statistics, exec_err, internal_err, plan_err}; use datafusion::datasource::{TableProvider, TableType}; use datafusion::error::DataFusionError; use datafusion::execution::{SendableRecordBatchStream, TaskContext}; @@ -34,6 +34,10 @@ use std::time::Duration; pub struct RowGeneratorWorkUnit { #[prost(uint64, tag = "1")] n_rows: u64, + #[prost(uint64, tag = "2")] + context_id: u64, + #[prost(uint64, optional, tag = "3")] + context: Option, } /// A scripted operation that the [`RowGeneratorFeedProvider`] performs on the @@ -118,6 +122,7 @@ impl RowGeneratorFeedProvider { struct FeedStreamState { iter: std::vec::IntoIter, counter: Count, + context: Option, done: bool, } @@ -139,6 +144,7 @@ impl WorkUnitFeedProvider for RowGeneratorFeedProvider { let state = FeedStreamState { iter: ops.into_iter(), counter, + context: Some(partition as u64), done: false, }; let stream = futures::stream::unfold(state, |mut state| async move { @@ -150,7 +156,12 @@ impl WorkUnitFeedProvider for RowGeneratorFeedProvider { match op { WorkUnitOp::Rows(n) => { state.counter.add(1); - return Some((Ok(RowGeneratorWorkUnit { n_rows: n }), state)); + let message = RowGeneratorWorkUnit { + n_rows: n, + context_id: 1, + context: state.context.take(), + }; + return Some((Ok(message), state)); } WorkUnitOp::Wait(d) => { tokio::time::sleep(d).await; @@ -409,9 +420,21 @@ impl ExecutionPlan for RowGeneratorExec { let schema = self.schema(); let tag = self.tag.clone(); let projection = self.projection.clone(); + let mut context = None; let stream = work_unit_feed.map(move |msg_result| { let msg = msg_result?; + if msg.context_id != 1 { + return exec_err!("unknown row-generator context id {}", msg.context_id); + } + if let Some(definition) = msg.context + && context.replace(definition).is_some() + { + return exec_err!("row-generator context was redefined"); + } + if context.is_none() { + return exec_err!("row-generator work unit references undefined context"); + } let n_rows = msg.n_rows as usize; // Build all columns, then select only the projected ones. let all_columns: Vec> = vec![ diff --git a/tests/work_unit_feed.rs b/tests/work_unit_feed.rs index 144684134..68689f96d 100644 --- a/tests/work_unit_feed.rs +++ b/tests/work_unit_feed.rs @@ -81,6 +81,31 @@ mod tests { Ok(()) } + /// Stateful work units define per-partition context once and reference it from later + /// messages. This verifies that coordinator batching and gRPC routing preserve each + /// partition's message order through the remote feed. + #[tokio::test] + async fn preserves_stateful_work_unit_context_per_partition() + -> Result<(), Box> { + let (_, results) = run_query( + r#" + SELECT * FROM test_work_unit( + 'source', 2, + 'rows(1),rows(1)', 'rows(1),rows(1)', + 'rows(1),rows(1)', 'rows(1),rows(1)' + ) + "#, + ) + .await?; + + let data_rows = results + .lines() + .filter(|line| line.starts_with("| source ")) + .count(); + assert_eq!(data_rows, 8, "expected every stateful work unit to arrive"); + Ok(()) + } + /// Tests that empty work unit feeds (no work units) produce no rows for that partition, /// while other partitions still work correctly through the distributed path. #[tokio::test]