use super::config::Engine;
const READING: [&str; 9] = [
"SELECT", "WITH", "VALUES", "TABLE", "SHOW", "DESCRIBE", "DESC", "EXPLAIN", "PRAGMA",
];
const WRITING: [&str; 5] = ["INSERT", "UPDATE", "DELETE", "MERGE", "INTO"];
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Statement {
pub text: String,
pub start: usize,
pub end: usize,
}
pub fn split(sql: &str, engine: Engine) -> Vec<Statement> {
statements(sql, engine)
.into_iter()
.filter_map(|(_, start, end)| {
let text = sql[start..end].trim();
if text.is_empty() {
return None;
}
let offset = sql[start..end].find(text).unwrap_or_default();
Some(Statement {
text: text.to_string(),
start: start + offset,
end: start + offset + text.len(),
})
})
.collect()
}
pub fn at_cursor(sql: &str, cursor: usize, engine: Engine) -> Option<Statement> {
let statements = split(sql, engine);
statements
.iter()
.find(|statement| cursor >= statement.start && cursor <= statement.end)
.or_else(|| {
statements
.iter()
.rev()
.find(|statement| statement.end <= cursor)
})
.or_else(|| statements.first())
.cloned()
}
const EXPLAIN_OPTIONS: [&str; 18] = [
"ANALYZE",
"ANALYSE",
"VERBOSE",
"COSTS",
"SETTINGS",
"BUFFERS",
"WAL",
"TIMING",
"SUMMARY",
"FORMAT",
"TREE",
"JSON",
"TRADITIONAL",
"EXTENDED",
"PARTITIONS",
"XML",
"YAML",
"GENERIC_PLAN",
];
pub fn explained(sql: &str, engine: Engine) -> Option<(String, bool)> {
let characters: Vec<char> = sql.chars().collect();
let offsets: Vec<usize> = sql
.char_indices()
.map(|(offset, _)| offset)
.chain(std::iter::once(sql.len()))
.collect();
let byte = |index: usize| offsets.get(index).copied().unwrap_or(sql.len());
let mut index = skip_trivia(&characters, 0, engine);
let (word, next) = word_at(&characters, index)?;
let first = word.to_ascii_uppercase();
if first != "EXPLAIN" && first != "ANALYZE" {
return None;
}
let mut analyzes = first == "ANALYZE";
index = next;
loop {
index = skip_trivia(&characters, index, engine);
match characters.get(index) {
Some('(') => {
let (found, next) = skip_parens(&characters, index, engine);
analyzes |= found;
index = next;
}
Some('=') => index += 1,
Some(_) => {
if let Some((word, next)) = word_at(&characters, index) {
let upper = word.to_ascii_uppercase();
if EXPLAIN_OPTIONS.contains(&upper.as_str()) {
analyzes |= is_analyze(&upper);
index = next;
continue;
}
}
return Some((sql[byte(index)..].trim().to_string(), analyzes));
}
None => return Some((String::new(), analyzes)),
}
}
}
pub fn first_write(sql: &str, engine: Engine) -> Option<String> {
statements(sql, engine)
.into_iter()
.find_map(|(words, _, _)| {
if reads(&words) {
return None;
}
Some(
words
.first()
.cloned()
.unwrap_or_else(|| "statement".to_string()),
)
})
}
pub fn leading_words(sql: &str, limit: usize, engine: Engine) -> Vec<String> {
statements(sql, engine)
.into_iter()
.next()
.map(|(words, _, _)| words.into_iter().take(limit).collect())
.unwrap_or_default()
}
fn transaction_control(words: &[String]) -> bool {
let word = |index: usize| words.get(index).map(String::as_str).unwrap_or("");
let control = match word(0) {
"BEGIN" | "COMMIT" | "END" | "ROLLBACK" | "ABORT" | "SAVEPOINT" | "RELEASE" => true,
"START" | "SET" => word(1) == "TRANSACTION",
_ => false,
};
control
&& !words
.iter()
.any(|word| word == "WRITE" || word == "PREPARED" || WRITING.contains(&word.as_str()))
}
fn reads(words: &[String]) -> bool {
let Some(first) = words.first() else {
return true;
};
if transaction_control(words) {
return true;
}
if !READING.contains(&first.as_str()) {
return false;
}
if words
.iter()
.any(|word| WRITING.contains(&word.as_str()) || is_analyze(word))
{
return false;
}
if first == "PRAGMA" {
if words.iter().any(|word| word == "=") {
return false;
}
if let Some(paren) = words.iter().position(|word| word == "(") {
let name = paren
.checked_sub(1)
.and_then(|index| words.get(index))
.map(String::as_str)
.unwrap_or("");
return READING_PRAGMAS.contains(&name);
}
}
true
}
fn is_analyze(word: &str) -> bool {
word.eq_ignore_ascii_case("ANALYZE") || word.eq_ignore_ascii_case("ANALYSE")
}
const READING_PRAGMAS: [&str; 9] = [
"TABLE_INFO",
"TABLE_XINFO",
"TABLE_LIST",
"INDEX_LIST",
"INDEX_INFO",
"INDEX_XINFO",
"FOREIGN_KEY_LIST",
"FOREIGN_KEY_CHECK",
"INTEGRITY_CHECK",
];
fn statements(sql: &str, engine: Engine) -> Vec<(Vec<String>, usize, usize)> {
let mut statements = Vec::new();
let mut words = Vec::new();
let mut word = String::new();
let characters: Vec<char> = sql.chars().collect();
let offsets: Vec<usize> = sql
.char_indices()
.map(|(offset, _)| offset)
.chain(std::iter::once(sql.len()))
.collect();
let byte = |index: usize| offsets.get(index).copied().unwrap_or(sql.len());
let mut start = 0;
let mut index = 0;
macro_rules! end_word {
() => {
if !word.is_empty() {
words.push(std::mem::take(&mut word).to_ascii_uppercase());
}
};
}
while index < characters.len() {
let character = characters[index];
let next = characters.get(index + 1).copied();
match character {
'-' if next == Some('-') => {
end_word!();
index = skip_until(&characters, index, "\n");
}
'#' if engine == Engine::MySql => {
end_word!();
index = skip_until(&characters, index, "\n");
}
'/' if next == Some('*') => {
end_word!();
index = skip_until(&characters, index + 2, "*/");
}
'\'' | '"' | '`' => {
end_word!();
index = skip_quoted(&characters, index, character, engine);
}
'$' if dollar_tag(&characters, index).is_some() => {
end_word!();
let tag = dollar_tag(&characters, index).expect("checked just above");
index = skip_until(&characters, index + tag.len(), &tag);
}
';' => {
end_word!();
statements.push((std::mem::take(&mut words), byte(start), byte(index)));
index += 1;
start = index;
}
'=' => {
end_word!();
words.push("=".to_string());
index += 1;
}
'(' if words.first().is_some_and(|first| first == "PRAGMA") => {
end_word!();
words.push("(".to_string());
index += 1;
}
character if character.is_alphanumeric() || character == '_' => {
word.push(character);
index += 1;
}
_ => {
end_word!();
index += 1;
}
}
}
end_word!();
if start < characters.len() {
statements.push((words, byte(start), sql.len()));
}
statements
}
fn skip_trivia(characters: &[char], mut index: usize, engine: Engine) -> usize {
loop {
while index < characters.len() && characters[index].is_whitespace() {
index += 1;
}
match characters
.get(index)
.copied()
.zip(characters.get(index + 1).copied())
{
Some(('-', '-')) => index = skip_until(characters, index, "\n"),
Some(('/', '*')) => index = skip_until(characters, index + 2, "*/"),
Some(('#', _)) if engine == Engine::MySql => {
index = skip_until(characters, index, "\n")
}
_ => return index,
}
}
}
fn word_at(characters: &[char], start: usize) -> Option<(String, usize)> {
let mut index = start;
let mut word = String::new();
while let Some(character) = characters.get(index) {
if character.is_alphanumeric() || *character == '_' {
word.push(*character);
index += 1;
} else {
break;
}
}
if word.is_empty() {
None
} else {
Some((word, index))
}
}
fn skip_parens(characters: &[char], start: usize, engine: Engine) -> (bool, usize) {
let mut depth = 0usize;
let mut analyzes = false;
let mut index = start;
while index < characters.len() {
match characters[index] {
'(' => depth += 1,
')' => {
depth -= 1;
if depth == 0 {
return (analyzes, index + 1);
}
}
quote @ ('\'' | '"' | '`') => {
index = skip_quoted(characters, index, quote, engine);
continue;
}
character if character.is_alphanumeric() || character == '_' => {
let (word, next) = word_at(characters, index).expect("a word starts here");
analyzes |= is_analyze(&word);
index = next;
continue;
}
_ => {}
}
index += 1;
}
(analyzes, index)
}
pub(crate) fn backslash_escapes(
characters: &[char],
from: usize,
quote: char,
engine: Engine,
) -> bool {
match engine {
Engine::MySql => quote != '`',
Engine::Postgres => {
quote == '\''
&& from
.checked_sub(1)
.and_then(|index| characters.get(index))
.is_some_and(|prefix| matches!(prefix, 'E' | 'e'))
&& from
.checked_sub(2)
.and_then(|index| characters.get(index))
.is_none_or(|before| !(before.is_alphanumeric() || *before == '_'))
}
Engine::Sqlite => false,
}
}
fn skip_until(characters: &[char], from: usize, terminator: &str) -> usize {
let terminator: Vec<char> = terminator.chars().collect();
let mut index = from;
while index < characters.len() {
if characters[index..].starts_with(terminator.as_slice()) {
return index + terminator.len();
}
index += 1;
}
characters.len()
}
fn skip_quoted(characters: &[char], from: usize, quote: char, engine: Engine) -> usize {
let escapes = backslash_escapes(characters, from, quote, engine);
let mut index = from + 1;
while index < characters.len() {
if escapes && characters[index] == '\\' {
index += 2;
continue;
}
if characters[index] == quote {
if characters.get(index + 1) == Some("e) {
index += 2;
continue;
}
return index + 1;
}
index += 1;
}
characters.len()
}
fn dollar_tag(characters: &[char], index: usize) -> Option<String> {
let mut tag = String::from("$");
for character in &characters[index + 1..] {
match character {
'$' => {
tag.push('$');
return Some(tag);
}
character if character.is_alphanumeric() || *character == '_' => tag.push(*character),
_ => return None,
}
}
None
}
#[cfg(test)]
mod tests {
use super::*;
fn split(sql: &str) -> Vec<Statement> {
super::split(sql, Engine::MySql)
}
fn at_cursor(sql: &str, cursor: usize) -> Option<Statement> {
super::at_cursor(sql, cursor, Engine::MySql)
}
fn explained(sql: &str) -> Option<(String, bool)> {
super::explained(sql, Engine::MySql)
}
fn first_write(sql: &str) -> Option<String> {
super::first_write(sql, Engine::MySql)
}
fn leading_words(sql: &str, limit: usize) -> Vec<String> {
super::leading_words(sql, limit, Engine::MySql)
}
#[test]
fn a_backslash_ends_nothing_on_sqlite_or_postgres() {
let sql = "select * from files where dir = 'C:\\';\ndelete from files where id = 3;";
for engine in [Engine::Sqlite, Engine::Postgres] {
let statements = super::split(sql, engine);
assert_eq!(statements.len(), 2, "{engine:?}");
let first = super::at_cursor(sql, 3, engine).unwrap();
assert_eq!(super::first_write(&first.text, engine), None, "{engine:?}");
assert_eq!(
super::first_write(sql, engine).as_deref(),
Some("DELETE"),
"{engine:?}"
);
}
assert_eq!(super::split(sql, Engine::MySql).len(), 1);
let escaped = "select E'it\\'s';\nselect 2;";
assert_eq!(super::split(escaped, Engine::Postgres).len(), 2);
let named = "select date 'x\\';\nselect 2;";
assert_eq!(super::split(named, Engine::Postgres).len(), 2);
}
#[test]
fn a_hash_is_a_comment_only_on_mysql() {
let sql = "select data #>> '{a}' from t;\nselect 2;";
assert_eq!(super::split(sql, Engine::Postgres).len(), 2);
let into = "select d #>> '{a}' as x into t2 from t";
assert_eq!(
super::first_write(into, Engine::Postgres).as_deref(),
Some("SELECT")
);
assert_eq!(
super::split("select 1; # note; select 2", Engine::MySql).len(),
2
);
}
#[test]
fn analyse_is_analyze() {
assert_eq!(
explained("explain analyse delete from items"),
Some(("delete from items".to_string(), true))
);
assert_eq!(
explained("explain (analyse) delete from items"),
Some(("delete from items".to_string(), true))
);
assert!(first_write("explain analyse create table t as select 1").is_some());
}
#[test]
fn a_pragma_set_with_parentheses_is_a_write() {
for sql in [
"PRAGMA user_version(5)",
"pragma journal_mode(delete)",
"PRAGMA main.writable_schema(1)",
] {
assert_eq!(first_write(sql).as_deref(), Some("PRAGMA"), "{sql}");
}
for sql in [
"PRAGMA table_info(items)",
"pragma main.index_list('items')",
"PRAGMA journal_mode",
] {
assert_eq!(first_write(sql), None, "{sql}");
}
}
#[test]
fn a_buffer_splits_into_its_statements() {
let sql = "select 1;\nselect 2\n";
let statements = split(sql);
assert_eq!(
statements
.iter()
.map(|statement| statement.text.as_str())
.collect::<Vec<_>>(),
["select 1", "select 2"]
);
assert_eq!(&sql[statements[1].start..statements[1].end], "select 2");
assert_eq!(split("select 'a; b'").len(), 1);
assert!(split(" ; \n ;").is_empty());
}
#[test]
fn leading_words_skip_comments_and_quotes() {
assert_eq!(
leading_words("/* first */ -- second\n lock Tables `t` write", 3),
["LOCK", "TABLES", "WRITE"]
);
assert_eq!(leading_words("vacuum; select 1", 5), ["VACUUM"]);
assert!(leading_words(" -- nothing\n", 2).is_empty());
}
#[test]
fn the_caret_picks_the_statement_it_is_in() {
let sql = "select 1;\nselect 2;\nselect 3;";
let at = |cursor| at_cursor(sql, cursor).map(|statement| statement.text);
assert_eq!(at(0).as_deref(), Some("select 1"));
assert_eq!(at(3).as_deref(), Some("select 1"));
assert_eq!(at(9).as_deref(), Some("select 1"));
assert_eq!(at(12).as_deref(), Some("select 2"));
assert_eq!(at(sql.len()).as_deref(), Some("select 3"));
assert_eq!(at_cursor("", 0), None);
}
#[test]
fn plain_reads_are_reads() {
for sql in [
"select * from items",
"SELECT 1; select 2;",
"show tables",
"explain select * from items",
"pragma table_info(items)",
"values (1), (2)",
" \n-- a comment\nselect 1",
"",
";",
] {
assert_eq!(first_write(sql), None, "{sql} should read");
}
}
#[test]
fn a_with_select_is_a_read() {
for sql in [
"with cte1 as (select a, b from table1), cte2 as (select c, d from table2) select b, d from cte1 join cte2 where cte1.a = cte2.c",
"WITH RECURSIVE countdown(n) AS (SELECT 1 UNION ALL SELECT n + 1 FROM countdown WHERE n < 3) SELECT n FROM countdown",
"with x as (select 1) values (2)",
] {
assert_eq!(first_write(sql), None, "{sql} should read");
}
assert_eq!(
first_write("with moved as (update items set x = 1 returning *) select * from moved")
.as_deref(),
Some("WITH")
);
assert_eq!(
first_write("with n as (select 1) insert into items select * from n").as_deref(),
Some("WITH")
);
}
#[test]
fn writes_are_named() {
assert_eq!(
first_write("insert into items values (1)").as_deref(),
Some("INSERT")
);
assert_eq!(
first_write("select 1; drop table items").as_deref(),
Some("DROP")
);
assert_eq!(
first_write("SET GLOBAL max_connections = 10").as_deref(),
Some("SET")
);
assert_eq!(
first_write("alter table items add column x int").as_deref(),
Some("ALTER")
);
assert_eq!(first_write("truncate items").as_deref(), Some("TRUNCATE"));
}
#[test]
fn transaction_control_is_not_a_write() {
for sql in [
"begin",
"BEGIN TRANSACTION",
"begin immediate",
"begin isolation level serializable",
"start transaction",
"START TRANSACTION READ ONLY",
"start transaction with consistent snapshot",
"commit",
"commit work",
"end",
"rollback",
"abort",
"rollback to savepoint a",
"savepoint a",
"release savepoint a",
"release a",
"set transaction isolation level repeatable read",
] {
assert_eq!(first_write(sql), None, "{sql} should not write");
}
}
#[test]
fn transaction_control_that_could_write_still_counts() {
assert_eq!(first_write("begin read write").as_deref(), Some("BEGIN"));
assert!(first_write("start transaction read write").is_some());
assert!(first_write("set transaction read write").is_some());
assert!(first_write("commit prepared 'x'").is_some());
assert!(first_write("set session characteristics as transaction read only").is_some());
assert!(first_write("set search_path = app").is_some());
}
#[test]
fn a_write_hiding_behind_a_reading_keyword_is_found() {
assert!(
first_write("with gone as (delete from items returning *) select * from gone")
.is_some()
);
assert!(first_write("explain analyze delete from items").is_some());
assert!(first_write("pragma journal_mode = wal").is_some());
}
#[test]
fn a_select_into_is_not_a_plain_read() {
assert_eq!(
first_write("select * into new_table from items").as_deref(),
Some("SELECT")
);
assert_eq!(
first_write("select * from items into outfile '/tmp/dump.csv'").as_deref(),
Some("SELECT")
);
assert!(first_write("SELECT id INTO DUMPFILE '/tmp/id' FROM items LIMIT 1").is_some());
}
#[test]
fn quoted_text_is_not_sql() {
assert_eq!(first_write("select 'drop table items' as warning"), None);
assert_eq!(first_write(r#"select "drop" from items"#), None);
assert_eq!(first_write("select $tag$ delete from items $tag$"), None);
assert_eq!(first_write("select 'a; drop table items'"), None);
}
#[test]
fn a_comment_cannot_hide_a_write() {
assert_eq!(first_write("select 1 -- drop table items"), None);
assert_eq!(first_write("/* drop table items */ select 1"), None);
assert_eq!(
first_write("/* comment */ delete from items").as_deref(),
Some("DELETE")
);
}
#[test]
fn an_explain_header_is_split_off_what_it_explains() {
let split = |sql: &str| explained(sql);
assert_eq!(split("select 1"), None);
assert_eq!(
split("explain select * from items"),
Some(("select * from items".to_string(), false))
);
assert_eq!(
split("EXPLAIN ANALYZE select * from items"),
Some(("select * from items".to_string(), true))
);
assert_eq!(
split("explain (analyze, buffers, format json) select 1"),
Some(("select 1".to_string(), true))
);
assert_eq!(
split("explain (format json) select 1"),
Some(("select 1".to_string(), false))
);
assert_eq!(
split("explain format=tree select 1"),
Some(("select 1".to_string(), false))
);
assert_eq!(
split("analyze format=json select 1"),
Some(("select 1".to_string(), true))
);
assert_eq!(
split("explain analyze delete from items"),
Some(("delete from items".to_string(), true))
);
assert_eq!(
split("explain delete from items"),
Some(("delete from items".to_string(), false))
);
assert_eq!(split("explain"), Some((String::new(), false)));
}
#[test]
fn an_unknown_statement_counts_as_a_write() {
assert!(first_write("call do_something()").is_some());
assert!(first_write("lock tables items write").is_some());
}
}