#![forbid(unsafe_code)]
use datafusion::prelude::SessionContext;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct UseTarget {
pub catalog: Option<String>,
pub schema: Option<String>,
}
fn unquote(ident: &str) -> String {
let t = ident.trim();
for q in ['`', '"'] {
if let Some(inner) = t.strip_prefix(q).and_then(|s| s.strip_suffix(q)) {
return inner.to_string();
}
}
t.to_string()
}
pub fn parse_use(query: &str) -> Option<UseTarget> {
let q = query.trim().trim_end_matches(';').trim();
let (head, remainder) = q.split_once(char::is_whitespace)?;
if !head.eq_ignore_ascii_case("USE") {
return None;
}
let remainder = remainder.trim();
if remainder.is_empty() {
return None;
}
let (is_catalog, name_part) = match remainder.split_once(char::is_whitespace) {
Some((w, r)) if w.eq_ignore_ascii_case("CATALOG") => (true, r.trim()),
Some((w, r))
if w.eq_ignore_ascii_case("SCHEMA")
|| w.eq_ignore_ascii_case("DATABASE")
|| w.eq_ignore_ascii_case("NAMESPACE") =>
{
(false, r.trim())
}
_ => (false, remainder),
};
let name = unquote(name_part);
if name.is_empty() {
return None;
}
if is_catalog {
return Some(UseTarget {
catalog: Some(name),
schema: None,
});
}
if let Some((cat, sch)) = name.split_once('.') {
return Some(UseTarget {
catalog: Some(unquote(cat)),
schema: Some(unquote(sch)),
});
}
Some(UseTarget {
catalog: None,
schema: Some(name),
})
}
pub fn apply_use(ctx: &SessionContext, query: &str) -> Option<Result<(), String>> {
let target = parse_use(query)?;
let state_ref = ctx.state_ref();
let mut state = state_ref.write();
let opts = state.config_mut().options_mut();
if let Some(catalog) = target.catalog {
opts.catalog.default_catalog = catalog;
}
if let Some(schema) = target.schema {
opts.catalog.default_schema = schema;
}
Some(Ok(()))
}
pub fn rewrite_show_databases(query: &str) -> Option<String> {
let q = query.trim().trim_end_matches(';').trim();
let upper = q.to_ascii_uppercase();
let is_show = upper.starts_with("SHOW DATABASES") || upper.starts_with("SHOW SCHEMAS");
if !is_show {
return None;
}
let like_clause = if let Some(idx) = upper.find(" LIKE ") {
let pat = q[idx + 6..].trim().trim_end_matches(';').trim();
Some(format!(" WHERE schema_name LIKE {pat}"))
} else {
None
};
Some(format!(
"SELECT schema_name AS namespace FROM information_schema.schemata{} ORDER BY namespace",
like_clause.unwrap_or_default()
))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_use_forms() {
assert_eq!(
parse_use("USE analytics"),
Some(UseTarget {
catalog: None,
schema: Some("analytics".into())
})
);
assert_eq!(
parse_use("USE SCHEMA sales"),
Some(UseTarget {
catalog: None,
schema: Some("sales".into())
})
);
assert_eq!(
parse_use("USE DATABASE sales;"),
Some(UseTarget {
catalog: None,
schema: Some("sales".into())
})
);
assert_eq!(
parse_use("USE CATALOG lakehouse"),
Some(UseTarget {
catalog: Some("lakehouse".into()),
schema: None
})
);
assert_eq!(
parse_use("USE lake.sales"),
Some(UseTarget {
catalog: Some("lake".into()),
schema: Some("sales".into())
})
);
assert_eq!(
parse_use("USE `my schema`").unwrap().schema.as_deref(),
Some("my schema")
);
assert_eq!(parse_use("SELECT 1"), None);
}
#[test]
fn show_databases_rewrite() {
assert!(
rewrite_show_databases("SHOW DATABASES")
.unwrap()
.contains("information_schema.schemata")
);
assert!(
rewrite_show_databases("SHOW SCHEMAS")
.unwrap()
.contains("AS namespace")
);
let with_like = rewrite_show_databases("SHOW DATABASES LIKE 'sal%'").unwrap();
assert!(with_like.contains("LIKE 'sal%'"));
assert_eq!(rewrite_show_databases("SHOW TABLES"), None);
}
#[tokio::test]
async fn use_changes_default_schema_end_to_end() {
let engine = crate::SqlEngine::new();
engine
.sql("USE information_schema")
.await
.expect("USE runs");
let batches = engine
.sql("SELECT count(*) AS c FROM tables")
.await
.expect("unqualified `tables` resolves via the new default schema")
.collect()
.await
.expect("collect");
let total: i64 = {
use arrow::array::Int64Array;
batches[0]
.column(0)
.as_any()
.downcast_ref::<Int64Array>()
.unwrap()
.value(0)
};
assert!(total > 0, "information_schema.tables should be non-empty");
}
#[tokio::test]
async fn show_databases_lists_schemas() {
let engine = crate::SqlEngine::new();
let batches = engine
.sql("SHOW DATABASES")
.await
.expect("SHOW DATABASES runs")
.collect()
.await
.expect("collect");
let total: usize = batches.iter().map(|b| b.num_rows()).sum();
assert!(total >= 1);
assert_eq!(batches[0].schema().field(0).name(), "namespace");
}
}