fix(sqlserver): support duplicate result columns

This commit is contained in:
二丫讲梵 2026-07-11 01:00:28 +08:00 committed by GitHub
parent 47ba8a2f3b
commit ee4658bac2
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
1 changed files with 132 additions and 12 deletions

View File

@ -1,3 +1,5 @@
use std::collections::HashSet;
use serde::{Deserialize, Serialize};
use crate::models::connection::DatabaseType;
@ -6,7 +8,7 @@ use crate::sql_dialect::{
firebird_rows_clause, pagination_strategy, quote_table_identifier, PaginationContext, TablePaginationStrategy,
};
use sqlparser::ast::{Expr, GroupByExpr, SelectItem, SetExpr, Statement};
use sqlparser::dialect::GenericDialect;
use sqlparser::dialect::{GenericDialect, MsSqlDialect};
use sqlparser::parser::Parser;
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
@ -166,7 +168,6 @@ pub fn build_paginated_query_sql(options: PaginatedQuerySqlOptions) -> QuerySqlB
if unsupported_pagination_type(options.database_type) {
return err("unsupported");
}
let safe_limit = options.limit.max(1);
let safe_offset = options.offset;
@ -215,6 +216,9 @@ pub fn build_count_query_sql(options: CountQuerySqlOptions) -> QuerySqlBuildResu
if options.database_type == Some(DatabaseType::Elasticsearch) {
return err("unsupported");
}
if options.database_type == Some(DatabaseType::SqlServer) && !sql_server_derived_table_projection_safe(&statement) {
return err("unsupported");
}
let alias = quote_table_identifier(options.database_type, "dbx_count");
let wrapped_sql = if options.database_type == Some(DatabaseType::SqlServer) {
@ -402,7 +406,8 @@ 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 Some(add_sql_server_existing_top_pagination(statement, limit, offset));
return sql_server_derived_table_projection_safe(statement)
.then(|| add_sql_server_existing_top_pagination(statement, limit, offset));
}
let order_by_index = find_top_level_trailing_order_by(statement);
@ -415,7 +420,7 @@ fn add_sql_server_offset_fetch(statement: &str, limit: usize, offset: usize) ->
}
let statement_without_order = order_by_index.map(|index| statement[..index].trim_end()).unwrap_or(statement);
if !sql_server_row_number_pagination_safe(statement_without_order) {
if !sql_server_derived_table_projection_safe(statement_without_order) {
return None;
}
@ -686,8 +691,8 @@ fn skip_sql_parenthesized(sql: &str, index: usize) -> usize {
sql.len()
}
fn sql_server_row_number_pagination_safe(statement: &str) -> bool {
let dialect = GenericDialect {};
fn sql_server_derived_table_projection_safe(statement: &str) -> bool {
let dialect = MsSqlDialect {};
let Ok(statements) = Parser::parse_sql(&dialect, statement) else {
return false;
};
@ -698,15 +703,33 @@ fn sql_server_row_number_pagination_safe(statement: &str) -> bool {
return false;
};
select.projection.iter().all(|item| match item {
SelectItem::Wildcard(_) | SelectItem::QualifiedWildcard(_, _) => false,
SelectItem::UnnamedExpr(expr) => sql_server_derived_projection_has_name(expr),
SelectItem::ExprWithAlias { .. } | SelectItem::ExprWithAliases { .. } => true,
if matches!(select.projection.as_slice(), [SelectItem::Wildcard(_)]) {
return select.from.len() == 1 && select.from[0].joins.is_empty();
}
let mut column_names = HashSet::with_capacity(select.projection.len());
select.projection.iter().all(|item| {
let Some(name) = sql_server_derived_projection_name(item) else {
return false;
};
column_names.insert(name.to_lowercase())
})
}
fn sql_server_derived_projection_has_name(expr: &Expr) -> bool {
matches!(expr, Expr::Identifier(_) | Expr::CompoundIdentifier(_))
fn sql_server_derived_projection_name(item: &SelectItem) -> Option<&str> {
match item {
SelectItem::ExprWithAlias { alias, .. } => Some(&alias.value),
SelectItem::UnnamedExpr(Expr::Identifier(identifier)) if !identifier.value.starts_with('@') => {
Some(&identifier.value)
}
SelectItem::UnnamedExpr(Expr::CompoundIdentifier(identifiers)) => {
identifiers.last().map(|identifier| identifier.value.as_str())
}
SelectItem::UnnamedExpr(_)
| SelectItem::ExprWithAliases { .. }
| SelectItem::QualifiedWildcard(_, _)
| SelectItem::Wildcard(_) => None,
}
}
fn add_sql_server_top(sql: &str, limit: usize) -> String {
@ -1351,6 +1374,31 @@ mod tests {
assert_eq!(result.sql.unwrap(), "SELECT TOP (100) COUNT(*) FROM TicketInfo");
}
#[test]
fn uses_sqlserver_top_for_unnamed_expression_on_first_page() {
let result = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql: "SELECT id + 1 FROM TicketInfo".to_string(),
database_type: Some(DatabaseType::SqlServer),
limit: 100,
offset: 0,
});
assert!(result.ok);
assert_eq!(result.sql.unwrap(), "SELECT TOP (100) id + 1 FROM TicketInfo");
}
#[test]
fn rejects_sqlserver_unnamed_expression_for_later_pages() {
let result = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql: "SELECT id + 1 FROM TicketInfo".to_string(),
database_type: Some(DatabaseType::SqlServer),
limit: 100,
offset: 100,
});
assert_eq!(result, err("unsupported"));
}
#[test]
fn uses_sqlserver_top_for_distinct_queries_without_order_by() {
let result = build_paginated_query_sql(PaginatedQuerySqlOptions {
@ -1458,6 +1506,78 @@ mod tests {
assert_eq!(plan.page_offset, Some(100));
}
#[test]
fn sqlserver_unsafe_projections_execute_original_sql_without_wrappers() {
let queries = [
"SELECT TOP 100 AAA, * FROM BBB",
"SELECT TOP 100 AAA, bbb AS aaa FROM BBB",
"SELECT TOP 100 a.id, b.id FROM AAA a JOIN BBB b ON a.id = b.id",
"SELECT TOP 100 a.*, b.* FROM AAA a JOIN BBB b ON a.id = b.id",
"SELECT TOP 100 a.*, * FROM AAA a",
"SELECT TOP 100 AAA + 1 FROM BBB",
];
for sql in queries {
for offset in [0, 100] {
let plan = build_query_pagination_execution_plan(QueryPaginationExecutionPlanOptions {
sql: sql.to_string(),
query_base_sql: sql.to_string(),
database_type: Some(DatabaseType::SqlServer),
pagination: QueryPagination { limit: 100, offset, session_id: None },
use_agent_cursor: false,
first_page_uses_actual_sql: false,
});
assert_eq!(plan.sql_to_execute, sql);
assert!(plan.page_sql.is_none());
assert!(plan.count_sql.is_none());
assert_eq!(plan.page_limit, None);
assert_eq!(plan.page_offset, None);
}
}
}
#[test]
fn sqlserver_unsafe_projection_is_not_wrapped_for_count() {
let result = build_count_query_sql(CountQuerySqlOptions {
original_sql: "SELECT TOP 100 AAA, * FROM BBB".to_string(),
database_type: Some(DatabaseType::SqlServer),
});
assert_eq!(result, err("unsupported"));
}
#[test]
fn sqlserver_unsafe_projection_is_not_paginated_directly() {
let result = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql: "SELECT TOP 100 AAA, * FROM BBB".to_string(),
database_type: Some(DatabaseType::SqlServer),
limit: 100,
offset: 100,
});
assert_eq!(result, err("unsupported"));
}
#[test]
fn sqlserver_unique_top_projection_keeps_server_pagination() {
let sql = "SELECT TOP 500 [id], [order_no] FROM [sales].[orders_10k]".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!(plan.sql_to_execute.contains("ROW_NUMBER()"));
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));
assert!(plan.count_sql.is_some());
}
#[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]";