Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
157 changes: 152 additions & 5 deletions iceberg/src/codec.rs
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
use std::fmt;
use std::sync::Arc;

use datafusion::arrow::datatypes::SchemaRef;
Expand All @@ -19,22 +20,43 @@ use prost::Message;
use crate::proto::generated::iceberg as pb;
use crate::{IcebergDataSource, IcebergWorkUnitFeed};

type RuntimeResolver = dyn Fn(&TaskContext) -> Result<iceberg::Runtime> + Send + Sync;

/// Physical plan codec for [`IcebergDataSource`].
#[derive(Debug, Clone)]
#[derive(Clone)]
pub struct IcebergCodec {
storage_factory: Arc<dyn StorageFactory>,
iceberg_runtime: iceberg::Runtime,
runtime_resolver: Arc<RuntimeResolver>,
}

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 {
/// Creates a codec using the storage and runtime configured on this process.
pub fn new(
storage_factory: Arc<dyn StorageFactory>,
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<dyn StorageFactory>,
runtime_resolver: impl Fn(&TaskContext) -> Result<iceberg::Runtime> + Send + Sync + 'static,
) -> Self {
Self {
storage_factory,
iceberg_runtime,
runtime_resolver: Arc::new(runtime_resolver),
}
}
}
Expand Down Expand Up @@ -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,
}))
}
Expand Down Expand Up @@ -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]
Expand Down Expand Up @@ -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::<iceberg::Runtime>()
.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<TokioRuntime> {
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<PhysicalPlanNode> {
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<dyn ExecutionPlan>) -> Result<Arc<dyn ExecutionPlan>> {
if let Some(exec) = plan.downcast_ref::<DataSourceExec>()
&& exec
Expand Down
Loading