fix(sqlserver): handle spatial query results

This commit is contained in:
t8y2 2026-06-05 11:11:14 +08:00
parent 5e37a518f2
commit 7283f2b34a
1 changed files with 432 additions and 10 deletions

View File

@ -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<String>,
system_type_name: Option<String>,
user_type_schema: Option<String>,
user_type_name: Option<String>,
}
async fn sqlserver_driver_result<T, E, F>(future: F) -> Result<T, String>
where
F: Future<Output = Result<T, E>>,
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<Vec<SqlServerDescribedColumn>, 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<Option<String>, 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<String> {
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::<Vec<_>>();
let source_alias_list =
source_columns.iter().map(|name| quote_sqlserver_identifier(name)).collect::<Vec<_>>().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::<Vec<_>>()
.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<String> {
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<SqlServerToken> {
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<usize> {
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<char> {
if index >= sql.len() {
None
} else {
sql[index..].chars().next()
}
}
fn push_sqlserver_result_set(results: &mut Vec<QueryResult>, result: Option<SqlServerResultSet>, 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<usize>,
) -> Result<Vec<QueryResult>, 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<String> {
#[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();