use crate::sql::expression::identifier::SimpleIdentifier;
use crate::sql::query::{SQLQuery, TableReference};
use crate::write_utils::{blank_line_or_space, maybe_newline, maybe_pad, newline_or_space, Indent};
use ordermap::OrderMap;
use std::fmt;
use std::fmt::{Display, Formatter};
use thiserror::Error;
#[derive(Debug, Error)]
pub enum CTEMergeError {
#[error(
"CTE name conflict: '{0}' already exists. Auto-generated CTE names should not shadow existing CTEs."
)]
NameConflict(SimpleIdentifier),
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct CTE {
pub expressions: OrderMap<SimpleIdentifier, Box<SQLQuery>>,
}
impl CTE {
pub fn with(mut self, name: SimpleIdentifier, query: SQLQuery) -> Self {
self.expressions.insert(name, Box::new(query));
self
}
pub fn merge(mut self, other: CTE) -> Result<Self, CTEMergeError> {
for (name, query) in other.expressions {
if self.expressions.contains_key(&name) {
return Err(CTEMergeError::NameConflict(name));
}
self.expressions.insert(name, query);
}
Ok(self)
}
pub fn get_names(&self) -> Vec<&SimpleIdentifier> {
self.expressions.keys().collect()
}
pub fn get_table_references(&self) -> Vec<TableReference> {
self.expressions
.values()
.flat_map(|q| q.get_table_references().into_iter())
.collect()
}
pub fn fmt_indented(&self, f: &mut Formatter<'_>, indentation: Indent) -> fmt::Result {
let expr_count = self.expressions.len();
if expr_count == 0 {
return Ok(());
}
write!(f, "WITH")?;
newline_or_space(f, indentation)?;
for (i, (k, v)) in self.expressions.iter().enumerate() {
maybe_pad(f, indentation.nested())?;
k.fmt_indented(f, indentation.nested())?;
write!(f, " AS (")?;
maybe_newline(f, indentation)?;
maybe_pad(f, indentation.nested().nested())?;
v.fmt_indented(f, indentation.nested().nested())?;
maybe_newline(f, indentation)?;
maybe_pad(f, indentation.nested())?;
write!(f, ")")?;
if i != expr_count - 1 {
write!(f, ",")?;
}
blank_line_or_space(f, indentation)?;
}
Ok(())
}
}
impl Display for CTE {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
self.fmt_indented(f, Indent::default())
}
}