1use super::{subquery_table, unsupported, Binder, BoundSource, RecursiveBody, SourceRows};
19use crate::ast::{self, CompoundOp, JoinKind, SelectId};
20use crate::catalog_view::TableInfo;
21use crate::diagnostic::{ParseError, ParseErrorKind};
22use crate::lexer::Span;
23
24pub const FIRST_ANONYMOUS_SHARED: usize = 1 << 20;
27
28#[derive(Clone, Debug, PartialEq, Eq)]
36pub struct CteBinding {
37 pub folded: Vec<u8>,
39 pub name: Vec<u8>,
41 pub columns: Vec<Vec<u8>>,
43 pub select: SelectId,
45 pub recursive: bool,
47 pub materialized: Option<bool>,
49}
50
51#[derive(Clone, Debug)]
53pub(super) struct RecursiveTarget {
54 pub(super) folded: Vec<u8>,
56 id: usize,
58 table: TableInfo,
60 referenced: bool,
62}
63
64impl Binder<'_> {
65 pub(crate) fn push_ctes(&mut self, with: &ast::With) -> Result<bool, ParseError> {
67 if with.ctes.is_empty() {
68 return Ok(false);
69 }
70 let mut bindings = Vec::with_capacity(with.ctes.len());
71 for cte in &with.ctes {
72 let folded = self.ast.folded(cte.name);
76 if bindings
77 .iter()
78 .any(|held: &CteBinding| held.folded.as_slice() == folded)
79 {
80 return Err(super::refused(
81 format!(
82 "duplicate WITH table name: {}",
83 String::from_utf8_lossy(self.ast.text(cte.name))
84 ),
85 crate::lexer::Span::default(),
86 ));
87 }
88 bindings.push(CteBinding {
89 folded: self.ast.folded(cte.name).to_vec(),
90 name: self.ast.text(cte.name).to_vec(),
91 columns: cte
92 .columns
93 .iter()
94 .map(|name| self.ast.text(*name).to_vec())
95 .collect(),
96 select: cte.select,
97 recursive: with.recursive,
98 materialized: cte.materialized,
99 });
100 }
101 self.ctes.push(bindings);
102 Ok(true)
103 }
104
105 pub(crate) fn pop_ctes(&mut self) {
107 self.ctes.pop();
108 }
109
110 pub(super) fn find_cte(&self, folded: &[u8]) -> Option<CteBinding> {
112 for level in self.ctes.iter().rev() {
113 if let Some(found) = level.iter().find(|cte| cte.folded == folded) {
114 return Some(found.clone());
115 }
116 }
117 None
118 }
119
120 pub(super) fn name_is_used_twice(&self, folded: &[u8]) -> bool {
128 self.name_uses(folded) > 1
129 }
130
131 pub(super) fn name_uses(&self, folded: &[u8]) -> u32 {
138 let mut uses = 0u32;
139 for at in 0..self.ast.from_term_count() {
140 let Some(term) = self.ast.from_term(ast::FromTermId(at as u32)) else {
141 continue;
142 };
143 if let ast::FromSource::Table {
144 database: None,
145 name,
146 ..
147 } = &term.source
148 {
149 if self.ast.folded(*name) == folded {
150 uses = uses.saturating_add(1);
151 }
152 }
153 }
154 uses
155 }
156
157 pub(super) fn share_last_source(&mut self, cte: &CteBinding) {
168 if cte.materialized == Some(false) || !self.name_is_used_twice(&cte.folded) {
169 return;
170 }
171 let key = match self.shared_ctes.iter().position(|(arena, select)| {
172 *arena == self.ast as *const _ as usize && *select == cte.select
173 }) {
174 Some(key) => key,
175 None => {
176 self.shared_ctes
177 .push((self.ast as *const _ as usize, cte.select));
178 self.shared_ctes.len().saturating_sub(1)
179 }
180 };
181 let Some(source) = self.sources.last_mut() else {
182 return;
183 };
184 let SourceRows::Subquery(block) = &mut source.rows else {
185 return;
186 };
187 if !block.correlations.is_empty() {
188 return;
189 }
190 let mut volatile = false;
191 let mut probe = (**block).clone();
192 crate::rewrite::rewrite_select(&mut probe, &mut |expr: &mut super::BoundExpr| {
193 if crate::plan::calls_a_volatile_function(expr) {
194 volatile = true;
195 }
196 });
197 if volatile {
198 block.shared = Some(key);
199 }
200 }
201
202 pub(super) fn share_uncorrelated_sources(&mut self, block: &mut super::BoundSelect) {
212 for source in &mut block.sources {
213 let SourceRows::Subquery(inner) = &mut source.rows else {
214 continue;
215 };
216 if inner.correlations.is_empty() {
217 if inner.shared.is_none() {
218 inner.shared = Some(FIRST_ANONYMOUS_SHARED + self.shared_anonymous);
219 self.shared_anonymous = self.shared_anonymous.saturating_add(1);
220 }
221 } else {
222 self.share_uncorrelated_sources(inner);
223 }
224 }
225 for (_, arm) in &mut block.compounds {
226 self.share_uncorrelated_sources(arm);
227 }
228 }
229
230 pub(super) fn select_names_itself(&self, select: ast::SelectId, folded: &[u8]) -> bool {
245 let Some(query) = self.ast.select(select) else {
246 return false;
247 };
248 if query
249 .with
250 .ctes
251 .iter()
252 .any(|inner| self.ast.folded(inner.name) == folded)
253 {
254 return false;
255 }
256 if self.core_names_cte(query.first, folded) {
257 return true;
258 }
259 query
260 .compounds
261 .iter()
262 .any(|(_, arm)| self.core_names_cte(*arm, folded))
263 }
264
265 pub(super) fn core_names_cte(&self, core: ast::SelectCoreId, folded: &[u8]) -> bool {
270 let Some(arm) = self.ast.core(core) else {
271 return false;
272 };
273 let ast::SelectBody::Select { from, .. } = &arm.body else {
274 return false;
275 };
276 self.terms_name_cte(from, folded)
277 }
278
279 pub(super) fn terms_name_cte(&self, terms: &[ast::FromTermId], folded: &[u8]) -> bool {
284 terms.iter().any(|id| match self.ast.from_term(*id) {
285 Some(term) => match &term.source {
286 ast::FromSource::Table { database, name, .. } => {
287 database.is_none() && self.ast.folded(*name) == folded
288 }
289 ast::FromSource::Subquery(select) => self.select_names_itself(*select, folded),
290 ast::FromSource::Join(inner) => self.terms_name_cte(inner, folded),
291 },
292 None => false,
293 })
294 }
295
296 pub(super) fn push_recursive_self(
298 &mut self,
299 position: usize,
300 alias: Option<ast::NameId>,
301 join: JoinKind,
302 ) -> Result<(), ParseError> {
303 let Some(target) = self.recursing.get_mut(position) else {
304 return Err(unsupported("unknown recursive reference", Span::default()));
305 };
306 target.referenced = true;
307 let cte = target.id;
308 let table = target.table.clone();
309 let alias = match alias {
310 Some(alias) => self.ast.text(alias).to_vec(),
311 None => table.name.clone(),
312 };
313 let id = self.sources.len();
314 self.sources.push(BoundSource {
315 index_hint: crate::bind::IndexChoice::Any,
316 id,
317 rows: SourceRows::RecursiveSelf { cte },
318 table: std::rc::Rc::new(table),
319 alias,
320 join,
321 constraint: None,
322 suppressed: Vec::new(),
323 index_exprs: Vec::new(),
324 written_schema: None,
325 derived: Default::default(),
326 });
327 if let Some(scope) = self.scopes.last_mut() {
328 scope.push(id);
329 }
330 Ok(())
331 }
332
333 pub(super) fn bind_recursive_cte(
342 &mut self,
343 cte: &CteBinding,
344 alias: Vec<u8>,
345 join: JoinKind,
346 span: Span,
347 ) -> Result<(), ParseError> {
348 let Some(select) = self.ast.select(cte.select) else {
349 return Err(unsupported("missing select", span));
350 };
351 if select.compounds.is_empty() {
352 return self.bind_subquery_term(
353 cte.select,
354 Some(alias),
355 cte.columns.clone(),
356 join,
357 span,
358 );
359 }
360 let arms: Vec<(CompoundOp, ast::SelectCoreId)> = select.compounds.clone();
361 let order_by = select.order_by.clone();
362 let limit = select.limit;
363 let offset = select.offset;
364 let first = select.first;
365
366 let id = self.sources.len();
367 self.sources.push(BoundSource {
371 index_hint: crate::bind::IndexChoice::Any,
372 id,
373 rows: SourceRows::Table,
374 table: std::rc::Rc::new(TableInfo::subquery(alias.clone(), 0, Vec::new())),
375 alias: alias.clone(),
376 join,
377 constraint: None,
378 suppressed: Vec::new(),
379 index_exprs: Vec::new(),
380 written_schema: None,
381 derived: Default::default(),
382 });
383
384 let seed = self.bind_isolated_arm(first)?;
385 let table = subquery_table(&alias, &cte.columns, &seed);
386 if !cte.columns.is_empty() && cte.columns.len() != seed.columns.len() {
387 return Err(super::refusal::named_column_count(
388 &alias,
389 seed.columns.len(),
390 cte.columns.len(),
391 span,
392 ));
393 }
394 self.recursing.push(RecursiveTarget {
395 folded: cte.folded.clone(),
396 id,
397 table: table.clone(),
398 referenced: false,
399 });
400 let mut seeds = vec![(CompoundOp::UnionAll, seed)];
401 let mut steps = Vec::new();
402 let mut outcome = Ok(());
403 for (op, arm) in &arms {
404 if !matches!(op, CompoundOp::Union | CompoundOp::UnionAll) {
405 outcome = Err(ParseError::new(
406 ParseErrorKind::Unsupported("recursive query does not use UNION or UNION ALL"),
407 span,
408 ));
409 break;
410 }
411 if let Some(target) = self.recursing.last_mut() {
412 target.referenced = false;
413 }
414 let bound = match self.bind_isolated_arm(*arm) {
415 Ok(bound) => bound,
416 Err(reason) => {
417 outcome = Err(reason);
418 break;
419 }
420 };
421 let referenced = self
422 .recursing
423 .last()
424 .is_some_and(|target| target.referenced);
425 if referenced {
426 steps.push((*op, bound));
427 } else {
428 seeds.push((*op, bound));
429 }
430 }
431 self.recursing.pop();
432 outcome?;
433 let seed_columns = seeds
437 .first()
438 .map_or_else(Vec::new, |(_, seed)| seed.columns.clone());
439 let other_arms: Vec<(ast::CompoundOp, crate::bind::BoundSelect)> =
441 seeds.iter().skip(1).chain(steps.iter()).cloned().collect();
442 let order_by = self.bind_compound_order_by(&order_by, &seed_columns, &other_arms)?;
443 let limit = limit.map(|expr| self.bind_expr(expr)).transpose()?;
444 let offset = offset.map(|expr| self.bind_expr(expr)).transpose()?;
445 let mut source = BoundSource {
446 index_hint: crate::bind::IndexChoice::Any,
447 id,
448 rows: SourceRows::Recursive(Box::new(RecursiveBody {
449 seeds,
450 steps,
451 order_by,
452 limit,
453 offset,
454 })),
455 table: std::rc::Rc::new(table),
456 alias,
457 join,
458 constraint: None,
459 suppressed: Vec::new(),
460 index_exprs: Vec::new(),
461 written_schema: None,
462 derived: Default::default(),
463 };
464 if let SourceRows::Recursive(body) = &mut source.rows {
465 if body.steps.is_empty() {
466 let mut arms = core::mem::take(&mut body.seeds);
469 if arms.is_empty() {
470 return Err(unsupported("missing select core", span));
471 }
472 let mut head = arms.remove(0).1;
473 head.compounds = arms;
474 head.order_by = core::mem::take(&mut body.order_by);
475 head.limit = body.limit.take();
476 head.offset = body.offset.take();
477 source.rows = SourceRows::Subquery(Box::new(head));
478 }
479 }
480 if let Some(slot) = self.sources.get_mut(id) {
481 *slot = source;
482 }
483 if let Some(scope) = self.scopes.last_mut() {
484 scope.push(id);
485 }
486 Ok(())
487 }
488}
489
490impl<'a> Binder<'a> {
491 pub(super) fn bind_cte_term(
499 &mut self,
500 cte: CteBinding,
501 folded: &[u8],
502 alias: Option<ast::NameId>,
503 join: JoinKind,
504 span: Span,
505 ) -> Result<(), ParseError> {
506 let alias = match alias {
507 Some(alias) => self.ast.text(alias).to_vec(),
508 None => cte.name.clone(),
509 };
510 if self.binding_ctes.contains(&cte.select) {
513 let _ = span;
515 return Err(ParseError::new(
516 ParseErrorKind::Refused(format!(
517 "circular reference: {}",
518 String::from_utf8_lossy(&cte.name)
519 )),
520 Span::default(),
521 ));
522 }
523 self.binding_ctes.push(cte.select);
524 let outcome = if cte.recursive || self.select_names_itself(cte.select, folded) {
529 self.bind_recursive_cte(&cte, alias, join, span)
530 } else {
531 let bound =
532 self.bind_subquery_term(cte.select, Some(alias), cte.columns.clone(), join, span);
533 if bound.is_ok() {
534 self.share_last_source(&cte);
535 }
536 bound
537 };
538 self.binding_ctes.pop();
539 let own = match self.ast.select(cte.select) {
542 Some(query) => core::iter::once(query.first)
543 .chain(query.compounds.iter().map(|(_, arm)| *arm))
544 .filter(|arm| self.core_names_cte(*arm, &cte.folded))
545 .count() as u32,
546 None => 0,
547 };
548 let uses = self.name_uses(&cte.folded).saturating_sub(own);
549 let id = self.scope().last().copied();
552 if let Some(source) = id.and_then(|id| self.sources.get_mut(id)) {
553 source.derived = super::DerivedNote {
554 cte: true,
555 materialized: cte.materialized,
556 uses,
557 name: cte.name.clone(),
558 ..super::DerivedNote::default()
559 };
560 }
561 outcome
562 }
563
564 pub(super) fn find_term_table(
574 &self,
575 database: Option<ast::NameId>,
576 database_name: Option<Vec<u8>>,
577 folded: &[u8],
578 ) -> (Option<&'a TableInfo>, Option<Vec<u8>>) {
579 let found = self.catalog.find_table(database_name.as_deref(), folded);
580 if found.is_none() && database.is_none() && database_name.is_some() {
581 let eponymous = self
582 .catalog
583 .find_table(None, folded)
584 .filter(|table| table.kind == crate::catalog_view::TableKind::Virtual);
585 if eponymous.is_some() {
586 return (eponymous, None);
587 }
588 }
589 (found, database_name)
590 }
591}