1use super::{BindingContext, QueryPlan, RowSchema, SQLError, SQLParam, SchemaScope};
10use crate::ast::ColumnType;
11use crate::routines::RoutineResolution;
12
13pub fn hide_recursive_generated_schema(schema: &RowSchema, visible: usize) -> RowSchema {
14 let visible = visible.min(schema.len());
15 let base = RowSchema::with_types(
16 schema.columns()[..visible].to_vec(),
17 schema.column_types()[..visible].to_vec(),
18 );
19 let generated = schema
20 .columns()
21 .iter()
22 .enumerate()
23 .skip(visible)
24 .map(|(position, column)| {
25 (
26 crate::ColumnIdentity::unqualified(column),
27 schema.column_type(position).cloned(),
28 )
29 })
30 .collect::<Vec<_>>();
31 RowSchema::with_typed_virtual_identities(&base, &generated)
32}
33
34pub fn extend_cte_generated_schema(
35 routines: &dyn RoutineResolution,
36 cte: &crate::plan::CtePlan,
37 schema: RowSchema,
38 params: &[SQLParam],
39) -> Result<RowSchema, SQLError> {
40 extend_cte_generated_schema_mode(routines, cte, schema, params, true)
41}
42
43fn extend_cte_generated_schema_mode(
44 routines: &dyn RoutineResolution,
45 cte: &crate::plan::CtePlan,
46 schema: RowSchema,
47 params: &[SQLParam],
48 reject_output_conflicts: bool,
49) -> Result<RowSchema, SQLError> {
50 if cte.search.is_none() && cte.cycle.is_none() {
51 return Ok(schema);
52 }
53 let mut columns = schema.columns().to_vec();
54 let mut types = schema.column_types().to_vec();
55 let base_columns = columns.clone();
56 let require_columns = |requested: &[String], kind: &str| -> Result<(), SQLError> {
57 let mut seen = std::collections::BTreeSet::new();
58 for column in requested {
59 if !seen.insert(column) {
60 return Err(SQLError::Routine {
61 sqlstate: "42701".into(),
62 message: format!("{kind} column \"{column}\" specified more than once"),
63 });
64 }
65 if !base_columns.iter().any(|candidate| candidate == column) {
66 return Err(SQLError::Routine {
67 sqlstate: "42601".into(),
68 message: format!("{kind} column \"{column}\" not in WITH query column list"),
69 });
70 }
71 }
72 Ok(())
73 };
74 let reject_conflict = |name: &str, kind: &str| -> Result<(), SQLError> {
75 if reject_output_conflicts && base_columns.iter().any(|column| column == name) {
76 return Err(SQLError::Routine {
77 sqlstate: "42601".into(),
78 message: format!(
79 "{kind} column name \"{name}\" already used in WITH query column list"
80 ),
81 });
82 }
83 Ok(())
84 };
85 if let Some(search) = &cte.search {
86 require_columns(&search.columns, "search")?;
87 reject_conflict(&search.sequence_column, "search sequence")?;
88 columns.push(search.sequence_column.clone());
89 types.push(Some(if search.breadth_first {
90 ColumnType::Record
91 } else {
92 ColumnType::Array(Box::new(ColumnType::Record))
93 }));
94 }
95 if let Some(cycle) = &cte.cycle {
96 require_columns(&cycle.columns, "cycle")?;
97 if cycle.mark_column == cycle.path_column {
98 return Err(SQLError::Routine {
99 sqlstate: "42601".into(),
100 message: "cycle mark column name and cycle path column name are the same".into(),
101 });
102 }
103 reject_conflict(&cycle.mark_column, "cycle mark")?;
104 reject_conflict(&cycle.path_column, "cycle path")?;
105 let empty = RowSchema::default();
106 let mark_type = match (
107 crate::scalar_type_with_resolver(&cycle.mark_value, &empty, params, routines)?,
108 crate::scalar_type_with_resolver(&cycle.mark_default, &empty, params, routines)?,
109 ) {
110 (Some(left), Some(right)) => Some(crate::common_type(&left, &right)?),
111 (left @ Some(_), None) | (None, left @ Some(_)) => left,
112 (None, None) => None,
113 };
114 if let Some(mark_type) = mark_type.as_ref() {
115 crate::equality_operand_type(mark_type, mark_type)?;
116 }
117 for column in &cycle.columns {
118 if let Some(column_type) = base_columns
119 .iter()
120 .position(|candidate| candidate == column)
121 .and_then(|position| schema.column_type(position))
122 {
123 crate::equality_operand_type(column_type, column_type)?;
124 }
125 }
126 columns.push(cycle.mark_column.clone());
127 types.push(mark_type);
128 columns.push(cycle.path_column.clone());
129 types.push(Some(ColumnType::Array(Box::new(ColumnType::Record))));
130 }
131 Ok(RowSchema::with_types(columns, types))
132}
133
134pub fn extend_recursive_cte_binding_schema(
135 routines: &dyn RoutineResolution,
136 cte: &crate::plan::CtePlan,
137 schema: RowSchema,
138 params: &[SQLParam],
139) -> Result<RowSchema, SQLError> {
140 let extended = extend_cte_generated_schema_mode(routines, cte, schema.clone(), params, false)?;
141 let generated = extended
142 .columns()
143 .iter()
144 .enumerate()
145 .skip(schema.len())
146 .map(|(position, column)| {
147 (
148 crate::ColumnIdentity::unqualified(column),
149 extended.column_type(position).cloned(),
150 )
151 })
152 .collect::<Vec<_>>();
153 Ok(RowSchema::with_typed_conflicting_virtual_identities(
154 &schema, &generated,
155 ))
156}
157
158pub fn analyze_recursive_control_step(
159 routines: &dyn RoutineResolution,
160 cte: &crate::plan::CtePlan,
161 step: &QueryPlan,
162 base_schema: RowSchema,
163 params: &[SQLParam],
164 ctes: &BindingContext,
165) -> Result<(), SQLError> {
166 let provisional = extend_recursive_cte_binding_schema(routines, cte, base_schema, params)?;
167 let mut scope = SchemaScope::for_analysis(ctes)?;
168 scope.ctes.insert(cte.name.clone(), provisional);
169 scope.bind_query(routines, step, params, None).map(|_| ())
170}