diff --git a/crates/integrations/datafusion/src/sql_context.rs b/crates/integrations/datafusion/src/sql_context.rs index d13bb93bf..0fb1213fc 100644 --- a/crates/integrations/datafusion/src/sql_context.rs +++ b/crates/integrations/datafusion/src/sql_context.rs @@ -58,6 +58,7 @@ use datafusion::common::tree_node::{TreeNode, TreeNodeRecursion}; use datafusion::common::TableReference; use datafusion::datasource::{MemTable, TableProvider}; use datafusion::error::{DataFusionError, Result as DFResult}; +use datafusion::execution::context::SQLOptions; use datafusion::execution::runtime_env::RuntimeEnv; use datafusion::execution::SessionStateBuilder; use datafusion::logical_expr::{Expr as LogicalExpr, LogicalPlan, Volatility}; @@ -447,6 +448,11 @@ impl SQLContext { /// Execute a SQL statement. Paimon database and table extensions are handled /// directly; everything else is delegated to DataFusion. pub async fn sql(&self, sql: &str) -> DFResult { + self.sql_with_options(sql, SQLOptions::new()).await + } + + /// Execute a SQL statement with options for statements delegated to DataFusion. + pub async fn sql_with_options(&self, sql: &str, options: SQLOptions) -> DFResult { let is_create_table = looks_like_create_table(sql); let enable_ident_normalization = self.ctx.enable_ident_normalization(); let (rewritten_sql, partition_keys) = if is_create_table { @@ -456,7 +462,7 @@ impl SQLContext { }; if contains_time_travel_keyword(&rewritten_sql) { // Time-travel queries are not DDL; skip our own parsing and handle directly. - return self.handle_time_travel_query(&rewritten_sql).await; + return self.handle_time_travel_query(&rewritten_sql, options).await; } if let Some(show_partitions) = crate::format_partition_ddl::parse_show_partitions(&rewritten_sql)? @@ -522,7 +528,7 @@ impl SQLContext { Statement::Use(Use::Object(name)) => self.handle_use_database(name).await, Statement::CreateTable(create_table) => { if create_table.temporary { - self.handle_create_temp_table(create_table).await + self.handle_create_temp_table(create_table, options).await } else { let (catalog, _catalog_name, _) = self.resolve_catalog_and_table(&create_table.name)?; @@ -538,7 +544,7 @@ impl SQLContext { Statement::ShowCreate { obj_type: ShowCreateObject::Table, obj_name, - } => self.handle_show_create_table(sql, obj_name).await, + } => self.handle_show_create_table(sql, obj_name, options).await, Statement::AlterTable(alter_table) => { if alter_table.location.is_some() && alter_table.operations.iter().any(|operation| { @@ -572,7 +578,7 @@ impl SQLContext { if insert.overwrite && insert.partitioned.as_ref().is_some_and(|p| !p.is_empty()) => { - self.handle_insert_overwrite_partition(insert, enable_ident_normalization) + self.handle_insert_overwrite_partition(insert, enable_ident_normalization, options) .await } Statement::Set(Set::SingleAssignment { @@ -596,7 +602,7 @@ impl SQLContext { .insert(paimon_key.to_string(), value); return ok_result(&self.ctx); } - self.ctx.sql(sql).await + self.ctx.sql_with_options(sql, options).await } Statement::Reset(ResetStatement { reset: Reset::ConfigurationParameter(name), @@ -607,7 +613,7 @@ impl SQLContext { self.dynamic_options.write().unwrap().remove(paimon_key); return ok_result(&self.ctx); } - self.ctx.sql(sql).await + self.ctx.sql_with_options(sql, options).await } Statement::Truncate(truncate) => { self.handle_truncate_table(truncate, enable_ident_normalization) @@ -625,15 +631,15 @@ impl SQLContext { Statement::CreateView(create_view) => { if create_view.temporary { // Temporary views are always handled by us (Paimon catalog temp storage) - self.handle_create_view(create_view).await + self.handle_create_view(create_view, options).await } else { // Non-temporary views: only intercept if the target catalog is Paimon let view_name = create_view.name.to_string(); let table_ref: TableReference = view_name.as_str().into(); if self.is_paimon_catalog_ref(&table_ref) { - self.handle_create_view(create_view).await + self.handle_create_view(create_view, options).await } else { - self.ctx.sql(sql).await + self.ctx.sql_with_options(sql, options).await } } } @@ -641,7 +647,7 @@ impl SQLContext { if self.is_paimon_function_name(&create_function.name) { self.handle_create_function(create_function).await } else { - self.ctx.sql(sql).await + self.ctx.sql_with_options(sql, options).await } } Statement::Drop { @@ -686,7 +692,7 @@ impl SQLContext { self.resolve_catalog_and_table(&names[0])?; self.handle_drop_table(&catalog, names, *if_exists).await } else { - self.ctx.sql(sql).await + self.ctx.sql_with_options(sql, options).await } } else { let targets_paimon_catalog = names.iter().any(|name| { @@ -694,7 +700,7 @@ impl SQLContext { self.is_paimon_catalog_ref(&table_ref) }); if !targets_paimon_catalog { - return self.ctx.sql(sql).await; + return self.ctx.sql_with_options(sql, options).await; } let [name] = names.as_slice() else { return Err(DataFusionError::Plan( @@ -752,9 +758,11 @@ impl SQLContext { ¤t_database, ) .await?; - self.ctx.sql(&expanded.to_string()).await + self.ctx + .sql_with_options(&expanded.to_string(), options) + .await } - _ => self.ctx.sql(sql).await, + _ => self.ctx.sql_with_options(sql, options).await, } } @@ -766,7 +774,11 @@ impl SQLContext { /// 3. For each table, create a `PaimonTableProvider` with the appropriate scan options /// (merged with session-scoped dynamic options) /// 4. Register them as UUID-named temp tables, execute the rewritten SQL, then deregister - async fn handle_time_travel_query(&self, sql: &str) -> DFResult { + async fn handle_time_travel_query( + &self, + sql: &str, + sql_options: SQLOptions, + ) -> DFResult { use crate::table::PaimonTableProvider; use paimon::spec::{SCAN_TIMESTAMP_MILLIS_OPTION, SCAN_VERSION_OPTION}; @@ -897,7 +909,7 @@ impl SQLContext { ¤t_database, ) .await?; - self.ctx.sql(&expanded).await + self.ctx.sql_with_options(&expanded, sql_options).await } /// Parse a timestamp string to milliseconds since epoch (using local timezone). @@ -1043,7 +1055,11 @@ impl SQLContext { ok_result(&self.ctx) } - async fn handle_create_temp_table(&self, ct: &CreateTable) -> DFResult { + async fn handle_create_temp_table( + &self, + ct: &CreateTable, + sql_options: SQLOptions, + ) -> DFResult { let table_ref: TableReference = ct.name.to_string().as_str().into(); if ct.if_not_exists && self.temp_table_exist(table_ref.clone())? { @@ -1075,7 +1091,7 @@ impl SQLContext { if let Some(query) = &ct.query { // CREATE TEMPORARY TABLE ... AS SELECT ... let query_sql = query.to_string(); - let df = self.ctx.sql(&query_sql).await?; + let df = self.ctx.sql_with_options(&query_sql, sql_options).await?; let schema = df.schema().inner().clone(); let batches = df.collect().await?; @@ -1200,11 +1216,18 @@ impl SQLContext { ok_result(&self.ctx) } - async fn handle_show_create_table(&self, sql: &str, name: &ObjectName) -> DFResult { + async fn handle_show_create_table( + &self, + sql: &str, + name: &ObjectName, + sql_options: SQLOptions, + ) -> DFResult { let (catalog, catalog_name, identifier) = self.resolve_catalog_and_table(name)?; let table = match catalog.get_table(&identifier).await { Ok(table) => table, - Err(paimon::Error::TableNotExist { .. }) => return self.ctx.sql(sql).await, + Err(paimon::Error::TableNotExist { .. }) => { + return self.ctx.sql_with_options(sql, sql_options).await; + } Err(e) => return Err(to_datafusion_error(e)), }; crate::table_loader::ensure_paimon_served(&table, &identifier)?; @@ -1569,6 +1592,7 @@ impl SQLContext { &self, insert: &Insert, enable_ident_normalization: bool, + sql_options: SQLOptions, ) -> DFResult { self.ensure_no_time_travel_for_write("INSERT OVERWRITE")?; let table_name = match &insert.table { @@ -1600,7 +1624,10 @@ impl SQLContext { let source = insert.source.as_ref().ok_or_else(|| { DataFusionError::Plan("INSERT OVERWRITE requires a source query".into()) })?; - let df = self.ctx.sql(&source.to_string()).await?; + let df = self + .ctx + .sql_with_options(&source.to_string(), sql_options) + .await?; let all_fields = table.schema().fields(); let non_static_fields: Vec<&PaimonDataField> = all_fields @@ -1779,7 +1806,11 @@ impl SQLContext { ok_result(&self.ctx) } - async fn handle_create_view(&self, create_view: &CreateView) -> DFResult { + async fn handle_create_view( + &self, + create_view: &CreateView, + sql_options: SQLOptions, + ) -> DFResult { if create_view.materialized { return Err(DataFusionError::Plan( "CREATE MATERIALIZED VIEW is not supported".to_string(), @@ -1792,7 +1823,7 @@ impl SQLContext { let view_name = create_view.name.to_string(); let table_ref: TableReference = view_name.as_str().into(); let (catalog, database, name) = self.resolve_temp_table_name(table_ref)?; - let df = self.ctx.sql(&query_sql).await?; + let df = self.ctx.sql_with_options(&query_sql, sql_options).await?; let logical_plan = df.logical_plan().clone(); if create_view.if_not_exists && self.temp_table_exist(format!("{catalog}.{database}.{name}"))? diff --git a/crates/integrations/datafusion/tests/sql_context_tests.rs b/crates/integrations/datafusion/tests/sql_context_tests.rs index d33342ad5..16a66f299 100644 --- a/crates/integrations/datafusion/tests/sql_context_tests.rs +++ b/crates/integrations/datafusion/tests/sql_context_tests.rs @@ -24,6 +24,7 @@ use async_trait::async_trait; use datafusion::arrow::array::{Array, Int64Array}; use datafusion::catalog::CatalogProvider; use datafusion::datasource::MemTable; +use datafusion::execution::context::SQLOptions; use paimon::catalog::{list_partitions_from_file_system, Identifier}; use paimon::spec::{ ArrayType, BinaryType, BlobType, CharType, DataType, FloatType, IntType, @@ -1523,6 +1524,59 @@ async fn test_ddl_context_delegates_select() { assert_eq!(total_rows, 0, "Empty table should return 0 rows"); } +#[tokio::test] +async fn test_sql_with_options_propagates_to_datafusion() { + let ctx = SQLContext::new(); + let options = SQLOptions::new().with_allow_statements(false); + + for sql in [ + "SET datafusion.execution.batch_size = 1024", + "RESET datafusion.execution.batch_size", + ] { + let error = ctx.sql_with_options(sql, options).await.unwrap_err(); + assert!(error.to_string().contains("Statement not supported")); + } +} + +#[tokio::test] +async fn test_sql_with_options_propagates_to_copy() { + let temp_dir = TempDir::new().unwrap(); + let destination = temp_dir.path().join("copy.parquet"); + let ctx = SQLContext::new(); + let options = SQLOptions::new().with_allow_dml(false); + let sql = format!("COPY (VALUES (1)) TO '{}'", destination.display()); + + let error = ctx.sql_with_options(&sql, options).await.unwrap_err(); + + assert!( + error.to_string().contains("DML not supported: COPY"), + "unexpected error: {error}" + ); + assert!(!destination.exists()); +} + +#[tokio::test] +async fn test_sql_with_options_propagates_to_explain() { + let temp_dir = TempDir::new().unwrap(); + let ctx = SQLContext::new(); + let options = SQLOptions::new().with_allow_dml(false); + + for (explain, file_name) in [ + ("EXPLAIN", "explain-copy.parquet"), + ("EXPLAIN ANALYZE", "explain-analyze-copy.parquet"), + ] { + let destination = temp_dir.path().join(file_name); + let sql = format!("{explain} COPY (VALUES (1)) TO '{}'", destination.display()); + let error = ctx.sql_with_options(&sql, options).await.unwrap_err(); + + assert!( + error.to_string().contains("DML not supported: COPY"), + "unexpected error: {error}" + ); + assert!(!destination.exists()); + } +} + // ======================= MULTI-CATALOG ======================= #[tokio::test]