1use super::{
19 from_aggregate_rel, from_cast, from_cross_rel, from_exchange_rel, from_fetch_rel,
20 from_field_reference, from_filter_rel, from_if_then, from_join_rel, from_literal,
21 from_nested, from_project_rel, from_read_rel, from_scalar_function, from_set_rel,
22 from_singular_or_list, from_sort_rel, from_subquery, from_substrait_rel,
23 from_substrait_rex, from_window_function,
24};
25use crate::extensions::Extensions;
26use crate::logical_plan::consumer::{
27 field_from_substrait_type_without_names, from_lambda,
28};
29use async_trait::async_trait;
30use datafusion::arrow::datatypes::{DataType, FieldRef};
31use datafusion::catalog::TableProvider;
32use datafusion::common::datatype::FieldExt;
33use datafusion::common::{
34 DFSchema, ScalarValue, TableReference, not_impl_err, substrait_err,
35};
36use datafusion::execution::{FunctionRegistry, SessionState};
37use datafusion::logical_expr::expr::LambdaVariable;
38use datafusion::logical_expr::{Expr, Extension, LogicalPlan};
39use std::collections::VecDeque;
40use std::sync::{Arc, RwLock};
41use substrait::proto::expression as substrait_expression;
42use substrait::proto::expression::{
43 Enum, FieldReference, IfThen, Literal, MultiOrList, Nested, ScalarFunction,
44 SingularOrList, SwitchExpression, WindowFunction,
45};
46use substrait::proto::{self, Type};
47use substrait::proto::{
48 AggregateRel, ConsistentPartitionWindowRel, CrossRel, DynamicParameter, ExchangeRel,
49 Expression, ExtensionLeafRel, ExtensionMultiRel, ExtensionSingleRel, FetchRel,
50 FilterRel, JoinRel, ProjectRel, ReadRel, Rel, SetRel, SortRel, r#type,
51};
52
53#[async_trait]
54pub trait SubstraitConsumer: Send + Sync + Sized {
194 async fn resolve_table_ref(
195 &self,
196 table_ref: &TableReference,
197 ) -> datafusion::common::Result<Option<Arc<dyn TableProvider>>>;
198
199 fn get_extensions(&self) -> &Extensions;
205 fn get_function_registry(&self) -> &impl FunctionRegistry;
206
207 async fn consume_rel(&self, rel: &Rel) -> datafusion::common::Result<LogicalPlan> {
215 from_substrait_rel(self, rel).await
216 }
217
218 async fn consume_read(
219 &self,
220 rel: &ReadRel,
221 ) -> datafusion::common::Result<LogicalPlan> {
222 from_read_rel(self, rel).await
223 }
224
225 async fn consume_filter(
226 &self,
227 rel: &FilterRel,
228 ) -> datafusion::common::Result<LogicalPlan> {
229 from_filter_rel(self, rel).await
230 }
231
232 async fn consume_fetch(
233 &self,
234 rel: &FetchRel,
235 ) -> datafusion::common::Result<LogicalPlan> {
236 from_fetch_rel(self, rel).await
237 }
238
239 async fn consume_aggregate(
240 &self,
241 rel: &AggregateRel,
242 ) -> datafusion::common::Result<LogicalPlan> {
243 from_aggregate_rel(self, rel).await
244 }
245
246 async fn consume_sort(
247 &self,
248 rel: &SortRel,
249 ) -> datafusion::common::Result<LogicalPlan> {
250 from_sort_rel(self, rel).await
251 }
252
253 async fn consume_join(
254 &self,
255 rel: &JoinRel,
256 ) -> datafusion::common::Result<LogicalPlan> {
257 from_join_rel(self, rel).await
258 }
259
260 async fn consume_project(
261 &self,
262 rel: &ProjectRel,
263 ) -> datafusion::common::Result<LogicalPlan> {
264 from_project_rel(self, rel).await
265 }
266
267 async fn consume_set(&self, rel: &SetRel) -> datafusion::common::Result<LogicalPlan> {
268 from_set_rel(self, rel).await
269 }
270
271 async fn consume_cross(
272 &self,
273 rel: &CrossRel,
274 ) -> datafusion::common::Result<LogicalPlan> {
275 from_cross_rel(self, rel).await
276 }
277
278 async fn consume_consistent_partition_window(
279 &self,
280 _rel: &ConsistentPartitionWindowRel,
281 ) -> datafusion::common::Result<LogicalPlan> {
282 not_impl_err!("Consistent Partition Window Rel not supported")
283 }
284
285 async fn consume_exchange(
286 &self,
287 rel: &ExchangeRel,
288 ) -> datafusion::common::Result<LogicalPlan> {
289 from_exchange_rel(self, rel).await
290 }
291
292 async fn consume_expression(
300 &self,
301 expr: &Expression,
302 input_schema: &DFSchema,
303 ) -> datafusion::common::Result<Expr> {
304 from_substrait_rex(self, expr, input_schema).await
305 }
306
307 async fn consume_literal(&self, expr: &Literal) -> datafusion::common::Result<Expr> {
308 from_literal(self, expr).await
309 }
310
311 async fn consume_field_reference(
312 &self,
313 expr: &FieldReference,
314 input_schema: &DFSchema,
315 ) -> datafusion::common::Result<Expr> {
316 from_field_reference(self, expr, input_schema).await
317 }
318
319 async fn consume_scalar_function(
320 &self,
321 expr: &ScalarFunction,
322 input_schema: &DFSchema,
323 ) -> datafusion::common::Result<Expr> {
324 from_scalar_function(self, expr, input_schema).await
325 }
326
327 async fn consume_window_function(
328 &self,
329 expr: &WindowFunction,
330 input_schema: &DFSchema,
331 ) -> datafusion::common::Result<Expr> {
332 from_window_function(self, expr, input_schema).await
333 }
334
335 async fn consume_if_then(
336 &self,
337 expr: &IfThen,
338 input_schema: &DFSchema,
339 ) -> datafusion::common::Result<Expr> {
340 from_if_then(self, expr, input_schema).await
341 }
342
343 async fn consume_switch(
344 &self,
345 _expr: &SwitchExpression,
346 _input_schema: &DFSchema,
347 ) -> datafusion::common::Result<Expr> {
348 not_impl_err!("Switch expression not supported")
349 }
350
351 async fn consume_singular_or_list(
352 &self,
353 expr: &SingularOrList,
354 input_schema: &DFSchema,
355 ) -> datafusion::common::Result<Expr> {
356 from_singular_or_list(self, expr, input_schema).await
357 }
358
359 async fn consume_multi_or_list(
360 &self,
361 _expr: &MultiOrList,
362 _input_schema: &DFSchema,
363 ) -> datafusion::common::Result<Expr> {
364 not_impl_err!("Multi Or List expression not supported")
365 }
366
367 async fn consume_cast(
368 &self,
369 expr: &substrait_expression::Cast,
370 input_schema: &DFSchema,
371 ) -> datafusion::common::Result<Expr> {
372 from_cast(self, expr, input_schema).await
373 }
374
375 async fn consume_subquery(
376 &self,
377 expr: &substrait_expression::Subquery,
378 input_schema: &DFSchema,
379 ) -> datafusion::common::Result<Expr> {
380 from_subquery(self, expr, input_schema).await
381 }
382
383 async fn consume_nested(
384 &self,
385 expr: &Nested,
386 input_schema: &DFSchema,
387 ) -> datafusion::common::Result<Expr> {
388 from_nested(self, expr, input_schema).await
389 }
390
391 async fn consume_enum(
392 &self,
393 _expr: &Enum,
394 _input_schema: &DFSchema,
395 ) -> datafusion::common::Result<Expr> {
396 not_impl_err!("Enum expression not supported")
397 }
398
399 async fn consume_dynamic_parameter(
400 &self,
401 expr: &DynamicParameter,
402 _input_schema: &DFSchema,
403 ) -> datafusion::common::Result<Expr> {
404 let id = format!("${}", expr.parameter_reference + 1);
405 let field = expr
406 .r#type
407 .as_ref()
408 .map(|t| {
409 super::from_substrait_type_without_names(self, t).map(|dt| {
410 Arc::new(datafusion::arrow::datatypes::Field::new(&id, dt, true))
411 })
412 })
413 .transpose()?;
414 Ok(Expr::Placeholder(
415 datafusion::logical_expr::expr::Placeholder::new_with_field(id, field),
416 ))
417 }
418
419 async fn consume_lambda(
420 &self,
421 expr: &proto::expression::Lambda,
422 input_schema: &DFSchema,
423 ) -> datafusion::common::Result<Expr> {
424 from_lambda(self, expr, input_schema).await
425 }
426
427 fn push_outer_schema(&self, _schema: Arc<DFSchema>) {}
434
435 fn pop_outer_schema(&self) {}
437
438 fn get_outer_schema(&self, _steps_out: usize) -> Option<Arc<DFSchema>> {
444 None
445 }
446
447 async fn consume_extension_leaf(
453 &self,
454 rel: &ExtensionLeafRel,
455 ) -> datafusion::common::Result<LogicalPlan> {
456 if let Some(detail) = rel.detail.as_ref() {
457 return substrait_err!(
458 "Missing handler for ExtensionLeafRel: {}",
459 detail.type_url
460 );
461 }
462 substrait_err!("Missing handler for ExtensionLeafRel")
463 }
464
465 async fn consume_extension_single(
466 &self,
467 rel: &ExtensionSingleRel,
468 ) -> datafusion::common::Result<LogicalPlan> {
469 if let Some(detail) = rel.detail.as_ref() {
470 return substrait_err!(
471 "Missing handler for ExtensionSingleRel: {}",
472 detail.type_url
473 );
474 }
475 substrait_err!("Missing handler for ExtensionSingleRel")
476 }
477
478 async fn consume_extension_multi(
479 &self,
480 rel: &ExtensionMultiRel,
481 ) -> datafusion::common::Result<LogicalPlan> {
482 if let Some(detail) = rel.detail.as_ref() {
483 return substrait_err!(
484 "Missing handler for ExtensionMultiRel: {}",
485 detail.type_url
486 );
487 }
488 substrait_err!("Missing handler for ExtensionMultiRel")
489 }
490
491 fn consume_user_defined_type(
494 &self,
495 user_defined_type: &r#type::UserDefined,
496 ) -> datafusion::common::Result<DataType> {
497 substrait_err!(
498 "Missing handler for user-defined type: {}",
499 user_defined_type.type_reference
500 )
501 }
502
503 fn consume_user_defined_literal(
504 &self,
505 user_defined_literal: &proto::expression::literal::UserDefined,
506 ) -> datafusion::common::Result<ScalarValue> {
507 let type_ref = match user_defined_literal.type_anchor_type {
508 Some(
509 proto::expression::literal::user_defined::TypeAnchorType::TypeReference(
510 ref_val,
511 ),
512 ) => ref_val,
513 Some(
514 proto::expression::literal::user_defined::TypeAnchorType::TypeAliasReference(_),
515 ) => {
516 return not_impl_err!(
517 "Type alias references in user-defined literals are not yet supported"
518 )
519 }
520 None => 0,
521 };
522 substrait_err!("Missing handler for user-defined literals {}", type_ref)
523 }
524
525 fn push_lambda_parameters(
532 &self,
533 _lambda_parameters: &[Type],
534 _input_schema: &DFSchema,
535 ) -> datafusion::common::Result<Vec<String>> {
536 not_impl_err!("SubstraitConsumer::push_lambda_parameters")
537 }
538
539 fn pop_lambda_parameters(&self) {}
541
542 fn lambda_variable(
547 &self,
548 _steps_out: usize,
549 _field_idx: usize,
550 ) -> datafusion::common::Result<Expr> {
551 not_impl_err!("SubstraitConsumer::lambda_variable")
552 }
553}
554
555pub struct DefaultSubstraitConsumer<'a> {
559 pub(super) extensions: &'a Extensions,
560 pub(super) state: &'a SessionState,
561 outer_schemas: RwLock<Vec<Arc<DFSchema>>>,
562 lambda_consumer: DefaultSubstraitLambdaConsumer,
563}
564
565impl<'a> DefaultSubstraitConsumer<'a> {
566 pub fn new(extensions: &'a Extensions, state: &'a SessionState) -> Self {
567 DefaultSubstraitConsumer {
568 extensions,
569 state,
570 outer_schemas: RwLock::new(Vec::new()),
571 lambda_consumer: DefaultSubstraitLambdaConsumer::new(),
572 }
573 }
574}
575
576#[async_trait]
577impl SubstraitConsumer for DefaultSubstraitConsumer<'_> {
578 async fn resolve_table_ref(
579 &self,
580 table_ref: &TableReference,
581 ) -> datafusion::common::Result<Option<Arc<dyn TableProvider>>> {
582 let table = table_ref.table().to_string();
583 let schema = self.state.schema_for_ref(table_ref.clone())?;
584 let table_provider = schema.table(&table).await?;
585 Ok(table_provider)
586 }
587
588 fn get_extensions(&self) -> &Extensions {
589 self.extensions
590 }
591
592 fn get_function_registry(&self) -> &impl FunctionRegistry {
593 self.state
594 }
595
596 fn push_outer_schema(&self, schema: Arc<DFSchema>) {
597 self.outer_schemas.write().unwrap().push(schema);
598 }
599
600 fn pop_outer_schema(&self) {
601 self.outer_schemas.write().unwrap().pop();
602 }
603
604 fn get_outer_schema(&self, steps_out: usize) -> Option<Arc<DFSchema>> {
605 let schemas = self.outer_schemas.read().unwrap();
606 schemas
609 .len()
610 .checked_sub(steps_out)
611 .and_then(|idx| schemas.get(idx).cloned())
612 }
613
614 async fn consume_extension_leaf(
615 &self,
616 rel: &ExtensionLeafRel,
617 ) -> datafusion::common::Result<LogicalPlan> {
618 let Some(ext_detail) = &rel.detail else {
619 return substrait_err!("Unexpected empty detail in ExtensionLeafRel");
620 };
621 let plan = self
622 .state
623 .serializer_registry()
624 .deserialize_logical_plan(&ext_detail.type_url, &ext_detail.value)?;
625 Ok(LogicalPlan::Extension(Extension { node: plan }))
626 }
627
628 async fn consume_extension_single(
629 &self,
630 rel: &ExtensionSingleRel,
631 ) -> datafusion::common::Result<LogicalPlan> {
632 let Some(ext_detail) = &rel.detail else {
633 return substrait_err!("Unexpected empty detail in ExtensionSingleRel");
634 };
635 let plan = self
636 .state
637 .serializer_registry()
638 .deserialize_logical_plan(&ext_detail.type_url, &ext_detail.value)?;
639 let Some(input_rel) = &rel.input else {
640 return substrait_err!(
641 "ExtensionSingleRel missing input rel, try using ExtensionLeafRel instead"
642 );
643 };
644 let input_plan = self.consume_rel(input_rel).await?;
645 let plan = plan.with_exprs_and_inputs(plan.expressions(), vec![input_plan])?;
646 Ok(LogicalPlan::Extension(Extension { node: plan }))
647 }
648
649 async fn consume_extension_multi(
650 &self,
651 rel: &ExtensionMultiRel,
652 ) -> datafusion::common::Result<LogicalPlan> {
653 let Some(ext_detail) = &rel.detail else {
654 return substrait_err!("Unexpected empty detail in ExtensionMultiRel");
655 };
656 let plan = self
657 .state
658 .serializer_registry()
659 .deserialize_logical_plan(&ext_detail.type_url, &ext_detail.value)?;
660 let mut inputs = Vec::with_capacity(rel.inputs.len());
661 for input in &rel.inputs {
662 let input_plan = self.consume_rel(input).await?;
663 inputs.push(input_plan);
664 }
665 let plan = plan.with_exprs_and_inputs(plan.expressions(), inputs)?;
666 Ok(LogicalPlan::Extension(Extension { node: plan }))
667 }
668
669 fn push_lambda_parameters(
670 &self,
671 lambda_parameters: &[Type],
672 input_schema: &DFSchema,
673 ) -> datafusion::common::Result<Vec<String>> {
674 self.lambda_consumer
675 .push_lambda_parameters(self, lambda_parameters, input_schema)
676 }
677
678 fn pop_lambda_parameters(&self) {
679 self.lambda_consumer.pop_lambda_parameters()
680 }
681
682 fn lambda_variable(
683 &self,
684 steps_out: usize,
685 field_idx: usize,
686 ) -> datafusion::common::Result<Expr> {
687 self.lambda_consumer.lambda_variable(steps_out, field_idx)
688 }
689}
690
691pub struct DefaultSubstraitLambdaConsumer {
695 inner: RwLock<DefaultSubstraitLambdaConsumerInner>,
696}
697
698struct DefaultSubstraitLambdaConsumerInner {
699 lambda_parameters: VecDeque<Vec<FieldRef>>,
704 next_lambda_parameter: usize,
705}
706
707impl Default for DefaultSubstraitLambdaConsumer {
708 fn default() -> Self {
709 Self::new()
710 }
711}
712
713impl DefaultSubstraitLambdaConsumer {
714 pub fn new() -> Self {
715 Self {
716 inner: RwLock::new(DefaultSubstraitLambdaConsumerInner {
717 lambda_parameters: VecDeque::new(),
718 next_lambda_parameter: 0,
719 }),
720 }
721 }
722
723 pub fn push_lambda_parameters(
724 &self,
725 consumer: &impl SubstraitConsumer,
726 lambda_parameters: &[Type],
727 input_schema: &DFSchema,
728 ) -> datafusion::common::Result<Vec<String>> {
729 let mut inner = self.inner.write().unwrap();
730
731 let lambda_parameters = lambda_parameters
732 .iter()
733 .map(|ty| {
734 let (assigned_number, default_name) =
735 next_lambda_parameter_name(inner.next_lambda_parameter, input_schema);
736
737 inner.next_lambda_parameter = assigned_number + 1;
738
739 Ok(field_from_substrait_type_without_names(consumer, ty)?
740 .renamed(&default_name))
741 })
742 .collect::<datafusion::common::Result<Vec<_>>>()?;
743
744 let names = lambda_parameters.iter().map(|f| f.name().clone()).collect();
745
746 inner.lambda_parameters.push_front(lambda_parameters);
747
748 Ok(names)
749 }
750
751 pub fn pop_lambda_parameters(&self) {
752 self.inner.write().unwrap().lambda_parameters.pop_front();
753 }
754
755 pub fn lambda_variable(
756 &self,
757 steps_out: usize,
758 field_idx: usize,
759 ) -> datafusion::common::Result<Expr> {
760 let lambda_parameters = &self.inner.read().unwrap().lambda_parameters;
761
762 let Some(lambda_parameters) = lambda_parameters.get(steps_out) else {
763 return substrait_err!(
764 "No lambda at {steps_out} steps out, got only {}",
765 lambda_parameters.len()
766 );
767 };
768
769 let Some(var) = lambda_parameters.get(field_idx) else {
770 return substrait_err!(
771 "At lambda {steps_out} steps out, no field at index {field_idx}, got only {}",
772 lambda_parameters.len()
773 );
774 };
775
776 Ok(Expr::LambdaVariable(LambdaVariable::new(
777 var.name().clone(),
778 Some(Arc::clone(var)),
779 )))
780 }
781}
782
783fn next_lambda_parameter_name(
789 mut next_lambda_parameter: usize,
790 input_schema: &DFSchema,
791) -> (usize, String) {
792 loop {
793 let name = format!("p{next_lambda_parameter}");
794
795 if !input_schema.has_column_with_unqualified_name(&name) {
797 return (next_lambda_parameter, name);
798 }
799
800 next_lambda_parameter += 1;
801 }
802}
803
804#[cfg(test)]
805mod tests {
806 use super::*;
807 use crate::logical_plan::consumer::utils::tests::test_consumer;
808 use datafusion::arrow::datatypes::{Field, Schema};
809
810 fn make_schema(fields: &[(&str, DataType)]) -> Arc<DFSchema> {
811 let arrow_fields: Vec<Field> = fields
812 .iter()
813 .map(|(name, dt)| Field::new(*name, dt.clone(), true))
814 .collect();
815 Arc::new(
816 DFSchema::try_from(Schema::new(arrow_fields))
817 .expect("failed to create schema"),
818 )
819 }
820
821 #[test]
822 fn test_get_outer_schema_empty_stack() {
823 let consumer = test_consumer();
824
825 assert!(consumer.get_outer_schema(0).is_none());
827 assert!(consumer.get_outer_schema(1).is_none());
828 assert!(consumer.get_outer_schema(2).is_none());
829 }
830
831 #[test]
832 fn test_get_outer_schema_single_level() {
833 let consumer = test_consumer();
834
835 let schema_a = make_schema(&[("a", DataType::Int64)]);
836 consumer.push_outer_schema(Arc::clone(&schema_a));
837
838 let result = consumer.get_outer_schema(1).unwrap();
840 assert_eq!(result.fields().len(), 1);
841 assert_eq!(result.fields()[0].name(), "a");
842
843 assert!(consumer.get_outer_schema(0).is_none());
845 assert!(consumer.get_outer_schema(2).is_none());
846
847 consumer.pop_outer_schema();
848 assert!(consumer.get_outer_schema(1).is_none());
849 }
850
851 #[test]
852 fn test_get_outer_schema_nested() {
853 let consumer = test_consumer();
854
855 let schema_a = make_schema(&[("a", DataType::Int64)]);
856 let schema_b = make_schema(&[("b", DataType::Utf8)]);
857
858 consumer.push_outer_schema(Arc::clone(&schema_a));
859 consumer.push_outer_schema(Arc::clone(&schema_b));
860
861 let result = consumer.get_outer_schema(1).unwrap();
863 assert_eq!(result.fields()[0].name(), "b");
864
865 let result = consumer.get_outer_schema(2).unwrap();
867 assert_eq!(result.fields()[0].name(), "a");
868
869 assert!(consumer.get_outer_schema(3).is_none());
871
872 consumer.pop_outer_schema();
874 let result = consumer.get_outer_schema(1).unwrap();
875 assert_eq!(result.fields()[0].name(), "a");
876 assert!(consumer.get_outer_schema(2).is_none());
877 }
878}