From 7283f2b34af17b1b30cdef4ba0b5b35dccb8b8a2 Mon Sep 17 00:00:00 2001 From: t8y2 <1156263951@qq.com> Date: Fri, 5 Jun 2026 11:11:14 +0800 Subject: [PATCH] fix(sqlserver): handle spatial query results --- crates/dbx-core/src/db/sqlserver.rs | 442 +++++++++++++++++++++++++++- 1 file changed, 432 insertions(+), 10 deletions(-) diff --git a/crates/dbx-core/src/db/sqlserver.rs b/crates/dbx-core/src/db/sqlserver.rs index dc18fc0c3..096448bd2 100644 --- a/crates/dbx-core/src/db/sqlserver.rs +++ b/crates/dbx-core/src/db/sqlserver.rs @@ -1,8 +1,10 @@ use crate::query::MAX_ROWS; use crate::sql::starts_with_executable_sql_keyword; use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo}; -use futures::TryStreamExt; +use futures::{FutureExt, TryStreamExt}; use rust_decimal::Decimal; +use std::future::Future; +use std::panic::AssertUnwindSafe; use std::time::{Duration, Instant}; use tiberius::{AuthMethod, Client, ColumnData, Config, FromSql, QueryItem, QueryStream, SqlBrowser}; use tokio::net::TcpStream; @@ -140,6 +142,324 @@ struct SqlServerResultSet { truncated: bool, } +#[derive(Debug, Clone, PartialEq, Eq)] +struct SqlServerDescribedColumn { + name: Option, + system_type_name: Option, + user_type_schema: Option, + user_type_name: Option, +} + +async fn sqlserver_driver_result(future: F) -> Result +where + F: Future>, + E: ToString, +{ + match AssertUnwindSafe(future).catch_unwind().await { + Ok(result) => result.map_err(|e| e.to_string()), + Err(_) => { + Err("SQL Server driver could not decode this result set. Unsupported columns may need to be cast to text." + .to_string()) + } + } +} + +async fn describe_sqlserver_result_set( + client: &mut SqlServerClient, + sql: &str, +) -> Result, String> { + let describe_sql = "\ + SELECT name, system_type_name, user_type_schema, user_type_name \ + FROM sys.dm_exec_describe_first_result_set(@P1, NULL, 0) \ + WHERE error_number IS NULL AND is_hidden = 0 \ + ORDER BY column_ordinal"; + let stream = sqlserver_driver_result(client.query(describe_sql, &[&sql])).await?; + let rows = sqlserver_driver_result(stream.into_first_result()).await?; + + Ok(rows + .iter() + .map(|row| SqlServerDescribedColumn { + name: row.try_get::<&str, _>(0).ok().flatten().map(str::to_string), + system_type_name: row.try_get::<&str, _>(1).ok().flatten().map(str::to_string), + user_type_schema: row.try_get::<&str, _>(2).ok().flatten().map(str::to_string), + user_type_name: row.try_get::<&str, _>(3).ok().flatten().map(str::to_string), + }) + .collect()) +} + +async fn spatial_safe_sqlserver_query(client: &mut SqlServerClient, sql: &str) -> Result, String> { + if !is_single_sqlserver_select(sql) { + return Ok(None); + } + let columns = describe_sqlserver_result_set(client, sql).await?; + Ok(build_spatial_safe_sqlserver_query(sql, &columns)) +} + +fn build_spatial_safe_sqlserver_query(sql: &str, columns: &[SqlServerDescribedColumn]) -> Option { + if columns.is_empty() || !columns.iter().any(is_sqlserver_spatial_column) { + return None; + } + let statement = normalized_sqlserver_select_statement(sql)?; + let source_alias = quote_sqlserver_identifier("dbx_spatial_source"); + let source_columns = (0..columns.len()).map(sqlserver_source_column_name).collect::>(); + let source_alias_list = + source_columns.iter().map(|name| quote_sqlserver_identifier(name)).collect::>().join(", "); + let select_list = columns + .iter() + .enumerate() + .map(|(index, column)| { + let output_name = sqlserver_output_column_name(column, index); + let quoted_output = quote_sqlserver_identifier(&output_name); + let source_column = quote_sqlserver_identifier(&source_columns[index]); + let value_ref = format!("{source_alias}.{source_column}"); + if is_sqlserver_spatial_column(column) { + format!("{quoted_output} = CASE WHEN {value_ref} IS NULL THEN NULL ELSE {value_ref}.STAsText() END") + } else { + format!("{quoted_output} = {value_ref}") + } + }) + .collect::>() + .join(", "); + + Some(format!("SELECT {select_list} FROM ({statement}) AS {source_alias}({source_alias_list})")) +} + +fn is_sqlserver_spatial_column(column: &SqlServerDescribedColumn) -> bool { + [&column.system_type_name, &column.user_type_name].into_iter().flatten().any(|name| { + let normalized = name.trim().trim_matches(['[', ']']).to_ascii_lowercase(); + normalized == "geometry" + || normalized == "geography" + || normalized.ends_with(".geometry") + || normalized.ends_with(".geography") + }) +} + +fn normalized_sqlserver_select_statement(sql: &str) -> Option { + let statement = trim_sqlserver_statement(sql); + let trimmed = statement.trim_start(); + if trimmed.is_empty() || !trimmed.get(..6).is_some_and(|prefix| prefix.eq_ignore_ascii_case("SELECT")) { + return None; + } + if has_top_level_select_into(trimmed) { + return None; + } + + let mut statement = trimmed.to_string(); + if let Some(order_by) = find_top_level_trailing_order_by(&statement) { + if !has_top_level_select_top(&statement) + && !has_top_level_offset_after(&statement, order_by) + && !has_top_level_for_xml(&statement) + { + statement.push_str(" OFFSET 0 ROWS"); + } + } + Some(statement) +} + +fn trim_sqlserver_statement(sql: &str) -> String { + let mut statement = sql.trim(); + while let Some(stripped) = statement.strip_suffix(';') { + statement = stripped.trim_end(); + } + statement.to_string() +} + +fn is_single_sqlserver_select(sql: &str) -> bool { + let statements = crate::sql::split_sql_statements(sql); + if statements.len() != 1 { + return false; + } + let statement = statements[0].trim_start(); + statement.get(..6).is_some_and(|prefix| prefix.eq_ignore_ascii_case("SELECT")) +} + +fn sqlserver_source_column_name(index: usize) -> String { + format!("dbx_col_{}", index + 1) +} + +fn sqlserver_output_column_name(column: &SqlServerDescribedColumn, index: usize) -> String { + column + .name + .as_deref() + .map(str::trim) + .filter(|name| !name.is_empty()) + .map(str::to_string) + .unwrap_or_else(|| format!("column_{}", index + 1)) +} + +fn quote_sqlserver_identifier(identifier: &str) -> String { + format!("[{}]", identifier.replace(']', "]]")) +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct SqlServerToken { + text: String, + start: usize, +} + +fn top_level_sqlserver_tokens(sql: &str) -> Vec { + let mut tokens = Vec::new(); + let mut i = 0; + let mut depth = 0usize; + + while i < sql.len() { + let ch = next_char(sql, i); + let next = next_char_at(sql, i + ch.len_utf8()); + + if ch == '-' && next == Some('-') { + i += 2; + 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() { + let current = next_char(sql, i); + let following = next_char_at(sql, i + current.len_utf8()); + if current == '*' && following == Some('/') { + i += 2; + break; + } + i += current.len_utf8(); + } + continue; + } + if matches!(ch, '\'' | '"') { + i = skip_sqlserver_quoted(sql, i, ch); + continue; + } + if ch == '[' { + i = skip_sqlserver_bracket_identifier(sql, i); + continue; + } + if ch == '(' { + depth += 1; + i += ch.len_utf8(); + continue; + } + if ch == ')' { + depth = depth.saturating_sub(1); + i += ch.len_utf8(); + continue; + } + if depth == 0 && is_sqlserver_token_start(ch) { + let start = i; + i += ch.len_utf8(); + while i < sql.len() && is_sqlserver_token_part(next_char(sql, i)) { + i += next_char(sql, i).len_utf8(); + } + tokens.push(SqlServerToken { text: sql[start..i].to_ascii_uppercase(), start }); + continue; + } + i += ch.len_utf8(); + } + + tokens +} + +fn has_top_level_select_into(sql: &str) -> bool { + let tokens = top_level_sqlserver_tokens(sql); + let Some(select_index) = tokens.iter().position(|token| token.text == "SELECT") else { + return false; + }; + let from_index = tokens + .iter() + .enumerate() + .find(|(index, token)| *index > select_index && token.text == "FROM") + .map(|(index, _)| index) + .unwrap_or(tokens.len()); + tokens[select_index + 1..from_index].iter().any(|token| token.text == "INTO") +} + +fn has_top_level_select_top(sql: &str) -> bool { + let tokens = top_level_sqlserver_tokens(sql); + let Some(select_index) = tokens.iter().position(|token| token.text == "SELECT") else { + return false; + }; + let from_index = tokens + .iter() + .enumerate() + .find(|(index, token)| *index > select_index && token.text == "FROM") + .map(|(index, _)| index) + .unwrap_or(tokens.len()); + tokens[select_index + 1..from_index].iter().any(|token| token.text == "TOP") +} + +fn has_top_level_for_xml(sql: &str) -> bool { + let tokens = top_level_sqlserver_tokens(sql); + tokens.windows(2).any(|tokens| tokens[0].text == "FOR" && tokens[1].text == "XML") +} + +fn has_top_level_offset_after(sql: &str, start: usize) -> bool { + top_level_sqlserver_tokens(sql).into_iter().any(|token| token.start > start && token.text == "OFFSET") +} + +fn find_top_level_trailing_order_by(sql: &str) -> Option { + let tokens = top_level_sqlserver_tokens(sql); + for index in (0..tokens.len().saturating_sub(1)).rev() { + if tokens[index].text == "ORDER" && tokens.get(index + 1).is_some_and(|token| token.text == "BY") { + return Some(tokens[index].start); + } + } + None +} + +fn skip_sqlserver_quoted(sql: &str, pos: usize, quote: char) -> usize { + let mut i = pos + quote.len_utf8(); + while i < sql.len() { + let ch = next_char(sql, i); + let next = next_char_at(sql, i + ch.len_utf8()); + if ch == quote { + if next == Some(quote) { + i += ch.len_utf8() + quote.len_utf8(); + continue; + } + return i + ch.len_utf8(); + } + i += ch.len_utf8(); + } + sql.len() +} + +fn skip_sqlserver_bracket_identifier(sql: &str, pos: usize) -> usize { + let mut i = pos + 1; + while i < sql.len() { + let ch = next_char(sql, i); + let next = next_char_at(sql, i + ch.len_utf8()); + if ch == ']' { + if next == Some(']') { + i += 2; + continue; + } + return i + 1; + } + i += ch.len_utf8(); + } + sql.len() +} + +fn is_sqlserver_token_start(ch: char) -> bool { + ch.is_ascii_alphabetic() || ch == '_' +} + +fn is_sqlserver_token_part(ch: char) -> bool { + ch.is_ascii_alphanumeric() || matches!(ch, '_' | '$' | '#') +} + +fn next_char(sql: &str, index: usize) -> char { + sql[index..].chars().next().unwrap_or('\0') +} + +fn next_char_at(sql: &str, index: usize) -> Option { + if index >= sql.len() { + None + } else { + sql[index..].chars().next() + } +} + fn push_sqlserver_result_set(results: &mut Vec, result: Option, start: Instant) { if let Some(result) = result { if result.rows.is_empty() && result.columns.is_empty() { @@ -591,11 +911,15 @@ pub async fn execute_query_with_max_rows( let start = Instant::now(); if starts_with_executable_sql_keyword(sql, &["SELECT", "EXEC", "WITH", "TABLE"]) { - let stream = client.query(sql, &[]).await.map_err(|e| e.to_string())?; - collect_first_result_limited(stream, start, max_rows).await + let query_sql = match spatial_safe_sqlserver_query(client, sql).await { + Ok(Some(sql)) => sql, + Ok(None) | Err(_) => sql.to_string(), + }; + let stream = sqlserver_driver_result(client.query(query_sql.as_str(), &[])).await?; + sqlserver_driver_result(collect_first_result_limited(stream, start, max_rows)).await } else if requires_simple_query_batch(sql) || is_transaction_control(sql) { - let stream = client.simple_query(sql).await.map_err(|e| e.to_string())?; - let _ = collect_result_sets_limited(stream, start, max_rows).await?; + let stream = sqlserver_driver_result(client.simple_query(sql)).await?; + let _ = sqlserver_driver_result(collect_result_sets_limited(stream, start, max_rows)).await?; Ok(QueryResult { columns: vec![], rows: vec![], @@ -606,7 +930,7 @@ pub async fn execute_query_with_max_rows( has_more: false, }) } else { - let result = client.execute(sql, &[]).await.map_err(|e| e.to_string())?; + let result = sqlserver_driver_result(client.execute(sql, &[])).await?; Ok(QueryResult { columns: vec![], rows: vec![], @@ -629,8 +953,16 @@ pub async fn execute_batch_with_max_rows( max_rows: Option, ) -> Result, String> { let start = Instant::now(); - let stream = client.simple_query(sql).await.map_err(|e| e.to_string())?; - let mut results = collect_result_sets_limited(stream, start, max_rows).await?; + if is_single_sqlserver_select(sql) { + if let Ok(Some(query_sql)) = spatial_safe_sqlserver_query(client, sql).await { + let stream = sqlserver_driver_result(client.query(query_sql.as_str(), &[])).await?; + return sqlserver_driver_result(collect_first_result_limited(stream, start, max_rows)) + .await + .map(|result| vec![result]); + } + } + let stream = sqlserver_driver_result(client.simple_query(sql)).await?; + let mut results = sqlserver_driver_result(collect_result_sets_limited(stream, start, max_rows)).await?; if results.is_empty() { results.push(QueryResult { @@ -724,8 +1056,9 @@ fn first_sql_tokens(sql: &str, limit: usize) -> Vec { #[cfg(test)] mod tests { use super::{ - requires_simple_query_batch, sqlserver_cell_to_json, sqlserver_columns_sql, sqlserver_indexes_sql, - sqlserver_list_objects_sql, SqlServerResultSet, + build_spatial_safe_sqlserver_query, is_sqlserver_spatial_column, requires_simple_query_batch, + sqlserver_cell_to_json, sqlserver_columns_sql, sqlserver_indexes_sql, sqlserver_list_objects_sql, + SqlServerDescribedColumn, SqlServerResultSet, }; use chrono::NaiveDate; use std::time::Instant; @@ -833,6 +1166,95 @@ mod tests { assert_eq!(sqlserver_cell_to_json(&cell), serde_json::json!("2026-05-13 09:08:07.123")); } + #[test] + fn sqlserver_detects_geometry_result_columns() { + assert!(is_sqlserver_spatial_column(&SqlServerDescribedColumn { + name: Some("polygon".to_string()), + system_type_name: Some("geometry".to_string()), + user_type_schema: Some("sys".to_string()), + user_type_name: Some("geometry".to_string()), + })); + assert!(is_sqlserver_spatial_column(&SqlServerDescribedColumn { + name: Some("shape".to_string()), + system_type_name: Some("geography".to_string()), + user_type_schema: Some("sys".to_string()), + user_type_name: Some("geography".to_string()), + })); + assert!(!is_sqlserver_spatial_column(&SqlServerDescribedColumn { + name: Some("name".to_string()), + system_type_name: Some("varchar(30)".to_string()), + user_type_schema: None, + user_type_name: None, + })); + } + + #[test] + fn sqlserver_wraps_geometry_columns_as_text() { + let rewritten = build_spatial_safe_sqlserver_query( + "SELECT * FROM dbo.tLandPolygon;", + &[ + SqlServerDescribedColumn { + name: Some("landId".to_string()), + system_type_name: Some("varchar(30)".to_string()), + user_type_schema: None, + user_type_name: None, + }, + SqlServerDescribedColumn { + name: Some("polygon".to_string()), + system_type_name: Some("geometry".to_string()), + user_type_schema: Some("sys".to_string()), + user_type_name: Some("geometry".to_string()), + }, + ], + ) + .unwrap(); + + assert_eq!( + rewritten, + "SELECT [landId] = [dbx_spatial_source].[dbx_col_1], [polygon] = CASE WHEN [dbx_spatial_source].[dbx_col_2] IS NULL THEN NULL ELSE [dbx_spatial_source].[dbx_col_2].STAsText() END FROM (SELECT * FROM dbo.tLandPolygon) AS [dbx_spatial_source]([dbx_col_1], [dbx_col_2])" + ); + } + + #[test] + fn sqlserver_does_not_wrap_non_spatial_columns() { + assert_eq!( + build_spatial_safe_sqlserver_query( + "SELECT landId FROM dbo.tLandPolygon", + &[SqlServerDescribedColumn { + name: Some("landId".to_string()), + system_type_name: Some("varchar(30)".to_string()), + user_type_schema: None, + user_type_name: None, + }] + ), + None + ); + } + + #[test] + fn sqlserver_preserves_order_by_when_wrapping_geometry_columns() { + let rewritten = build_spatial_safe_sqlserver_query( + "SELECT landId, polygon FROM dbo.tLandPolygon ORDER BY landId DESC", + &[ + SqlServerDescribedColumn { + name: Some("landId".to_string()), + system_type_name: Some("varchar(30)".to_string()), + user_type_schema: None, + user_type_name: None, + }, + SqlServerDescribedColumn { + name: Some("polygon".to_string()), + system_type_name: Some("geometry".to_string()), + user_type_schema: Some("sys".to_string()), + user_type_name: Some("geometry".to_string()), + }, + ], + ) + .unwrap(); + + assert!(rewritten.contains("ORDER BY landId DESC OFFSET 0 ROWS")); + } + #[test] fn sqlserver_keeps_empty_result_sets_when_metadata_exists() { let mut results = Vec::new();