Skip to content
Open
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
73 changes: 73 additions & 0 deletions datafusion/core/tests/sql/select.rs
Original file line number Diff line number Diff line change
Expand Up @@ -387,6 +387,79 @@ async fn test_query_parameters_with_metadata() -> Result<()> {
Ok(())
}

#[tokio::test]
async fn test_limit_offset_parameters_named() -> Result<()> {
let ctx = SessionContext::new();
let df = ctx
.sql(
"SELECT value FROM (VALUES (10), (20), (30)) AS t(value) \
ORDER BY value LIMIT $rows OFFSET $skip",
)
.await?;
assert_eq!(
df.logical_plan().get_parameter_types()?,
HashMap::from([
("$rows".to_string(), Some(DataType::Int64)),
("$skip".to_string(), Some(DataType::Int64)),
])
);
let results = df
.with_param_values(vec![
("rows", ScalarValue::Int64(Some(1))),
("skip", ScalarValue::Int64(Some(1))),
])?
.collect()
.await?;
datafusion::assert_batches_eq!(
[
"+-------+",
"| value |",
"+-------+",
"| 20 |",
"+-------+"
],
&results
);
Ok(())
}

#[tokio::test]
async fn test_limit_offset_parameters_keep_field_metadata() -> Result<()> {
let ctx = SessionContext::new();
let metadata = HashMap::from([("some_key".to_string(), "some_value".to_string())]);
let schema = Arc::new(Schema::new(vec![
Field::new("value", DataType::Int32, false).with_metadata(metadata.clone()),
]));
ctx.register_batch(
"t",
RecordBatch::try_new(schema, vec![Arc::new(Int32Array::from(vec![1, 1]))])?,
)?;

for clause in ["LIMIT", "OFFSET"] {
let sql = format!("SELECT $1 AS value FROM t WHERE value = $1 {clause} $1");
let df = ctx.sql(&sql).await?;
let results = df
.with_param_values(ParamValues::List(vec![ScalarAndMetadata::new(
ScalarValue::Int32(Some(1)),
Some(metadata.clone().into()),
)]))?
.collect()
.await?;
datafusion::assert_batches_eq!(
[
"+-------+",
"| value |",
"+-------+",
"| 1 |",
"+-------+"
],
&results
);
}

Ok(())
}

#[tokio::test]
async fn test_version_function() {
let expected_version = format!(
Expand Down
22 changes: 18 additions & 4 deletions datafusion/expr/src/logical_plan/plan.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1836,14 +1836,21 @@ impl LogicalPlan {
.collect())
}

/// Walk the logical plan, find any `Placeholder` tokens, and return a map of their IDs and FieldRefs
/// Walk the logical plan, find any `Placeholder` tokens, and return a map of their IDs and FieldRefs.
/// Bare `LIMIT`/`OFFSET` parameters default to `Int64` if no occurrence provides a type.
pub fn get_parameter_fields(
&self,
) -> Result<HashMap<String, Option<FieldRef>>, DataFusionError> {
let mut param_types: HashMap<String, Option<FieldRef>> = HashMap::new();
let mut row_count_parameters: HashSet<String> = HashSet::new();

self.apply_with_subqueries(|plan| {
plan.apply_expressions(|expr| {
if matches!(plan, LogicalPlan::Limit(_))
&& let Expr::Placeholder(Placeholder { id, field: None }) = expr
{
row_count_parameters.insert(id.clone());
}
expr.apply(|expr| {
if let Expr::Placeholder(Placeholder { id, field }) = expr {
let prev = param_types.get(id);
Expand All @@ -1860,15 +1867,22 @@ impl LogicalPlan {
param_types.insert(id.clone(), Some(Arc::clone(field)));
}
_ => {
param_types.insert(id.clone(), None);
param_types.entry(id.clone()).or_insert(None);
}
}
}
Ok(TreeNodeRecursion::Continue)
})
})
})
.map(|_| param_types)
})?;

for id in row_count_parameters {
param_types
.entry(id)
.or_default()
.get_or_insert_with(|| Arc::new(Field::new("", DataType::Int64, true)));
}
Ok(param_types)
}

// ------------
Expand Down
40 changes: 36 additions & 4 deletions datafusion/optimizer/src/analyzer/type_coercion.rs
Original file line number Diff line number Diff line change
Expand Up @@ -100,6 +100,25 @@ impl AnalyzerRule for TypeCoercion {
fn analyze(&self, plan: LogicalPlan, config: &ConfigOptions) -> Result<LogicalPlan> {
static EMPTY_SCHEMA: LazyLock<DFSchema> = LazyLock::new(DFSchema::empty);

// Record bare row-count parameter types before coercion makes them indistinguishable
// from explicit casts. A type inferred elsewhere takes precedence over the default.
let parameter_fields = plan.get_parameter_fields()?;
let plan = plan
.transform_up_with_subqueries(|plan| {
if !matches!(plan, LogicalPlan::Limit(_)) {
return Ok(Transformed::no(plan));
}
plan.map_expressions(|expr| match expr {
Expr::Placeholder(mut placeholder) if placeholder.field.is_none() => {
placeholder.field =
parameter_fields.get(&placeholder.id).cloned().flatten();
Ok(Transformed::yes(Expr::Placeholder(placeholder)))
}
expr => Ok(Transformed::no(expr)),
})
})?
.data;

// recurse
let transformed_plan = plan
.transform_up_with_subqueries(|plan| analyze_internal(&EMPTY_SCHEMA, plan))?
Expand Down Expand Up @@ -1631,10 +1650,10 @@ mod test {
use arrow::datatypes::{DataType, Field, Schema, SchemaBuilder, TimeUnit};
use insta::assert_snapshot;

use crate::analyzer::Analyzer;
use crate::analyzer::type_coercion::{
TypeCoercion, TypeCoercionRewriter, coerce_case_expression,
};
use crate::analyzer::{Analyzer, AnalyzerRule};
use crate::assert_analyzed_plan_with_config_eq_snapshot;
use datafusion_common::config::ConfigOptions;
use datafusion_common::tree_node::{TransformedResult, TreeNode};
Expand All @@ -1646,9 +1665,9 @@ mod test {
use datafusion_expr::test::function_stub::avg_udaf;
use datafusion_expr::{
AccumulatorFactoryFunction, AggregateUDF, BinaryExpr, Case, ColumnarValue, Expr,
ExprSchemable, Filter, LogicalPlan, Operator, ScalarFunctionArgs, ScalarUDF,
ScalarUDFImpl, Signature, SimpleAggregateUDF, Subquery, Union, Volatility, cast,
col, create_udaf, is_true, lit,
ExprSchemable, Filter, Limit, LogicalPlan, Operator, ScalarFunctionArgs,
ScalarUDF, ScalarUDFImpl, Signature, SimpleAggregateUDF, Subquery, Union,
Volatility, cast, col, create_udaf, is_true, lit, placeholder,
};
use datafusion_functions_aggregate::average::AvgAccumulator;

Expand Down Expand Up @@ -1733,6 +1752,19 @@ mod test {
Ok(())
}

#[test]
fn untyped_limit_parameter_remains_inferable_after_coercion() -> Result<()> {
let plan = LogicalPlan::Limit(Limit {
input: empty(),
fetch: Some(Box::new(placeholder("$1"))),
skip: None,
});
let analyzed = TypeCoercion::new().analyze(plan, &ConfigOptions::default())?;

assert_eq!(analyzed.get_parameter_types()?["$1"], Some(DataType::Int64));
Ok(())
}

#[test]
fn simple_case() -> Result<()> {
let expr = col("a").lt(lit(2_u32));
Expand Down
69 changes: 69 additions & 0 deletions datafusion/optimizer/tests/optimizer_integration.rs
Original file line number Diff line number Diff line change
Expand Up @@ -670,6 +670,75 @@ fn optimize_plan(plan: LogicalPlan) -> Result<LogicalPlan> {
optimizer.optimize(plan, &config, observe)
}

#[test]
fn limit_offset_parameters_keep_type_after_optimization() -> Result<()> {
for sql in [
"SELECT 20 AS value LIMIT $1",
"SELECT 20 AS value OFFSET $1",
] {
let optimized = test_sql(sql)?;
assert_eq!(
optimized.get_parameter_types()?,
HashMap::from([("$1".to_string(), Some(DataType::Int64))]),
"{sql}"
);
}
Ok(())
}

#[test]
fn limit_offset_parameters_leave_cast_inputs_unresolved_after_optimization() -> Result<()>
{
for sql in [
"SELECT $1 AS value",
"SELECT $1 AS value LIMIT CAST($1 AS INT)",
"SELECT $1 AS value LIMIT CAST($1 AS BIGINT)",
"SELECT $1 AS value FROM (VALUES (1), (2)) AS t(v) OFFSET CAST($1 AS BIGINT)",
"SELECT $1 AS value FROM (SELECT 1 LIMIT CAST($1 AS BIGINT)) AS t",
] {
let optimized = test_sql(sql)?;
assert_eq!(
optimized.get_parameter_types()?,
HashMap::from([("$1".to_string(), None)]),
"{sql}"
);
}
Ok(())
}

#[test]
fn limit_parameter_binding_with_inferred_type() -> Result<()> {
let statements = Parser::parse_sql(
&GenericDialect {},
"SELECT $1 + CAST(1 AS INT) AS value LIMIT $1",
)?;
let context_provider = MyContextProvider::default();
let plan =
SqlToRel::new(&context_provider).sql_statement_to_plan(statements[0].clone())?;
let analyzed =
Analyzer::new().execute_and_check(plan, &ConfigOptions::default(), |_, _| {})?;
let bound = analyzed.with_param_values(vec![ScalarValue::Int32(Some(1))])?;
let optimized = optimize_plan(bound)?;
let limit = match &optimized {
LogicalPlan::Limit(limit) => limit,
LogicalPlan::Projection(projection) => {
let LogicalPlan::Limit(limit) = projection.input.as_ref() else {
panic!("expected LIMIT under projection, got {optimized:?}");
};
limit
}
_ => panic!("expected LIMIT plan, got {optimized:?}"),
};
assert!(
matches!(
limit.get_fetch_type()?,
datafusion_expr::FetchType::Literal(Some(1))
),
"unexpected optimized LIMIT: {optimized}"
);
Ok(())
}

/// Extension node that does NOT implement `necessary_children_exprs`.
/// Used to test that the optimizer still processes subtrees below such nodes.
#[derive(Debug, Hash, PartialEq, Eq)]
Expand Down
14 changes: 8 additions & 6 deletions datafusion/sql/src/statement.rs
Original file line number Diff line number Diff line change
Expand Up @@ -880,16 +880,18 @@ impl<S: ContextProvider> SqlToRel<'_, S> {

if fields.is_empty() {
let map_types = plan.get_parameter_fields()?;
let param_types: Vec<_> = (1..=map_types.len())
.filter_map(|i| {
let param_types: Option<Vec<FieldRef>> = (1..=map_types.len())
.map(|i| {
let key = format!("${i}");
map_types.get(&key).and_then(|opt| opt.clone())
})
.collect();
fields.extend(param_types.iter().cloned());
planner_context.with_prepare_param_data_types(
param_types.into_iter().map(Some).collect(),
);
if let Some(param_types) = param_types {
fields.extend(param_types.iter().cloned());
planner_context.with_prepare_param_data_types(
param_types.into_iter().map(Some).collect(),
);
}
}

Ok(LogicalPlan::Statement(PlanStatement::Prepare(Prepare {
Expand Down
Loading
Loading