1use ahash::AHashSet;
32use std::sync::{Arc, OnceLock};
33
34use radixdb_core::{CompactArc, CompactVec, StringMap};
35use radixdb_core::{DataType, Error, Result, Row, RowVec, Value};
36use radixdb_functions::FunctionRegistry;
37use radixdb_sql::ast::*;
38use radixdb_sql::token::{Position, Token, TokenType};
39use radixdb_storage::traits::QueryResult;
40
41use super::aggregation::{AggregationExecutorExt, AggregationHost};
42use super::context::ExecutionContext;
43use super::expression::{compile_expression_with_context, ExpressionEval};
44use super::pipeline::paging::evaluate_page_expression;
45use super::pipeline::set::merge_set_type;
46use super::query_classification::{get_classification, QueryClassification};
47use super::subquery::{SubqueryExecutorExt, SubqueryHost};
48use super::utils::build_column_index_map;
49use super::utils::RetainedRowsBudget;
50use super::window::{WindowExecutorExt, WindowHost};
51
52pub type CteData = (
55 CompactArc<Vec<String>>,
56 CompactArc<Vec<(i64, Row)>>,
57 Arc<OnceLock<CompactArc<Vec<Row>>>>,
58);
59
60pub type CteDataMap = StringMap<CteData>;
64
65#[derive(Clone)]
71pub struct CteRegistry {
72 data: Arc<CteDataMap>,
75}
76
77impl Default for CteRegistry {
78 fn default() -> Self {
79 Self::new()
80 }
81}
82
83impl CteRegistry {
84 pub fn new() -> Self {
86 Self {
87 data: Arc::new(StringMap::new()),
88 }
89 }
90
91 pub fn store(&mut self, name: &str, columns: Vec<String>, rows: RowVec) {
100 let name_lower = name.to_lowercase();
101 let rows_vec: Vec<(i64, Row)> = rows.into_iter().collect();
103 Arc::make_mut(&mut self.data).insert(
104 name_lower,
105 (
106 CompactArc::new(columns),
107 CompactArc::new(rows_vec),
108 Arc::new(OnceLock::new()),
109 ),
110 );
111 }
112
113 pub fn store_arc(
118 &mut self,
119 name: &str,
120 columns: CompactArc<Vec<String>>,
121 rows: CompactArc<Vec<(i64, Row)>>,
122 materialized_rows: Arc<OnceLock<CompactArc<Vec<Row>>>>,
123 ) {
124 let name_lower = name.to_lowercase();
125 Arc::make_mut(&mut self.data).insert(name_lower, (columns, rows, materialized_rows));
126 }
127
128 pub fn get(&self, name: &str) -> Option<&CteData> {
130 self.data.get(&name.to_lowercase())
131 }
132
133 pub fn data(&self) -> Arc<CteDataMap> {
138 self.data.clone()
139 }
140
141 pub fn iter(&self) -> impl Iterator<Item = (&String, &CteData)> {
143 self.data.iter()
144 }
145}
146
147pub trait CteHost: AggregationHost + SubqueryHost + WindowHost {
151 fn cte_function_registry(&self) -> &FunctionRegistry;
152 fn cte_execute_select(
153 &self,
154 statement: &SelectStatement,
155 context: &ExecutionContext,
156 ) -> Result<Box<dyn QueryResult>>;
157}
158
159pub struct CteExecutor<'a, H: CteHost + ?Sized> {
162 host: &'a H,
163}
164
165impl<'a, H: CteHost + ?Sized> CteExecutor<'a, H> {
166 fn new(host: &'a H) -> Self {
167 Self { host }
168 }
169}
170
171pub trait CteExecutorExt: CteHost {
173 fn execute_select_with_ctes(
174 &self,
175 statement: &SelectStatement,
176 context: &ExecutionContext,
177 ) -> Result<Box<dyn QueryResult>> {
178 CteExecutor::new(self).execute_select_with_ctes(statement, context)
179 }
180
181 fn execute_query_on_cte_result(
182 &self,
183 statement: &SelectStatement,
184 context: &ExecutionContext,
185 columns: Vec<String>,
186 rows: RowVec,
187 ) -> Result<(Vec<String>, RowVec)> {
188 CteExecutor::new(self).execute_query_on_cte_result(statement, context, columns, rows)
189 }
190
191 fn execute_query_on_cte_result_inner(
192 &self,
193 statement: &SelectStatement,
194 context: &ExecutionContext,
195 columns: Vec<String>,
196 rows: RowVec,
197 skip_order_limit: bool,
198 ) -> Result<(Vec<String>, RowVec, bool)> {
199 CteExecutor::new(self).execute_query_on_cte_result_inner(
200 statement,
201 context,
202 columns,
203 rows,
204 skip_order_limit,
205 )
206 }
207
208 fn has_cte(&self, statement: &SelectStatement) -> bool {
209 CteExecutor::new(self).has_cte(statement)
210 }
211
212 fn try_inline_ctes(
213 &self,
214 statement: &SelectStatement,
215 with_clause: &WithClause,
216 ) -> Option<SelectStatement> {
217 CteExecutor::new(self).try_inline_ctes(statement, with_clause)
218 }
219}
220
221impl<T: CteHost + ?Sized> CteExecutorExt for T {}
222
223fn materialize_result(mut result: Box<dyn QueryResult>) -> Result<RowVec> {
224 let mut rows = result
225 .estimated_count()
226 .map_or_else(RowVec::new, RowVec::with_capacity);
227 let mut row_id = 0i64;
228 while result.next() {
229 rows.push((row_id, result.take_row()));
230 row_id += 1;
231 }
232 if let Some(error) = result.last_error() {
233 return Err(error);
234 }
235 Ok(rows)
236}
237
238impl<H: CteHost + ?Sized> CteExecutor<'_, H> {
239 pub(crate) fn execute_select_with_ctes(
241 &self,
242 stmt: &SelectStatement,
243 ctx: &ExecutionContext,
244 ) -> Result<Box<dyn QueryResult>> {
245 let with_clause = match &stmt.with {
247 Some(with) => with,
248 None => return self.host.cte_execute_select(stmt, ctx),
249 };
250
251 if let Some(inlined_stmt) = self.try_inline_ctes(stmt, with_clause) {
257 return self.host.cte_execute_select(&inlined_stmt, ctx);
259 }
260
261 let mut cte_registry = CteRegistry::new();
263
264 for cte in &with_clause.ctes {
266 let (columns, rows) = if cte.is_recursive {
268 let aliases = if cte.column_names.is_empty() {
270 None
271 } else {
272 Some(cte.column_names.as_slice())
273 };
274 self.execute_recursive_cte_with_columns(
275 &cte.name.value,
276 &cte.query,
277 ctx,
278 &mut cte_registry,
279 aliases,
280 )?
281 } else {
282 self.execute_cte_query(&cte.query, ctx, &mut cte_registry)?
283 };
284
285 let columns = if !cte.column_names.is_empty() {
287 cte.column_names
288 .iter()
289 .enumerate()
290 .map(|(i, alias)| {
291 if i < columns.len() {
292 alias.value.to_string()
293 } else {
294 columns
295 .get(i)
296 .cloned()
297 .unwrap_or_else(|| format!("col{}", i))
298 }
299 })
300 .collect()
301 } else {
302 columns
303 };
304
305 cte_registry.store(&cte.name.value, columns, rows);
307 }
308
309 self.execute_main_query_with_ctes(stmt, ctx, &mut cte_registry)
311 }
312
313 fn execute_cte_query(
315 &self,
316 stmt: &SelectStatement,
317 ctx: &ExecutionContext,
318 cte_registry: &mut CteRegistry,
319 ) -> Result<(Vec<String>, RowVec)> {
320 let ctx_with_ctes = ctx.with_cte_data(cte_registry.data());
321 let mut statement = stmt.clone();
322 statement.with = None;
323 let result = self.host.cte_execute_select(&statement, &ctx_with_ctes)?;
324 let columns = result.columns().to_vec();
325 let rows = materialize_result(result)?;
326
327 Ok((columns, rows))
328 }
329
330 fn execute_recursive_cte_with_columns(
332 &self,
333 cte_name: &str,
334 stmt: &SelectStatement,
335 ctx: &ExecutionContext,
336 cte_registry: &mut CteRegistry,
337 column_aliases: Option<&[Identifier]>,
338 ) -> Result<(Vec<String>, RowVec)> {
339 use radixdb_sql::ast::SetOperationType;
340
341 const MAX_ITERATIONS: usize = 10000;
343
344 if stmt.set_operations.is_empty() {
346 return Err(Error::InvalidArgument(
347 "Recursive CTE must have UNION ALL between anchor and recursive members"
348 .to_string(),
349 ));
350 }
351
352 for set_op in &stmt.set_operations {
354 if !matches!(set_op.operation, SetOperationType::UnionAll) {
355 return Err(Error::InvalidArgument(
356 "Recursive CTE only supports UNION ALL (not UNION)".to_string(),
357 ));
358 }
359 }
360
361 let anchor_stmt = SelectStatement {
363 token: stmt.token.clone(),
364 distinct: stmt.distinct,
365 distinct_on: stmt.distinct_on.clone(),
366 columns: stmt.columns.clone(),
367 with: None,
368 table_expr: stmt.table_expr.clone(),
369 where_clause: stmt.where_clause.clone(),
370 group_by: stmt.group_by.clone(),
371 having: stmt.having.clone(),
372 window_defs: stmt.window_defs.clone(),
373 order_by: vec![], limit: None,
375 offset: None,
376 set_operations: vec![],
377 };
378
379 let result = self.host.cte_execute_select(&anchor_stmt, ctx)?;
380 let anchor_columns = result.columns().to_vec();
381
382 if let Some(aliases) = column_aliases {
383 if aliases.len() != anchor_columns.len() {
384 return Err(Error::InvalidArgument(format!(
385 "recursive CTE {cte_name} declares {} columns but anchor returns {}",
386 aliases.len(),
387 anchor_columns.len()
388 )));
389 }
390 }
391
392 let columns: Vec<String> = if let Some(aliases) = column_aliases {
394 aliases
395 .iter()
396 .enumerate()
397 .map(|(i, alias)| {
398 if i < anchor_columns.len() {
399 alias.value.to_string()
400 } else {
401 anchor_columns
402 .get(i)
403 .cloned()
404 .unwrap_or_else(|| format!("col{}", i))
405 }
406 })
407 .collect()
408 } else {
409 anchor_columns
410 };
411
412 let mut all_rows = materialize_result(result)?;
413 if all_rows.iter().any(|(_, row)| row.len() != columns.len()) {
414 return Err(Error::InvalidArgument(format!(
415 "recursive CTE {cte_name} anchor row width does not match its {} columns",
416 columns.len()
417 )));
418 }
419
420 if all_rows.is_empty() {
422 return Ok((columns, all_rows));
423 }
424
425 let mut target_types = vec![DataType::Null; columns.len()];
426 for (_, row) in all_rows.iter() {
427 for (column, value) in row.iter().enumerate() {
428 target_types[column] = merge_set_type(target_types[column], value.data_type())
429 .map_err(|error| {
430 Error::InvalidArgument(format!(
431 "recursive CTE {cte_name} anchor types are incompatible: {error}"
432 ))
433 })?;
434 }
435 }
436 for (_, row) in all_rows.iter_mut() {
437 for (value, target_type) in row.iter_mut().zip(&target_types) {
438 if value.data_type() != *target_type {
439 *value = value.try_coerce_to_type(*target_type).map_err(|error| {
440 Error::InvalidArgument(format!(
441 "recursive CTE {cte_name} anchor type is incompatible: {error}"
442 ))
443 })?;
444 }
445 }
446 }
447
448 let mut retained = RetainedRowsBudget::new("recursive CTE");
450 for (_, row) in all_rows.iter() {
451 retained.admit(row)?; retained.admit(row)?; }
454 let mut working_rows = all_rows.clone();
455
456 let mut converged = false;
458 for _iteration in 0..MAX_ITERATIONS {
459 ctx.check_cancelled()?;
460 if working_rows.is_empty() {
461 converged = true;
462 break;
463 }
464
465 let mut temp_registry = CteRegistry::new();
467 for (name, (cols, rows, materialized_rows)) in cte_registry.iter() {
469 temp_registry.store_arc(
470 name,
471 cols.clone(),
472 CompactArc::clone(rows),
473 Arc::clone(materialized_rows),
474 );
475 }
476 let working_rows_arc = CompactArc::new(working_rows.into_iter().collect());
479 temp_registry.store_arc(
480 cte_name,
481 CompactArc::new(columns.clone()),
482 CompactArc::clone(&working_rows_arc),
483 Arc::new(OnceLock::new()),
484 );
485
486 let mut new_rows = RowVec::new();
488 for set_op in &stmt.set_operations {
489 let mut recursive_result =
491 self.execute_cte_query(&set_op.right, ctx, &mut temp_registry)?;
492
493 if recursive_result.0.len() != columns.len() {
494 return Err(Error::InvalidArgument(format!(
495 "recursive CTE {cte_name} member returns {} columns but anchor returns {}",
496 recursive_result.0.len(),
497 columns.len()
498 )));
499 }
500
501 let mut merged_types = target_types.clone();
502 for (_, row) in recursive_result.1.iter() {
503 if row.len() != columns.len() {
504 return Err(Error::InvalidArgument(format!(
505 "recursive CTE {cte_name} member row width {} does not match {}",
506 row.len(),
507 columns.len()
508 )));
509 }
510 for (column, value) in row.iter().enumerate() {
511 merged_types[column] =
512 merge_set_type(merged_types[column], value.data_type()).map_err(
513 |error| {
514 Error::InvalidArgument(format!(
515 "recursive CTE {cte_name} member type is incompatible: {error}"
516 ))
517 },
518 )?;
519 }
520 }
521
522 if merged_types != target_types {
523 for (_, row) in all_rows.iter_mut().chain(new_rows.iter_mut()) {
524 for (value, target_type) in row.iter_mut().zip(&merged_types) {
525 if value.data_type() != *target_type {
526 *value =
527 value.try_coerce_to_type(*target_type).map_err(|error| {
528 Error::InvalidArgument(format!(
529 "recursive CTE {cte_name} type migration failed: {error}"
530 ))
531 })?;
532 }
533 }
534 }
535 target_types = merged_types;
536 }
537
538 let base_id = new_rows.len() as i64;
540 for (i, (_, mut row)) in recursive_result.1.drain(..).enumerate() {
541 if i & 0xff == 0 {
542 ctx.check_cancelled()?;
543 }
544 if row.len() != columns.len() {
545 return Err(Error::InvalidArgument(format!(
546 "recursive CTE {cte_name} member row width {} does not match {}",
547 row.len(),
548 columns.len()
549 )));
550 }
551 for (value, target_type) in row.iter_mut().zip(&target_types) {
552 if value.data_type() != *target_type {
553 *value = value.try_coerce_to_type(*target_type).map_err(|error| {
554 Error::InvalidArgument(format!(
555 "recursive CTE {cte_name} member type is incompatible with anchor: {error}"
556 ))
557 })?;
558 }
559 }
560 retained.admit(&row)?;
561 new_rows.push((base_id + i as i64, row));
562 }
563 }
564
565 drop(temp_registry);
568 for (_, row) in working_rows_arc.iter() {
569 retained.release(row);
570 }
571
572 if new_rows.is_empty() {
573 converged = true;
574 break;
575 }
576
577 let base_id = all_rows.len() as i64;
579 for (i, (_, row)) in new_rows.iter().enumerate() {
580 retained.admit(row)?;
581 all_rows.push((base_id + i as i64, row.clone()));
582 }
583
584 working_rows = new_rows;
586 }
587
588 if !converged {
589 return Err(Error::InvalidArgument(format!(
590 "recursive CTE {cte_name} exceeded {MAX_ITERATIONS} iterations"
591 )));
592 }
593
594 let classification = get_classification(stmt);
595 all_rows =
596 self.apply_order_by_limit_offset(stmt, ctx, &classification, all_rows, &columns)?;
597
598 Ok((columns, all_rows))
599 }
600
601 fn execute_main_query_with_ctes(
603 &self,
604 stmt: &SelectStatement,
605 ctx: &ExecutionContext,
606 cte_registry: &mut CteRegistry,
607 ) -> Result<Box<dyn QueryResult>> {
608 let ctx_with_ctes = ctx.with_cte_data(cte_registry.data());
609 let mut statement = stmt.clone();
610 statement.with = None;
611 self.host.cte_execute_select(&statement, &ctx_with_ctes)
615 }
616
617 #[allow(dead_code)]
619 pub(crate) fn execute_query_on_cte_result(
620 &self,
621 stmt: &SelectStatement,
622 ctx: &ExecutionContext,
623 cte_columns: Vec<String>,
624 cte_rows: RowVec,
625 ) -> Result<(Vec<String>, RowVec)> {
626 let (cols, rows, _applied) =
627 self.execute_query_on_cte_result_inner(stmt, ctx, cte_columns, cte_rows, false)?;
628 Ok((cols, rows))
629 }
630
631 pub(crate) fn execute_query_on_cte_result_inner(
636 &self,
637 stmt: &SelectStatement,
638 ctx: &ExecutionContext,
639 cte_columns: Vec<String>,
640 cte_rows: RowVec,
641 skip_order_limit: bool,
642 ) -> Result<(Vec<String>, RowVec, bool)> {
643 let classification = get_classification(stmt);
645
646 let filtered_rows = if let Some(ref where_clause) = stmt.where_clause {
648 let processed_where = if classification.where_has_subqueries {
651 self.host.process_where_subqueries(where_clause, ctx)?
652 } else {
653 (**where_clause).clone()
654 };
655
656 let mut eval = ExpressionEval::compile_with_options(
658 &processed_where,
659 &cte_columns,
660 None,
661 ctx.outer_columns(),
662 None,
663 self.host.cte_function_registry(),
664 )?
665 .with_context(ctx);
666
667 let mut result = RowVec::new();
668 let mut row_id = 0i64;
669 for (_, row) in cte_rows {
670 if eval.eval_bool_checked(&row)? {
671 result.push((row_id, row));
672 row_id += 1;
673 }
674 }
675 result
676 } else {
677 cte_rows
678 };
679
680 if classification.has_aggregation {
682 let result = self.host.execute_select_with_aggregation(
683 stmt,
684 ctx,
685 filtered_rows,
686 &cte_columns,
687 )?;
688 let columns = result.columns().to_vec();
689 let mut rows = materialize_result(result)?;
690
691 if !skip_order_limit {
692 rows =
693 self.apply_order_by_limit_offset(stmt, ctx, &classification, rows, &columns)?;
694 }
695 return Ok((columns, rows, !skip_order_limit));
696 }
697
698 if classification.has_window_functions {
700 let result = self.host.execute_select_with_window_functions(
701 stmt,
702 ctx,
703 &filtered_rows,
704 &cte_columns,
705 )?;
706 let columns = result.columns().to_vec();
707 let mut rows = materialize_result(result)?;
708
709 if !skip_order_limit {
710 rows =
711 self.apply_order_by_limit_offset(stmt, ctx, &classification, rows, &columns)?;
712 }
713 return Ok((columns, rows, !skip_order_limit));
714 }
715
716 let processed_columns = self
718 .host
719 .try_process_select_subqueries(&stmt.columns, ctx)?;
720 let columns_to_use = processed_columns.as_ref().unwrap_or(&stmt.columns);
721
722 let output_columns =
724 self.resolve_cte_output_columns_from_exprs(columns_to_use, &cte_columns)?;
725
726 let needs_projection = self.needs_projection_for_columns(columns_to_use);
727
728 if skip_order_limit {
729 let needs_source_for_order = classification.has_order_by
734 && self.order_by_needs_source_columns(
735 &stmt.order_by,
736 &output_columns,
737 &cte_columns,
738 );
739
740 if needs_source_for_order {
741 let mut combined_columns = output_columns.clone();
744 for src_col in &cte_columns {
745 if !combined_columns
746 .iter()
747 .any(|c| c.eq_ignore_ascii_case(src_col))
748 {
749 combined_columns.push(src_col.clone());
750 }
751 }
752
753 let result_rows = if needs_projection {
755 let projected = self.project_cte_rows_from_columns(
756 columns_to_use,
757 &filtered_rows,
758 &cte_columns,
759 ctx,
760 )?;
761 let extra_src_indices: Vec<usize> = cte_columns
763 .iter()
764 .enumerate()
765 .filter(|(_, c)| {
766 !output_columns.iter().any(|oc| oc.eq_ignore_ascii_case(c))
767 })
768 .map(|(i, _)| i)
769 .collect();
770
771 projected
772 .into_iter()
773 .zip(filtered_rows.iter())
774 .map(|((id, proj_row), (_, src_row))| {
775 let mut vals = proj_row.into_values();
776 for &idx in &extra_src_indices {
777 vals.push(
778 src_row.get(idx).cloned().unwrap_or(Value::null_unknown()),
779 );
780 }
781 (id, Row::from_values(vals))
782 })
783 .collect()
784 } else {
785 filtered_rows
786 };
787 return Ok((combined_columns, result_rows, false));
788 }
789
790 let result_rows = if needs_projection {
791 self.project_cte_rows_from_columns(
792 columns_to_use,
793 &filtered_rows,
794 &cte_columns,
795 ctx,
796 )?
797 } else {
798 filtered_rows
799 };
800 return Ok((output_columns, result_rows, false));
801 }
802
803 let needs_pre_sort = classification.has_order_by
806 && self.order_by_needs_source_columns(&stmt.order_by, &output_columns, &cte_columns);
807
808 let mut result_rows = if needs_pre_sort {
810 if needs_projection {
811 let sorted =
813 self.apply_order_by_to_rows(filtered_rows, &stmt.order_by, &cte_columns)?;
814 self.project_cte_rows_from_columns(columns_to_use, &sorted, &cte_columns, ctx)?
815 } else {
816 self.apply_order_by_to_rows(filtered_rows, &stmt.order_by, &cte_columns)?
817 }
818 } else {
819 let mut rows = if needs_projection {
821 self.project_cte_rows_from_columns(
822 columns_to_use,
823 &filtered_rows,
824 &cte_columns,
825 ctx,
826 )?
827 } else {
828 filtered_rows
829 };
830
831 if classification.has_order_by {
832 rows = self.apply_order_by_to_rows(rows, &stmt.order_by, &output_columns)?;
833 }
834 rows
835 };
836
837 if classification.has_offset {
839 if let Some(ref offset_expr) = stmt.offset {
840 let offset = evaluate_page_expression(offset_expr, ctx, "OFFSET")?;
841 if offset > 0 && offset < result_rows.len() {
843 result_rows.drain(..offset);
844 } else if offset >= result_rows.len() {
845 result_rows.clear();
846 }
847 }
848 }
849
850 if classification.has_limit {
851 if let Some(ref limit_expr) = stmt.limit {
852 let limit = evaluate_page_expression(limit_expr, ctx, "LIMIT")?;
853 if limit < result_rows.len() {
854 result_rows.truncate(limit);
855 }
856 }
857 }
858
859 Ok((output_columns, result_rows, true))
860 }
861
862 fn extract_cte_name_for_lookup(&self, expr: &Expression) -> Option<String> {
864 match expr {
865 Expression::CteReference(cte_ref) => Some(cte_ref.name.value.to_string()),
866 Expression::TableSource(simple_table_source) => {
867 Some(simple_table_source.name.value.to_string())
868 }
869 Expression::Identifier(id) => Some(id.value.to_string()),
870 _ => None,
871 }
872 }
873
874 fn resolve_cte_output_columns_from_exprs(
876 &self,
877 columns: &[Expression],
878 cte_columns: &[String],
879 ) -> Result<Vec<String>> {
880 let mut output_columns = Vec::new();
881
882 for (i, col_expr) in columns.iter().enumerate() {
883 match col_expr {
884 Expression::Star(_) | Expression::QualifiedStar(_) => {
885 output_columns.extend(cte_columns.iter().cloned());
886 }
887 Expression::Identifier(id) => {
888 output_columns.push(id.value.to_string());
889 }
890 Expression::Aliased(aliased) => {
891 output_columns.push(aliased.alias.value.to_string());
892 }
893 _ => {
894 output_columns.push(format!("expr{}", i + 1));
895 }
896 }
897 }
898
899 if output_columns.is_empty() {
900 output_columns = cte_columns.to_vec();
901 }
902
903 Ok(output_columns)
904 }
905
906 fn needs_projection_for_columns(&self, columns: &[Expression]) -> bool {
908 if columns.is_empty() {
909 return false;
910 }
911
912 if columns.len() == 1 {
914 if let Expression::Star(_) = &columns[0] {
915 return false;
916 }
917 }
918
919 true
920 }
921
922 fn project_cte_rows_from_columns(
924 &self,
925 columns: &[Expression],
926 rows: &RowVec,
927 cte_columns: &[String],
928 ctx: &ExecutionContext,
929 ) -> Result<RowVec> {
930 use super::expression::{ExecuteContext, ExprVM, SharedProgram};
931
932 let col_index_map = build_column_index_map(cte_columns);
933
934 enum CompiledColumn {
937 Star,
938 Identifier(usize),
939 Compiled(SharedProgram),
940 }
941
942 let compiled_columns: Vec<CompiledColumn> = columns
943 .iter()
944 .map(|col_expr| match col_expr {
945 Expression::Star(_) => Ok(CompiledColumn::Star),
946 Expression::Identifier(id) => {
947 let idx = col_index_map
948 .get(id.value_lower.as_str())
949 .copied()
950 .ok_or_else(|| Error::ColumnNotFound(id.value.to_string()))?;
951 Ok(CompiledColumn::Identifier(idx))
952 }
953 Expression::Aliased(aliased) => {
954 let program = compile_expression_with_context(
955 &aliased.expression,
956 cte_columns,
957 ctx.outer_columns(),
958 self.host.cte_function_registry(),
959 )?;
960 Ok(CompiledColumn::Compiled(program))
961 }
962 _ => {
963 let program = compile_expression_with_context(
964 col_expr,
965 cte_columns,
966 ctx.outer_columns(),
967 self.host.cte_function_registry(),
968 )?;
969 Ok(CompiledColumn::Compiled(program))
970 }
971 })
972 .collect::<Result<Vec<_>>>()?;
973
974 let mut vm = ExprVM::new();
976 let mut result_rows = RowVec::with_capacity(rows.len());
977
978 for (row_id, (_, row)) in rows.iter().enumerate() {
979 let mut values: CompactVec<Value> =
981 CompactVec::with_capacity(columns.len().max(row.len()));
982 let mut exec_ctx = ExecuteContext::new(row)
984 .with_params(ctx.params())
985 .with_named_params(ctx.named_params())
986 .with_transaction_id(ctx.transaction_id())
987 .with_stored_function_invoker(ctx.stored_function_invoker());
988 if let Some(outer_row) = ctx.outer_row() {
989 exec_ctx = exec_ctx.with_outer_row(outer_row);
990 }
991
992 for compiled in &compiled_columns {
993 match compiled {
994 CompiledColumn::Star => {
995 values.extend(row.iter().cloned());
997 }
998 CompiledColumn::Identifier(idx) => {
999 values.push(row.get(*idx).cloned().unwrap_or_else(Value::null_unknown));
1000 }
1001 CompiledColumn::Compiled(program) => {
1002 values.push(vm.execute_cow(program, &exec_ctx)?);
1003 }
1004 }
1005 }
1006
1007 result_rows.push((row_id as i64, Row::from_compact_vec(values)));
1008 }
1009
1010 Ok(result_rows)
1011 }
1012
1013 pub(crate) fn has_cte(&self, stmt: &SelectStatement) -> bool {
1015 stmt.with.is_some()
1016 }
1017
1018 fn apply_order_by_to_rows(
1020 &self,
1021 mut rows: RowVec,
1022 order_by: &[radixdb_sql::ast::OrderByExpression],
1023 columns: &[String],
1024 ) -> Result<RowVec> {
1025 if order_by.is_empty() || rows.is_empty() {
1026 return Ok(rows);
1027 }
1028
1029 let col_index_map = build_column_index_map(columns);
1031
1032 let order_specs: Vec<(Option<usize>, bool, Option<bool>)> = order_by
1034 .iter()
1035 .map(|ob| {
1036 let col_idx = match &ob.expression {
1037 Expression::Identifier(id) => {
1038 col_index_map.get(id.value_lower.as_str()).copied()
1039 }
1040 Expression::QualifiedIdentifier(qi) => {
1041 let full_name =
1043 format!("{}.{}", qi.qualifier, qi.name.value).to_lowercase();
1044 col_index_map
1045 .get(&full_name)
1046 .or_else(|| col_index_map.get(qi.name.value_lower.as_str()))
1047 .copied()
1048 }
1049 Expression::IntegerLiteral(lit) => {
1050 let pos = lit.value as usize;
1052 if pos > 0 && pos <= columns.len() {
1053 Some(pos - 1)
1054 } else {
1055 None
1056 }
1057 }
1058 _ => None,
1059 };
1060 (col_idx, ob.ascending, ob.nulls_first)
1061 })
1062 .collect();
1063
1064 rows.sort_by(|(_, a), (_, b)| {
1067 for (col_idx, ascending, nulls_first) in &order_specs {
1068 if let Some(idx) = col_idx {
1069 let a_val = a.get(*idx);
1070 let b_val = b.get(*idx);
1071
1072 let a_is_null = a_val.is_none() || a_val.map(|v| v.is_null()).unwrap_or(true);
1074 let b_is_null = b_val.is_none() || b_val.map(|v| v.is_null()).unwrap_or(true);
1075
1076 if a_is_null || b_is_null {
1078 if a_is_null && b_is_null {
1079 continue; }
1081 let nulls_come_first = nulls_first.unwrap_or(!*ascending);
1083 let cmp = if a_is_null {
1084 if nulls_come_first {
1085 std::cmp::Ordering::Less
1086 } else {
1087 std::cmp::Ordering::Greater
1088 }
1089 } else if nulls_come_first {
1090 std::cmp::Ordering::Greater
1091 } else {
1092 std::cmp::Ordering::Less
1093 };
1094 return cmp;
1095 }
1096
1097 let cmp = match (a_val, b_val) {
1099 (Some(av), Some(bv)) => {
1100 av.partial_cmp(bv).unwrap_or(std::cmp::Ordering::Equal)
1101 }
1102 _ => std::cmp::Ordering::Equal,
1103 };
1104
1105 let cmp = if !*ascending { cmp.reverse() } else { cmp };
1106
1107 if cmp != std::cmp::Ordering::Equal {
1108 return cmp;
1109 }
1110 }
1111 }
1112 std::cmp::Ordering::Equal
1113 });
1114
1115 Ok(rows)
1116 }
1117
1118 fn order_by_needs_source_columns(
1122 &self,
1123 order_by: &[OrderByExpression],
1124 output_columns: &[String],
1125 source_columns: &[String],
1126 ) -> bool {
1127 let output_lower: AHashSet<String> =
1128 output_columns.iter().map(|c| c.to_lowercase()).collect();
1129
1130 for ob in order_by {
1131 if self.expr_references_source_not_output(&ob.expression, &output_lower, source_columns)
1132 {
1133 return true;
1134 }
1135 }
1136 false
1137 }
1138
1139 fn expr_references_source_not_output(
1143 &self,
1144 expr: &Expression,
1145 output_lower: &AHashSet<String>,
1146 source_columns: &[String],
1147 ) -> bool {
1148 let check = |e: &Expression| {
1149 self.expr_references_source_not_output(e, output_lower, source_columns)
1150 };
1151 let is_source_not_output = |name: &str| {
1152 !output_lower.contains(name)
1153 && source_columns.iter().any(|c| c.eq_ignore_ascii_case(name))
1154 };
1155
1156 match expr {
1157 Expression::Identifier(id) => is_source_not_output(id.value_lower.as_str()),
1159 Expression::QualifiedIdentifier(qi) => {
1160 is_source_not_output(qi.name.value_lower.as_str())
1161 }
1162
1163 Expression::IntegerLiteral(_)
1165 | Expression::FloatLiteral(_)
1166 | Expression::StringLiteral(_)
1167 | Expression::BooleanLiteral(_)
1168 | Expression::NullLiteral(_)
1169 | Expression::IntervalLiteral(_)
1170 | Expression::Parameter(_)
1171 | Expression::Star(_)
1172 | Expression::QualifiedStar(_)
1173 | Expression::Default(_) => false,
1174
1175 Expression::Prefix(p) => check(&p.right),
1177 Expression::Infix(inf) => check(&inf.left) || check(&inf.right),
1178 Expression::FunctionCall(fc) => fc.arguments.iter().any(&check),
1179 Expression::Cast(c) => check(&c.expr),
1180 Expression::Aliased(a) => check(&a.expression),
1181 Expression::Case(case) => {
1182 case.value.as_ref().is_some_and(|v| check(v))
1183 || case
1184 .when_clauses
1185 .iter()
1186 .any(|w| check(&w.condition) || check(&w.then_result))
1187 || case.else_value.as_ref().is_some_and(|e| check(e))
1188 }
1189 Expression::Between(b) => check(&b.expr) || check(&b.lower) || check(&b.upper),
1190 Expression::In(i) => check(&i.left),
1191 Expression::Like(l) => check(&l.left) || check(&l.pattern),
1192 Expression::Distinct(d) => check(&d.expr),
1193 Expression::Window(w) => w.function.arguments.iter().any(check),
1194
1195 _ => true,
1197 }
1198 }
1199
1200 fn apply_order_by_limit_offset(
1204 &self,
1205 stmt: &SelectStatement,
1206 ctx: &ExecutionContext,
1207 classification: &QueryClassification,
1208 mut rows: RowVec,
1209 columns: &[String],
1210 ) -> Result<RowVec> {
1211 if classification.has_order_by {
1212 rows = self.apply_order_by_to_rows(rows, &stmt.order_by, columns)?;
1213 }
1214
1215 if classification.has_offset {
1216 if let Some(ref offset_expr) = stmt.offset {
1217 let offset = evaluate_page_expression(offset_expr, ctx, "OFFSET")?;
1218 if offset > 0 && offset < rows.len() {
1219 rows.drain(..offset);
1220 } else if offset >= rows.len() {
1221 rows.clear();
1222 }
1223 }
1224 }
1225
1226 if classification.has_limit {
1227 if let Some(ref limit_expr) = stmt.limit {
1228 let limit = evaluate_page_expression(limit_expr, ctx, "LIMIT")?;
1229 if limit < rows.len() {
1230 rows.truncate(limit);
1231 }
1232 }
1233 }
1234
1235 Ok(rows)
1236 }
1237
1238 fn should_use_limit_pushdown_instead(
1247 &self,
1248 stmt: &SelectStatement,
1249 with_clause: &WithClause,
1250 ) -> bool {
1251 if stmt.limit.is_none() || !stmt.order_by.is_empty() {
1253 return false;
1254 }
1255
1256 let join_source = match &stmt.table_expr {
1258 Some(expr) => match expr.as_ref() {
1259 Expression::JoinSource(js) => js,
1260 _ => return false,
1261 },
1262 None => return false,
1263 };
1264
1265 let join_type = join_source.join_type.to_uppercase();
1266 let is_inner_join = join_type == "INNER" || join_type.is_empty() || join_type == "JOIN";
1267 let is_left_join = join_type == "LEFT" || join_type == "LEFT OUTER";
1268 let is_right_join = join_type == "RIGHT" || join_type == "RIGHT OUTER";
1269
1270 if !is_inner_join && !is_left_join && !is_right_join {
1271 return false;
1272 }
1273
1274 let cte_names: AHashSet<String> = with_clause
1276 .ctes
1277 .iter()
1278 .filter(|c| !c.is_recursive && !c.query.group_by.columns.is_empty())
1279 .map(|c| c.name.value_lower.to_string())
1280 .collect();
1281
1282 if cte_names.is_empty() {
1283 return false;
1284 }
1285
1286 let left_cte = self
1287 .extract_cte_name_for_lookup(&join_source.left)
1288 .filter(|n| cte_names.contains(&n.to_lowercase()));
1289 let right_cte = self
1290 .extract_cte_name_for_lookup(&join_source.right)
1291 .filter(|n| cte_names.contains(&n.to_lowercase()));
1292
1293 match (&left_cte, &right_cte) {
1297 (Some(_), None) if is_inner_join || is_right_join => true,
1298 (None, Some(_)) if is_inner_join || is_left_join => true,
1299 _ => false,
1300 }
1301 }
1302
1303 pub(crate) fn try_inline_ctes(
1311 &self,
1312 stmt: &SelectStatement,
1313 with_clause: &WithClause,
1314 ) -> Option<SelectStatement> {
1315 let simple_projection = stmt.columns.len() == 1
1320 && matches!(
1321 stmt.columns[0],
1322 Expression::Star(_) | Expression::QualifiedStar(_)
1323 );
1324 if !simple_projection
1325 || stmt.where_clause.is_some()
1326 || stmt.having.is_some()
1327 || !stmt.group_by.columns.is_empty()
1328 || !stmt.window_defs.is_empty()
1329 || !stmt.order_by.is_empty()
1330 || stmt.limit.is_some()
1331 || stmt.offset.is_some()
1332 || !stmt.set_operations.is_empty()
1333 || stmt.distinct
1334 || !stmt.distinct_on.is_empty()
1335 {
1336 return None;
1337 }
1338
1339 let table_expr = stmt.table_expr.as_ref()?;
1341
1342 if self.should_use_limit_pushdown_instead(stmt, with_clause) {
1348 return None;
1349 }
1350
1351 let cte_names_lower: Vec<(String, &CommonTableExpression)> = with_clause
1354 .ctes
1355 .iter()
1356 .map(|cte| {
1357 if cte.is_recursive || !cte.column_names.is_empty() {
1359 return Err(());
1360 }
1361 Ok((cte.name.value_lower.to_string(), cte))
1362 })
1363 .collect::<std::result::Result<Vec<_>, _>>()
1364 .ok()?;
1365
1366 let cte_defs: StringMap<&CommonTableExpression> = cte_names_lower.iter().cloned().collect();
1368 let cte_name_set: AHashSet<&str> = cte_defs.keys().map(|s| s.as_str()).collect();
1369
1370 for (_, cte) in &cte_names_lower {
1373 for other_cte_name in &cte_name_set {
1374 if self.query_references_cte(&cte.query, other_cte_name) {
1375 return None;
1377 }
1378 }
1379 }
1380
1381 let mut table_ref_counts: StringMap<usize> =
1384 cte_defs.keys().map(|name| (name.clone(), 0)).collect();
1385 let mut where_ref_counts: StringMap<usize> = table_ref_counts.clone();
1386
1387 self.count_cte_references_in_expr(table_expr, &mut table_ref_counts);
1389
1390 if let Some(ref where_clause) = stmt.where_clause {
1392 self.count_cte_references_in_expr(where_clause, &mut where_ref_counts);
1393 }
1394
1395 for name in cte_defs.keys() {
1399 let table_refs = table_ref_counts.get(name).copied().unwrap_or(0);
1400 let where_refs = where_ref_counts.get(name).copied().unwrap_or(0);
1401
1402 if where_refs > 0 {
1404 return None;
1405 }
1406
1407 if table_refs > 1 {
1409 return None;
1410 }
1411 }
1413
1414 let any_refs = table_ref_counts.values().any(|&count| count > 0);
1418 if !any_refs {
1419 return None;
1421 }
1422
1423 let inlined_expr = self.try_inline_cte_references(table_expr, &cte_defs)?;
1425
1426 Some(SelectStatement {
1427 token: stmt.token.clone(),
1428 distinct: stmt.distinct,
1429 distinct_on: stmt.distinct_on.clone(),
1430 columns: stmt.columns.clone(),
1431 with: None, table_expr: Some(Box::new(inlined_expr)),
1433 where_clause: stmt.where_clause.clone(),
1434 group_by: stmt.group_by.clone(),
1435 having: stmt.having.clone(),
1436 window_defs: stmt.window_defs.clone(),
1437 order_by: stmt.order_by.clone(),
1438 limit: stmt.limit.clone(),
1439 offset: stmt.offset.clone(),
1440 set_operations: stmt.set_operations.clone(),
1441 })
1442 }
1443
1444 fn query_references_cte(&self, stmt: &SelectStatement, cte_name: &str) -> bool {
1446 if let Some(ref table_expr) = stmt.table_expr {
1448 if self.expr_references_cte(table_expr, cte_name) {
1449 return true;
1450 }
1451 }
1452
1453 if let Some(ref where_clause) = stmt.where_clause {
1455 if self.expr_references_cte(where_clause, cte_name) {
1456 return true;
1457 }
1458 }
1459
1460 false
1461 }
1462
1463 fn expr_references_cte(&self, expr: &Expression, cte_name: &str) -> bool {
1465 match expr {
1466 Expression::CteReference(cte_ref) => cte_ref.name.value.eq_ignore_ascii_case(cte_name),
1467 Expression::TableSource(ts) => ts.name.value.eq_ignore_ascii_case(cte_name),
1468 Expression::Identifier(id) => id.value.eq_ignore_ascii_case(cte_name),
1469 Expression::JoinSource(js) => {
1470 self.expr_references_cte(&js.left, cte_name)
1471 || self.expr_references_cte(&js.right, cte_name)
1472 }
1473 Expression::SubquerySource(sq) => self.query_references_cte(&sq.subquery, cte_name),
1474 Expression::ScalarSubquery(sq) => self.query_references_cte(&sq.subquery, cte_name),
1475 Expression::In(in_expr) => {
1476 if let Expression::ScalarSubquery(sq) = &*in_expr.right {
1478 self.query_references_cte(&sq.subquery, cte_name)
1479 } else {
1480 false
1481 }
1482 }
1483 Expression::Exists(ex) => self.query_references_cte(&ex.subquery, cte_name),
1484 Expression::Infix(infix) => {
1485 self.expr_references_cte(&infix.left, cte_name)
1486 || self.expr_references_cte(&infix.right, cte_name)
1487 }
1488 _ => false,
1489 }
1490 }
1491
1492 fn count_cte_references_in_stmt(
1494 &self,
1495 stmt: &SelectStatement,
1496 ref_counts: &mut StringMap<usize>,
1497 ) {
1498 if let Some(ref table_expr) = stmt.table_expr {
1500 self.count_cte_references_in_expr(table_expr, ref_counts);
1501 }
1502
1503 if let Some(ref where_clause) = stmt.where_clause {
1505 self.count_cte_references_in_expr(where_clause, ref_counts);
1506 }
1507
1508 for col in &stmt.columns {
1510 self.count_cte_references_in_expr(col, ref_counts);
1511 }
1512 }
1513
1514 fn count_cte_references_in_expr(&self, expr: &Expression, ref_counts: &mut StringMap<usize>) {
1516 match expr {
1517 Expression::CteReference(cte_ref) => {
1518 let name: &str = cte_ref.name.value_lower.as_str();
1519 if let Some(count) = ref_counts.get_mut(name) {
1520 *count += 1;
1521 }
1522 }
1523 Expression::TableSource(ts) => {
1524 let name: &str = ts.name.value_lower.as_str();
1525 if let Some(count) = ref_counts.get_mut(name) {
1526 *count += 1;
1527 }
1528 }
1529 Expression::Identifier(id) => {
1530 let name: &str = id.value_lower.as_str();
1531 if let Some(count) = ref_counts.get_mut(name) {
1532 *count += 1;
1533 }
1534 }
1535 Expression::JoinSource(js) => {
1536 self.count_cte_references_in_expr(&js.left, ref_counts);
1537 self.count_cte_references_in_expr(&js.right, ref_counts);
1538 }
1539 Expression::SubquerySource(sq) => {
1540 self.count_cte_references_in_stmt(&sq.subquery, ref_counts);
1541 }
1542 Expression::ScalarSubquery(sq) => {
1543 self.count_cte_references_in_stmt(&sq.subquery, ref_counts);
1544 }
1545 Expression::In(in_expr) => {
1546 self.count_cte_references_in_expr(&in_expr.left, ref_counts);
1547 if let Expression::ScalarSubquery(sq) = &*in_expr.right {
1549 self.count_cte_references_in_stmt(&sq.subquery, ref_counts);
1550 }
1551 }
1552 Expression::Exists(ex) => {
1553 self.count_cte_references_in_stmt(&ex.subquery, ref_counts);
1554 }
1555 Expression::Infix(infix) => {
1556 self.count_cte_references_in_expr(&infix.left, ref_counts);
1557 self.count_cte_references_in_expr(&infix.right, ref_counts);
1558 }
1559 Expression::Aliased(aliased) => {
1560 self.count_cte_references_in_expr(&aliased.expression, ref_counts);
1561 }
1562 _ => {}
1563 }
1564 }
1565
1566 fn try_inline_cte_references(
1569 &self,
1570 expr: &Expression,
1571 cte_defs: &StringMap<&CommonTableExpression>,
1572 ) -> Option<Expression> {
1573 match expr {
1574 Expression::CteReference(cte_ref) => {
1575 let name = &cte_ref.name.value_lower;
1577 cte_defs.get(name.as_str()).map(|cte| {
1578 let alias = cte_ref
1580 .alias
1581 .clone()
1582 .unwrap_or_else(|| cte_ref.name.clone());
1583 Expression::SubquerySource(Box::new(SubqueryTableSource {
1584 token: Token::new(TokenType::Punctuator, "(", Position::new(0, 0, 0)),
1585 subquery: cte.query.clone(),
1586 alias: Some(alias),
1587 }))
1588 })
1589 }
1590 Expression::TableSource(ts) => {
1591 let name = &ts.name.value_lower;
1593 cte_defs.get(name.as_str()).map(|cte| {
1594 let alias = ts.alias.clone().unwrap_or_else(|| ts.name.clone());
1596 Expression::SubquerySource(Box::new(SubqueryTableSource {
1597 token: Token::new(TokenType::Punctuator, "(", Position::new(0, 0, 0)),
1598 subquery: cte.query.clone(),
1599 alias: Some(alias),
1600 }))
1601 })
1602 }
1603 Expression::JoinSource(js) => {
1604 let left_changed = self.try_inline_cte_references(&js.left, cte_defs);
1605 let right_changed = self.try_inline_cte_references(&js.right, cte_defs);
1606
1607 if left_changed.is_some() || right_changed.is_some() {
1609 let left = left_changed.unwrap_or_else(|| (*js.left).clone());
1610 let right = right_changed.unwrap_or_else(|| (*js.right).clone());
1611 Some(Expression::JoinSource(Box::new(JoinTableSource {
1612 token: js.token.clone(),
1613 left: Box::new(left),
1614 right: Box::new(right),
1615 join_type: js.join_type.clone(),
1616 condition: js.condition.clone(),
1617 using_columns: js.using_columns.clone(),
1618 })))
1619 } else {
1620 None
1621 }
1622 }
1623 Expression::SubquerySource(sq) => {
1624 if let Some(ref table_expr) = sq.subquery.table_expr {
1626 if let Some(inlined) = self.try_inline_cte_references(table_expr, cte_defs) {
1627 let mut new_subquery = (*sq.subquery).clone();
1628 new_subquery.table_expr = Some(Box::new(inlined));
1629 return Some(Expression::SubquerySource(Box::new(SubqueryTableSource {
1630 token: sq.token.clone(),
1631 subquery: Box::new(new_subquery),
1632 alias: sq.alias.clone(),
1633 })));
1634 }
1635 }
1636 None
1637 }
1638 _ => None, }
1640 }
1641}