diff --git a/iceberg/src/codec.rs b/iceberg/src/codec.rs index 173e261b6..35bdef978 100644 --- a/iceberg/src/codec.rs +++ b/iceberg/src/codec.rs @@ -1,3 +1,4 @@ +use std::fmt; use std::sync::Arc; use datafusion::arrow::datatypes::SchemaRef; @@ -19,11 +20,21 @@ use prost::Message; use crate::proto::generated::iceberg as pb; use crate::{IcebergDataSource, IcebergWorkUnitFeed}; +type RuntimeResolver = dyn Fn(&TaskContext) -> Result + Send + Sync; + /// Physical plan codec for [`IcebergDataSource`]. -#[derive(Debug, Clone)] +#[derive(Clone)] pub struct IcebergCodec { storage_factory: Arc, - iceberg_runtime: iceberg::Runtime, + runtime_resolver: Arc, +} + +impl fmt::Debug for IcebergCodec { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("IcebergCodec") + .field("storage_factory", &self.storage_factory) + .finish_non_exhaustive() + } } impl IcebergCodec { @@ -31,10 +42,21 @@ impl IcebergCodec { pub fn new( storage_factory: Arc, iceberg_runtime: iceberg::Runtime, + ) -> Self { + Self::new_with_runtime_resolver(storage_factory, move |_| Ok(iceberg_runtime.clone())) + } + + /// Selects a runtime from the [`TaskContext`] for each decoded scan. + /// + /// The resolver is not called during encoding; its errors fail decoding. + /// Callers must keep the selected Tokio runtimes alive through execution. + pub fn new_with_runtime_resolver( + storage_factory: Arc, + runtime_resolver: impl Fn(&TaskContext) -> Result + Send + Sync + 'static, ) -> Self { Self { storage_factory, - iceberg_runtime, + runtime_resolver: Arc::new(runtime_resolver), } } } @@ -99,7 +121,7 @@ impl PhysicalExtensionCodec for IcebergCodec { column_stats: None, table_snapshot: None, iceberg_file_io, - iceberg_runtime: self.iceberg_runtime.clone(), + iceberg_runtime: (self.runtime_resolver)(ctx)?, feed, })) } @@ -142,10 +164,16 @@ impl PhysicalExtensionCodec for IcebergCodec { #[cfg(test)] mod tests { - use datafusion::common::Statistics; + use datafusion::common::{Statistics, exec_err}; use datafusion::datasource::source::DataSource; + use datafusion::prelude::SessionConfig; + use datafusion_distributed::{DistributedCodec, DistributedExt}; + use datafusion_proto::physical_plan::AsExecutionPlan; + use datafusion_proto::protobuf::PhysicalPlanNode; + use tokio::runtime::{Builder, Handle, Runtime as TokioRuntime}; use super::*; + use crate::common::df_err; use crate::test_utils::IcebergTestHarness; #[tokio::test] @@ -189,6 +217,125 @@ mod tests { Ok(()) } + #[test] + fn fixed_runtime_is_preserved() -> Result<()> { + let io = runtime()?; + let cpu = runtime()?; + let codec = IcebergCodec::new( + Arc::new(OpenDalResolvingStorageFactory::new()), + iceberg::Runtime::new_with_split(&io, &cpu), + ); + io.block_on(async { + let proto = encoded_scan(&codec).await?; + let ctx = task_context(iceberg::Runtime::new(&io)); + assert_decoded_runtime(&proto, &codec, &ctx, &io, &cpu).await + }) + } + + #[test] + fn default_runtime_is_captured_at_construction() -> Result<()> { + let startup = runtime()?; + let worker = runtime()?; + let codec = { + let _guard = startup.enter(); + IcebergCodec::default() + }; + worker.block_on(async { + let proto = encoded_scan(&codec).await?; + assert_decoded_runtime(&proto, &codec, &TaskContext::default(), &startup, &startup) + .await + }) + } + + #[test] + fn resolves_each_decode_with_query_cpu_and_shared_io() -> Result<()> { + let codec = IcebergCodec::new_with_runtime_resolver( + Arc::new(OpenDalResolvingStorageFactory::new()), + |ctx| { + Ok(ctx + .session_config() + .get_extension::() + .expect("worker runtime extension") + .as_ref() + .clone()) + }, + ); + let config = SessionConfig::new().with_distributed_user_codec(codec); + let codec = DistributedCodec::new_combined_with_user(&config); + let io = runtime()?; + let cpu_a = runtime()?; + let cpu_b = runtime()?; + io.block_on(async { + let proto = encoded_scan(&codec).await?; + let ctx_a = task_context(iceberg::Runtime::new_with_split(&io, &cpu_a)); + let ctx_b = task_context(iceberg::Runtime::new_with_split(&io, &cpu_b)); + assert_decoded_runtime(&proto, &codec, &ctx_a, &io, &cpu_a).await?; + assert_decoded_runtime(&proto, &codec, &ctx_b, &io, &cpu_b).await + }) + } + + #[tokio::test] + async fn propagates_runtime_resolver_errors_only_on_decode() -> Result<()> { + let codec = IcebergCodec::new_with_runtime_resolver( + Arc::new(OpenDalResolvingStorageFactory::new()), + |_| exec_err!("query runtime unavailable"), + ); + let proto = encoded_scan(&codec).await?; + let error = proto + .try_into_physical_plan(&TaskContext::default(), &codec) + .unwrap_err(); + assert!( + error.to_string().contains("query runtime unavailable"), + "{error}" + ); + Ok(()) + } + + fn runtime() -> Result { + Ok(Builder::new_multi_thread() + .worker_threads(1) + .enable_all() + .build()?) + } + + fn task_context(runtime: iceberg::Runtime) -> TaskContext { + TaskContext::default() + .with_session_config(SessionConfig::new().with_extension(Arc::new(runtime))) + } + + async fn encoded_scan(codec: &dyn PhysicalExtensionCodec) -> Result { + let harness = IcebergTestHarness::new().await?; + PhysicalPlanNode::try_from_physical_plan(harness.scan().await?, codec) + } + + async fn assert_decoded_runtime( + proto: &PhysicalPlanNode, + codec: &dyn PhysicalExtensionCodec, + ctx: &TaskContext, + io: &TokioRuntime, + cpu: &TokioRuntime, + ) -> Result<()> { + let decoded = proto.try_into_physical_plan(ctx, codec)?; + let runtime = &iceberg_source(&decoded)?.iceberg_runtime; + assert_eq!( + runtime + .io() + .spawn(async { Handle::current().id() }) + .await + .map_err(df_err)?, + io.handle().id(), + ); + assert_eq!( + runtime + .cpu() + .spawn(async { Handle::current().id() }) + .await + .map_err(df_err)?, + cpu.handle().id(), + ); + Ok(()) + } + fn iceberg_plan(plan: &Arc) -> Result> { if let Some(exec) = plan.downcast_ref::() && exec