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
524 changes: 446 additions & 78 deletions datafusion/substrait/src/logical_plan/consumer/expr/literal.rs

Large diffs are not rendered by default.

179 changes: 143 additions & 36 deletions datafusion/substrait/src/logical_plan/consumer/rel/read_rel.rs
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@
// under the License.

use crate::logical_plan::consumer::SubstraitConsumer;
use crate::logical_plan::consumer::from_substrait_literal;
use crate::logical_plan::consumer::from_substrait_literal_with_expected_field;
use crate::logical_plan::consumer::from_substrait_named_struct;
use crate::logical_plan::consumer::utils::ensure_schema_compatibility;
use datafusion::common::{
Expand Down Expand Up @@ -157,21 +157,16 @@ pub async fn from_read_rel(
}

let mut row_exprs = vec![];
let mut name_idx = 0;
for expression in &row.fields {
// Top-level names are provided through schema
// Each expression consumes at least one name, and Literals may consume additional names.
name_idx += 1;
for (expression, expected_field) in
row.fields.iter().zip(substrait_schema.fields())
{
let expr = match expression.rex_type.as_ref() {
Some(substrait::proto::expression::RexType::Literal(lit)) => {
// Values literals need 'named_struct.names' so nested struct fields keep their names from the ReadRel base schema.
// This is important for nested struct fields to retain their names.
Expr::Literal(
from_substrait_literal(
from_substrait_literal_with_expected_field(
consumer,
lit,
&named_struct.names,
&mut name_idx,
expected_field,
)?,
None,
)
Expand All @@ -184,19 +179,11 @@ pub async fn from_read_rel(
};
row_exprs.push(expr);
}

if name_idx != named_struct.names.len() {
return substrait_err!(
"Names list must match exactly to nested schema, but found {} uses for {} names",
name_idx,
named_struct.names.len()
);
}
exprs.push(row_exprs);
}
exprs
} else {
convert_literal_rows(consumer, vt, named_struct)?
convert_literal_rows(consumer, vt, &substrait_schema)?
};

Ok(LogicalPlan::Values(Values {
Expand Down Expand Up @@ -254,38 +241,35 @@ pub async fn from_read_rel(
/// Converts Substrait literal rows from a VirtualTable into DataFusion expressions.
///
/// This function processes the deprecated `values` field of VirtualTable, converting
/// each literal value into a `Expr::Literal` while tracking and validating the name
/// indices against the provided named struct schema.
/// each literal value into a `Expr::Literal` using the matching field from the schema.
fn convert_literal_rows(
consumer: &impl SubstraitConsumer,
vt: &substrait::proto::read_rel::VirtualTable,
named_struct: &substrait::proto::NamedStruct,
schema: &DFSchema,
) -> datafusion::common::Result<Vec<Vec<Expr>>> {
#[expect(deprecated)]
vt.values
.iter()
.map(|row| {
let mut name_idx = 0;
if row.fields.len() != schema.fields().len() {
return substrait_err!(
"Field count mismatch: expected {} fields but found {} in virtual table row",
schema.fields().len(),
row.fields.len()
);
}
let lits = row
.fields
.iter()
.map(|lit| {
name_idx += 1; // top-level names are provided through schema
Ok(Expr::Literal(from_substrait_literal(
.zip(schema.fields())
.map(|(lit, expected_field)| {
Ok(Expr::Literal(from_substrait_literal_with_expected_field(
consumer,
lit,
&named_struct.names,
&mut name_idx,
expected_field,
)?, None))
})
.collect::<datafusion::common::Result<_>>()?;
if name_idx != named_struct.names.len() {
return substrait_err!(
"Names list must match exactly to nested schema, but found {} uses for {} names",
name_idx,
named_struct.names.len()
);
}
Ok(lits)
})
.collect::<datafusion::common::Result<_>>()
Expand Down Expand Up @@ -365,3 +349,126 @@ fn apply_projection(
_ => plan_err!("DataFrame passed to apply_projection must be a TableScan"),
}
}

#[cfg(test)]
mod tests {
use super::*;
use crate::logical_plan::consumer::utils::tests::test_consumer;
use crate::logical_plan::producer::{
DefaultSubstraitProducer, to_substrait_named_struct,
};
use datafusion::arrow::array::Array;
use datafusion::arrow::datatypes::{DataType, Field, Schema};
use datafusion::common::ScalarValue;
use datafusion::prelude::SessionContext;
use substrait::proto::expression::Literal;
use substrait::proto::expression::RexType;
use substrait::proto::expression::literal::{
List, LiteralType, Struct as LiteralStruct,
};
use substrait::proto::expression::nested::Struct as ExpressionStruct;
use substrait::proto::read_rel::VirtualTable;

fn list_literal() -> Literal {
Literal {
nullable: false,
type_variation_reference: 0,
literal_type: Some(LiteralType::List(List {
values: vec![Literal {
nullable: false,
type_variation_reference: 0,
literal_type: Some(LiteralType::I32(1)),
}],
})),
}
}

fn list_schema() -> datafusion::common::Result<DFSchema> {
let list_type =
DataType::List(Arc::new(Field::new_list_field(DataType::Int32, false)));
DFSchema::try_from(Schema::new(vec![Field::new("list", list_type, false)]))
}

#[test]
fn deprecated_literal_rows_use_expected_fields() -> datafusion::common::Result<()> {
let schema = list_schema()?;
let list_type = schema.field(0).data_type();
#[expect(deprecated)]
let virtual_table = VirtualTable {
values: vec![LiteralStruct {
fields: vec![list_literal()],
}],
..Default::default()
};

let rows = convert_literal_rows(&test_consumer(), &virtual_table, &schema)?;
let Expr::Literal(ScalarValue::List(list), _) = &rows[0][0] else {
panic!("expected list literal")
};
assert_eq!(list.data_type(), list_type);

#[expect(deprecated)]
let mismatched_table = VirtualTable {
values: vec![LiteralStruct { fields: vec![] }],
..Default::default()
};
let err = convert_literal_rows(&test_consumer(), &mismatched_table, &schema)
.unwrap_err();
assert!(
err.to_string().contains("Field count mismatch"),
"got: {err}"
);

Ok(())
}

#[tokio::test]
async fn expression_rows_use_expected_fields() -> datafusion::common::Result<()> {
let schema = list_schema()?;
let state = SessionContext::new().state();
let mut producer = DefaultSubstraitProducer::new(&state);
let base_schema =
to_substrait_named_struct(&mut producer, &DFSchemaRef::new(schema.clone()))?;
let expression = Expression {
rex_type: Some(RexType::Literal(list_literal())),
};
let virtual_table = VirtualTable {
expressions: vec![ExpressionStruct {
fields: vec![expression],
}],
..Default::default()
};
let read = ReadRel {
base_schema: Some(base_schema.clone()),
read_type: Some(ReadType::VirtualTable(virtual_table)),
..Default::default()
};

let plan = from_read_rel(&test_consumer(), &read).await?;
let LogicalPlan::Values(values) = plan else {
panic!("expected Values plan")
};
let Expr::Literal(ScalarValue::List(list), _) = &values.values[0][0] else {
panic!("expected list literal")
};
assert_eq!(list.data_type(), schema.field(0).data_type());

let mismatched_read = ReadRel {
base_schema: Some(base_schema),
read_type: Some(ReadType::VirtualTable(VirtualTable {
expressions: vec![ExpressionStruct { fields: vec![] }],
..Default::default()
})),
..Default::default()
};
let err = from_read_rel(&test_consumer(), &mismatched_read)
.await
.unwrap_err();
assert!(
err.to_string().contains("Field count mismatch"),
"got: {err}"
);

Ok(())
}
}
8 changes: 4 additions & 4 deletions datafusion/substrait/src/logical_plan/consumer/types.rs
Original file line number Diff line number Diff line change
Expand Up @@ -179,9 +179,7 @@ pub fn from_substrait_type(
})?;
let field = Arc::new(Field::new_list_field(
from_substrait_type(consumer, inner_type, dfs_names, name_idx)?,
// We ignore Substrait's nullability here to match to_substrait_literal
// which always creates nullable lists
true,
type_is_nullable(inner_type)?,
));
match list.type_variation_reference {
DEFAULT_CONTAINER_TYPE_VARIATION_REF => Ok(DataType::List(field)),
Expand All @@ -198,6 +196,7 @@ pub fn from_substrait_type(
let value_type = map.value.as_ref().ok_or_else(|| {
substrait_datafusion_err!("Map type must have value type")
})?;
let value_nullable = type_is_nullable(value_type)?;
let key_type =
from_substrait_type(consumer, key_type, dfs_names, name_idx)?;
let value_type =
Expand All @@ -206,7 +205,8 @@ pub fn from_substrait_type(
match map.type_variation_reference {
DEFAULT_MAP_TYPE_VARIATION_REF => {
let key_field = Arc::new(Field::new("key", key_type, false));
let value_field = Arc::new(Field::new("value", value_type, true));
let value_field =
Arc::new(Field::new("value", value_type, value_nullable));
Ok(DataType::Map(
Arc::new(Field::new_struct(
"entries",
Expand Down
96 changes: 94 additions & 2 deletions datafusion/substrait/src/logical_plan/producer/expr/literal.rs
Original file line number Diff line number Diff line change
Expand Up @@ -406,10 +406,15 @@ fn convert_array_to_literal_list<T: OffsetSizeTrait>(
#[cfg(test)]
mod tests {
use super::*;
use crate::logical_plan::consumer::from_substrait_literal_without_names;
use crate::logical_plan::consumer::tests::test_consumer;
use crate::logical_plan::consumer::{
from_substrait_literal_with_expected_field, from_substrait_literal_without_names,
};
use crate::logical_plan::producer::DefaultSubstraitProducer;
use datafusion::arrow::array::{Int64Builder, MapBuilder, StringBuilder};
use datafusion::arrow::array::{
AsArray, Int64Builder, MapArray, MapBuilder, StringBuilder, new_empty_array,
};
use datafusion::arrow::buffer::OffsetBuffer;
use datafusion::arrow::datatypes::{
DataType, Field, IntervalDayTime, IntervalMonthDayNano,
};
Expand Down Expand Up @@ -548,6 +553,79 @@ mod tests {
Ok(())
}

#[test]
fn round_trip_literals_with_expected_field() -> Result<()> {
let required_struct = ScalarStructBuilder::new()
.with_scalar(
Field::new("required", DataType::Int32, false),
ScalarValue::Int32(Some(1)),
)
.build()?;
let state = SessionContext::default().state();
let mut producer = DefaultSubstraitProducer::new(&state);
let substrait_literal = to_substrait_literal(&mut producer, &required_struct)?;
let schema_less =
from_substrait_literal_without_names(&test_consumer(), &substrait_literal)?;
let DataType::Struct(fields) = schema_less.data_type() else {
panic!("expected struct literal")
};
assert!(fields[0].is_nullable());

round_trip_literal_with_expected_field(required_struct.clone())?;

let required_list = ScalarValue::List(ScalarValue::new_list(
std::slice::from_ref(&required_struct),
&required_struct.data_type(),
false,
));
round_trip_literal_with_expected_field(required_list)?;

let empty_required_list = ScalarValue::List(ScalarValue::new_list(
&[],
&required_struct.data_type(),
false,
));
round_trip_literal_with_expected_field(empty_required_list)?;

let required_entry = ScalarStructBuilder::new()
.with_scalar(
Field::new("key", DataType::Utf8, false),
ScalarValue::Utf8(Some("key".to_string())),
)
.with_scalar(
Field::new("value", DataType::Int64, false),
ScalarValue::Int64(Some(1)),
)
.build()?;
let entries = ScalarValue::iter_to_array([required_entry])?
.as_struct()
.to_owned();
let entries_field =
Arc::new(Field::new("entries", entries.data_type().clone(), false));
let required_map = ScalarValue::Map(Arc::new(MapArray::new(
Arc::clone(&entries_field),
OffsetBuffer::new(vec![0, 1].into()),
entries,
None,
false,
)));
round_trip_literal_with_expected_field(required_map)?;

let empty_entries = new_empty_array(entries_field.data_type())
.as_struct()
.to_owned();
let empty_required_map = ScalarValue::Map(Arc::new(MapArray::new(
entries_field,
OffsetBuffer::new(vec![0, 0].into()),
empty_entries,
None,
false,
)));
round_trip_literal_with_expected_field(empty_required_map)?;

Ok(())
}

fn round_trip_literal(scalar: ScalarValue) -> Result<()> {
println!("Checking round trip of {scalar:?}");
let state = SessionContext::default().state();
Expand All @@ -558,4 +636,18 @@ mod tests {
assert_eq!(scalar, roundtrip_scalar);
Ok(())
}

fn round_trip_literal_with_expected_field(scalar: ScalarValue) -> Result<()> {
let state = SessionContext::default().state();
let mut producer = DefaultSubstraitProducer::new(&state);
let substrait_literal = to_substrait_literal(&mut producer, &scalar)?;
let expected_field = Field::new("expected", scalar.data_type(), false);
let roundtrip_scalar = from_substrait_literal_with_expected_field(
&test_consumer(),
&substrait_literal,
&expected_field,
)?;
assert_eq!(scalar, roundtrip_scalar);
Ok(())
}
}
Loading
Loading