use crate::reader::Reader;
use crate::{naming, parser::SourceTree, GgsqlError, Result};
use std::collections::HashSet;
use tree_sitter::Node;
#[derive(Debug, Clone)]
pub struct CteDefinition {
pub name: String,
pub body: String,
pub column_aliases: Vec<String>,
}
pub fn extract_ctes(source_tree: &SourceTree) -> Vec<CteDefinition> {
let root = source_tree.root();
source_tree
.find_nodes(&root, "(cte_definition) @cte")
.into_iter()
.filter_map(|node| parse_cte_definition(&node, source_tree.source))
.collect()
}
fn parse_cte_definition(node: &Node, source: &str) -> Option<CteDefinition> {
let mut name: Option<String> = None;
let mut column_aliases: Vec<String> = Vec::new();
let mut body_start: Option<usize> = None;
let mut body_end: Option<usize> = None;
let mut cursor = node.walk();
for child in node.children(&mut cursor) {
match child.kind() {
"identifier" => {
if name.is_none() {
name = Some(get_node_text(&child, source).to_string());
} else {
column_aliases.push(get_node_text(&child, source).to_string());
}
}
"select_statement" | "subquery_body" | "with_statement" => {
body_start = Some(child.start_byte());
body_end = Some(child.end_byte());
}
_ => {}
}
}
match (name, body_start, body_end) {
(Some(n), Some(start), Some(end)) => {
let body = source[start..end].to_string();
Some(CteDefinition {
name: n,
body,
column_aliases,
})
}
_ => None,
}
}
pub(crate) fn get_node_text<'a>(node: &Node, source: &'a str) -> &'a str {
&source[node.start_byte()..node.end_byte()]
}
pub fn transform_cte_references(sql: &str, cte_names: &HashSet<String>) -> String {
if cte_names.is_empty() {
return sql.to_string();
}
let Ok(sites) = crate::parser::extract_table_ref_sites(sql) else {
return sql.to_string();
};
let string_ranges = crate::parser::string_literal_ranges(sql);
let in_string = |pos: usize| string_ranges.iter().any(|&(s, e)| pos >= s && pos < e);
let temp_of = |raw: &str| -> Option<String> {
let name = naming::unquote_ident(raw);
cte_names
.iter()
.find(|c| naming::unquote_ident(c).eq_ignore_ascii_case(&name))
.map(|c| naming::quote_ident(&naming::cte_table(c)))
};
let mut replacements: Vec<(usize, usize, String)> = Vec::new();
for site in &sites {
if let Some(temp) = temp_of(&site.raw) {
replacements.push((site.start, site.end, temp));
}
}
let site_starts: HashSet<usize> = sites.iter().map(|s| s.start).collect();
for cte in cte_names {
let temp = naming::quote_ident(&naming::cte_table(cte));
let bare = naming::unquote_ident(cte);
let pattern = format!(r"((?i:{}))\s*\.", regex::escape(&bare));
let Ok(re) = regex::Regex::new(&pattern) else {
continue;
};
for caps in re.captures_iter(sql) {
let name = caps.get(1).unwrap();
let start = name.start();
if site_starts.contains(&start) || in_string(start) {
continue;
}
let ok_prefix = sql[..start]
.chars()
.next_back()
.is_none_or(|c| !(c.is_alphanumeric() || c == '_' || c == '.' || c == '"'));
if ok_prefix {
replacements.push((name.start(), name.end(), temp.clone()));
}
}
}
let mut result = sql.to_string();
replacements.sort_by_key(|(start, _, _)| std::cmp::Reverse(*start));
replacements.dedup_by_key(|(start, _, _)| *start);
for (start, end, replacement) in replacements {
result.replace_range(start..end, &replacement);
}
result
}
fn unquoted_dot_positions(raw: &str) -> Vec<usize> {
let mut positions = Vec::new();
let mut in_quote = false;
for (i, b) in raw.bytes().enumerate() {
match b {
b'"' => in_quote = !in_quote,
b'.' if !in_quote => positions.push(i),
_ => {}
}
}
positions
}
fn identifier_components(raw: &str) -> Vec<&str> {
let mut parts = Vec::new();
let mut start = 0;
for pos in unquoted_dot_positions(raw) {
parts.push(raw[start..pos].trim());
start = pos + 1;
}
parts.push(raw[start..].trim());
parts
}
fn last_identifier_component(raw: &str) -> &str {
match unquoted_dot_positions(raw).last() {
Some(&pos) => raw[pos + 1..].trim(),
None => raw.trim(),
}
}
pub fn transform_source_references(sql: &str, reader: &dyn Reader) -> Result<String> {
if !reader.caches_sources() {
return Ok(sql.to_string());
}
let sites = match crate::parser::extract_table_ref_sites(sql) {
Ok(sites) => sites,
Err(_) => return Ok(sql.to_string()),
};
let is_cache_resolvable = |raw: &str| {
raw.starts_with('\'')
|| raw.starts_with("ggsql:")
|| naming::is_internal_table(&naming::unquote_ident(raw))
};
let has_cache_ref = sites.iter().any(|s| is_cache_resolvable(&s.raw));
let primary_sites: Vec<&crate::parser::TableRefSite> = sites
.iter()
.filter(|s| !is_cache_resolvable(&s.raw))
.collect();
if !has_cache_ref || primary_sites.is_empty() {
return Ok(sql.to_string());
}
let last_of = |raw: &str| -> String { last_identifier_component(raw).to_string() };
let mut staged_for: std::collections::HashMap<String, String> =
std::collections::HashMap::new();
for site in &primary_sites {
if !staged_for.contains_key(&site.raw) {
let t = naming::staged_source_table(staged_for.len());
reader.materialize_table(&t, &[], &format!("SELECT * FROM {}", site.raw))?;
staged_for.insert(site.raw.clone(), t);
}
}
let mut last_counts: std::collections::HashMap<String, usize> =
std::collections::HashMap::new();
let mut seen_unaliased: HashSet<&str> = HashSet::new();
for site in &primary_sites {
if !site.has_alias && seen_unaliased.insert(site.raw.as_str()) {
*last_counts.entry(last_of(&site.raw)).or_default() += 1;
}
}
let last_collides = |raw: &str| last_counts.get(&last_of(raw)).copied().unwrap_or(0) > 1;
let string_ranges = crate::parser::string_literal_ranges(sql);
let in_string = |pos: usize| string_ranges.iter().any(|&(s, e)| pos >= s && pos < e);
let site_starts: HashSet<usize> = primary_sites.iter().map(|s| s.start).collect();
let mut replacements: Vec<(usize, usize, String)> = Vec::new();
for site in &primary_sites {
let quoted = naming::quote_ident(&staged_for[&site.raw]);
let replacement = if site.has_alias || last_collides(&site.raw) {
quoted
} else {
format!("{} AS {}", quoted, last_of(&site.raw))
};
replacements.push((site.start, site.end, replacement));
}
for raw in staged_for.keys() {
if unquoted_dot_positions(raw).is_empty() {
continue;
}
let target = if last_collides(raw) {
naming::quote_ident(&staged_for[raw])
} else {
last_of(raw)
};
let pattern = format!(
r"({})\s*\.",
identifier_components(raw)
.iter()
.map(|c| {
if c.starts_with('"') {
regex::escape(c)
} else {
format!("(?i:{})", regex::escape(c))
}
})
.collect::<Vec<_>>()
.join(r"\s*\.\s*")
);
let Ok(re) = regex::Regex::new(&pattern) else {
continue;
};
for caps in re.captures_iter(sql) {
let qualifier = caps.get(1).unwrap();
let start = qualifier.start();
if site_starts.contains(&start) || in_string(start) {
continue;
}
let ok_prefix = sql[..start]
.chars()
.next_back()
.is_none_or(|c| !(c.is_alphanumeric() || c == '_' || c == '.' || c == '"'));
if ok_prefix {
replacements.push((start, qualifier.end(), target.clone()));
}
}
}
let mut result = sql.to_string();
replacements.sort_by_key(|(start, _, _)| std::cmp::Reverse(*start));
for (start, end, replacement) in replacements {
result.replace_range(start..end, &replacement);
}
Ok(result)
}
pub fn materialize_ctes(ctes: &[CteDefinition], reader: &dyn Reader) -> Result<HashSet<String>> {
let mut materialized = HashSet::new();
for cte in ctes {
let transformed_body = transform_cte_references(&cte.body, &materialized);
let transformed_body = transform_source_references(&transformed_body, reader)?;
let temp_table_name = naming::cte_table(&cte.name);
reader
.materialize_table(&temp_table_name, &cte.column_aliases, &transformed_body)
.map_err(|e| {
GgsqlError::ReaderError(format!("Failed to materialize CTE '{}': {}", cte.name, e))
})?;
materialized.insert(cte.name.clone());
}
Ok(materialized)
}
pub fn split_with_query(source_tree: &SourceTree) -> Option<(String, String)> {
let root = source_tree.root();
let with_node = source_tree.find_node(&root, "(with_statement) @with")?;
let mut cursor = with_node.walk();
let mut last_cte_end: Option<usize> = None;
let mut tail_node = None;
let mut seen_cte = false;
for child in with_node.children(&mut cursor) {
match child.kind() {
"cte_definition" => {
seen_cte = true;
last_cte_end = Some(child.end_byte());
}
"select_statement" if seen_cte => {
tail_node = Some((child, false));
break;
}
"from_statement" if seen_cte => {
tail_node = Some((child, true));
break;
}
_ => {}
}
}
let cte_prefix = source_tree.source[with_node.start_byte()..last_cte_end?].to_string();
let (node, is_from) = tail_node?;
let trailing = if is_from {
format!("SELECT * {}", source_tree.get_text(&node))
} else {
source_tree.get_text(&node)
};
Some((cte_prefix, trailing))
}
pub fn extract_side_effects(source_tree: &SourceTree) -> Vec<String> {
let root = source_tree.root();
let side_effect_stmts = r#"
(sql_statement
[(create_statement)
(insert_statement)
(update_statement)
(delete_statement)] @stmt)
"#;
source_tree
.find_texts(&root, side_effect_stmts)
.into_iter()
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.collect()
}
pub fn transform_global_sql(
source_tree: &SourceTree,
materialized_ctes: &HashSet<String>,
reader: &dyn Reader,
) -> Result<Option<String>> {
let root = source_tree.root();
let select_sql = split_with_query(source_tree)
.map(|(_, select)| select)
.or_else(|| {
source_tree.find_text(&root, "(sql_statement (select_statement) @select)")
});
if let Some(select_sql) = select_sql {
let select_sql = transform_cte_references(&select_sql, materialized_ctes);
return Ok(Some(transform_source_references(&select_sql, reader)?));
}
if !has_executable_sql(source_tree) {
return Ok(None);
}
let viz_from_query = source_tree
.find_text(
&root,
r#"(visualise_statement (visualise_from source: (_) @source))"#,
)
.map(|table| {
let q = format!("SELECT * FROM {}", table);
let q = transform_cte_references(&q, materialized_ctes);
transform_source_references(&q, reader)
})
.transpose()?;
if viz_from_query.is_some() || !extract_side_effects(source_tree).is_empty() {
Ok(viz_from_query)
} else {
source_tree
.extract_sql()
.map(|s| {
let s = transform_cte_references(&s, materialized_ctes);
transform_source_references(&s, reader)
})
.transpose()
}
}
pub fn has_executable_sql(source_tree: &SourceTree) -> bool {
let root = source_tree.root();
let direct_statements = r#"
(sql_statement
[(select_statement)
(create_statement)
(insert_statement)
(update_statement)
(delete_statement)
(from_statement)] @stmt)
"#;
if source_tree.find_node(&root, direct_statements).is_some() {
return true;
}
if split_with_query(source_tree).is_some() {
return true;
}
let visualise_from = r#"
(visualise_statement
(visualise_from) @from)
"#;
if source_tree.find_node(&root, visualise_from).is_some() {
return true;
}
false
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_extract_ctes_single() {
let sql = "WITH sales AS (SELECT * FROM raw_sales) SELECT * FROM sales";
let source_tree = SourceTree::new(sql).unwrap();
let ctes = extract_ctes(&source_tree);
assert_eq!(ctes.len(), 1);
assert_eq!(ctes[0].name, "sales");
assert!(ctes[0].body.contains("SELECT * FROM raw_sales"));
}
#[test]
fn test_extract_ctes_multiple() {
let sql = "WITH
sales AS (SELECT * FROM raw_sales),
targets AS (SELECT * FROM goals)
SELECT * FROM sales";
let source_tree = SourceTree::new(sql).unwrap();
let ctes = extract_ctes(&source_tree);
assert_eq!(ctes.len(), 2);
assert_eq!(ctes[0].name, "sales");
assert_eq!(ctes[1].name, "targets");
}
#[test]
fn test_extract_ctes_with_column_aliases() {
let sql = "WITH t(value, label) AS (SELECT * FROM (VALUES (70, 'Target'))) SELECT * FROM t";
let source_tree = SourceTree::new(sql).unwrap();
let ctes = extract_ctes(&source_tree);
assert_eq!(ctes.len(), 1);
assert_eq!(ctes[0].name, "t");
assert_eq!(ctes[0].column_aliases, vec!["value", "label"]);
}
#[test]
fn test_extract_ctes_without_column_aliases() {
let sql = "WITH sales AS (SELECT * FROM raw_sales) SELECT * FROM sales";
let source_tree = SourceTree::new(sql).unwrap();
let ctes = extract_ctes(&source_tree);
assert_eq!(ctes.len(), 1);
assert_eq!(ctes[0].name, "sales");
assert!(ctes[0].column_aliases.is_empty());
}
#[test]
fn test_extract_ctes_none() {
let sql = "SELECT * FROM sales WHERE year = 2024";
let source_tree = SourceTree::new(sql).unwrap();
let ctes = extract_ctes(&source_tree);
assert!(ctes.is_empty());
}
#[test]
fn test_transform_cte_references() {
let test_cases: Vec<(
&str,
Vec<&str>,
Vec<&str>, // strings that should be in result
Option<&str>, // exact match (if result should equal this)
)> = vec![
(
"SELECT * FROM sales WHERE year = 2024",
vec!["sales"],
vec!["FROM \"__ggsql_cte_sales_", "__\" WHERE year = 2024"],
None,
),
(
"SELECT sales.date, targets.revenue FROM sales JOIN targets ON sales.id = targets.id",
vec!["sales", "targets"],
vec![
"FROM \"__ggsql_cte_sales_",
"JOIN \"__ggsql_cte_targets_",
"__ggsql_cte_sales_", "__ggsql_cte_targets_", ],
None,
),
(
"WHERE sales.date > '2024-01-01' AND sales.revenue > 100",
vec!["sales"],
vec!["__ggsql_cte_sales_"],
None,
),
(
"SELECT * FROM other_table",
vec!["sales"],
vec![],
Some("SELECT * FROM other_table"),
),
(
"SELECT * FROM sales",
vec![],
vec![],
Some("SELECT * FROM sales"),
),
(
"SELECT wholesale.date FROM wholesale",
vec!["sales"],
vec![],
Some("SELECT wholesale.date FROM wholesale"),
),
];
for (sql, cte_names_vec, expected_contains, exact_match) in test_cases {
let cte_names: HashSet<String> = cte_names_vec.iter().map(|s| s.to_string()).collect();
let result = transform_cte_references(sql, &cte_names);
if let Some(expected) = exact_match {
assert_eq!(result, expected, "SQL '{}' should remain unchanged", sql);
} else {
for expected in &expected_contains {
assert!(
result.contains(expected),
"Result '{}' should contain '{}' for SQL '{}'",
result,
expected,
sql
);
}
if !cte_names_vec.is_empty() {
assert!(
result.contains(naming::session_id()),
"Result should contain session UUID"
);
}
}
}
}
#[test]
fn test_transform_cte_references_comma_join_second_position() {
let ctes: HashSet<String> = ["cte"].iter().map(|s| s.to_string()).collect();
let out = transform_cte_references("SELECT * FROM base, cte WHERE base.k = cte.k", &ctes);
assert!(
!out.contains("FROM base, cte "),
"cte table ref not rewritten: {out}"
);
assert!(out.contains("__ggsql_cte_cte_"));
assert_eq!(out.matches("__ggsql_cte_cte_").count(), 2);
}
#[test]
fn test_transform_cte_references_preserves_string_literals() {
let ctes: HashSet<String> = ["cte"].iter().map(|s| s.to_string()).collect();
let out = transform_cte_references("SELECT cte.k, 'cte.k' AS lit FROM cte", &ctes);
assert!(out.contains("'cte.k'"), "literal was corrupted: {out}");
assert_eq!(out.matches("__ggsql_cte_cte_").count(), 2);
}
#[test]
fn test_transform_cte_references_whitespace_around_dot() {
let ctes: HashSet<String> = ["cte"].iter().map(|s| s.to_string()).collect();
let out = transform_cte_references("SELECT cte . v FROM cte", &ctes);
assert!(!out.contains("cte . v"), "qualifier not rewritten: {out}");
assert_eq!(out.matches("__ggsql_cte_cte_").count(), 2);
}
#[test]
fn test_transform_cte_references_case_insensitive() {
let ctes: HashSet<String> = ["cte"].iter().map(|s| s.to_string()).collect();
let out = transform_cte_references("SELECT CTE.v FROM CTE", &ctes);
assert_eq!(out.matches("__ggsql_cte_cte_").count(), 2);
}
struct MockReader {
caches: bool,
staged: std::cell::RefCell<Vec<(String, String)>>,
}
impl MockReader {
fn new(caches: bool) -> Self {
Self {
caches,
staged: std::cell::RefCell::new(Vec::new()),
}
}
}
impl Reader for MockReader {
fn execute_sql(&self, _sql: &str) -> Result<crate::DataFrame> {
unreachable!("staging must not touch execute_sql in these tests")
}
fn register(&self, _name: &str, _df: crate::DataFrame, _replace: bool) -> Result<()> {
Ok(())
}
fn execute(&self, _query: &str) -> Result<crate::reader::Spec> {
unreachable!()
}
fn caches_sources(&self) -> bool {
self.caches
}
fn materialize_table(
&self,
name: &str,
_column_aliases: &[String],
body_sql: &str,
) -> Result<()> {
self.staged
.borrow_mut()
.push((name.to_string(), body_sql.to_string()));
Ok(())
}
}
#[test]
fn test_transform_source_references_quoted_primary_ref() {
let reader = MockReader::new(true);
let sql = "SELECT * FROM \"__ggsql_cte_t__\" JOIN \"my base\" ON 1 = 1";
let out = transform_source_references(sql, &reader).unwrap();
let staged = reader.staged.borrow();
assert_eq!(staged.len(), 1);
assert_eq!(staged[0].1, "SELECT * FROM \"my base\"");
assert!(out.contains("__ggsql_staged_0_"));
assert!(!out.contains("JOIN \"my base\""));
}
#[test]
fn test_transform_source_references_non_caching_reader_unchanged() {
let reader = MockReader::new(false);
let sql = "SELECT * FROM \"__ggsql_cte_t__\" JOIN base ON base.k = 1";
let out = transform_source_references(sql, &reader).unwrap();
assert_eq!(out, sql);
assert!(reader.staged.borrow().is_empty());
}
#[test]
fn test_transform_source_references_all_primary_unchanged() {
let reader = MockReader::new(true);
let out = transform_source_references("SELECT * FROM a JOIN base ON a.k = base.k", &reader)
.unwrap();
assert_eq!(out, "SELECT * FROM a JOIN base ON a.k = base.k");
assert!(reader.staged.borrow().is_empty());
}
#[test]
fn test_transform_source_references_stages_mixed_body() {
let reader = MockReader::new(true);
let sql = "SELECT * FROM \"__ggsql_cte_t__\" JOIN base ON base.k = 1";
let out = transform_source_references(sql, &reader).unwrap();
let staged = reader.staged.borrow();
assert_eq!(staged.len(), 1, "base should be staged exactly once");
assert!(staged[0].0.starts_with("__ggsql_staged_0_"));
assert_eq!(staged[0].1, "SELECT * FROM base");
assert!(out.contains("__ggsql_staged_0_"));
assert!(out.contains("\"__ggsql_cte_t__\"")); assert!(!out.contains("JOIN base"));
}
#[test]
fn test_transform_source_references_reversed_from_join() {
let reader = MockReader::new(true);
let sql = "SELECT * FROM base JOIN \"__ggsql_cte_t__\" ON base.k = 1";
let out = transform_source_references(sql, &reader).unwrap();
assert_eq!(reader.staged.borrow().len(), 1);
assert!(out.contains("FROM \"__ggsql_staged_0_"));
assert!(!out.contains("FROM base"));
}
#[test]
fn test_transform_source_references_case_insensitive_qualifier() {
let reader = MockReader::new(true);
let sql = "SELECT MYSCHEMA.BASE.w FROM \"__ggsql_cte_t__\" \
JOIN myschema.base ON MYSCHEMA.BASE.k = 1";
let out = transform_source_references(sql, &reader).unwrap();
assert_eq!(reader.staged.borrow().len(), 1);
assert!(
!out.contains("MYSCHEMA"),
"case-variant qualifier not rewritten: {out}"
);
}
#[test]
fn test_last_identifier_component() {
assert_eq!(last_identifier_component("base"), "base");
assert_eq!(last_identifier_component("schema.base"), "base");
assert_eq!(last_identifier_component("cat.schema.base"), "base");
assert_eq!(last_identifier_component("\"schema\".\"base\""), "\"base\"");
assert_eq!(last_identifier_component("schema.\"base\""), "\"base\"");
assert_eq!(last_identifier_component("\"my.base\""), "\"my.base\"");
}
#[test]
fn test_transform_source_references_fully_quoted_qualified() {
let reader = MockReader::new(true);
let sql = "SELECT \"myschema\".\"base\".w FROM \"__ggsql_cte_t__\" \
JOIN \"myschema\".\"base\" ON \"myschema\".\"base\".k = 1";
let out = transform_source_references(sql, &reader).unwrap();
let staged = reader.staged.borrow();
assert_eq!(staged.len(), 1);
assert_eq!(staged[0].1, "SELECT * FROM \"myschema\".\"base\"");
assert!(out.contains("AS \"base\""));
assert!(!out.contains("\"myschema\".\"base\""));
}
#[test]
fn test_transform_source_references_whitespace_around_dot() {
let reader = MockReader::new(true);
let sql = "SELECT myschema .\nbase . w FROM \"__ggsql_cte_t__\" \
JOIN myschema . base ON myschema . base . k = 1";
let out = transform_source_references(sql, &reader).unwrap();
let staged = reader.staged.borrow();
assert_eq!(staged.len(), 1);
assert!(!out.contains("myschema"), "qualifier not rewritten: {out}");
assert!(out.contains("AS base"));
}
#[test]
fn test_transform_source_references_schema_qualified() {
let reader = MockReader::new(true);
let sql = "SELECT * FROM \"__ggsql_cte_t__\" JOIN myschema.base ON base.k = 1";
let out = transform_source_references(sql, &reader).unwrap();
let staged = reader.staged.borrow();
assert_eq!(staged.len(), 1);
assert_eq!(staged[0].1, "SELECT * FROM myschema.base");
assert!(out.contains("__ggsql_staged_0_"));
assert!(!out.contains("JOIN myschema.base"));
}
#[test]
fn test_transform_source_references_same_ref_staged_once() {
let reader = MockReader::new(true);
let sql = "SELECT * FROM base JOIN \"__ggsql_cte_t__\" ON base.k = base.j";
let _ = transform_source_references(sql, &reader).unwrap();
assert_eq!(reader.staged.borrow().len(), 1);
}
#[test]
fn test_transform_source_references_comma_join() {
let reader = MockReader::new(true);
let sql = "SELECT * FROM \"__ggsql_cte_t__\", base";
let out = transform_source_references(sql, &reader).unwrap();
let staged = reader.staged.borrow();
assert_eq!(staged.len(), 1);
assert_eq!(staged[0].1, "SELECT * FROM base");
assert!(out.contains("__ggsql_staged_0_"));
assert!(out.contains("\"__ggsql_cte_t__\""));
}
#[test]
fn test_transform_source_references_preserves_string_literals() {
let reader = MockReader::new(true);
let sql = "SELECT * FROM \"__ggsql_cte_t__\" JOIN myschema.base \
ON note = 'myschema.base.k'";
let out = transform_source_references(sql, &reader).unwrap();
assert_eq!(reader.staged.borrow().len(), 1);
assert!(out.contains("'myschema.base.k'"));
assert!(!out.contains("JOIN myschema.base "));
}
#[test]
fn test_transform_source_references_builtin_not_staged() {
let reader = MockReader::new(true);
let sql = "SELECT * FROM \"__ggsql_cte_t__\" JOIN ggsql:penguins ON 1 = 1";
let out = transform_source_references(sql, &reader).unwrap();
assert!(reader.staged.borrow().is_empty());
assert_eq!(out, sql);
}
#[test]
fn test_split_with_query_basic() {
let sql = "WITH cte AS (SELECT * FROM x) SELECT * FROM cte";
let source_tree = SourceTree::new(sql).unwrap();
let (prefix, select) = split_with_query(&source_tree).unwrap();
assert_eq!(prefix, "WITH cte AS (SELECT * FROM x)");
assert_eq!(select, "SELECT * FROM cte");
}
#[test]
fn test_split_with_query_multiple_ctes() {
let sql = "WITH a AS (SELECT 1), b AS (SELECT 2) SELECT * FROM a JOIN b";
let source_tree = SourceTree::new(sql).unwrap();
let (prefix, select) = split_with_query(&source_tree).unwrap();
assert_eq!(prefix, "WITH a AS (SELECT 1), b AS (SELECT 2)");
assert_eq!(select, "SELECT * FROM a JOIN b");
}
#[test]
fn test_split_with_query_nested_subquery() {
let sql = "WITH cte AS (SELECT * FROM (SELECT 1)) SELECT * FROM cte";
let source_tree = SourceTree::new(sql).unwrap();
let (prefix, select) = split_with_query(&source_tree).unwrap();
assert_eq!(prefix, "WITH cte AS (SELECT * FROM (SELECT 1))");
assert_eq!(select, "SELECT * FROM cte");
}
#[test]
fn test_split_with_query_string_with_select_keyword() {
let sql = "WITH cte AS (SELECT 'SELECT' AS col) SELECT * FROM cte";
let source_tree = SourceTree::new(sql).unwrap();
let (prefix, select) = split_with_query(&source_tree).unwrap();
assert_eq!(prefix, "WITH cte AS (SELECT 'SELECT' AS col)");
assert_eq!(select, "SELECT * FROM cte");
}
#[test]
fn test_split_with_query_string_with_parens() {
let sql = "WITH cte AS (SELECT '()' AS col) SELECT * FROM cte";
let source_tree = SourceTree::new(sql).unwrap();
let (prefix, select) = split_with_query(&source_tree).unwrap();
assert_eq!(prefix, "WITH cte AS (SELECT '()' AS col)");
assert_eq!(select, "SELECT * FROM cte");
}
#[test]
fn test_split_with_query_not_a_with() {
let sql = "SELECT * FROM x";
let source_tree = SourceTree::new(sql).unwrap();
assert!(split_with_query(&source_tree).is_none());
}
#[test]
fn test_split_with_query_no_trailing_select() {
let sql = "WITH cte AS (SELECT 1) VISUALISE DRAW point";
let source_tree = SourceTree::new(sql).unwrap();
assert!(split_with_query(&source_tree).is_none());
}
#[test]
fn test_split_with_query_stat_transform_output() {
let sql = "WITH __stat_src__ AS (SELECT x FROM data), \
__binned__ AS (SELECT x, COUNT(*) AS count FROM __stat_src__ GROUP BY x) \
SELECT *, count * 1.0 / SUM(count) OVER () AS density FROM __binned__";
let source_tree = SourceTree::new(sql).unwrap();
let (prefix, select) = split_with_query(&source_tree).unwrap();
assert!(prefix.starts_with("WITH __stat_src__"));
assert!(prefix.contains("__binned__"));
assert!(prefix.ends_with(")"));
assert!(select.starts_with("SELECT *"));
assert!(select.contains("density"));
}
}