From 1497ae1009665f78f3b45cfc895b3d76c53827c0 Mon Sep 17 00:00:00 2001 From: zipg Date: Fri, 24 Jul 2026 11:50:20 +0800 Subject: [PATCH] fix(sqlserver): handle duplicate result column names --- crates/dbx-core/src/db/sqlserver.rs | 154 ++++++++++++++++++++++------ 1 file changed, 125 insertions(+), 29 deletions(-) diff --git a/crates/dbx-core/src/db/sqlserver.rs b/crates/dbx-core/src/db/sqlserver.rs index 1439555a6..cc3f278ca 100644 --- a/crates/dbx-core/src/db/sqlserver.rs +++ b/crates/dbx-core/src/db/sqlserver.rs @@ -5,6 +5,9 @@ use crate::types::{ TriggerInfo, }; use futures::{FutureExt, TryStreamExt}; +use sqlparser::ast::{Expr, SelectItem, SetExpr, Statement}; +use sqlparser::dialect::MsSqlDialect; +use sqlparser::parser::Parser; use std::future::Future; use std::panic::AssertUnwindSafe; use std::sync::{Arc as StdArc, Mutex as StdMutex}; @@ -44,7 +47,7 @@ const SQLSERVER_RESULT_TYPE_PROBE_SQL: &str = "\ DECLARE @dbx_probe_sql nvarchar(max); \ BEGIN TRY \ SET @dbx_probe_sql = N'SELECT TOP (0) * INTO ' + QUOTENAME(@dbx_probe_table) + \ - N' FROM (' + @P3 + N') AS dbx_probe_source'; \ + N' FROM ' + @P3; \ EXEC sys.sp_executesql @dbx_probe_sql; \ SELECT c.name, TYPE_NAME(c.system_type_id) AS system_type_name, \ SCHEMA_NAME(t.schema_id) AS user_type_schema, t.name AS user_type_name \ @@ -419,6 +422,12 @@ struct SqlServerDescribedColumn { user_type_name: Option, } +#[derive(Debug, PartialEq, Eq)] +struct SqlServerLegacyProbe { + source_sql: String, + output_names: Option>>, +} + async fn sqlserver_driver_result(future: F) -> Result where F: Future>, @@ -440,23 +449,24 @@ pub fn is_driver_panic_error(error: &str) -> bool { async fn describe_sqlserver_result_set( client: &mut SqlServerClient, sql: &str, - legacy_sql: &str, + legacy_probe: &SqlServerLegacyProbe, ) -> Result, String> { - describe_sqlserver_result_set_with_mode(client, sql, legacy_sql, false).await + describe_sqlserver_result_set_with_mode(client, sql, legacy_probe, false).await } async fn describe_sqlserver_result_set_with_mode( client: &mut SqlServerClient, sql: &str, - legacy_sql: &str, + legacy_probe: &SqlServerLegacyProbe, force_legacy: bool, ) -> Result, String> { // SQL Server 2008 has no first-result-set DMV. Keep one probe round trip by // selecting the modern DMV path server-side and using metadata-only execution otherwise. let force_legacy = i32::from(force_legacy); - let mut stream = - sqlserver_driver_result(client.query(SQLSERVER_RESULT_TYPE_PROBE_SQL, &[&sql, &force_legacy, &legacy_sql])) - .await?; + let mut stream = sqlserver_driver_result( + client.query(SQLSERVER_RESULT_TYPE_PROBE_SQL, &[&sql, &force_legacy, &legacy_probe.source_sql]), + ) + .await?; let mut active_result_index = None; let mut uses_describe_dmv = None; let mut rows = Vec::new(); @@ -486,7 +496,15 @@ async fn describe_sqlserver_result_set_with_mode( if uses_describe_dmv.is_none() { return Err("SQL Server result type probe did not report its compatibility mode".to_string()); } - Ok(rows.iter().map(sqlserver_described_column_from_row).collect()) + let mut columns = rows.iter().map(sqlserver_described_column_from_row).collect::>(); + if uses_describe_dmv == Some(false) { + if let Some(output_names) = &legacy_probe.output_names { + for (column, output_name) in columns.iter_mut().zip(output_names) { + column.name.clone_from(output_name); + } + } + } + Ok(columns) } fn sqlserver_described_column_from_row(row: &Row) -> SqlServerDescribedColumn { @@ -506,13 +524,55 @@ async fn sqlserver_unsafe_type_query(client: &mut SqlServerClient, sql: &str) -> if !is_single_sqlserver_select(sql) { return Ok(None); } - let Some(legacy_sql) = normalized_sqlserver_select_statement(sql) else { + let Some(legacy_probe) = sqlserver_legacy_probe(sql) else { return Ok(None); }; - let columns = describe_sqlserver_result_set(client, sql, &legacy_sql).await?; + let columns = describe_sqlserver_result_set(client, sql, &legacy_probe).await?; Ok(build_sqlserver_unsafe_type_query(sql, &columns)) } +fn sqlserver_legacy_probe(sql: &str) -> Option { + let statement = normalized_sqlserver_select_statement(sql)?; + let output_names = sqlserver_projection_output_names(&statement); + let source_alias = quote_sqlserver_identifier("dbx_probe_source"); + let source_sql = if let Some(names) = &output_names { + let aliases = (0..names.len()) + .map(sqlserver_source_column_name) + .map(|name| quote_sqlserver_identifier(&name)) + .collect::>() + .join(", "); + format!("({statement}) AS {source_alias}({aliases})") + } else { + format!("({statement}) AS {source_alias}") + }; + Some(SqlServerLegacyProbe { source_sql, output_names }) +} + +fn sqlserver_projection_output_names(statement: &str) -> Option>> { + let statements = Parser::parse_sql(&MsSqlDialect {}, statement).ok()?; + let [Statement::Query(query)] = statements.as_slice() else { + return None; + }; + let SetExpr::Select(select) = query.body.as_ref() else { + return None; + }; + select + .projection + .iter() + .map(|item| match item { + SelectItem::ExprWithAlias { alias, .. } => Some(Some(alias.value.clone())), + SelectItem::UnnamedExpr(Expr::Identifier(identifier)) => { + Some((!identifier.value.starts_with('@')).then(|| identifier.value.clone())) + } + SelectItem::UnnamedExpr(Expr::CompoundIdentifier(identifiers)) => { + Some(identifiers.last().map(|identifier| identifier.value.clone())) + } + SelectItem::UnnamedExpr(_) => Some(None), + SelectItem::ExprWithAliases { .. } | SelectItem::QualifiedWildcard(_, _) | SelectItem::Wildcard(_) => None, + }) + .collect() +} + fn build_sqlserver_unsafe_type_query(sql: &str, columns: &[SqlServerDescribedColumn]) -> Option { if columns.is_empty() || !columns.iter().any(is_sqlserver_unsafe_column) { return None; @@ -2232,7 +2292,7 @@ mod tests { query_result_with_server_messages, requires_simple_query_batch, sqlserver_batch_can_use_execute, sqlserver_cell_to_json, sqlserver_columns_sql, sqlserver_completion_assistant_sql, sqlserver_dml_output_returns_rows, sqlserver_hidden_schema_names, sqlserver_indexes_sql, - sqlserver_list_objects_sql, sqlserver_list_schemas_sql, sqlserver_list_tables_sql, + sqlserver_legacy_probe, sqlserver_list_objects_sql, sqlserver_list_schemas_sql, sqlserver_list_tables_sql, sqlserver_schema_name_predicate, sqlserver_table_comment_sql, sqlserver_visible_object_predicate, strip_dbx_sqlserver_row_number_column, SqlServerDescribedColumn, SqlServerResultSet, SQLSERVER_RESULT_TYPE_PROBE_SQL, @@ -3116,11 +3176,40 @@ mod tests { assert!(SQLSERVER_RESULT_TYPE_PROBE_SQL.contains("sys.dm_exec_describe_first_result_set")); assert!(SQLSERVER_RESULT_TYPE_PROBE_SQL.contains("##dbx_result_type_probe_")); assert!(SQLSERVER_RESULT_TYPE_PROBE_SQL.contains("SELECT TOP (0) * INTO")); + assert!(SQLSERVER_RESULT_TYPE_PROBE_SQL.contains("N' FROM ' + @P3")); assert!(SQLSERVER_RESULT_TYPE_PROBE_SQL.contains("FROM tempdb.sys.columns")); assert!(!SQLSERVER_RESULT_TYPE_PROBE_SQL.contains("FMTONLY")); assert_eq!(SQLSERVER_RESULT_TYPE_PROBE_SQL.matches("SELECT @dbx_use_describe_dmv").count(), 1); } + #[test] + fn sqlserver_legacy_probe_uses_unique_internal_names_for_duplicate_outputs() { + let probe = sqlserver_legacy_probe( + "SELECT a.HJRQ, b.HJRQ, a.id + b.id AS total, GETDATE() FROM dbo.a a JOIN dbo.b b ON b.id = a.id", + ) + .unwrap(); + + assert_eq!( + probe.source_sql, + "(SELECT a.HJRQ, b.HJRQ, a.id + b.id AS total, GETDATE() FROM dbo.a a JOIN dbo.b b ON b.id = a.id) AS [dbx_probe_source]([dbx_col_1], [dbx_col_2], [dbx_col_3], [dbx_col_4])" + ); + assert_eq!( + probe.output_names, + Some(vec![Some("HJRQ".to_string()), Some("HJRQ".to_string()), Some("total".to_string()), None]) + ); + } + + #[test] + fn sqlserver_legacy_probe_keeps_wildcards_on_existing_server_validation_path() { + let probe = sqlserver_legacy_probe("SELECT a.*, b.HJRQ FROM dbo.a a JOIN dbo.b b ON b.id = a.id").unwrap(); + + assert_eq!( + probe.source_sql, + "(SELECT a.*, b.HJRQ FROM dbo.a a JOIN dbo.b b ON b.id = a.id) AS [dbx_probe_source]" + ); + assert_eq!(probe.output_names, None); + } + #[test] fn sqlserver_legacy_probe_errors_are_blocking_only_when_marked_unsafe() { assert!(is_blocking_sqlserver_unsafe_probe_error( @@ -3228,17 +3317,16 @@ mod tests { client.simple_query(setup).await.unwrap().into_results().await.unwrap(); let sql = "SELECT id, payload FROM #dbx_issue_4002"; - let ordinary_columns = super::describe_sqlserver_result_set_with_mode( - &mut client, - "SELECT 42 AS answer", - "SELECT 42 AS answer", - true, - ) - .await - .unwrap(); + let ordinary_probe = super::sqlserver_legacy_probe("SELECT 42 AS answer").unwrap(); + let ordinary_columns = + super::describe_sqlserver_result_set_with_mode(&mut client, "SELECT 42 AS answer", &ordinary_probe, true) + .await + .unwrap(); assert_eq!(ordinary_columns[0].system_type_name.as_deref(), Some("int")); - let legacy_columns = super::describe_sqlserver_result_set_with_mode(&mut client, sql, sql, true).await.unwrap(); + let legacy_probe = super::sqlserver_legacy_probe(sql).unwrap(); + let legacy_columns = + super::describe_sqlserver_result_set_with_mode(&mut client, sql, &legacy_probe, true).await.unwrap(); assert_eq!(legacy_columns.len(), 2); assert!(is_sqlserver_variant_column(&legacy_columns[1])); @@ -3252,15 +3340,23 @@ mod tests { let variant = super::execute_query(&mut client, sql).await.unwrap(); assert_eq!(variant.rows, vec![vec![serde_json::json!(1), serde_json::json!("legacy")]]); - let boundary_error = super::describe_sqlserver_result_set_with_mode( - &mut client, - "SELECT 1 AS duplicate_name, 2 AS duplicate_name", - "SELECT 1 AS duplicate_name, 2 AS duplicate_name", - true, - ) - .await - .unwrap_err(); - assert!(is_blocking_sqlserver_unsafe_probe_error(&boundary_error)); + let duplicate_sql = "SELECT id AS HJRQ, payload AS HJRQ FROM #dbx_issue_4002"; + let duplicate_probe = super::sqlserver_legacy_probe(duplicate_sql).unwrap(); + let duplicate_columns = + super::describe_sqlserver_result_set_with_mode(&mut client, duplicate_sql, &duplicate_probe, true) + .await + .unwrap(); + assert_eq!(duplicate_columns.len(), 2); + assert_eq!(duplicate_columns[0].name.as_deref(), Some("HJRQ")); + assert_eq!(duplicate_columns[1].name.as_deref(), Some("HJRQ")); + assert!(is_sqlserver_variant_column(&duplicate_columns[1])); + + let duplicate_rewritten = build_sqlserver_unsafe_type_query(duplicate_sql, &duplicate_columns).unwrap(); + let duplicate_rows = client.query(duplicate_rewritten, &[]).await.unwrap().into_first_result().await.unwrap(); + assert_eq!(duplicate_rows[0].columns()[0].name(), "HJRQ"); + assert_eq!(duplicate_rows[0].columns()[1].name(), "HJRQ"); + assert_eq!(duplicate_rows[0].get::(0), Some(1)); + assert_eq!(duplicate_rows[0].get::<&str, _>(1), Some("legacy")); let continued = super::execute_query(&mut client, "SELECT CAST(7 AS int) AS still_connected").await.unwrap(); assert_eq!(continued.rows, vec![vec![serde_json::json!(7)]]);