fix(sql_safety): 修正子查询边界误判

This commit is contained in:
Illuminated2020 2026-05-10 17:19:15 +08:00
parent 1909856467
commit 99e7997620
1 changed files with 72 additions and 9 deletions

View File

@ -150,14 +150,22 @@ fn is_ddl_token(token: &str) -> bool {
matches!(token, "CREATE" | "ALTER" | "DROP" | "TRUNCATE" | "RENAME")
}
#[derive(Debug, Clone, PartialEq, Eq)]
struct SqlToken {
text: String,
depth: usize,
}
fn has_unfiltered_destructive_write(sql: &str) -> bool {
executable_statements(sql).into_iter().any(|statement| {
scanned_executable_statements(sql).into_iter().any(|statement| {
statement.iter().enumerate().any(|(index, token)| {
if !matches!(token.as_str(), "DELETE" | "UPDATE") {
if !matches!(token.text.as_str(), "DELETE" | "UPDATE") {
return false;
}
!statement[index + 1..].iter().any(|token| matches!(token.as_str(), "WHERE" | "LIMIT"))
!statement[index + 1..]
.iter()
.any(|boundary| boundary.depth == token.depth && matches!(boundary.text.as_str(), "WHERE" | "LIMIT"))
})
})
}
@ -167,14 +175,27 @@ fn executable_tokens(sql: &str) -> Vec<String> {
}
fn executable_statements(sql: &str) -> Vec<Vec<String>> {
scanned_executable_statements(sql)
.into_iter()
.map(|statement| statement.into_iter().map(|token| token.text).collect())
.collect()
}
fn scanned_executable_statements(sql: &str) -> Vec<Vec<SqlToken>> {
let mut statements = Vec::new();
let mut current = Vec::new();
scan_executable_tokens(sql, &mut current, &mut statements);
let mut depth = 0;
scan_executable_tokens(sql, &mut current, &mut statements, &mut depth);
push_statement(&mut current, &mut statements);
statements
}
fn scan_executable_tokens(sql: &str, current: &mut Vec<String>, statements: &mut Vec<Vec<String>>) {
fn scan_executable_tokens(
sql: &str,
current: &mut Vec<SqlToken>,
statements: &mut Vec<Vec<SqlToken>>,
depth: &mut usize,
) {
let bytes = sql.as_bytes();
let mut i = 0;
@ -185,7 +206,21 @@ fn scan_executable_tokens(sql: &str, current: &mut Vec<String>, statements: &mut
}
if bytes[i] == b';' {
push_statement(current, statements);
if *depth == 0 {
push_statement(current, statements);
}
i += 1;
continue;
}
if bytes[i] == b'(' {
*depth += 1;
i += 1;
continue;
}
if bytes[i] == b')' {
*depth = depth.saturating_sub(1);
i += 1;
continue;
}
@ -202,7 +237,7 @@ fn scan_executable_tokens(sql: &str, current: &mut Vec<String>, statements: &mut
if i + 2 < bytes.len() && bytes[i + 2] == b'!' {
let content_start = i + 3;
let content_end = block_comment_end(bytes, content_start);
scan_executable_tokens(&sql[content_start..content_end], current, statements);
scan_executable_tokens(&sql[content_start..content_end], current, statements, depth);
i = (content_end + 2).min(bytes.len());
} else {
i += 2;
@ -248,7 +283,7 @@ fn scan_executable_tokens(sql: &str, current: &mut Vec<String>, statements: &mut
while i < bytes.len() && (bytes[i].is_ascii_alphanumeric() || bytes[i] == b'_') {
i += 1;
}
current.push(sql[start..i].to_ascii_uppercase());
current.push(SqlToken { text: sql[start..i].to_ascii_uppercase(), depth: *depth });
continue;
}
@ -256,7 +291,7 @@ fn scan_executable_tokens(sql: &str, current: &mut Vec<String>, statements: &mut
}
}
fn push_statement(current: &mut Vec<String>, statements: &mut Vec<Vec<String>>) {
fn push_statement(current: &mut Vec<SqlToken>, statements: &mut Vec<Vec<SqlToken>>) {
if !current.is_empty() {
statements.push(std::mem::take(current));
}
@ -380,6 +415,34 @@ mod tests {
assert_eq!(risk_for("TRUNCATE TABLE users", RiskContext::new("dev")).risk_level, RiskLevel::Critical);
}
#[test]
fn destructive_writes_only_count_top_level_where_or_limit_as_boundaries() {
assert_eq!(
risk_for(
"DELETE FROM users USING (SELECT id FROM archived WHERE stale = true) old",
RiskContext::new("dev")
)
.risk_level,
RiskLevel::Critical
);
assert_eq!(
risk_for(
"UPDATE users SET active = false FROM (SELECT id FROM flags LIMIT 10) flags",
RiskContext::new("dev")
)
.risk_level,
RiskLevel::Critical
);
assert_eq!(
risk_for(
"DELETE FROM users WHERE id IN (SELECT user_id FROM archived WHERE stale = true)",
RiskContext::new("dev")
)
.risk_level,
RiskLevel::Medium
);
}
#[test]
fn postgresql_dollar_quotes_do_not_contribute_tokens() {
assert_eq!(classify_sql("SELECT $$ DELETE FROM users $$"), OperationClass::Read);