use std::collections::HashSet;
use sqlparser::dialect::SQLiteDialect;
use sqlparser::keywords::Keyword;
use sqlparser::tokenizer::{Location, Token, TokenWithSpan, Tokenizer};
use thiserror::Error;
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize)]
pub struct SourcePosition {
pub line: usize,
pub column: usize,
pub offset: usize,
pub length: usize,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize)]
pub struct Placeholder {
pub name: String,
pub position: SourcePosition,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize)]
pub struct TableRef {
pub name: String,
pub position: SourcePosition,
}
#[derive(Debug, Clone, Default, PartialEq, Eq, serde::Serialize)]
pub struct ScanResult {
pub placeholders: Vec<Placeholder>,
pub tables: Vec<TableRef>,
}
#[derive(Debug, Error)]
pub enum ScanError {
#[error("SQL tokenizer error: {0}")]
Tokenize(String),
}
pub fn scan_sql(sql: &str) -> Result<ScanResult, ScanError> {
let dialect = SQLiteDialect {};
let tokens = Tokenizer::new(&dialect, sql)
.tokenize_with_location()
.map_err(|e| ScanError::Tokenize(e.to_string()))?;
let line_starts = compute_line_starts(sql);
let significant: Vec<&TokenWithSpan> = tokens
.iter()
.filter(|t| !matches!(t.token, Token::Whitespace(_)))
.collect();
Ok(ScanResult {
placeholders: scan_placeholders(&significant, &line_starts),
tables: scan_tables(&significant, &line_starts),
})
}
pub fn undeclared_tables<'a>(
result: &'a ScanResult,
declared_labels: &[String],
) -> Vec<&'a TableRef> {
let declared: HashSet<String> = declared_labels.iter().map(|l| l.to_lowercase()).collect();
result
.tables
.iter()
.filter(|t| !declared.contains(&t.name.to_lowercase()))
.collect()
}
fn compute_line_starts(sql: &str) -> Vec<usize> {
let mut starts = vec![0usize];
for (idx, ch) in sql.chars().enumerate() {
if ch == '\n' {
starts.push(idx + 1);
}
}
starts
}
fn to_offset(line_starts: &[usize], loc: Location) -> usize {
let line_idx = (loc.line as usize).saturating_sub(1);
let line_start = line_starts
.get(line_idx)
.copied()
.unwrap_or_else(|| line_starts.last().copied().unwrap_or(0));
line_start + (loc.column as usize).saturating_sub(1)
}
fn source_position(line_starts: &[usize], start: Location, end: Location) -> SourcePosition {
let offset = to_offset(line_starts, start);
let end_offset = to_offset(line_starts, end);
SourcePosition {
line: start.line as usize,
column: start.column as usize,
offset,
length: end_offset.saturating_sub(offset),
}
}
fn scan_placeholders(tokens: &[&TokenWithSpan], line_starts: &[usize]) -> Vec<Placeholder> {
let mut seen = HashSet::new();
let mut out = Vec::new();
for pair in tokens.windows(2) {
let (colon, word) = (pair[0], pair[1]);
if !matches!(colon.token, Token::Colon) {
continue;
}
if colon.span.end != word.span.start {
continue;
}
let Token::Word(w) = &word.token else {
continue;
};
if w.quote_style.is_some() {
continue;
}
if !seen.insert(w.value.clone()) {
continue;
}
out.push(Placeholder {
name: w.value.clone(),
position: source_position(line_starts, colon.span.start, word.span.end),
});
}
out
}
fn ends_from_list(keyword: Keyword) -> bool {
matches!(
keyword,
Keyword::WHERE
| Keyword::GROUP
| Keyword::HAVING
| Keyword::ORDER
| Keyword::LIMIT
| Keyword::UNION
| Keyword::INTERSECT
| Keyword::EXCEPT
| Keyword::WINDOW
)
}
fn scan_tables(tokens: &[&TokenWithSpan], line_starts: &[usize]) -> Vec<TableRef> {
let cte_names = collect_cte_names(tokens);
let mut paren_depth: usize = 0;
let mut from_active: Vec<bool> = vec![false];
let mut seen = HashSet::new();
let mut out = Vec::new();
for (i, tok) in tokens.iter().enumerate() {
match &tok.token {
Token::LParen => {
paren_depth += 1;
from_active.push(false);
}
Token::RParen => {
if paren_depth > 0 {
paren_depth -= 1;
from_active.pop();
}
}
Token::Word(w) if w.keyword == Keyword::FROM => {
from_active[paren_depth] = true;
}
Token::Word(w) if ends_from_list(w.keyword) => {
from_active[paren_depth] = false;
}
Token::Word(w) if w.keyword == Keyword::NoKeyword => {
let is_candidate = i > 0
&& match &tokens[i - 1].token {
Token::Word(pw) => {
pw.keyword == Keyword::FROM || pw.keyword == Keyword::JOIN
}
Token::Comma => from_active[paren_depth],
_ => false,
};
if !is_candidate {
continue;
}
let excluded = matches!(
tokens.get(i + 1).map(|t| &t.token),
Some(Token::LParen) | Some(Token::Period)
);
if excluded {
continue;
}
if cte_names.contains(&w.value.to_lowercase()) {
continue;
}
if !seen.insert(w.value.to_lowercase()) {
continue;
}
out.push(TableRef {
name: w.value.clone(),
position: source_position(line_starts, tok.span.start, tok.span.end),
});
}
_ => {}
}
}
out
}
fn collect_cte_names(tokens: &[&TokenWithSpan]) -> HashSet<String> {
let mut names = HashSet::new();
let mut i = 0usize;
while i < tokens.len() {
let is_with = matches!(&tokens[i].token, Token::Word(w) if w.keyword == Keyword::WITH);
if !is_with {
i += 1;
continue;
}
i += 1;
if matches!(tokens.get(i).map(|t| &t.token), Some(Token::Word(w)) if w.keyword == Keyword::RECURSIVE)
{
i += 1;
}
while let Some(Token::Word(name)) = tokens.get(i).map(|t| &t.token) {
names.insert(name.value.to_lowercase());
i += 1;
if matches!(tokens.get(i).map(|t| &t.token), Some(Token::LParen)) {
i = skip_balanced_parens(tokens, i);
}
let is_as = matches!(tokens.get(i).map(|t| &t.token), Some(Token::Word(w)) if w.keyword == Keyword::AS);
if !is_as {
break;
}
i += 1;
if !matches!(tokens.get(i).map(|t| &t.token), Some(Token::LParen)) {
break;
}
i = skip_balanced_parens(tokens, i);
if matches!(tokens.get(i).map(|t| &t.token), Some(Token::Comma)) {
i += 1;
continue; }
break;
}
}
names
}
fn skip_balanced_parens(tokens: &[&TokenWithSpan], open_paren_idx: usize) -> usize {
let mut depth = 0i32;
let mut i = open_paren_idx;
loop {
match tokens.get(i).map(|t| &t.token) {
Some(Token::LParen) => depth += 1,
Some(Token::RParen) => {
depth -= 1;
if depth == 0 {
return i + 1;
}
}
Some(_) => {}
None => return tokens.len(),
}
i += 1;
}
}
#[cfg(test)]
mod tests {
use super::*;
fn names(result: &ScanResult) -> Vec<&str> {
result.tables.iter().map(|t| t.name.as_str()).collect()
}
fn placeholder_names(result: &ScanResult) -> Vec<&str> {
result
.placeholders
.iter()
.map(|p| p.name.as_str())
.collect()
}
#[test]
fn simple_from_reports_table_and_position() {
let result = scan_sql("SELECT * FROM v").unwrap();
assert_eq!(names(&result), vec!["v"]);
assert_eq!(result.tables[0].position.line, 1);
assert_eq!(result.tables[0].position.column, 15);
assert_eq!(result.tables[0].position.offset, 14);
assert_eq!(result.tables[0].position.length, 1);
}
#[test]
fn join_variants_report_both_sides() {
let result = scan_sql("SELECT * FROM a LEFT JOIN b ON a.id = b.id").unwrap();
assert_eq!(names(&result), vec!["a", "b"]);
}
#[test]
fn comma_list_reports_every_item_ignoring_aliases() {
let result = scan_sql("SELECT * FROM a, b c, d AS e").unwrap();
assert_eq!(names(&result), vec!["a", "b", "d"]);
}
#[test]
fn cte_body_reported_but_cte_reference_excluded() {
let result = scan_sql("WITH cte AS (SELECT * FROM t) SELECT * FROM cte").unwrap();
assert_eq!(names(&result), vec!["t"]);
}
#[test]
fn subquery_in_from_reports_nothing() {
let result = scan_sql("SELECT * FROM (SELECT 1) sub").unwrap();
assert!(result.tables.is_empty());
}
#[test]
fn table_valued_function_reports_nothing() {
let result = scan_sql("SELECT * FROM json_each(x)").unwrap();
assert!(result.tables.is_empty());
}
#[test]
fn quoted_identifier_reports_unquoted_name() {
let result = scan_sql(r#"SELECT * FROM "Quoted""#).unwrap();
assert_eq!(names(&result), vec!["Quoted"]);
}
#[test]
fn qualified_name_reports_nothing() {
let result = scan_sql("SELECT * FROM main.t").unwrap();
assert!(result.tables.is_empty());
}
#[test]
fn dedup_is_case_insensitive_and_keeps_first_spelling() {
let result = scan_sql("SELECT * FROM v v2 JOIN V").unwrap();
assert_eq!(names(&result), vec!["v"]);
}
#[test]
fn placeholder_ignores_literal_and_comment_reports_real_one() {
let result = scan_sql("SELECT * WHERE x = ':notaparam' AND y = :ward -- :comment").unwrap();
assert_eq!(placeholder_names(&result), vec!["ward"]);
}
#[test]
fn double_colon_cast_is_not_a_placeholder() {
let result = scan_sql("SELECT x::int").unwrap();
assert!(result.placeholders.is_empty());
}
#[test]
fn placeholder_dedup_keeps_first_occurrence() {
let result = scan_sql("SELECT * WHERE a = :p AND b = :p").unwrap();
assert_eq!(result.placeholders.len(), 1);
}
#[test]
fn multibyte_text_offsets_count_chars_not_bytes() {
let sql = "SELECT * WHERE name = 'café' AND x = :p";
let result = scan_sql(sql).unwrap();
assert_eq!(placeholder_names(&result), vec!["p"]);
let pos = result.placeholders[0].position;
assert_eq!(pos.offset, sql.chars().count() - 2);
}
#[test]
fn multiline_sql_reports_correct_line_and_column() {
let sql = "SELECT *\nFROM v\nWHERE x = :p";
let result = scan_sql(sql).unwrap();
assert_eq!(result.tables[0].position.line, 2);
assert_eq!(result.tables[0].position.column, 6);
assert_eq!(result.placeholders[0].position.line, 3);
assert_eq!(result.placeholders[0].position.column, 11);
}
#[test]
fn tokenizer_failure_is_reported_as_scan_error() {
let err = scan_sql("SELECT 'unterminated").unwrap_err();
assert!(matches!(err, ScanError::Tokenize(_)));
}
#[test]
fn undeclared_tables_helper_is_case_insensitive_and_ordered() {
let result = scan_sql("SELECT * FROM a JOIN B JOIN c").unwrap();
let declared = vec!["A".to_string(), "c".to_string()];
let undeclared = undeclared_tables(&result, &declared);
assert_eq!(
undeclared
.iter()
.map(|t| t.name.as_str())
.collect::<Vec<_>>(),
vec!["B"]
);
}
}