fix(sqlserver): paginate joins with duplicate result columns

Closes #5284
This commit is contained in:
t8y2 2026-08-04 12:14:49 +08:00
parent b6e1d6d9c4
commit f8c71e61dc
No known key found for this signature in database
3 changed files with 217 additions and 31 deletions

View File

@ -471,9 +471,11 @@ async fn collect_first_result_limited(
mut stream: QueryStream<'_>,
start: Instant,
max_rows: Option<usize>,
result_offset: usize,
spatial_columns: &[SqlServerSpatialColumn],
) -> Result<QueryResult, String> {
let row_limit = query_result_row_limit(max_rows);
let mut remaining_offset = result_offset;
let mut columns: Vec<String> = vec![];
let mut column_types: Vec<String> = vec![];
let mut rows: Vec<Vec<serde_json::Value>> = Vec::new();
@ -491,6 +493,10 @@ async fn collect_first_result_limited(
}
QueryItem::Metadata(_) => {}
QueryItem::Row(row) if row.result_index() == 0 => {
if remaining_offset > 0 {
remaining_offset -= 1;
continue;
}
if rows.len() < row_limit {
let (values, srids) =
row_to_json_with_spatial_metadata(&row, spatial_columns, |column_index, srid| {
@ -529,6 +535,7 @@ struct SqlServerResultSet {
column_types: Vec<String>,
rows: Vec<Vec<serde_json::Value>>,
truncated: bool,
remaining_offset: usize,
}
pub struct SqlServerStreamExportSummary {
@ -1226,20 +1233,25 @@ async fn collect_result_sets_limited(
mut stream: QueryStream<'_>,
start: Instant,
max_rows: Option<usize>,
result_offset: usize,
) -> Result<Vec<QueryResult>, String> {
let row_limit = query_result_row_limit(max_rows);
let mut results = Vec::new();
let mut current: Option<SqlServerResultSet> = None;
let mut saw_result_set = false;
while let Some(item) = stream.try_next().await.map_err(|e| e.to_string())? {
match item {
QueryItem::Metadata(metadata) => {
push_sqlserver_result_set(&mut results, current.take(), start);
let remaining_offset = if saw_result_set { 0 } else { result_offset };
saw_result_set = true;
current = Some(SqlServerResultSet {
columns: columns_from_metadata(&metadata),
column_types: column_types_from_metadata(&metadata),
rows: Vec::new(),
truncated: false,
remaining_offset,
});
}
QueryItem::Row(row) => {
@ -1248,7 +1260,13 @@ async fn collect_result_sets_limited(
column_types: row.columns().iter().map(sqlserver_column_type_name).collect(),
rows: Vec::new(),
truncated: false,
remaining_offset: if saw_result_set { 0 } else { result_offset },
});
saw_result_set = true;
if result.remaining_offset > 0 {
result.remaining_offset -= 1;
continue;
}
if result.rows.len() < row_limit {
result.rows.push(row_to_json(&row));
} else {
@ -2478,6 +2496,7 @@ pub async fn execute_query_with_max_rows(
max_rows: Option<usize>,
) -> Result<QueryResult, String> {
let start = Instant::now();
let result_offset = crate::query_result_sql::sqlserver_result_offset(sql);
if starts_with_executable_sql_keyword(sql, &["SELECT", "EXEC", "WITH", "TABLE"])
|| sqlserver_dml_output_returns_rows(sql)
@ -2490,7 +2509,14 @@ pub async fn execute_query_with_max_rows(
};
let (result, messages) = capture_sqlserver_messages(async {
let stream = sqlserver_driver_result(client.query(query.sql.as_str(), &[])).await?;
sqlserver_driver_result(collect_first_result_limited(stream, start, max_rows, &query.spatial_columns)).await
sqlserver_driver_result(collect_first_result_limited(
stream,
start,
max_rows,
result_offset,
&query.spatial_columns,
))
.await
})
.await;
let mut result = query_result_with_server_messages(result?, messages);
@ -2531,6 +2557,7 @@ pub async fn execute_batch_with_max_rows(
max_rows: Option<usize>,
) -> Result<Vec<QueryResult>, String> {
let start = Instant::now();
let result_offset = crate::query_result_sql::sqlserver_result_offset(sql);
if sqlserver_batch_can_use_execute(sql) {
let (result, messages) = capture_sqlserver_messages(sqlserver_driver_result(client.execute(sql, &[]))).await;
let result = result?;
@ -2562,6 +2589,7 @@ pub async fn execute_batch_with_max_rows(
stream,
start,
max_rows,
result_offset,
&query.spatial_columns,
))
.await
@ -2591,9 +2619,10 @@ pub async fn execute_simple_batch_with_max_rows(
max_rows: Option<usize>,
) -> Result<Vec<QueryResult>, String> {
let start = Instant::now();
let result_offset = crate::query_result_sql::sqlserver_result_offset(sql);
let (results, messages) = capture_sqlserver_messages(async {
let stream = sqlserver_driver_result(client.simple_query(sql)).await?;
sqlserver_driver_result(collect_result_sets_limited(stream, start, max_rows)).await
sqlserver_driver_result(collect_result_sets_limited(stream, start, max_rows, result_offset)).await
})
.await;
let mut results = results?;
@ -2629,9 +2658,10 @@ async fn execute_simple_batch_first_result_with_max_rows(
max_rows: Option<usize>,
) -> Result<QueryResult, String> {
let start = Instant::now();
let result_offset = crate::query_result_sql::sqlserver_result_offset(sql);
let (result, messages) = capture_sqlserver_messages(async {
let stream = sqlserver_driver_result(client.simple_query(sql)).await?;
sqlserver_driver_result(collect_first_result_limited(stream, start, max_rows, &[])).await
sqlserver_driver_result(collect_first_result_limited(stream, start, max_rows, result_offset, &[])).await
})
.await;
let mut result = query_result_with_server_messages(result?, messages);
@ -3214,7 +3244,7 @@ mod tests {
let first_result = source.split("async fn execute_simple_batch_first_result_with_max_rows").nth(1).unwrap();
let first_result = first_result.split("fn strip_dbx_sqlserver_row_number_column").next().unwrap();
assert!(first_result.contains("collect_first_result_limited(stream, start, max_rows, &[])"));
assert!(first_result.contains("collect_first_result_limited(stream, start, max_rows, result_offset, &[])"));
assert!(!first_result.contains("collect_result_sets_limited"));
}
@ -4106,6 +4136,7 @@ mod tests {
column_types: vec![],
rows: vec![],
truncated: false,
remaining_offset: 0,
}),
Instant::now(),
);
@ -4120,7 +4151,13 @@ mod tests {
let mut results = Vec::new();
super::push_sqlserver_result_set(
&mut results,
Some(SqlServerResultSet { columns: vec![], column_types: vec![], rows: vec![], truncated: false }),
Some(SqlServerResultSet {
columns: vec![],
column_types: vec![],
rows: vec![],
truncated: false,
remaining_offset: 0,
}),
Instant::now(),
);

View File

@ -439,8 +439,11 @@ 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 sql_server_derived_table_projection_safe(statement)
.then(|| add_sql_server_existing_top_pagination(statement, limit, offset));
return if sql_server_derived_table_projection_safe(statement) {
Some(add_sql_server_existing_top_pagination(statement, limit, offset))
} else {
(offset > 0).then(|| add_sql_server_rowcount_pagination(statement, limit, offset))
};
}
let order_by_index = find_top_level_trailing_order_by(statement);
@ -454,7 +457,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_derived_table_projection_safe(statement_without_order) {
return None;
return Some(add_sql_server_rowcount_pagination(statement, limit, offset));
}
let row_number_order = order_by_index
@ -466,6 +469,33 @@ fn add_sql_server_offset_fetch(statement: &str, limit: usize, offset: usize) ->
))
}
const SQLSERVER_RESULT_OFFSET_PREFIX: &str = "/*__dbx_result_offset=";
const SQLSERVER_RESULT_OFFSET_SUFFIX: &str = "__*/";
fn add_sql_server_rowcount_pagination(statement: &str, limit: usize, offset: usize) -> String {
let row_count = offset.saturating_add(limit);
let escaped_statement = statement.replace('\'', "''");
// Keep duplicate result-column names intact while bounding the server response
// on every SQL Server version supported by DBX. The dynamic batch scopes
// SET ROWCOUNT to this execution instead of leaking it into the tab session.
format!(
"EXEC sys.sp_executesql N'SET ROWCOUNT {row_count}; {escaped_statement}'; {SQLSERVER_RESULT_OFFSET_PREFIX}{offset}{SQLSERVER_RESULT_OFFSET_SUFFIX}"
)
}
pub(crate) fn sqlserver_result_offset(sql: &str) -> usize {
let sql = sql.trim_end();
if !sql.starts_with("EXEC sys.sp_executesql N'SET ROWCOUNT ") || !sql.ends_with(SQLSERVER_RESULT_OFFSET_SUFFIX) {
return 0;
}
let Some(marker_index) = sql.rfind(SQLSERVER_RESULT_OFFSET_PREFIX) else {
return 0;
};
let value_start = marker_index + SQLSERVER_RESULT_OFFSET_PREFIX.len();
let value_end = sql.len() - SQLSERVER_RESULT_OFFSET_SUFFIX.len();
sql[value_start..value_end].parse().unwrap_or(0)
}
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 {
@ -1472,7 +1502,7 @@ mod tests {
}
#[test]
fn rejects_sqlserver_unnamed_expression_for_later_pages() {
fn paginates_sqlserver_unnamed_expression_with_rowcount() {
let result = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql: "SELECT id + 1 FROM TicketInfo".to_string(),
database_type: Some(DatabaseType::SqlServer),
@ -1480,7 +1510,9 @@ mod tests {
offset: 100,
});
assert_eq!(result, err("unsupported"));
let sql = result.sql.expect("build unnamed expression page");
assert!(sql.starts_with("EXEC sys.sp_executesql N'SET ROWCOUNT 200; SELECT id + 1 FROM TicketInfo'"));
assert_eq!(sqlserver_result_offset(&sql), 100);
}
#[test]
@ -1591,7 +1623,7 @@ mod tests {
}
#[test]
fn sqlserver_unsafe_projections_execute_original_sql_without_wrappers() {
fn sqlserver_unsafe_top_projections_use_rowcount_for_later_pages() {
let queries = [
"SELECT TOP 100 AAA, * FROM BBB",
"SELECT TOP 100 AAA, bbb AS aaa FROM BBB",
@ -1602,22 +1634,32 @@ mod tests {
];
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,
});
let first_page = 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: 0, session_id: None },
use_agent_cursor: false,
first_page_uses_actual_sql: false,
});
assert_eq!(first_page.sql_to_execute, sql);
assert!(first_page.page_sql.is_none());
assert!(first_page.count_sql.is_none());
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);
}
let later_page = 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: 100, session_id: None },
use_agent_cursor: false,
first_page_uses_actual_sql: false,
});
assert!(later_page.sql_to_execute.starts_with("EXEC sys.sp_executesql N'SET ROWCOUNT 200; "));
assert_eq!(sqlserver_result_offset(&later_page.sql_to_execute), 100);
assert_eq!(later_page.page_sql, Some(later_page.sql_to_execute.clone()));
assert!(later_page.count_sql.is_none());
assert_eq!(later_page.page_limit, Some(100));
assert_eq!(later_page.page_offset, Some(100));
}
}
@ -1632,15 +1674,42 @@ mod tests {
}
#[test]
fn sqlserver_unsafe_projection_is_not_paginated_directly() {
fn sqlserver_join_wildcard_uses_bounded_rowcount_pagination() {
let sql = "SELECT * FROM WZ_CKGL_WZLLDSQ_DETAIL d LEFT JOIN WZ_CKGL_WZLLDSQ_MAIN AS m ON m.ID = d.ParentID";
let result = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql: "SELECT TOP 100 AAA, * FROM BBB".to_string(),
original_sql: sql.to_string(),
database_type: Some(DatabaseType::SqlServer),
limit: 500,
offset: 500,
});
assert!(result.ok);
let sql = result.sql.unwrap();
assert_eq!(
sql,
"EXEC sys.sp_executesql N'SET ROWCOUNT 1000; SELECT * FROM WZ_CKGL_WZLLDSQ_DETAIL d LEFT JOIN WZ_CKGL_WZLLDSQ_MAIN AS m ON m.ID = d.ParentID'; /*__dbx_result_offset=500__*/"
);
assert_eq!(sqlserver_result_offset(&sql), 500);
}
#[test]
fn sqlserver_result_offset_ignores_user_sql_markers() {
assert_eq!(sqlserver_result_offset("SELECT 1 /*__dbx_result_offset=500__*/"), 0);
}
#[test]
fn sqlserver_rowcount_pagination_escapes_string_literals() {
let result = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql: "SELECT * FROM detail d JOIN parent p ON p.id = d.parent_id WHERE p.label = N'O''Brien'"
.to_string(),
database_type: Some(DatabaseType::SqlServer),
limit: 100,
offset: 100,
});
assert_eq!(result, err("unsupported"));
let sql = result.sql.expect("build rowcount pagination SQL");
assert!(sql.contains("WHERE p.label = N''O''''Brien''"));
assert_eq!(sqlserver_result_offset(&sql), 100);
}
#[test]
@ -1998,7 +2067,7 @@ WHERE u.id = picked.id;
}
#[test]
fn rejects_sqlserver_wildcard_later_pages_to_avoid_duplicate_columns() {
fn paginates_sqlserver_wildcard_later_pages_without_derived_tables() {
let result = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql:
"SELECT b.ProjectType,* FROM VesselBusinessOpportunity a LEFT JOIN JDDR_sys_BasicConfig_ProjectInfo_Data b ON a.ProjectID = b.ID"
@ -2008,7 +2077,9 @@ WHERE u.id = picked.id;
offset: 100,
});
assert_eq!(result, err("unsupported"));
let sql = result.sql.expect("build wildcard page");
assert!(sql.starts_with("EXEC sys.sp_executesql N'SET ROWCOUNT 200; SELECT b.ProjectType,*"));
assert_eq!(sqlserver_result_offset(&sql), 100);
}
#[test]

View File

@ -1,4 +1,6 @@
use dbx_core::db::sqlserver;
use dbx_core::models::connection::DatabaseType;
use dbx_core::query_result_sql::{build_paginated_query_sql, PaginatedQuerySqlOptions};
use std::time::Duration;
async fn connect_sqlserver() -> sqlserver::SqlServerClient {
@ -71,3 +73,79 @@ async fn sqlserver_single_result_query_drains_later_results() {
.expect("execute query after draining later results");
assert_eq!(follow_up.rows[0][0].as_i64(), Some(2));
}
#[tokio::test]
#[ignore = "requires DBX_TEST_SQLSERVER_HOST and DBX_TEST_SQLSERVER_PASSWORD"]
async fn sqlserver_duplicate_join_columns_page_without_leaking_rowcount() {
let mut client = connect_sqlserver().await;
sqlserver::execute_simple_batch_with_max_rows(
&mut client,
"CREATE TABLE #dbx_page_detail (id INT NOT NULL PRIMARY KEY, parent_id INT NOT NULL); \
CREATE TABLE #dbx_page_main (id INT NOT NULL PRIMARY KEY, label NVARCHAR(20) NOT NULL); \
WITH numbers AS ( \
SELECT TOP (1200) ROW_NUMBER() OVER (ORDER BY (SELECT NULL)) AS value \
FROM sys.all_objects a CROSS JOIN sys.all_objects b \
) \
INSERT INTO #dbx_page_main (id, label) SELECT value, CONCAT(N'label-', value) FROM numbers; \
INSERT INTO #dbx_page_detail (id, parent_id) SELECT id, id FROM #dbx_page_main;",
None,
)
.await
.expect("create pagination fixtures");
let reported_query = "SELECT * FROM #dbx_page_detail d LEFT JOIN #dbx_page_main m ON m.id = d.parent_id";
let reported_page = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql: reported_query.to_string(),
database_type: Some(DatabaseType::SqlServer),
limit: 500,
offset: 500,
});
let reported_results = sqlserver::execute_batch_with_max_rows(
&mut client,
reported_page.sql.as_deref().expect("build reported pagination SQL"),
Some(500),
)
.await
.expect("execute reported duplicate-column page");
assert_eq!(reported_results.first().expect("reported page result").rows.len(), 500);
let query = format!("{reported_query} ORDER BY d.id");
let page = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql: query,
database_type: Some(DatabaseType::SqlServer),
limit: 500,
offset: 500,
});
let page_sql = page.sql.expect("build duplicate-column pagination SQL");
let results = sqlserver::execute_batch_with_max_rows(&mut client, &page_sql, Some(500))
.await
.expect("execute duplicate-column page");
let result = results.first().expect("page result");
assert_eq!(result.rows.len(), 500);
assert_eq!(result.rows.first().and_then(|row| row[0].as_i64()), Some(501));
assert_eq!(result.rows.last().and_then(|row| row[0].as_i64()), Some(1000));
let failing_page = build_paginated_query_sql(PaginatedQuerySqlOptions {
original_sql:
"SELECT * FROM #dbx_page_detail d LEFT JOIN #dbx_missing_page_table m ON m.id = d.parent_id ORDER BY d.id"
.to_string(),
database_type: Some(DatabaseType::SqlServer),
limit: 1,
offset: 1,
});
let error = sqlserver::execute_batch_with_max_rows(
&mut client,
failing_page.sql.as_deref().expect("build failing pagination SQL"),
Some(1),
)
.await
.expect_err("missing joined table should fail");
assert!(error.to_ascii_lowercase().contains("dbx_missing_page_table"));
let follow_up =
sqlserver::execute_query(&mut client, "SELECT value FROM (VALUES (1), (2)) rows(value) ORDER BY value")
.await
.expect("execute query after ROWCOUNT page");
assert_eq!(follow_up.rows.len(), 2);
}