fix(query): insert pagination before FOR UPDATE locking clauses

This commit is contained in:
AiLi 2026-08-07 02:41:15 +08:00 committed by GitHub
parent 59187f2e41
commit 87fd52de87
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
1 changed files with 271 additions and 7 deletions

View File

@ -224,6 +224,19 @@ pub fn build_count_query_sql(options: CountQuerySqlOptions) -> QuerySqlBuildResu
if unsupported_pagination_type(options.database_type) {
return err("unsupported");
}
// A locking clause does not affect cardinality and cannot appear inside
// every dialect's derived-table count query. PostgreSQL permits pagination
// after the lock clause; decline counting that uncommon order rather than
// accidentally dropping the user's explicit LIMIT/OFFSET.
let tokens = top_level_sql_tokens(&statement);
let statement = if let Some(index) = locking_clause_index(&tokens) {
if has_pagination_clause_after(&tokens, index) {
return err("locking");
}
statement[..index].trim_end().to_string()
} else {
statement
};
// ES SQL can't wrap a SELECT in `SELECT COUNT(*) FROM (...)` — the
// driver already reports the true match count via affected_rows.
if matches!(options.database_type, Some(DatabaseType::Elasticsearch | DatabaseType::Easysearch)) {
@ -952,12 +965,13 @@ fn add_questdb_limit(statement: &str, limit: usize, offset: usize) -> String {
}
return format!("{statement};");
}
if offset > 0 {
let limit_sql = if offset > 0 {
let upper_bound = offset + limit;
append_sql_suffix(statement, &format!("LIMIT {offset}, {upper_bound};"))
format!("LIMIT {offset}, {upper_bound}")
} else {
append_sql_suffix(statement, &format!("LIMIT {limit};"))
}
format!("LIMIT {limit}")
};
append_or_insert_before_locking(statement, &limit_sql)
}
fn has_top_level_limit(sql: &str) -> bool {
@ -1057,7 +1071,7 @@ fn add_fetch_first_limit(statement: &str, limit: usize, offset: usize) -> String
return format!("{statement};");
}
let offset_sql = if offset > 0 { format!(" OFFSET {offset} ROWS") } else { String::new() };
append_sql_suffix(statement, &format!("{offset_sql} FETCH FIRST {limit} ROWS ONLY;"))
append_or_insert_before_locking(statement, &format!("{offset_sql} FETCH FIRST {limit} ROWS ONLY"))
}
fn add_firebird_rows_limit(statement: &str, limit: usize, offset: usize) -> String {
@ -1114,7 +1128,7 @@ fn add_standard_limit(
if database_type == Some(DatabaseType::ClickHouse) {
return add_clickhouse_limit(statement, &limit_sql);
}
append_sql_suffix(statement, &format!("{limit_sql};"))
append_or_insert_before_locking(statement, &limit_sql)
}
fn add_outer_standard_limit(
@ -1138,7 +1152,45 @@ fn add_clickhouse_limit(statement: &str, limit_sql: &str) -> String {
return format!("{statement_before_settings} {limit_sql} {settings_clause};");
}
append_sql_suffix(statement, &format!("{limit_sql};"))
append_or_insert_before_locking(statement, limit_sql)
}
/// Insert pagination before a top-level locking clause; SQL dialects require
/// LIMIT/FETCH to precede FOR UPDATE, FOR SHARE, or LOCK IN SHARE MODE.
fn append_or_insert_before_locking(statement: &str, clause: &str) -> String {
let clause = clause.trim();
if let Some(index) = locking_clause_index(&top_level_sql_tokens(statement)) {
let before = statement[..index].trim_end();
let after = statement[index..].trim_start();
let separator = if sql_suffix_needs_newline(before) { "\n" } else { " " };
return format!("{before}{separator}{clause} {after};");
}
append_sql_suffix(statement, &format!("{clause};"))
}
const LOCKING_CLAUSE_PATTERNS: &[&[&str]] = &[
&["FOR", "UPDATE"],
&["FOR", "SHARE"],
&["FOR", "KEY", "SHARE"],
&["FOR", "NO", "KEY", "UPDATE"],
&["LOCK", "IN", "SHARE", "MODE"],
];
fn locking_clause_index(tokens: &[SqlToken]) -> Option<usize> {
tokens.iter().enumerate().find_map(|(index, token)| {
LOCKING_CLAUSE_PATTERNS
.iter()
.any(|pattern| token_sequence_matches(&tokens[index..], pattern))
.then_some(token.start)
})
}
fn token_sequence_matches(tokens: &[SqlToken], expected: &[&str]) -> bool {
tokens.len() >= expected.len() && tokens.iter().zip(expected).all(|(token, expected)| token.text == *expected)
}
fn has_pagination_clause_after(tokens: &[SqlToken], index: usize) -> bool {
tokens.iter().any(|token| token.start > index && matches!(token.text.as_str(), "LIMIT" | "OFFSET" | "FETCH"))
}
fn clickhouse_settings_clause_index(statement: &str) -> Option<usize> {
@ -1340,6 +1392,14 @@ fn top_level_sql_tokens(sql: &str) -> Vec<SqlToken> {
continue;
}
if ch == '#' {
i += 1;
while i < sql.len() && next_char(sql, i) != '\n' {
i += next_char(sql, i).len_utf8();
}
continue;
}
if ch == '/' && next == Some('*') {
i += 2;
while i < sql.len() {
@ -1354,6 +1414,15 @@ fn top_level_sql_tokens(sql: &str) -> Vec<SqlToken> {
continue;
}
// PostgreSQL dollar-quoted bodies may contain arbitrary SQL keywords.
// Skip them before scanning for top-level clauses such as FOR UPDATE.
if ch == '$' {
if let Some(end) = skip_sql_dollar_quoted(sql, i) {
i = end;
continue;
}
}
if matches!(ch, '\'' | '"' | '`') {
i = skip_sql_quoted(sql, i, ch);
continue;
@ -1392,6 +1461,22 @@ fn top_level_sql_tokens(sql: &str) -> Vec<SqlToken> {
tokens
}
fn skip_sql_dollar_quoted(sql: &str, pos: usize) -> Option<usize> {
let tag_end_offset = sql.get(pos + 1..)?.find('$')?;
let tag_end = pos + 1 + tag_end_offset;
let tag = &sql[pos + 1..tag_end];
let valid_tag = tag.is_empty()
|| (tag.chars().next().is_some_and(|ch| ch.is_ascii_alphabetic() || ch == '_')
&& tag.chars().all(|ch| ch.is_ascii_alphanumeric() || ch == '_'));
if !valid_tag {
return None;
}
let delimiter = &sql[pos..=tag_end];
let content_start = tag_end + 1;
sql.get(content_start..)?.find(delimiter).map(|closing_offset| content_start + closing_offset + delimiter.len())
}
fn skip_sql_quoted(sql: &str, pos: usize, quote: char) -> usize {
let mut i = pos + quote.len_utf8();
while i < sql.len() {
@ -2147,6 +2232,185 @@ WHERE u.id = picked.id;
assert_eq!(plan.page_offset, Some(0));
}
#[test]
fn mysql_for_update_places_limit_before_locking_clause() {
let result = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql: "SELECT * FROM `test`\nwhere id=1 for update".to_string(),
database_type: Some(DatabaseType::Mysql),
limit: 100,
offset: 0,
});
assert!(result.ok);
assert_eq!(result.sql.unwrap(), "SELECT * FROM `test`\nwhere id=1 LIMIT 100 for update;");
}
#[test]
fn locking_query_plan_keeps_server_pagination_and_count() {
let sql = "SELECT * FROM `test`\nwhere id=1 for update".to_string();
let plan = build_query_pagination_execution_plan(QueryPaginationExecutionPlanOptions {
sql: sql.clone(),
query_base_sql: sql.clone(),
database_type: Some(DatabaseType::Mysql),
pagination: QueryPagination { limit: 100, offset: 0, session_id: None },
use_agent_cursor: false,
first_page_uses_actual_sql: false,
});
assert_eq!(plan.sql_to_execute, "SELECT * FROM `test`\nwhere id=1 LIMIT 100 for update;");
assert_eq!(plan.page_sql, Some("SELECT * FROM `test`\nwhere id=1 LIMIT 100 for update;".to_string()));
assert_eq!(plan.page_limit, Some(100));
assert_eq!(plan.page_offset, Some(0));
assert_eq!(
plan.count_sql,
Some("SELECT COUNT(*) AS dbx_total_rows FROM (SELECT * FROM `test`\nwhere id=1) `dbx_count`;".to_string())
);
assert!(!plan.use_agent_result_session);
}
#[test]
fn locking_query_later_page_places_offset_before_locking_clause() {
let sql = "SELECT * FROM t WHERE deleted = 0 FOR UPDATE SKIP LOCKED".to_string();
let plan = build_query_pagination_execution_plan(QueryPaginationExecutionPlanOptions {
sql: sql.clone(),
query_base_sql: sql.clone(),
database_type: Some(DatabaseType::Mysql),
pagination: QueryPagination { limit: 100, offset: 100, session_id: None },
use_agent_cursor: false,
first_page_uses_actual_sql: false,
});
assert_eq!(
plan.sql_to_execute,
"SELECT * FROM t WHERE deleted = 0 LIMIT 100 OFFSET 100 FOR UPDATE SKIP LOCKED;"
);
assert_eq!(plan.page_limit, Some(100));
assert_eq!(plan.page_offset, Some(100));
}
#[test]
fn locking_query_count_removes_locking_clause() {
let result = build_count_query_sql(CountQuerySqlOptions {
original_sql: "SELECT * FROM t WHERE deleted = 0 FOR UPDATE".to_string(),
database_type: Some(DatabaseType::Mysql),
});
assert_eq!(
result.sql,
Some("SELECT COUNT(*) AS dbx_total_rows FROM (SELECT * FROM t WHERE deleted = 0) `dbx_count`;".to_string())
);
}
#[test]
fn locking_query_count_preserves_postgres_limit_after_lock_by_declining_rewrite() {
let result = build_count_query_sql(CountQuerySqlOptions {
original_sql: "SELECT * FROM t FOR UPDATE LIMIT 10".to_string(),
database_type: Some(DatabaseType::Postgres),
});
assert_eq!(result, err("locking"));
}
#[test]
fn places_limit_before_supported_top_level_locking_clause_variants() {
for (sql, expected) in [
("SELECT * FROM t FOR UPDATE", "SELECT * FROM t LIMIT 100 FOR UPDATE;"),
("SELECT * FROM t FOR SHARE", "SELECT * FROM t LIMIT 100 FOR SHARE;"),
("SELECT * FROM t FOR KEY SHARE", "SELECT * FROM t LIMIT 100 FOR KEY SHARE;"),
("SELECT * FROM t FOR NO KEY UPDATE", "SELECT * FROM t LIMIT 100 FOR NO KEY UPDATE;"),
("SELECT * FROM t LOCK IN SHARE MODE", "SELECT * FROM t LIMIT 100 LOCK IN SHARE MODE;"),
] {
let result = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql: sql.to_string(),
database_type: Some(DatabaseType::Postgres),
limit: 100,
offset: 0,
});
assert_eq!(result.sql.as_deref(), Some(expected), "{sql}");
}
}
#[test]
fn nested_for_update_does_not_block_outer_limit_append() {
let result = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql: "SELECT * FROM (SELECT id FROM t FOR UPDATE) locked".to_string(),
database_type: Some(DatabaseType::Mysql),
limit: 100,
offset: 0,
});
assert!(result.ok);
assert_eq!(result.sql.unwrap(), "SELECT * FROM (SELECT id FROM t FOR UPDATE) locked LIMIT 100;");
}
#[test]
fn for_xml_is_not_treated_as_locking_clause() {
let result = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql: "SELECT id, name FROM users FOR XML PATH('row')".to_string(),
database_type: Some(DatabaseType::Mysql),
limit: 100,
offset: 0,
});
assert!(result.ok);
assert_eq!(result.sql.unwrap(), "SELECT id, name FROM users FOR XML PATH('row') LIMIT 100;");
}
#[test]
fn ordinary_select_still_appends_limit() {
let result = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql: "SELECT * FROM t WHERE deleted = 0".to_string(),
database_type: Some(DatabaseType::Mysql),
limit: 100,
offset: 0,
});
assert!(result.ok);
assert_eq!(result.sql.unwrap(), "SELECT * FROM t WHERE deleted = 0 LIMIT 100;");
}
#[test]
fn locking_keywords_inside_postgres_dollar_quote_are_not_rewritten() {
let sql = "SELECT $$FOR UPDATE$$ AS message";
let paginated = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql: sql.to_string(),
database_type: Some(DatabaseType::Postgres),
limit: 100,
offset: 0,
});
let counted = build_count_query_sql(CountQuerySqlOptions {
original_sql: sql.to_string(),
database_type: Some(DatabaseType::Postgres),
});
assert_eq!(paginated.sql.as_deref(), Some("SELECT $$FOR UPDATE$$ AS message LIMIT 100;"));
assert_eq!(
counted.sql.as_deref(),
Some("SELECT COUNT(*) AS dbx_total_rows FROM (SELECT $$FOR UPDATE$$ AS message) \"dbx_count\";")
);
}
#[test]
fn locking_keywords_inside_mysql_hash_comment_are_not_rewritten() {
let sql = "SELECT * FROM t\n# FOR UPDATE LIMIT 1";
let paginated = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql: sql.to_string(),
database_type: Some(DatabaseType::Mysql),
limit: 100,
offset: 0,
});
let counted = build_count_query_sql(CountQuerySqlOptions {
original_sql: sql.to_string(),
database_type: Some(DatabaseType::Mysql),
});
assert_eq!(paginated.sql.as_deref(), Some("SELECT * FROM t\n# FOR UPDATE LIMIT 1\nLIMIT 100;"));
assert_eq!(
counted.sql.as_deref(),
Some("SELECT COUNT(*) AS dbx_total_rows FROM (SELECT * FROM t\n# FOR UPDATE LIMIT 1\n) `dbx_count`;")
);
}
#[test]
fn clickhouse_scalar_with_select_can_be_counted() {
let result = build_count_query_sql(CountQuerySqlOptions {