Skip to content
Merged
Show file tree
Hide file tree
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
77 changes: 54 additions & 23 deletions crates/integrations/datafusion/src/sql_context.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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};
Expand Down Expand Up @@ -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<DataFrame> {
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<DataFrame> {
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 {
Expand All @@ -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)?
Expand Down Expand Up @@ -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)?;
Expand All @@ -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| {
Expand Down Expand Up @@ -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 {
Expand All @@ -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),
Expand All @@ -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)
Expand All @@ -625,23 +631,23 @@ 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
}
}
}
Statement::CreateFunction(create_function) => {
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 {
Expand Down Expand Up @@ -686,15 +692,15 @@ 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| {
let table_ref: TableReference = name.to_string().as_str().into();
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(
Expand Down Expand Up @@ -752,9 +758,11 @@ impl SQLContext {
&current_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,
}
}

Expand All @@ -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<DataFrame> {
async fn handle_time_travel_query(
&self,
sql: &str,
sql_options: SQLOptions,
) -> DFResult<DataFrame> {
use crate::table::PaimonTableProvider;
use paimon::spec::{SCAN_TIMESTAMP_MILLIS_OPTION, SCAN_VERSION_OPTION};

Expand Down Expand Up @@ -897,7 +909,7 @@ impl SQLContext {
&current_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).
Expand Down Expand Up @@ -1043,7 +1055,11 @@ impl SQLContext {
ok_result(&self.ctx)
}

async fn handle_create_temp_table(&self, ct: &CreateTable) -> DFResult<DataFrame> {
async fn handle_create_temp_table(
&self,
ct: &CreateTable,
sql_options: SQLOptions,
) -> DFResult<DataFrame> {
let table_ref: TableReference = ct.name.to_string().as_str().into();

if ct.if_not_exists && self.temp_table_exist(table_ref.clone())? {
Expand Down Expand Up @@ -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?;

Expand Down Expand Up @@ -1200,11 +1216,18 @@ impl SQLContext {
ok_result(&self.ctx)
}

async fn handle_show_create_table(&self, sql: &str, name: &ObjectName) -> DFResult<DataFrame> {
async fn handle_show_create_table(
&self,
sql: &str,
name: &ObjectName,
sql_options: SQLOptions,
) -> DFResult<DataFrame> {
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)?;
Expand Down Expand Up @@ -1569,6 +1592,7 @@ impl SQLContext {
&self,
insert: &Insert,
enable_ident_normalization: bool,
sql_options: SQLOptions,
) -> DFResult<DataFrame> {
self.ensure_no_time_travel_for_write("INSERT OVERWRITE")?;
let table_name = match &insert.table {
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -1779,7 +1806,11 @@ impl SQLContext {
ok_result(&self.ctx)
}

async fn handle_create_view(&self, create_view: &CreateView) -> DFResult<DataFrame> {
async fn handle_create_view(
&self,
create_view: &CreateView,
sql_options: SQLOptions,
) -> DFResult<DataFrame> {
if create_view.materialized {
return Err(DataFusionError::Plan(
"CREATE MATERIALIZED VIEW is not supported".to_string(),
Expand All @@ -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}"))?
Expand Down
54 changes: 54 additions & 0 deletions crates/integrations/datafusion/tests/sql_context_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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]
Expand Down
Loading