const DANGEROUS_KEYWORDS: &[&str] = &[
"select", "insert", "update", "delete", "drop", "alter", "truncate",
"create", "grant", "revoke", "union", "into", "outfile", "load_file",
"sleep", "benchmark", "exec", "execute", "xp_cmdshell", "shutdown",
"attach", "pragma", "vacuum", "reindex",
];
fn is_ident_char(c: char) -> bool {
c.is_ascii_alphanumeric() || c == '_' || c == '$'
}
const ERR_PREFIX: &str = "Invalid SQL identifier";
pub fn validate_identifier(identifier: &str) -> Result<String, String> {
let id = identifier.trim();
if id.is_empty() {
return Err(format!("{}: empty", ERR_PREFIX));
}
let first = id.chars().next().unwrap_or('\0');
if first.is_ascii_digit() {
return Err(format!("{}: cannot start with a digit: `{}`", ERR_PREFIX, id));
}
if !id.chars().all(is_ident_char) {
return Err(format!("{}: contains invalid characters: `{}`", ERR_PREFIX, id));
}
Ok(id.to_string())
}
pub fn validate_qualified_identifier(identifier: &str) -> Result<String, String> {
let id = identifier.trim();
if id.is_empty() {
return Err(format!("{}: empty", ERR_PREFIX));
}
let ok = id.chars().all(|c| {
c.is_ascii_alphanumeric()
|| c == '_'
|| c == '$'
|| c == '.'
|| c == '('
|| c == ')'
|| c == '*'
|| c == ' '
});
if !ok {
return Err(format!(
"{}: contains invalid characters: `{}`",
ERR_PREFIX, id
));
}
if contains_injection_pattern(id).is_some() {
return Err(format!("{}: contains injection pattern: `{}`", ERR_PREFIX, id));
}
Ok(id.to_string())
}
pub fn is_reserved_word(identifier: &str) -> bool {
let upper = identifier.to_ascii_uppercase();
RESERVED_WORDS.contains(&upper.as_str())
}
const RESERVED_WORDS: &[&str] = &[
"SELECT", "FROM", "WHERE", "INSERT", "UPDATE", "DELETE", "CREATE", "DROP",
"ALTER", "TABLE", "INDEX", "AND", "OR", "NOT", "NULL", "IN", "LIKE",
"BETWEEN", "ORDER", "GROUP", "BY", "HAVING", "LIMIT", "OFFSET", "JOIN",
"INNER", "LEFT", "RIGHT", "FULL", "ON", "AS", "UNION", "ALL", "DISTINCT",
"PRIMARY", "KEY", "FOREIGN", "REFERENCES", "UNIQUE", "CHECK", "DEFAULT",
"CASE", "WHEN", "THEN", "ELSE", "END", "IS", "DESC", "ASC", "COUNT",
"SUM", "AVG", "MIN", "MAX", "VALUES", "SET", "TO", "EXISTS", "CASCADE",
];
pub fn quote_identifier(identifier: &str) -> Option<String> {
match validate_identifier(identifier) {
Ok(id) => {
if is_reserved_word(&id) {
Some(format!("`{}`", id))
} else {
Some(id)
}
}
Err(_) => None,
}
}
pub fn escape_string(input: &str) -> String {
input.replace('\'', "''")
}
pub fn contains_injection_pattern(sql: &str) -> Option<(usize, &str)> {
let lower = sql.to_ascii_lowercase();
let bytes: Vec<char> = lower.chars().collect();
let mut i = 0;
let mut first_token_done = false;
while i < bytes.len() {
if bytes[i].is_whitespace() {
i += 1;
continue;
}
if bytes[i] == '\'' {
i += 1;
while i < bytes.len() {
if bytes[i] == '\'' {
if i + 1 < bytes.len() && bytes[i + 1] == '\'' {
i += 2;
continue;
}
i += 1;
break;
}
if bytes[i] == ';' {
first_token_done = true;
break;
}
i += 1;
}
continue;
}
if bytes[i] == '-' && i + 1 < bytes.len() && bytes[i + 1] == '-' {
while i < bytes.len() && bytes[i] != '\n' {
i += 1;
}
continue;
}
if bytes[i] == '/' && i + 1 < bytes.len() && bytes[i + 1] == '*' {
i += 2;
while i + 1 < bytes.len() && !(bytes[i] == '*' && bytes[i + 1] == '/') {
i += 1;
}
i += 2;
continue;
}
if !is_ident_char(bytes[i]) {
if bytes[i] == ';' {
first_token_done = true;
}
i += 1;
continue;
}
let start = i;
while i < bytes.len() && is_ident_char(bytes[i]) {
i += 1;
}
let word: String = bytes[start..i].iter().collect();
let matched: Option<&str> = DANGEROUS_KEYWORDS.iter().copied().find(|kw| {
kw.len() == word.len() && kw.eq_ignore_ascii_case(&word)
});
if let Some(kw) = matched {
let always_suspicious = matches!(kw, "drop" | "alter" | "truncate" | "create" | "grant" | "revoke");
if always_suspicious || (first_token_done && kw_is_mid(kw)) {
return Some((start, sql.get(start..i).unwrap_or("")));
}
first_token_done = true;
} else {
if first_token_done
&& (word.eq_ignore_ascii_case("or") || word.eq_ignore_ascii_case("and"))
&& is_tautology_literal(&bytes, i)
{
return Some((start, sql.get(start..i).unwrap_or("")));
}
first_token_done = true;
}
}
None
}
fn is_tautology_literal(bytes: &[char], from: usize) -> bool {
let mut j = from;
while j < bytes.len() && bytes[j].is_whitespace() {
j += 1;
}
let start = j;
while j < bytes.len() && bytes[j].is_ascii_digit() {
j += 1;
}
if j == start {
return false;
}
while j < bytes.len() && bytes[j].is_whitespace() {
j += 1;
}
if j >= bytes.len() || bytes[j] != '=' {
return false;
}
j += 1;
while j < bytes.len() && bytes[j].is_whitespace() {
j += 1;
}
let end_start = j;
while j < bytes.len() && bytes[j].is_ascii_digit() {
j += 1;
}
j > end_start
}
fn kw_is_mid(kw: &str) -> bool {
matches!(
kw,
"select" | "insert" | "update" | "delete" | "union" | "into" | "outfile" | "load_file"
)
}
#[derive(Debug, Clone, Copy, Default)]
pub struct SqlSanitizer;
impl SqlSanitizer {
pub fn identifier(identifier: &str) -> String {
quote_identifier(identifier).unwrap_or_else(|| {
eprintln!(
"[torm::sql_safety] rejected unsafe identifier: {:?}",
identifier
);
String::new()
})
}
pub fn escape(unescaped: &str) -> String {
escape_string(unescaped)
}
pub fn quote(identifier: &str) -> Option<String> {
quote_identifier(identifier)
}
pub fn check(sql: &str) -> Option<(usize, &str)> {
contains_injection_pattern(sql)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_validate_identifier_ok() {
assert_eq!(validate_identifier("user_name"), Ok("user_name".to_string()));
assert_eq!(validate_identifier("_tmp"), Ok("_tmp".to_string()));
assert_eq!(validate_identifier("table2"), Ok("table2".to_string()));
}
#[test]
fn test_validate_identifier_rejects_injection() {
assert!(validate_identifier("name; DROP TABLE users").is_err());
assert!(validate_identifier("1abc").is_err());
assert!(validate_identifier("col' OR '1'='1").is_err());
assert!(validate_identifier("").is_err());
assert!(validate_identifier(" ").is_err());
assert!(validate_identifier("a b").is_err());
}
#[test]
fn test_quote_identifier() {
assert_eq!(quote_identifier("user_name"), Some("user_name".to_string()));
assert_eq!(quote_identifier("select"), Some("`select`".to_string()));
assert_eq!(quote_identifier("order"), Some("`order`".to_string()));
assert_eq!(quote_identifier("bad name"), None);
assert_eq!(quote_identifier("id; DROP"), None);
}
#[test]
fn test_escape_string() {
assert_eq!(escape_string("O'Reilly"), "O''Reilly");
assert_eq!(escape_string("plain"), "plain");
assert_eq!(escape_string("a''b"), "a''''b");
assert_eq!(escape_string(""), "");
}
#[test]
fn test_sanitizer_identifier() {
assert_eq!(SqlSanitizer::identifier("user_name"), "user_name");
assert_eq!(SqlSanitizer::identifier("select"), "`select`");
assert_eq!(SqlSanitizer::identifier("bad name"), "");
assert_eq!(SqlSanitizer::identifier("id; DROP TABLE users"), "");
}
#[test]
fn test_contains_injection_pattern() {
assert!(contains_injection_pattern("SELECT * FROM users WHERE id = ?").is_none());
assert!(contains_injection_pattern("SELECT * FROM users WHERE name = 'select'").is_none());
assert!(contains_injection_pattern("SELECT * FROM users WHERE name = 'It''s a drop test'").is_none());
assert!(contains_injection_pattern("'; DROP TABLE users; --").is_some());
assert!(contains_injection_pattern("UNION SELECT * FROM admin").is_some());
assert!(contains_injection_pattern("1 OR 1=1; DELETE FROM users").is_some());
}
#[test]
fn test_contains_tautology_pattern() {
assert!(contains_injection_pattern("x = 'a' OR 1=1 --").is_some());
assert!(contains_injection_pattern("password = 'x' OR 1=1").is_some());
assert!(contains_injection_pattern("name = 'a' AND 2=2").is_some());
assert!(contains_injection_pattern("SELECT * FROM users WHERE active = 1 OR deleted = 0").is_none());
assert!(contains_injection_pattern("SELECT * FROM users WHERE age = 18").is_none());
}
}