fix: 修复 SQL Server 查询结果分页重复第一页 (#1927)

Co-authored-by: staff <staff@qimaos-MacBook-Pro.local>
This commit is contained in:
zipg 2026-06-26 15:58:10 +08:00 committed by GitHub
parent 4e5cb62102
commit 2aba40ecb7
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
1 changed files with 338 additions and 5 deletions

View File

@ -306,10 +306,12 @@ fn single_selectable_statement(original_sql: &str) -> Result<String, ()> {
if statement.is_empty() {
return Err(());
}
if statement.len() != base_sql.trim_end_matches(';').trim().len() {
if !single_statement_matches_base_sql(&statement, base_sql) {
return Err(());
}
let upper = statement.trim_start_matches(';').trim_start().to_ascii_uppercase();
let statement_without_leading_comments =
strip_leading_statement_comments(statement.trim_start_matches(';').trim_start());
let upper = statement_without_leading_comments.to_ascii_uppercase();
if upper.starts_with("WITH") {
if !cte_main_statement_is_select(&statement) {
return Err(());
@ -324,6 +326,17 @@ fn single_selectable_statement(original_sql: &str) -> Result<String, ()> {
Ok(statement)
}
fn single_statement_matches_base_sql(statement: &str, base_sql: &str) -> bool {
let normalized_statement = statement.trim().trim_end_matches(';').trim();
let normalized_base = base_sql.trim().trim_end_matches(';').trim();
if normalized_statement.len() == normalized_base.len() {
return true;
}
let base_without_leading_comments =
strip_leading_statement_comments(normalized_base).trim().trim_end_matches(';').trim();
normalized_statement == base_without_leading_comments
}
fn starts_with_cte(sql: &str) -> bool {
sql.trim_start().trim_start_matches(';').trim_start().to_ascii_uppercase().starts_with("WITH")
}
@ -402,7 +415,7 @@ fn add_sql_server_offset_fetch(statement: &str, limit: usize, offset: usize) ->
return (offset == 0).then(|| statement.to_string());
}
if has_top_level_select_top(statement) {
return (offset == 0).then(|| statement.to_string());
return Some(add_sql_server_existing_top_pagination(statement, limit, offset));
}
let order_by_index = find_top_level_trailing_order_by(statement);
@ -428,6 +441,237 @@ fn add_sql_server_offset_fetch(statement: &str, limit: usize, offset: usize) ->
))
}
fn add_sql_server_existing_top_pagination(statement: &str, limit: usize, offset: usize) -> String {
let row_number_order = format!("ORDER BY {}", sql_server_default_pagination_order(statement));
if offset == 0 {
return format!("SELECT TOP ({limit}) * FROM ({statement}) [dbx_page] {row_number_order};");
}
let end = offset + limit;
format!(
"SELECT * FROM (SELECT dbx_page_source.*, ROW_NUMBER() OVER ({row_number_order}) AS [__dbx_row_num] FROM ({statement}) dbx_page_source) dbx_page WHERE [__dbx_row_num] > {offset} AND [__dbx_row_num] <= {end} ORDER BY [__dbx_row_num];"
)
}
fn sql_server_default_pagination_order(statement: &str) -> String {
first_simple_sqlserver_projection(statement).unwrap_or_else(|| "(SELECT NULL)".to_string())
}
fn first_simple_sqlserver_projection(statement: &str) -> Option<String> {
let sql = statement.trim();
let sql = &sql[skip_leading_sql_comments(sql, 0)..];
if sql.len() < 6 || !sql[..6].eq_ignore_ascii_case("SELECT") {
return None;
}
let mut index = 6;
index = skip_sql_whitespace(sql, index);
if let Some(next) = skip_sql_keyword(sql, index, "DISTINCT").or_else(|| skip_sql_keyword(sql, index, "ALL")) {
index = skip_sql_whitespace(sql, next);
}
if let Some(next) = skip_sql_keyword(sql, index, "TOP") {
index = skip_sqlserver_top_clause(sql, next);
}
let projection_start = skip_sql_whitespace(sql, index);
let projection_end = find_first_projection_end(sql, projection_start)?;
let projection = sql[projection_start..projection_end].trim();
if is_simple_sqlserver_order_projection(projection) {
Some(projection.to_string())
} else {
None
}
}
fn skip_leading_sql_comments(sql: &str, mut index: usize) -> usize {
loop {
index = skip_sql_whitespace(sql, index);
if sql[index..].starts_with("--") {
index += 2;
while index < sql.len() && next_char(sql, index) != '\n' {
index += next_char(sql, index).len_utf8();
}
continue;
}
if sql[index..].starts_with("/*") {
index += 2;
while index < sql.len() {
let ch = next_char(sql, index);
let next = next_char_at(sql, index + ch.len_utf8());
index += ch.len_utf8();
if ch == '*' && next == Some('/') {
index += 1;
break;
}
}
continue;
}
return index;
}
}
fn strip_leading_statement_comments(sql: &str) -> &str {
&sql[skip_leading_sql_comments(sql, 0)..]
}
fn skip_sqlserver_top_clause(sql: &str, index: usize) -> usize {
let mut cursor = skip_sql_whitespace(sql, index);
if next_char_at(sql, cursor) == Some('(') {
cursor = skip_sql_parenthesized(sql, cursor);
} else {
while cursor < sql.len() && !next_char(sql, cursor).is_whitespace() && next_char(sql, cursor) != ',' {
cursor += next_char(sql, cursor).len_utf8();
}
}
cursor = skip_sql_whitespace(sql, cursor);
if let Some(next) = skip_sql_keyword(sql, cursor, "PERCENT") {
cursor = skip_sql_whitespace(sql, next);
}
if let Some(next) = skip_sql_keyword(sql, cursor, "WITH") {
let after_with = skip_sql_whitespace(sql, next);
if let Some(after_ties) = skip_sql_keyword(sql, after_with, "TIES") {
cursor = skip_sql_whitespace(sql, after_ties);
}
}
cursor
}
fn find_first_projection_end(sql: &str, start: usize) -> Option<usize> {
let mut index = start;
let mut depth = 0usize;
while index < sql.len() {
let ch = next_char(sql, index);
if matches!(ch, '\'' | '"' | '`') {
index = skip_sql_quoted(sql, index, ch);
continue;
}
if ch == '[' {
index = skip_sql_bracket_identifier(sql, index);
continue;
}
if ch == '(' {
depth += 1;
index += ch.len_utf8();
continue;
}
if ch == ')' {
depth = depth.saturating_sub(1);
index += ch.len_utf8();
continue;
}
if depth == 0 && ch == ',' {
return Some(index);
}
if depth == 0 && sql_keyword_at(sql, index, "FROM") {
return Some(index);
}
index += ch.len_utf8();
}
None
}
fn is_simple_sqlserver_order_projection(projection: &str) -> bool {
if projection.is_empty() || projection == "*" {
return false;
}
let mut expect_part = true;
let mut saw_part = false;
let mut index = 0;
while index < projection.len() {
let ch = next_char(projection, index);
if ch.is_whitespace() {
return false;
}
if ch == '.' {
if expect_part {
return false;
}
expect_part = true;
index += 1;
continue;
}
if ch == '[' {
if !expect_part {
return false;
}
let next = skip_sql_bracket_identifier(projection, index);
if next <= index + 1 || next > projection.len() {
return false;
}
saw_part = true;
expect_part = false;
index = next;
continue;
}
if is_sql_token_start(ch) {
if !expect_part {
return false;
}
index += ch.len_utf8();
while index < projection.len() && is_sql_token_part(next_char(projection, index)) {
index += next_char(projection, index).len_utf8();
}
saw_part = true;
expect_part = false;
continue;
}
return false;
}
saw_part && !expect_part
}
fn skip_sql_whitespace(sql: &str, mut index: usize) -> usize {
while index < sql.len() && next_char(sql, index).is_whitespace() {
index += next_char(sql, index).len_utf8();
}
index
}
fn skip_sql_keyword(sql: &str, index: usize, keyword: &str) -> Option<usize> {
sql_keyword_at(sql, index, keyword).then_some(index + keyword.len())
}
fn sql_keyword_at(sql: &str, index: usize, keyword: &str) -> bool {
let Some(candidate) = sql.get(index..index + keyword.len()) else {
return false;
};
if !candidate.eq_ignore_ascii_case(keyword) {
return false;
}
let before_ok = index == 0 || !is_sql_token_part(next_char_before(sql, index));
let after = index + keyword.len();
let after_ok = after >= sql.len() || !is_sql_token_part(next_char(sql, after));
before_ok && after_ok
}
fn skip_sql_parenthesized(sql: &str, index: usize) -> usize {
let mut cursor = index;
let mut depth = 0usize;
while cursor < sql.len() {
let ch = next_char(sql, cursor);
if matches!(ch, '\'' | '"' | '`') {
cursor = skip_sql_quoted(sql, cursor, ch);
continue;
}
if ch == '[' {
cursor = skip_sql_bracket_identifier(sql, cursor);
continue;
}
if ch == '(' {
depth += 1;
} else if ch == ')' {
depth = depth.saturating_sub(1);
cursor += ch.len_utf8();
if depth == 0 {
return cursor;
}
continue;
}
cursor += ch.len_utf8();
}
sql.len()
}
fn sql_server_row_number_pagination_safe(statement: &str) -> bool {
let dialect = GenericDialect {};
let Ok(statements) = Parser::parse_sql(&dialect, statement) else {
@ -858,6 +1102,10 @@ fn next_char(sql: &str, index: usize) -> char {
sql[index..].chars().next().unwrap_or('\0')
}
fn next_char_before(sql: &str, index: usize) -> char {
sql[..index].chars().next_back().unwrap_or('\0')
}
fn next_char_at(sql: &str, index: usize) -> Option<char> {
if index >= sql.len() {
None
@ -1000,7 +1248,7 @@ mod tests {
}
#[test]
fn keeps_existing_sqlserver_top_clause() {
fn paginates_existing_sqlserver_top_clause_on_first_page() {
let result = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql: "SELECT TOP 1000 * FROM TicketInfo".to_string(),
database_type: Some(DatabaseType::SqlServer),
@ -1009,7 +1257,92 @@ mod tests {
});
assert!(result.ok);
assert_eq!(result.sql.unwrap(), "SELECT TOP 1000 * FROM TicketInfo");
assert_eq!(
result.sql.unwrap(),
"SELECT TOP (100) * FROM (SELECT TOP 1000 * FROM TicketInfo) [dbx_page] ORDER BY (SELECT NULL);"
);
}
#[test]
fn paginates_existing_sqlserver_top_clause_for_later_pages() {
let result = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql: "SELECT TOP 1000 * FROM TicketInfo".to_string(),
database_type: Some(DatabaseType::SqlServer),
limit: 100,
offset: 100,
});
assert!(result.ok);
assert_eq!(
result.sql.unwrap(),
"SELECT * FROM (SELECT dbx_page_source.*, ROW_NUMBER() OVER (ORDER BY (SELECT NULL)) AS [__dbx_row_num] FROM (SELECT TOP 1000 * FROM TicketInfo) dbx_page_source) dbx_page WHERE [__dbx_row_num] > 100 AND [__dbx_row_num] <= 200 ORDER BY [__dbx_row_num];"
);
}
#[test]
fn sqlserver_top_query_later_page_keeps_page_metadata() {
let sql = "SELECT TOP 1000 * FROM TicketInfo".to_string();
let plan = build_query_pagination_execution_plan(QueryPaginationExecutionPlanOptions {
sql: sql.clone(),
query_base_sql: sql,
database_type: Some(DatabaseType::SqlServer),
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 (SELECT dbx_page_source.*, ROW_NUMBER() OVER (ORDER BY (SELECT NULL)) AS [__dbx_row_num] FROM (SELECT TOP 1000 * FROM TicketInfo) dbx_page_source) dbx_page WHERE [__dbx_row_num] > 100 AND [__dbx_row_num] <= 200 ORDER BY [__dbx_row_num];"
);
assert_eq!(plan.page_sql, Some(plan.sql_to_execute.clone()));
assert_eq!(plan.page_limit, Some(100));
assert_eq!(plan.page_offset, Some(100));
}
#[test]
fn paginates_sqlserver_top_parenthesized_projection_query_by_first_column() {
let sql = "SELECT TOP (500) [id], [order_no], [store_id], [product_id], [customer_name], [quantity], [amount], [order_status], [created_at] FROM [sales].[orders_10k]";
let first_page = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql: sql.to_string(),
database_type: Some(DatabaseType::SqlServer),
limit: 100,
offset: 0,
});
let second_page = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql: sql.to_string(),
database_type: Some(DatabaseType::SqlServer),
limit: 100,
offset: 100,
});
assert!(first_page.ok);
assert!(second_page.ok);
assert_eq!(
first_page.sql.unwrap(),
"SELECT TOP (100) * FROM (SELECT TOP (500) [id], [order_no], [store_id], [product_id], [customer_name], [quantity], [amount], [order_status], [created_at] FROM [sales].[orders_10k]) [dbx_page] ORDER BY [id];"
);
assert_eq!(
second_page.sql.unwrap(),
"SELECT * FROM (SELECT dbx_page_source.*, ROW_NUMBER() OVER (ORDER BY [id]) AS [__dbx_row_num] FROM (SELECT TOP (500) [id], [order_no], [store_id], [product_id], [customer_name], [quantity], [amount], [order_status], [created_at] FROM [sales].[orders_10k]) dbx_page_source) dbx_page WHERE [__dbx_row_num] > 100 AND [__dbx_row_num] <= 200 ORDER BY [__dbx_row_num];"
);
}
#[test]
fn paginates_sqlserver_top_query_after_leading_comment_by_first_column() {
let sql = "-- 测试\nSELECT TOP (500) [id], [order_no], [store_id], [product_id], [customer_name], [quantity], [amount], [order_status], [created_at] FROM [sales].[orders_10k]";
let result = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql: sql.to_string(),
database_type: Some(DatabaseType::SqlServer),
limit: 100,
offset: 100,
});
assert!(result.ok);
assert_eq!(
result.sql.unwrap(),
"SELECT * FROM (SELECT dbx_page_source.*, ROW_NUMBER() OVER (ORDER BY [id]) AS [__dbx_row_num] FROM (SELECT TOP (500) [id], [order_no], [store_id], [product_id], [customer_name], [quantity], [amount], [order_status], [created_at] FROM [sales].[orders_10k]) dbx_page_source) dbx_page WHERE [__dbx_row_num] > 100 AND [__dbx_row_num] <= 200 ORDER BY [__dbx_row_num];"
);
}
#[test]