use super::{subquery_table, unsupported, Binder, BoundSource, RecursiveBody, SourceRows};
use crate::ast::{self, CompoundOp, JoinKind, SelectId};
use crate::catalog_view::TableInfo;
use crate::diagnostic::{ParseError, ParseErrorKind};
use crate::lexer::Span;
pub const FIRST_ANONYMOUS_SHARED: usize = 1 << 20;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct CteBinding {
pub folded: Vec<u8>,
pub name: Vec<u8>,
pub columns: Vec<Vec<u8>>,
pub select: SelectId,
pub recursive: bool,
pub materialized: Option<bool>,
pub level: usize,
}
#[derive(Clone, Debug)]
pub(super) struct RecursiveTarget {
pub(super) folded: Vec<u8>,
id: usize,
table: TableInfo,
referenced: bool,
}
impl Binder<'_> {
pub(crate) fn push_ctes(&mut self, with: &ast::With) -> Result<bool, ParseError> {
if with.ctes.is_empty() {
return Ok(false);
}
let mut bindings = Vec::with_capacity(with.ctes.len());
for cte in &with.ctes {
let folded = self.ast.folded(cte.name);
if bindings
.iter()
.any(|held: &CteBinding| held.folded.as_slice() == folded)
{
return Err(super::refused(
format!(
"duplicate WITH table name: {}",
String::from_utf8_lossy(self.ast.text(cte.name))
),
crate::lexer::Span::default(),
));
}
bindings.push(CteBinding {
folded: self.ast.folded(cte.name).to_vec(),
name: self.ast.text(cte.name).to_vec(),
columns: cte
.columns
.iter()
.map(|name| self.ast.text(*name).to_vec())
.collect(),
select: cte.select,
recursive: with.recursive,
materialized: cte.materialized,
level: self.ctes.len(),
});
}
self.ctes.push(bindings);
Ok(true)
}
pub(crate) fn pop_ctes(&mut self) {
self.ctes.pop();
}
pub(super) fn find_cte(&self, folded: &[u8]) -> Option<CteBinding> {
for level in self.ctes.iter().rev() {
if let Some(found) = level.iter().find(|cte| cte.folded == folded) {
return Some(found.clone());
}
}
None
}
pub(super) fn name_is_used_twice(&self, folded: &[u8]) -> bool {
self.name_uses(folded) > 1
}
pub(super) fn name_uses(&self, folded: &[u8]) -> u32 {
let mut uses = 0u32;
for at in 0..self.ast.from_term_count() {
let Some(term) = self.ast.from_term(ast::FromTermId(at as u32)) else {
continue;
};
if let ast::FromSource::Table {
database: None,
name,
..
} = &term.source
{
if self.ast.folded(*name) == folded {
uses = uses.saturating_add(1);
}
}
}
uses
}
pub(super) fn share_last_source(&mut self, cte: &CteBinding) {
if cte.materialized == Some(false) || !self.name_is_used_twice(&cte.folded) {
return;
}
let key = match self.shared_ctes.iter().position(|(arena, select)| {
*arena == self.ast as *const _ as usize && *select == cte.select
}) {
Some(key) => key,
None => {
self.shared_ctes
.push((self.ast as *const _ as usize, cte.select));
self.shared_ctes.len().saturating_sub(1)
}
};
let Some(source) = self.sources.last_mut() else {
return;
};
let SourceRows::Subquery(block) = &mut source.rows else {
return;
};
if !block.correlations.is_empty() {
return;
}
let mut volatile = false;
let mut probe = (**block).clone();
crate::rewrite::rewrite_select(&mut probe, &mut |expr: &mut super::BoundExpr| {
if crate::plan::calls_a_volatile_function(expr) {
volatile = true;
}
});
if volatile {
block.shared = Some(key);
}
}
pub(super) fn share_uncorrelated_sources(&mut self, block: &mut super::BoundSelect) {
for source in &mut block.sources {
let SourceRows::Subquery(inner) = &mut source.rows else {
continue;
};
if inner.correlations.is_empty() {
if inner.shared.is_none() {
inner.shared = Some(FIRST_ANONYMOUS_SHARED + self.shared_anonymous);
self.shared_anonymous = self.shared_anonymous.saturating_add(1);
}
} else {
self.share_uncorrelated_sources(inner);
}
}
for (_, arm) in &mut block.compounds {
self.share_uncorrelated_sources(arm);
}
}
pub(super) fn select_names_itself(&self, select: ast::SelectId, folded: &[u8]) -> bool {
let Some(query) = self.ast.select(select) else {
return false;
};
if query
.with
.ctes
.iter()
.any(|inner| self.ast.folded(inner.name) == folded)
{
return false;
}
if self.core_names_cte(query.first, folded) {
return true;
}
query
.compounds
.iter()
.any(|(_, arm)| self.core_names_cte(*arm, folded))
}
pub(super) fn core_names_cte(&self, core: ast::SelectCoreId, folded: &[u8]) -> bool {
let Some(arm) = self.ast.core(core) else {
return false;
};
let ast::SelectBody::Select { from, .. } = &arm.body else {
return false;
};
self.terms_name_cte(from, folded)
}
pub(super) fn terms_name_cte(&self, terms: &[ast::FromTermId], folded: &[u8]) -> bool {
terms.iter().any(|id| match self.ast.from_term(*id) {
Some(term) => match &term.source {
ast::FromSource::Table { database, name, .. } => {
database.is_none() && self.ast.folded(*name) == folded
}
ast::FromSource::Subquery(select) => self.select_names_itself(*select, folded),
ast::FromSource::Join(inner) => self.terms_name_cte(inner, folded),
},
None => false,
})
}
pub(super) fn push_recursive_self(
&mut self,
position: usize,
alias: Option<ast::NameId>,
join: JoinKind,
) -> Result<(), ParseError> {
let Some(target) = self.recursing.get_mut(position) else {
return Err(unsupported("unknown recursive reference", Span::default()));
};
target.referenced = true;
let cte = target.id;
let table = target.table.clone();
let alias = match alias {
Some(alias) => self.ast.text(alias).to_vec(),
None => table.name.clone(),
};
let id = self.sources.len();
self.sources.push(BoundSource {
index_hint: crate::bind::IndexChoice::Any,
id,
rows: SourceRows::RecursiveSelf { cte },
table: std::rc::Rc::new(table),
alias,
join,
constraint: None,
suppressed: Vec::new(),
index_exprs: Vec::new(),
written_schema: None,
derived: Default::default(),
});
if let Some(scope) = self.scopes.last_mut() {
scope.push(id);
}
Ok(())
}
pub(super) fn bind_recursive_cte(
&mut self,
cte: &CteBinding,
alias: Vec<u8>,
join: JoinKind,
span: Span,
) -> Result<(), ParseError> {
let Some(select) = self.ast.select(cte.select) else {
return Err(unsupported("missing select", span));
};
if select.compounds.is_empty() {
return self.bind_subquery_term(
cte.select,
Some(alias),
cte.columns.clone(),
join,
span,
);
}
let arms: Vec<(CompoundOp, ast::SelectCoreId)> = select.compounds.clone();
let order_by = select.order_by.clone();
let limit = select.limit;
let offset = select.offset;
let first = select.first;
let id = self.sources.len();
self.sources.push(BoundSource {
index_hint: crate::bind::IndexChoice::Any,
id,
rows: SourceRows::Table,
table: std::rc::Rc::new(TableInfo::subquery(alias.clone(), 0, Vec::new())),
alias: alias.clone(),
join,
constraint: None,
suppressed: Vec::new(),
index_exprs: Vec::new(),
written_schema: None,
derived: Default::default(),
});
let seed = self.bind_isolated_arm(first)?;
let table = subquery_table(&alias, &cte.columns, &seed);
named_columns_fit(&alias, &cte.columns, seed.columns.len(), span)?;
self.recursing.push(RecursiveTarget {
folded: cte.folded.clone(),
id,
table: table.clone(),
referenced: false,
});
let mut seeds = vec![(CompoundOp::UnionAll, seed)];
let mut steps = Vec::new();
let mut outcome = Ok(());
for (op, arm) in &arms {
if !matches!(op, CompoundOp::Union | CompoundOp::UnionAll) {
outcome = Err(ParseError::new(
ParseErrorKind::Unsupported("recursive query does not use UNION or UNION ALL"),
span,
));
break;
}
if let Some(target) = self.recursing.last_mut() {
target.referenced = false;
}
let bound = match self.bind_isolated_arm(*arm) {
Ok(bound) => bound,
Err(reason) => {
outcome = Err(reason);
break;
}
};
let referenced = self
.recursing
.last()
.is_some_and(|target| target.referenced);
if let Err(refusal) =
arm_refusal(&bound, table.columns.len(), referenced, *op, &alias, span)
{
outcome = Err(refusal);
break;
}
if referenced {
steps.push((*op, bound));
} else {
seeds.push((*op, bound));
}
}
self.recursing.pop();
outcome?;
let seed_columns = seeds
.first()
.map_or_else(Vec::new, |(_, seed)| seed.columns.clone());
let other_arms: Vec<(ast::CompoundOp, crate::bind::BoundSelect)> =
seeds.iter().skip(1).chain(steps.iter()).cloned().collect();
let order_by = self.bind_compound_order_by(&order_by, &seed_columns, &other_arms)?;
let limit = limit.map(|expr| self.bind_expr(expr)).transpose()?;
let offset = offset.map(|expr| self.bind_expr(expr)).transpose()?;
let mut source = BoundSource {
index_hint: crate::bind::IndexChoice::Any,
id,
rows: SourceRows::Recursive(Box::new(RecursiveBody {
seeds,
steps,
order_by,
limit,
offset,
})),
table: std::rc::Rc::new(table),
alias,
join,
constraint: None,
suppressed: Vec::new(),
index_exprs: Vec::new(),
written_schema: None,
derived: Default::default(),
};
if let SourceRows::Recursive(body) = &mut source.rows {
if body.steps.is_empty() {
let mut arms = core::mem::take(&mut body.seeds);
if arms.is_empty() {
return Err(unsupported("missing select core", span));
}
let mut head = arms.remove(0).1;
head.compounds = arms;
head.order_by = core::mem::take(&mut body.order_by);
head.limit = body.limit.take();
head.offset = body.offset.take();
source.rows = SourceRows::Subquery(Box::new(head));
}
}
if let Some(slot) = self.sources.get_mut(id) {
*slot = source;
}
if let Some(scope) = self.scopes.last_mut() {
scope.push(id);
}
Ok(())
}
}
impl<'a> Binder<'a> {
pub(super) fn bind_cte_term(
&mut self,
cte: CteBinding,
folded: &[u8],
alias: Option<ast::NameId>,
join: JoinKind,
span: Span,
) -> Result<(), ParseError> {
let alias = match alias {
Some(alias) => self.ast.text(alias).to_vec(),
None => cte.name.clone(),
};
if self.binding_ctes.contains(&cte.select) {
let _ = span;
return Err(ParseError::new(
ParseErrorKind::Refused(format!(
"circular reference: {}",
String::from_utf8_lossy(&cte.name)
)),
Span::default(),
));
}
self.binding_ctes.push(cte.select);
let hidden = self
.ctes
.split_off(cte.level.saturating_add(1).min(self.ctes.len()));
let outcome = if cte.recursive || self.select_names_itself(cte.select, folded) {
self.bind_recursive_cte(&cte, alias, join, span)
} else {
let bound =
self.bind_subquery_term(cte.select, Some(alias), cte.columns.clone(), join, span);
if bound.is_ok() {
self.share_last_source(&cte);
}
bound
};
self.ctes.extend(hidden);
self.binding_ctes.pop();
let own = match self.ast.select(cte.select) {
Some(query) => core::iter::once(query.first)
.chain(query.compounds.iter().map(|(_, arm)| *arm))
.filter(|arm| self.core_names_cte(*arm, &cte.folded))
.count() as u32,
None => 0,
};
let uses = self.name_uses(&cte.folded).saturating_sub(own);
let id = self.scope().last().copied();
if let Some(source) = id.and_then(|id| self.sources.get_mut(id)) {
source.derived = super::DerivedNote {
cte: true,
materialized: cte.materialized,
uses,
name: cte.name.clone(),
..super::DerivedNote::default()
};
}
outcome
}
pub(super) fn find_term_table(
&self,
database: Option<ast::NameId>,
database_name: Option<Vec<u8>>,
folded: &[u8],
) -> (Option<&'a TableInfo>, Option<Vec<u8>>) {
let found = self.catalog.find_table(database_name.as_deref(), folded);
if found.is_none() && database.is_none() && database_name.is_some() {
let eponymous = self
.catalog
.find_table(None, folded)
.filter(|table| table.kind == crate::catalog_view::TableKind::Virtual);
if eponymous.is_some() {
return (eponymous, None);
}
}
(found, database_name)
}
}
fn named_columns_fit<T>(
alias: &[u8],
named: &[T],
width: usize,
span: Span,
) -> Result<(), ParseError> {
if named.is_empty() || named.len() == width {
return Ok(());
}
Err(super::refusal::named_column_count(
alias,
width,
named.len(),
span,
))
}
fn arm_refusal(
arm: &crate::bind::BoundSelect,
width: usize,
recursive: bool,
op: CompoundOp,
alias: &[u8],
span: Span,
) -> Result<(), ParseError> {
if arm.columns.len() != width {
return Err(super::refusal::compound_width_mismatch(op, span));
}
match recursive
.then(|| recursive_step_refusal(arm, alias))
.flatten()
{
Some(reason) => Err(super::refused(reason, span)),
None => Ok(()),
}
}
fn recursive_step_refusal(step: &crate::bind::BoundSelect, alias: &[u8]) -> Option<String> {
if !step.aggregates.is_empty() || !step.group_by.is_empty() {
return Some("recursive aggregate queries not supported".to_string());
}
if !step.windows.is_empty() {
return Some("cannot use window functions in recursive queries".to_string());
}
let references = step
.sources
.iter()
.filter(|source| matches!(source.rows, SourceRows::RecursiveSelf { .. }))
.count();
(references > 1).then(|| {
format!(
"multiple references to recursive table: {}",
String::from_utf8_lossy(alias)
)
})
}