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;
#[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,
}
#[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 {
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,
});
}
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 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(),
});
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;
if !order_by.is_empty() || limit.is_some() || offset.is_some() {
return Err(ParseError::new(
ParseErrorKind::Unsupported(
"ORDER BY and LIMIT are not allowed on a recursive CTE",
),
span,
));
}
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(),
});
let seed = self.bind_isolated_arm(first)?;
let table = subquery_table(&alias, &cte.columns, &seed);
if !cte.columns.is_empty() && cte.columns.len() != seed.columns.len() {
return Err(ParseError::new(
ParseErrorKind::Unsupported("the named column list does not match the query"),
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 referenced {
steps.push((*op, bound));
} else {
seeds.push((*op, bound));
}
}
self.recursing.pop();
outcome?;
let mut source = BoundSource {
index_hint: crate::bind::IndexChoice::Any,
id,
rows: SourceRows::Recursive(Box::new(RecursiveBody { seeds, steps })),
table: std::rc::Rc::new(table),
alias,
join,
constraint: None,
suppressed: Vec::new(),
index_exprs: Vec::new(),
};
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;
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(())
}
}