use std::borrow::Cow;
use crate::error::Error;
use crate::expr::{Expr, IntoExpr};
use crate::writer::{Expression, SqlWriter};
use super::{MaybeAbsent, write_quoted_list};
#[derive(Debug, Clone, Default)]
pub struct Cte {
pub name: Cow<'static, str>,
pub columns: Vec<Cow<'static, str>>,
pub query: Option<Expr>,
pub materialized: Option<bool>,
pub search: CteSearch,
pub cycle: CteCycle,
}
impl Cte {
pub fn new(name: impl Into<Cow<'static, str>>, query: impl IntoExpr) -> Self {
Cte {
name: name.into(),
query: Some(query.into_expr()),
..Cte::default()
}
}
pub fn is_empty(&self) -> bool {
self.name.is_empty() && self.query.is_none()
}
}
impl Expression for Cte {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
if self.is_empty() {
return;
}
let Some(query) = &self.query else {
w.record_error(Error::Incomplete("the query of a CTE"));
return;
};
w.push_quoted(&[&self.name]);
write_quoted_list(w, &self.columns, " (", ", ", ")");
w.push_str(" AS ");
match self.materialized {
None => {}
Some(true) => w.push_str("MATERIALIZED "),
Some(false) => w.push_str("NOT MATERIALIZED "),
}
w.push_str("(");
w.write_expr(query);
w.push_str(")");
w.write_if(!self.search.is_empty(), " ", &self.search, "");
w.write_if(!self.cycle.is_empty(), " ", &self.cycle, "");
}
}
#[derive(Debug, Clone, Default)]
pub struct CteSearch {
pub order: SearchOrder,
pub columns: Vec<Cow<'static, str>>,
pub set: Cow<'static, str>,
}
impl CteSearch {
pub fn new(
order: SearchOrder,
columns: impl IntoIterator<Item = impl Into<Cow<'static, str>>>,
set: impl Into<Cow<'static, str>>,
) -> Self {
CteSearch {
order,
columns: columns.into_iter().map(Into::into).collect(),
set: set.into(),
}
}
pub fn is_empty(&self) -> bool {
self.columns.is_empty()
}
}
impl Expression for CteSearch {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
if self.is_empty() {
return;
}
if self.set.is_empty() {
w.record_error(Error::Incomplete("the SET column of a CTE SEARCH clause"));
return;
}
w.push_str("SEARCH ");
w.push_str(self.order.as_str());
w.push_str(" FIRST BY ");
write_quoted_list(w, &self.columns, "", ", ", "");
w.push_str(" SET ");
w.push_quoted(&[&self.set]);
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub enum SearchOrder {
#[default]
Breadth,
Depth,
}
impl SearchOrder {
pub fn as_str(self) -> &'static str {
match self {
SearchOrder::Breadth => "BREADTH",
SearchOrder::Depth => "DEPTH",
}
}
}
#[derive(Debug, Clone, Default)]
pub struct CteCycle {
pub columns: Vec<Cow<'static, str>>,
pub set: Cow<'static, str>,
pub using: Cow<'static, str>,
pub to: Option<Expr>,
pub default_val: Option<Expr>,
}
impl CteCycle {
pub fn new(
columns: impl IntoIterator<Item = impl Into<Cow<'static, str>>>,
set: impl Into<Cow<'static, str>>,
using: impl Into<Cow<'static, str>>,
) -> Self {
CteCycle {
columns: columns.into_iter().map(Into::into).collect(),
set: set.into(),
using: using.into(),
..CteCycle::default()
}
}
pub fn is_empty(&self) -> bool {
self.columns.is_empty()
}
}
impl Expression for CteCycle {
fn write_sql(&self, w: &mut SqlWriter<'_>) {
if self.is_empty() {
return;
}
if self.set.is_empty() || self.using.is_empty() {
w.record_error(Error::Incomplete(
"the SET and USING columns of a CTE CYCLE clause",
));
return;
}
if self.to.is_some() != self.default_val.is_some() {
w.record_error(Error::Incomplete(
"both TO and DEFAULT of a CTE CYCLE clause",
));
return;
}
w.push_str("CYCLE ");
write_quoted_list(w, &self.columns, "", ", ", "");
w.push_str(" SET ");
w.push_quoted(&[&self.set]);
if let (Some(to), Some(default_val)) = (&self.to, &self.default_val) {
w.push_str(" TO ");
w.write_expr(to);
w.push_str(" DEFAULT ");
w.write_expr(default_val);
}
w.push_str(" USING ");
w.push_quoted(&[&self.using]);
}
}
impl MaybeAbsent for Cte {
fn is_absent(&self) -> bool {
self.is_empty()
}
}
#[cfg(test)]
mod tests {
use keelson_sqlcheck::testing::assert_frag_sql;
use super::*;
use crate::dialect::testing::Numbered;
use crate::expr::arg;
use crate::value::Value;
use crate::writer::build;
const FRAME: &str = r#"WITH {} SELECT * FROM "c""#;
const RECURSIVE_FRAME: &str = r#"WITH RECURSIVE {} SELECT * FROM "c""#;
const AFTER_RECURSIVE_CTE: &str = concat!(
r#"WITH RECURSIVE "c" AS ("#,
r#"SELECT 1 AS "id" UNION ALL SELECT "id" + 1 FROM "c" WHERE "id" < 5"#,
r#") {} SELECT * FROM "c""#
);
fn sub() -> Expr {
Expr::join((
Expr::raw(r#"SELECT "id" FROM posts WHERE "id" ="#),
arg(1i32),
))
}
const SUB_SQL: &str = r#"SELECT "id" FROM posts WHERE "id" = $1"#;
fn recursive_sub() -> Expr {
Expr::raw(r#"SELECT 1 AS "id" UNION ALL SELECT "id" + 1 FROM "c" WHERE "id" < 5"#)
}
const RECURSIVE_SUB_SQL: &str =
r#"SELECT 1 AS "id" UNION ALL SELECT "id" + 1 FROM "c" WHERE "id" < 5"#;
fn sql(e: &impl Expression) -> String {
build(&Numbered, e).expect("render").0
}
#[test]
fn an_untouched_cte_writes_nothing() {
assert_eq!(build(&Numbered, &Cte::default()).unwrap().0, "");
assert!(Cte::default().is_empty());
}
#[test]
fn a_bare_cte_is_name_as_query() {
let (rendered, args) = build(&Numbered, &Cte::new("c", sub())).unwrap();
assert_frag_sql(FRAME, &rendered, &format!(r#""c" AS ({SUB_SQL})"#));
assert_eq!(args, vec![Value::I32(1)]);
}
#[test]
fn column_names_follow_the_cte_name() {
let two_cols = Expr::raw(r#"SELECT "id", "title" FROM posts"#);
let cte = Cte {
columns: vec!["id".into(), "data".into()],
..Cte::new("c", two_cols)
};
assert_frag_sql(
FRAME,
&sql(&cte),
r#""c" ("id", "data") AS (SELECT "id", "title" FROM posts)"#,
);
}
#[test]
fn materialisation_is_three_valued() {
let base = Cte::new("c", sub());
assert_frag_sql(FRAME, &sql(&base), &format!(r#""c" AS ({SUB_SQL})"#));
let yes = Cte {
materialized: Some(true),
..base.clone()
};
assert_frag_sql(
FRAME,
&sql(&yes),
&format!(r#""c" AS MATERIALIZED ({SUB_SQL})"#),
);
let no = Cte {
materialized: Some(false),
..base
};
assert_frag_sql(
FRAME,
&sql(&no),
&format!(r#""c" AS NOT MATERIALIZED ({SUB_SQL})"#),
);
}
#[test]
fn a_named_cte_with_no_query_is_a_recorded_failure() {
let cte = Cte {
name: "c".into(),
..Cte::default()
};
let err = build(&Numbered, &cte).unwrap_err();
assert!(
matches!(&err, Error::Incomplete(what) if what.contains("CTE")),
"got: {err}"
);
}
#[test]
fn search_and_cycle_follow_the_query_and_hinge_on_their_columns() {
let mut cte = Cte::new("c", recursive_sub());
cte.search = CteSearch::new(SearchOrder::Depth, ["id"], "ordercol");
cte.cycle = CteCycle {
set: "is_cycle".into(),
using: "path".into(),
..CteCycle::default()
};
assert_frag_sql(
RECURSIVE_FRAME,
&sql(&cte),
&format!(r#""c" AS ({RECURSIVE_SUB_SQL}) SEARCH DEPTH FIRST BY "id" SET "ordercol""#),
);
cte.cycle.columns = vec!["id".into()];
assert_frag_sql(
RECURSIVE_FRAME,
&sql(&cte),
&format!(
concat!(
r#""c" AS ({}) SEARCH DEPTH FIRST BY "id" SET "ordercol""#,
r#" CYCLE "id" SET "is_cycle" USING "path""#
),
RECURSIVE_SUB_SQL
),
);
}
#[test]
fn breadth_is_the_default_search_order() {
let search = CteSearch::new(SearchOrder::default(), ["id"], "seq");
assert_frag_sql(
AFTER_RECURSIVE_CTE,
&sql(&search),
r#"SEARCH BREADTH FIRST BY "id" SET "seq""#,
);
}
#[test]
fn a_search_clause_without_its_set_column_is_a_recorded_failure() {
let search = CteSearch {
columns: vec!["id".into()],
..CteSearch::default()
};
let err = build(&Numbered, &search).unwrap_err();
assert!(
matches!(&err, Error::Incomplete(what)
if what.contains("SET") && what.contains("SEARCH")),
"got: {err}"
);
}
#[test]
fn the_cycle_mark_values_are_written_as_one_optional_group() {
let mut cycle = CteCycle::new(["id"], "is_cycle", "path");
cycle.to = Some(Expr::literal("Y"));
cycle.default_val = Some(Expr::literal("N"));
let (rendered, args) = build(&Numbered, &cycle).unwrap();
assert_frag_sql(
AFTER_RECURSIVE_CTE,
&rendered,
r#"CYCLE "id" SET "is_cycle" TO 'Y' DEFAULT 'N' USING "path""#,
);
assert!(args.is_empty(), "a constant binds nothing");
cycle.default_val = None;
let err = build(&Numbered, &cycle).unwrap_err();
assert!(
matches!(&err, Error::Incomplete(what)
if what.contains("TO") && what.contains("DEFAULT") && what.contains("CYCLE")),
"got: {err}"
);
}
#[test]
fn a_cycle_clause_without_its_added_columns_is_a_recorded_failure() {
let cycle = CteCycle {
columns: vec!["id".into()],
..CteCycle::default()
};
let err = build(&Numbered, &cycle).unwrap_err();
assert!(
matches!(&err, Error::Incomplete(what)
if what.contains("SET") && what.contains("USING") && what.contains("CYCLE")),
"got: {err}"
);
}
}