diff --git a/crates/dbx-core/src/db/sqlserver.rs b/crates/dbx-core/src/db/sqlserver.rs index 9c24894ef..1b2c939fc 100644 --- a/crates/dbx-core/src/db/sqlserver.rs +++ b/crates/dbx-core/src/db/sqlserver.rs @@ -471,9 +471,11 @@ async fn collect_first_result_limited( mut stream: QueryStream<'_>, start: Instant, max_rows: Option, + result_offset: usize, spatial_columns: &[SqlServerSpatialColumn], ) -> Result { let row_limit = query_result_row_limit(max_rows); + let mut remaining_offset = result_offset; let mut columns: Vec = vec![]; let mut column_types: Vec = vec![]; let mut rows: Vec> = 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, rows: Vec>, 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, + result_offset: usize, ) -> Result, String> { let row_limit = query_result_row_limit(max_rows); let mut results = Vec::new(); let mut current: Option = 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, ) -> Result { 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, ) -> Result, 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, ) -> Result, 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, ) -> Result { 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(), ); diff --git a/crates/dbx-core/src/query_result_sql.rs b/crates/dbx-core/src/query_result_sql.rs index bcd6bc051..f0f538387 100644 --- a/crates/dbx-core/src/query_result_sql.rs +++ b/crates/dbx-core/src/query_result_sql.rs @@ -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] diff --git a/crates/dbx-core/tests/sqlserver_batch_results.rs b/crates/dbx-core/tests/sqlserver_batch_results.rs index 96a9830cc..0ce20094e 100644 --- a/crates/dbx-core/tests/sqlserver_batch_results.rs +++ b/crates/dbx-core/tests/sqlserver_batch_results.rs @@ -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); +}