fix(sqlserver): handle spatial query results
This commit is contained in:
parent
5e37a518f2
commit
7283f2b34a
|
|
@ -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();
|
||||
|
|
|
|||
Loading…
Reference in New Issue