1use radixdb_core::StringMap;
4use radixdb_sql::ast::{Expression, SelectStatement};
5
6use crate::utils::build_column_index_map;
7
8#[derive(Debug, Clone, PartialEq, Eq)]
10pub struct ProjectionScanPlan {
11 pub scan_indices: Vec<usize>,
12 pub scan_columns: Vec<String>,
13 pub output_indices_in_scan: Vec<usize>,
14 pub output_columns: Vec<String>,
15}
16
17#[derive(Debug, Clone, PartialEq, Eq)]
19pub struct NarrowKeyStreamPlan {
20 pub scan_indices: Vec<usize>,
21 pub scan_columns: Vec<String>,
22 pub key_index_in_scan: usize,
23}
24
25pub fn simple_projection_indices(
26 select: &[Expression],
27 all_columns: &[String],
28) -> Option<(Vec<usize>, Vec<String>)> {
29 if select.len() == 1
30 && matches!(
31 &select[0],
32 Expression::Star(_) | Expression::QualifiedStar(_)
33 )
34 {
35 return Some(((0..all_columns.len()).collect(), all_columns.to_vec()));
36 }
37
38 let columns = build_column_index_map(all_columns);
39 let mut indices = Vec::with_capacity(select.len());
40 let mut names = Vec::with_capacity(select.len());
41 for expression in select {
42 let (source, output) = match expression {
43 Expression::Identifier(identifier) => (
44 resolve_column(&columns, expression)?,
45 identifier.value.to_string(),
46 ),
47 Expression::QualifiedIdentifier(identifier) => (
48 resolve_column(&columns, expression)?,
49 identifier.name.value.to_string(),
50 ),
51 Expression::Aliased(aliased) => (
52 resolve_column(&columns, &aliased.expression)?,
53 aliased.alias.value.to_string(),
54 ),
55 _ => return None,
56 };
57 indices.push(source);
58 names.push(output);
59 }
60 Some((indices, names))
61}
62
63pub fn filtered_simple_projection(
64 filter: &Expression,
65 output_indices: &[usize],
66 output_columns: &[String],
67 all_columns: &[String],
68) -> Option<ProjectionScanPlan> {
69 if all_columns.is_empty() {
70 return None;
71 }
72 let columns = build_column_index_map(all_columns);
73 let mut filter_indices = Vec::new();
74 if !collect_expression_columns(filter, &columns, &mut filter_indices) {
75 return None;
76 }
77
78 let scan_indices = ordered_union_indices(
79 all_columns.len(),
80 output_indices.iter().copied().chain(filter_indices),
81 )?;
82 if scan_indices.len() == all_columns.len() {
83 return None;
84 }
85
86 let mut positions = vec![None; all_columns.len()];
87 for (position, source) in scan_indices.iter().copied().enumerate() {
88 positions[source] = Some(position);
89 }
90 let output_indices_in_scan = output_indices
91 .iter()
92 .map(|source| positions.get(*source).copied().flatten())
93 .collect::<Option<Vec<_>>>()?;
94 Some(plan_from_indices(
95 scan_indices,
96 all_columns,
97 output_indices_in_scan,
98 output_columns,
99 ))
100}
101
102pub fn narrow_key_stream(
103 filter: Option<&Expression>,
104 key_column: &str,
105 all_columns: &[String],
106) -> Option<NarrowKeyStreamPlan> {
107 if all_columns.is_empty() {
108 return None;
109 }
110 let columns = build_column_index_map(all_columns);
111 let key_index = *columns.get(&key_column.to_lowercase())?;
112 let mut required = vec![key_index];
113 if let Some(filter) = filter {
114 if !collect_expression_columns(filter, &columns, &mut required) {
115 return None;
116 }
117 }
118 let scan_indices = ordered_union_indices(all_columns.len(), required)?;
119 let key_index_in_scan = scan_indices.iter().position(|index| *index == key_index)?;
120 let scan_columns = scan_indices
121 .iter()
122 .map(|index| all_columns[*index].clone())
123 .collect();
124 Some(NarrowKeyStreamPlan {
125 scan_indices,
126 scan_columns,
127 key_index_in_scan,
128 })
129}
130
131pub fn filtered_expression_projection(
132 filter: &Expression,
133 select: &[Expression],
134 output_columns: &[String],
135 all_columns: &[String],
136) -> Option<ProjectionScanPlan> {
137 dependency_projection(
138 std::iter::once(filter).chain(select.iter()),
139 output_columns,
140 all_columns,
141 )
142}
143
144pub fn expression_projection(
145 select: &[Expression],
146 output_columns: &[String],
147 all_columns: &[String],
148) -> Option<ProjectionScanPlan> {
149 dependency_projection(select.iter(), output_columns, all_columns)
150}
151
152pub fn ordered_distinct_projection(
153 filter: Option<&Expression>,
154 statement: &SelectStatement,
155 output_columns: &[String],
156 all_columns: &[String],
157) -> Option<ProjectionScanPlan> {
158 let expressions = filter
159 .into_iter()
160 .chain(statement.columns.iter())
161 .chain(statement.order_by.iter().map(|order| &order.expression))
162 .chain(statement.distinct_on.iter());
163 dependency_projection(expressions, output_columns, all_columns)
164}
165
166fn dependency_projection<'a>(
167 expressions: impl IntoIterator<Item = &'a Expression>,
168 output_columns: &[String],
169 all_columns: &[String],
170) -> Option<ProjectionScanPlan> {
171 if all_columns.is_empty() {
172 return None;
173 }
174 let columns = build_column_index_map(all_columns);
175 let mut required = Vec::new();
176 for expression in expressions {
177 if !collect_expression_columns(expression, &columns, &mut required) {
178 return None;
179 }
180 }
181 let scan_indices = ordered_union_indices(all_columns.len(), required)?;
182 if scan_indices.len() == all_columns.len() {
183 return None;
184 }
185 Some(plan_from_indices(
186 scan_indices,
187 all_columns,
188 Vec::new(),
189 output_columns,
190 ))
191}
192
193fn plan_from_indices(
194 scan_indices: Vec<usize>,
195 all_columns: &[String],
196 output_indices_in_scan: Vec<usize>,
197 output_columns: &[String],
198) -> ProjectionScanPlan {
199 let scan_columns = scan_indices
200 .iter()
201 .map(|index| all_columns[*index].clone())
202 .collect();
203 ProjectionScanPlan {
204 scan_indices,
205 scan_columns,
206 output_indices_in_scan,
207 output_columns: output_columns.to_vec(),
208 }
209}
210
211fn ordered_union_indices(
212 column_count: usize,
213 indices: impl IntoIterator<Item = usize>,
214) -> Option<Vec<usize>> {
215 let mut needed = vec![false; column_count];
216 for index in indices {
217 *needed.get_mut(index)? = true;
218 }
219 Some(
220 needed
221 .iter()
222 .enumerate()
223 .filter_map(|(index, needed)| needed.then_some(index))
224 .collect(),
225 )
226}
227
228fn resolve_column(columns: &StringMap<usize>, expression: &Expression) -> Option<usize> {
229 match expression {
230 Expression::Identifier(identifier) => columns.get(identifier.value_lower.as_str()).copied(),
231 Expression::QualifiedIdentifier(identifier) => {
232 let qualified = format!(
233 "{}.{}",
234 identifier.qualifier.value_lower, identifier.name.value_lower
235 );
236 columns
237 .get(qualified.as_str())
238 .or_else(|| columns.get(identifier.name.value_lower.as_str()))
239 .copied()
240 }
241 _ => None,
242 }
243}
244
245fn collect_expression_columns(
246 expression: &Expression,
247 columns: &StringMap<usize>,
248 output: &mut Vec<usize>,
249) -> bool {
250 match expression {
251 Expression::Identifier(_) | Expression::QualifiedIdentifier(_) => {
252 if let Some(index) = resolve_column(columns, expression) {
253 output.push(index);
254 true
255 } else {
256 false
257 }
258 }
259 Expression::Aliased(value) => {
260 collect_expression_columns(&value.expression, columns, output)
261 }
262 Expression::FunctionCall(function) => {
263 function
264 .arguments
265 .iter()
266 .all(|argument| collect_expression_columns(argument, columns, output))
267 && function
268 .filter
269 .as_ref()
270 .is_none_or(|filter| collect_expression_columns(filter, columns, output))
271 && function
272 .order_by
273 .iter()
274 .all(|order| collect_expression_columns(&order.expression, columns, output))
275 }
276 Expression::Infix(value) => {
277 collect_expression_columns(&value.left, columns, output)
278 && collect_expression_columns(&value.right, columns, output)
279 }
280 Expression::Prefix(value) => collect_expression_columns(&value.right, columns, output),
281 Expression::Distinct(value) => collect_expression_columns(&value.expr, columns, output),
282 Expression::In(value) => {
283 collect_expression_columns(&value.left, columns, output)
284 && collect_expression_columns(&value.right, columns, output)
285 }
286 Expression::InHashSet(value) => collect_expression_columns(&value.column, columns, output),
287 Expression::Between(value) => {
288 collect_expression_columns(&value.expr, columns, output)
289 && collect_expression_columns(&value.lower, columns, output)
290 && collect_expression_columns(&value.upper, columns, output)
291 }
292 Expression::Like(value) => {
293 collect_expression_columns(&value.left, columns, output)
294 && collect_expression_columns(&value.pattern, columns, output)
295 && value
296 .escape
297 .as_ref()
298 .is_none_or(|escape| collect_expression_columns(escape, columns, output))
299 }
300 Expression::List(value) => value
301 .elements
302 .iter()
303 .all(|item| collect_expression_columns(item, columns, output)),
304 Expression::ExpressionList(value) => value
305 .expressions
306 .iter()
307 .all(|item| collect_expression_columns(item, columns, output)),
308 Expression::Case(value) => {
309 value
310 .value
311 .as_ref()
312 .is_none_or(|item| collect_expression_columns(item, columns, output))
313 && value.when_clauses.iter().all(|when| {
314 collect_expression_columns(&when.condition, columns, output)
315 && collect_expression_columns(&when.then_result, columns, output)
316 })
317 && value
318 .else_value
319 .as_ref()
320 .is_none_or(|item| collect_expression_columns(item, columns, output))
321 }
322 Expression::Cast(value) => collect_expression_columns(&value.expr, columns, output),
323 Expression::IntegerLiteral(_)
324 | Expression::FloatLiteral(_)
325 | Expression::StringLiteral(_)
326 | Expression::BooleanLiteral(_)
327 | Expression::NullLiteral(_)
328 | Expression::IntervalLiteral(_)
329 | Expression::BoundValue(_)
330 | Expression::Parameter(_)
331 | Expression::Default(_) => true,
332 Expression::Star(_)
333 | Expression::QualifiedStar(_)
334 | Expression::AllAny(_)
335 | Expression::Exists(_)
336 | Expression::ScalarSubquery(_)
337 | Expression::Window(_)
338 | Expression::TableSource(_)
339 | Expression::JoinSource(_)
340 | Expression::SubquerySource(_)
341 | Expression::ValuesSource(_)
342 | Expression::CteReference(_)
343 | Expression::FunctionTableSource(_) => false,
344 }
345}
346
347#[cfg(test)]
348mod tests {
349 use super::*;
350 use radixdb_sql::parse_sql;
351
352 fn select(sql: &str) -> SelectStatement {
353 let mut statements = parse_sql(sql).unwrap();
354 match statements.remove(0) {
355 radixdb_sql::ast::Statement::Select(statement) => statement,
356 _ => panic!("expected SELECT"),
357 }
358 }
359
360 #[test]
361 fn filtered_projection_reads_union_in_source_order() {
362 let statement = select("SELECT c, a FROM t WHERE b = 1");
363 let all = vec!["a".into(), "b".into(), "c".into(), "unused".into()];
364 let (output, names) = simple_projection_indices(&statement.columns, &all).unwrap();
365 let plan = filtered_simple_projection(
366 statement.where_clause.as_deref().unwrap(),
367 &output,
368 &names,
369 &all,
370 )
371 .unwrap();
372 assert_eq!(plan.scan_indices, [0, 1, 2]);
373 assert_eq!(plan.output_indices_in_scan, [2, 0]);
374 }
375
376 #[test]
377 fn constant_projection_can_request_exact_empty_scan() {
378 let statement = select("SELECT 42 FROM t");
379 let all = vec!["a".into(), "b".into()];
380 let plan = expression_projection(&statement.columns, &["expr1".into()], &all).unwrap();
381 assert!(plan.scan_indices.is_empty());
382 }
383}