Skip to content

Commit 92b9a70

Browse files
committed
ai draft
1 parent 9230808 commit 92b9a70

10 files changed

Lines changed: 296 additions & 67 deletions

File tree

kernel/src/actions/visitors.rs

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1786,7 +1786,7 @@ mod tests {
17861786
engine
17871787
.evaluation_handler()
17881788
.new_expression_evaluator(
1789-
get_commit_schema().clone(),
1789+
get_all_actions_schema().clone(),
17901790
expression.into(),
17911791
InCommitTimestampVisitor::schema().into(),
17921792
)

kernel/src/engine/arrow_expression/mod.rs

Lines changed: 93 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,7 @@ use evaluate_expression::{evaluate_expression, evaluate_predicate};
66
use itertools::Itertools;
77
use tracing::debug;
88

9-
use super::arrow_conversion::{TryFromKernel as _, TryIntoArrow as _};
9+
use super::arrow_conversion::{TryFromArrow as _, TryFromKernel as _, TryIntoArrow as _};
1010
use crate::arrow::array::{self, ArrayBuilder, ArrayRef, RecordBatch, StructArray};
1111
use crate::arrow::datatypes::{
1212
DataType as ArrowDataType, Field as ArrowField, Schema as ArrowSchema,
@@ -15,7 +15,7 @@ use crate::engine::arrow_data::{extract_record_batch, ArrowEngineData};
1515
use crate::engine::arrow_utils::apply_schema::{apply_schema, apply_schema_to};
1616
use crate::error::{DeltaResult, Error};
1717
use crate::expressions::{ArrayData, Expression, ExpressionRef, PredicateRef, Scalar};
18-
use crate::schema::{DataType, PrimitiveType, SchemaRef};
18+
use crate::schema::{DataType, PrimitiveType, SchemaRef, StructType};
1919
use crate::utils::require;
2020
use crate::{EngineData, EvaluationHandler, ExpressionEvaluator, PredicateEvaluator};
2121

@@ -254,7 +254,7 @@ impl EvaluationHandler for ArrowEvaluationHandler {
254254
output_type: DataType,
255255
) -> DeltaResult<Arc<dyn ExpressionEvaluator>> {
256256
Ok(Arc::new(DefaultExpressionEvaluator {
257-
_input_schema: schema,
257+
input_schema: schema,
258258
expression,
259259
output_type,
260260
}))
@@ -266,7 +266,7 @@ impl EvaluationHandler for ArrowEvaluationHandler {
266266
predicate: PredicateRef,
267267
) -> DeltaResult<Arc<dyn PredicateEvaluator>> {
268268
Ok(Arc::new(DefaultPredicateEvaluator {
269-
_input_schema: schema,
269+
input_schema: schema,
270270
predicate,
271271
}))
272272
}
@@ -341,7 +341,7 @@ impl EvaluationHandler for ArrowEvaluationHandler {
341341

342342
#[derive(Debug)]
343343
pub struct DefaultExpressionEvaluator {
344-
_input_schema: SchemaRef,
344+
input_schema: SchemaRef,
345345
expression: ExpressionRef,
346346
output_type: DataType,
347347
}
@@ -350,14 +350,7 @@ impl ExpressionEvaluator for DefaultExpressionEvaluator {
350350
fn evaluate(&self, batch: &dyn EngineData) -> DeltaResult<Box<dyn EngineData>> {
351351
debug!("Arrow evaluator evaluating: {:#?}", self.expression);
352352
let batch = extract_record_batch(batch)?;
353-
// TODO: make sure we have matching schemas for validation
354-
// if batch.schema().as_ref() != &input_schema {
355-
// return Err(Error::Generic(format!(
356-
// "input schema does not match batch schema: {:?} != {:?}",
357-
// input_schema,
358-
// batch.schema()
359-
// )));
360-
// };
353+
validate_input_schema(&self.input_schema, batch.schema().as_ref())?;
361354
let batch = match (self.expression.as_ref(), &self.output_type) {
362355
(Expression::StructPatch(patch), DataType::Struct(_)) if patch.is_empty() => {
363356
// Empty patch optimization: Skip expression evaluation and directly apply the
@@ -388,22 +381,15 @@ impl ExpressionEvaluator for DefaultExpressionEvaluator {
388381

389382
#[derive(Debug)]
390383
pub struct DefaultPredicateEvaluator {
391-
_input_schema: SchemaRef,
384+
input_schema: SchemaRef,
392385
predicate: PredicateRef,
393386
}
394387

395388
impl PredicateEvaluator for DefaultPredicateEvaluator {
396389
fn evaluate(&self, batch: &dyn EngineData) -> DeltaResult<Box<dyn EngineData>> {
397390
debug!("Arrow evaluator evaluating: {:#?}", self.predicate);
398391
let batch = extract_record_batch(batch)?;
399-
// TODO: make sure we have matching schemas for validation
400-
// if batch.schema().as_ref() != &input_schema {
401-
// return Err(Error::Generic(format!(
402-
// "input schema does not match batch schema: {:?} != {:?}",
403-
// input_schema,
404-
// batch.schema()
405-
// )));
406-
// };
392+
validate_input_schema(&self.input_schema, batch.schema().as_ref())?;
407393
let array = evaluate_predicate(&self.predicate, batch, false)?;
408394
let schema = ArrowSchema::new(vec![ArrowField::new(
409395
"output",
@@ -414,3 +400,88 @@ impl PredicateEvaluator for DefaultPredicateEvaluator {
414400
Ok(Box::new(ArrowEngineData::new(batch)))
415401
}
416402
}
403+
404+
fn validate_input_schema(input_schema: &SchemaRef, batch_schema: &ArrowSchema) -> DeltaResult<()> {
405+
let batch_schema = StructType::try_from_arrow(batch_schema)?;
406+
require!(
407+
input_schema.num_fields() == batch_schema.num_fields(),
408+
Error::generic(format!(
409+
"Input schema fields {:?} do not match batch schema fields {:?}",
410+
input_schema
411+
.fields()
412+
.map(|field| field.name())
413+
.collect_vec(),
414+
batch_schema
415+
.fields()
416+
.map(|field| field.name())
417+
.collect_vec()
418+
))
419+
);
420+
421+
for (input_field, batch_field) in input_schema.fields().zip(batch_schema.fields()) {
422+
require!(
423+
input_field.name() == batch_field.name(),
424+
Error::generic(format!(
425+
"Input schema field '{}' does not match batch schema field '{}'",
426+
input_field.name(),
427+
batch_field.name()
428+
))
429+
);
430+
require!(
431+
input_types_compatible(
432+
input_field.data_type(),
433+
batch_field.data_type(),
434+
input_field.is_nullable() || batch_field.is_nullable(),
435+
),
436+
Error::generic(format!(
437+
"Input schema type for '{}' does not match the batch schema type: {:?} != {:?}",
438+
input_field.name(),
439+
input_field.data_type(),
440+
batch_field.data_type()
441+
))
442+
);
443+
}
444+
Ok(())
445+
}
446+
447+
fn input_types_compatible(
448+
input_type: &DataType,
449+
batch_type: &DataType,
450+
allow_omitted_struct_fields: bool,
451+
) -> bool {
452+
match (input_type, batch_type) {
453+
(DataType::Primitive(input), DataType::Primitive(batch)) => input == batch,
454+
(DataType::Struct(input), DataType::Struct(batch)) => {
455+
if !allow_omitted_struct_fields && input.num_fields() != batch.num_fields() {
456+
return false;
457+
}
458+
459+
input.fields().all(|input_field| {
460+
batch.field(input_field.name()).is_none_or(|batch_field| {
461+
input_types_compatible(
462+
input_field.data_type(),
463+
batch_field.data_type(),
464+
input_field.is_nullable() || batch_field.is_nullable(),
465+
)
466+
})
467+
}) && batch.fields().all(|batch_field| {
468+
input.field(batch_field.name()).is_some() || allow_omitted_struct_fields
469+
})
470+
}
471+
(DataType::Array(input), DataType::Array(batch)) => input_types_compatible(
472+
input.element_type(),
473+
batch.element_type(),
474+
input.contains_null() || batch.contains_null(),
475+
),
476+
(DataType::Map(input), DataType::Map(batch)) => {
477+
input_types_compatible(input.key_type(), batch.key_type(), false)
478+
&& input_types_compatible(
479+
input.value_type(),
480+
batch.value_type(),
481+
input.value_contains_null() || batch.value_contains_null(),
482+
)
483+
}
484+
(DataType::Variant(input), DataType::Variant(batch)) => input == batch,
485+
_ => false,
486+
}
487+
}

kernel/src/engine/arrow_expression/tests.rs

Lines changed: 122 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1291,6 +1291,128 @@ fn test_evaluator_mixed_string_types_struct_expression() {
12911291
.unwrap();
12921292
}
12931293

1294+
#[derive(Clone, Copy, Debug)]
1295+
enum EvaluatorKind {
1296+
Expression,
1297+
Predicate,
1298+
}
1299+
1300+
#[derive(Clone, Copy, Debug)]
1301+
enum TopLevelSchemaMismatch {
1302+
ExtraField,
1303+
MissingField,
1304+
ReorderedFields,
1305+
RenamedField,
1306+
WrongType,
1307+
}
1308+
1309+
#[rstest]
1310+
#[case::extra_field(
1311+
TopLevelSchemaMismatch::ExtraField,
1312+
"Input schema fields [\"a\", \"b\"] do not match batch schema fields [\"a\", \"b\", \"c\"]"
1313+
)]
1314+
#[case::missing_field(
1315+
TopLevelSchemaMismatch::MissingField,
1316+
"Input schema fields [\"a\", \"b\"] do not match batch schema fields [\"a\"]"
1317+
)]
1318+
#[case::reordered_fields(
1319+
TopLevelSchemaMismatch::ReorderedFields,
1320+
"Input schema field 'a' does not match batch schema field 'b'"
1321+
)]
1322+
#[case::renamed_field(
1323+
TopLevelSchemaMismatch::RenamedField,
1324+
"Input schema field 'b' does not match batch schema field 'c'"
1325+
)]
1326+
#[case::wrong_type(
1327+
TopLevelSchemaMismatch::WrongType,
1328+
"Input schema type for 'a' does not match the batch schema type"
1329+
)]
1330+
fn evaluator_rejects_mismatched_top_level_schema(
1331+
#[values(EvaluatorKind::Expression, EvaluatorKind::Predicate)] evaluator_kind: EvaluatorKind,
1332+
#[case] mismatch: TopLevelSchemaMismatch,
1333+
#[case] expected_error: &str,
1334+
) {
1335+
let input_schema = schema_ref! {
1336+
nullable "a": INTEGER,
1337+
nullable "b": STRING,
1338+
};
1339+
let batch_schema = match mismatch {
1340+
TopLevelSchemaMismatch::ExtraField => Schema::new(vec![
1341+
Field::new("a", DataType::Int32, true),
1342+
Field::new("b", DataType::Utf8, true),
1343+
Field::new("c", DataType::Boolean, true),
1344+
]),
1345+
TopLevelSchemaMismatch::MissingField => {
1346+
Schema::new(vec![Field::new("a", DataType::Int32, true)])
1347+
}
1348+
TopLevelSchemaMismatch::ReorderedFields => Schema::new(vec![
1349+
Field::new("b", DataType::Utf8, true),
1350+
Field::new("a", DataType::Int32, true),
1351+
]),
1352+
TopLevelSchemaMismatch::RenamedField => Schema::new(vec![
1353+
Field::new("a", DataType::Int32, true),
1354+
Field::new("c", DataType::Utf8, true),
1355+
]),
1356+
TopLevelSchemaMismatch::WrongType => Schema::new(vec![
1357+
Field::new("a", DataType::Int64, true),
1358+
Field::new("b", DataType::Utf8, true),
1359+
]),
1360+
};
1361+
let batch = ArrowEngineData::new(RecordBatch::new_empty(Arc::new(batch_schema)));
1362+
let handler = ArrowEvaluationHandler;
1363+
let result = match evaluator_kind {
1364+
EvaluatorKind::Expression => handler
1365+
.new_expression_evaluator(input_schema, Arc::new(col!("a")), KernelDataType::INTEGER)
1366+
.unwrap()
1367+
.evaluate(&batch),
1368+
EvaluatorKind::Predicate => handler
1369+
.new_predicate_evaluator(input_schema, Arc::new(Predicate::TRUE))
1370+
.unwrap()
1371+
.evaluate(&batch),
1372+
};
1373+
1374+
assert_result_error_with_message(result, expected_error);
1375+
}
1376+
1377+
#[rstest]
1378+
fn evaluator_accepts_omitted_fields_in_nullable_struct(
1379+
#[values(false, true)] batch_is_sparse: bool,
1380+
) {
1381+
let sparse_struct = schema! {};
1382+
let rich_struct = schema! {
1383+
not_null "a": INTEGER,
1384+
};
1385+
let (input_struct, batch_struct) = if batch_is_sparse {
1386+
(rich_struct, sparse_struct)
1387+
} else {
1388+
(sparse_struct, rich_struct)
1389+
};
1390+
let input_schema = schema_ref! {
1391+
nullable "s": (input_struct),
1392+
};
1393+
let batch_schema = schema_ref! {
1394+
not_null "s": (batch_struct),
1395+
};
1396+
let batch_schema: Schema = batch_schema.as_ref().try_into_arrow().unwrap();
1397+
1398+
validate_input_schema(&input_schema, &batch_schema).unwrap();
1399+
}
1400+
1401+
#[test]
1402+
fn evaluator_rejects_conflicting_field_in_nullable_struct() {
1403+
let input_schema = schema_ref! {
1404+
nullable "s": { not_null "a": INTEGER },
1405+
};
1406+
let batch_schema = schema_ref! {
1407+
nullable "s": { nullable "a": STRING },
1408+
};
1409+
let batch_schema: Schema = batch_schema.as_ref().try_into_arrow().unwrap();
1410+
1411+
let result = validate_input_schema(&input_schema, &batch_schema);
1412+
1413+
assert_result_error_with_message(result, "Input schema type for 's' does not match");
1414+
}
1415+
12941416
// helper to build a RecordBatch via `create_many` and assert it equals `expected`
12951417
fn assert_create_many(rows: &[&[Scalar]], schema: SchemaRef, expected: RecordBatch) {
12961418
let handler = ArrowEvaluationHandler;

0 commit comments

Comments
 (0)