diff --git a/crates/dbx-core/src/db/sqlserver.rs b/crates/dbx-core/src/db/sqlserver.rs index d2f486c50..9ccba1aa6 100644 --- a/crates/dbx-core/src/db/sqlserver.rs +++ b/crates/dbx-core/src/db/sqlserver.rs @@ -5,7 +5,7 @@ use crate::types::{ TriggerInfo, }; use futures::{FutureExt, TryStreamExt}; -use sqlparser::ast::{Expr, SelectItem, SetExpr, Statement}; +use sqlparser::ast::{Expr, Ident, SelectItem, SetExpr, Statement}; use sqlparser::dialect::MsSqlDialect; use sqlparser::parser::Parser; use std::future::Future; @@ -428,6 +428,20 @@ struct SqlServerDescribedColumn { struct SqlServerLegacyProbe { source_sql: String, output_names: Option>>, + output_name_overrides: Vec, +} + +#[derive(Debug, PartialEq, Eq)] +struct SqlServerWildcardProjectionProbe { + statement: String, + output_name_overrides: Vec, +} + +#[derive(Debug, PartialEq, Eq)] +struct SqlServerProbeOutputNameOverride { + projection_ordinal: usize, + probe_name: String, + output_name: Option, } async fn sqlserver_driver_result(future: F) -> Result @@ -500,13 +514,33 @@ async fn describe_sqlserver_result_set_with_mode( } 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); + restore_sqlserver_legacy_probe_output_names(&mut columns, legacy_probe); + } + Ok(columns) +} + +fn restore_sqlserver_legacy_probe_output_names( + columns: &mut [SqlServerDescribedColumn], + legacy_probe: &SqlServerLegacyProbe, +) { + 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); + } + } else if !legacy_probe.output_name_overrides.is_empty() { + for column in columns { + let output_name = column.name.as_deref().and_then(|probe_name| { + legacy_probe + .output_name_overrides + .iter() + .find(|output_name_override| output_name_override.probe_name.eq_ignore_ascii_case(probe_name)) + .map(|output_name_override| output_name_override.output_name.clone()) + }); + if let Some(output_name) = output_name { + column.name = output_name; } } } - Ok(columns) } fn sqlserver_described_column_from_row(row: &Row) -> SqlServerDescribedColumn { @@ -534,20 +568,27 @@ async fn sqlserver_unsafe_type_query(client: &mut SqlServerClient, sql: &str) -> } fn sqlserver_legacy_probe(sql: &str) -> Option { + let nonce = uuid::Uuid::new_v4().simple().to_string(); + sqlserver_legacy_probe_with_nonce(sql, &nonce) +} + +fn sqlserver_legacy_probe_with_nonce(sql: &str, nonce: &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 (source_sql, output_name_overrides) = 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})") + (format!("({statement}) AS {source_alias}({aliases})"), Vec::new()) + } else if let Some(wildcard_probe) = sqlserver_wildcard_projection_probe(&statement, nonce) { + (format!("({}) AS {source_alias}", wildcard_probe.statement), wildcard_probe.output_name_overrides) } else { - format!("({statement}) AS {source_alias}") + (format!("({statement}) AS {source_alias}"), Vec::new()) }; - Some(SqlServerLegacyProbe { source_sql, output_names }) + Some(SqlServerLegacyProbe { source_sql, output_names, output_name_overrides }) } fn sqlserver_projection_output_names(statement: &str) -> Option>> { @@ -558,21 +599,58 @@ fn sqlserver_projection_output_names(statement: &str) -> Option Option> { + 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, + } +} + +fn sqlserver_probe_explicit_alias(nonce: &str, projection_ordinal: usize) -> String { + format!("__dbx_probe_{nonce}_explicit_{projection_ordinal}__") +} + +fn sqlserver_wildcard_projection_probe(statement: &str, nonce: &str) -> Option { + let mut statements = Parser::parse_sql(&MsSqlDialect {}, statement).ok()?; + let [Statement::Query(query)] = statements.as_mut_slice() else { + return None; + }; + let SetExpr::Select(select) = query.body.as_mut() else { + return None; + }; + if !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() + .any(|item| matches!(item, SelectItem::QualifiedWildcard(_, _) | SelectItem::Wildcard(_))) + { + return None; + } + + let mut output_name_overrides = Vec::new(); + for (index, item) in select.projection.iter_mut().enumerate() { + let (expr, output_name) = match item { + SelectItem::UnnamedExpr(expr) => (expr.clone(), sqlserver_projection_item_output_name(item)?), + SelectItem::ExprWithAlias { expr, alias } => (expr.clone(), Some(alias.value.clone())), + SelectItem::QualifiedWildcard(_, _) | SelectItem::Wildcard(_) => continue, + SelectItem::ExprWithAliases { .. } => return None, + }; + let projection_ordinal = index + 1; + let probe_name = sqlserver_probe_explicit_alias(nonce, projection_ordinal); + *item = SelectItem::ExprWithAlias { expr, alias: Ident::with_quote('[', probe_name.clone()) }; + output_name_overrides.push(SqlServerProbeOutputNameOverride { projection_ordinal, probe_name, output_name }); + } + + Some(SqlServerWildcardProjectionProbe { statement: query.to_string(), output_name_overrides }) } fn build_sqlserver_unsafe_type_query(sql: &str, columns: &[SqlServerDescribedColumn]) -> Option { @@ -2328,13 +2406,15 @@ mod tests { use super::{ build_sqlserver_unsafe_type_query, capture_sqlserver_messages, format_sqlserver_numeric, is_blocking_sqlserver_unsafe_probe_error, is_sqlserver_spatial_column, is_sqlserver_variant_column, - 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_filter_definition_error, sqlserver_hidden_schema_names, - sqlserver_indexes_sql, sqlserver_legacy_indexes_sql, sqlserver_legacy_probe, sqlserver_list_objects_sql, - sqlserver_list_schemas_sql, sqlserver_list_tables_sql, sqlserver_schema_name_predicate, + query_result_with_server_messages, requires_simple_query_batch, restore_sqlserver_legacy_probe_output_names, + sqlserver_batch_can_use_execute, sqlserver_cell_to_json, sqlserver_columns_sql, + sqlserver_completion_assistant_sql, sqlserver_dml_output_returns_rows, sqlserver_filter_definition_error, + sqlserver_hidden_schema_names, sqlserver_indexes_sql, sqlserver_legacy_indexes_sql, sqlserver_legacy_probe, + sqlserver_legacy_probe_with_nonce, sqlserver_list_objects_sql, sqlserver_list_schemas_sql, + sqlserver_list_tables_sql, sqlserver_probe_explicit_alias, 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, + SqlServerDescribedColumn, SqlServerProbeOutputNameOverride, SqlServerResultSet, + SQLSERVER_RESULT_TYPE_PROBE_SQL, }; use crate::types::{ CompletionAssistantMatchMode, CompletionAssistantObjectKind, CompletionAssistantRequest, QueryResult, @@ -3271,14 +3351,120 @@ mod tests { } #[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(); + fn sqlserver_legacy_probe_aliases_explicit_columns_next_to_wildcards() { + let nonce = "0123456789abcdef0123456789abcdef"; + let probe = + sqlserver_legacy_probe_with_nonce("SELECT a.*, b.HJRQ FROM dbo.a a JOIN dbo.b b ON b.id = a.id", nonce) + .unwrap(); + let probe_name = sqlserver_probe_explicit_alias(nonce, 2); - 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!(probe.source_sql.starts_with("(SELECT a.*,")); + assert!(probe.source_sql.contains(&format!("b.HJRQ AS [{probe_name}]"))); + assert!(probe.source_sql.ends_with(") AS [dbx_probe_source]")); assert_eq!(probe.output_names, None); + assert_eq!( + probe.output_name_overrides, + vec![SqlServerProbeOutputNameOverride { + projection_ordinal: 2, + probe_name, + output_name: Some("HJRQ".to_string()), + }] + ); + } + + #[test] + fn sqlserver_legacy_probe_keeps_quoted_columns_around_qualified_wildcard() { + let nonce = "fedcba9876543210fedcba9876543210"; + let probe = + sqlserver_legacy_probe_with_nonce("SELECT t.[Order], t.*, t.[After] FROM dbo.t AS t", nonce).unwrap(); + let first_probe_name = sqlserver_probe_explicit_alias(nonce, 1); + let third_probe_name = sqlserver_probe_explicit_alias(nonce, 3); + + assert!(probe.source_sql.contains(&format!("t.[Order] AS [{first_probe_name}], t.*"))); + assert!(probe.source_sql.contains(&format!("t.*, t.[After] AS [{third_probe_name}]"))); + assert!(probe.source_sql.contains("FROM dbo.t AS t")); + assert_eq!( + probe.output_name_overrides, + vec![ + SqlServerProbeOutputNameOverride { + projection_ordinal: 1, + probe_name: first_probe_name, + output_name: Some("Order".to_string()), + }, + SqlServerProbeOutputNameOverride { + projection_ordinal: 3, + probe_name: third_probe_name, + output_name: Some("After".to_string()), + }, + ] + ); + } + + #[test] + fn sqlserver_legacy_probe_supports_explicit_columns_followed_by_wildcard() { + let nonce = "11223344556677889900aabbccddeeff"; + let probe = sqlserver_legacy_probe_with_nonce("SELECT ybbz, cytzrq, jzrq, * FROM dbo.t", nonce).unwrap(); + let first_probe_name = sqlserver_probe_explicit_alias(nonce, 1); + let second_probe_name = sqlserver_probe_explicit_alias(nonce, 2); + let third_probe_name = sqlserver_probe_explicit_alias(nonce, 3); + + assert!(probe.source_sql.contains(&format!("ybbz AS [{first_probe_name}]"))); + assert!(probe.source_sql.contains(&format!("cytzrq AS [{second_probe_name}]"))); + assert!(probe.source_sql.contains(&format!("jzrq AS [{third_probe_name}]"))); + assert!(probe.source_sql.contains(", * FROM dbo.t")); + assert_eq!( + probe.output_name_overrides, + vec![ + SqlServerProbeOutputNameOverride { + projection_ordinal: 1, + probe_name: first_probe_name, + output_name: Some("ybbz".to_string()), + }, + SqlServerProbeOutputNameOverride { + projection_ordinal: 2, + probe_name: second_probe_name, + output_name: Some("cytzrq".to_string()), + }, + SqlServerProbeOutputNameOverride { + projection_ordinal: 3, + probe_name: third_probe_name, + output_name: Some("jzrq".to_string()), + }, + ] + ); + } + + #[test] + fn sqlserver_legacy_probe_does_not_collide_with_old_probe_alias_column() { + let nonce = "00112233445566778899aabbccddeeff"; + let old_probe_name = "__dbx_probe_explicit_1__"; + let probe = + sqlserver_legacy_probe_with_nonce("SELECT t.[__dbx_probe_explicit_1__], t.* FROM dbo.t AS t", nonce) + .unwrap(); + let generated_probe_name = sqlserver_probe_explicit_alias(nonce, 1); + + assert_ne!(generated_probe_name, old_probe_name); + assert!(probe.source_sql.contains(&format!("t.[{old_probe_name}] AS [{generated_probe_name}], t.*"))); + assert_eq!(probe.output_name_overrides[0].projection_ordinal, 1); + + let mut columns = vec![ + SqlServerDescribedColumn { + name: Some(generated_probe_name), + system_type_name: Some("int".to_string()), + user_type_schema: None, + user_type_name: None, + }, + SqlServerDescribedColumn { + name: Some(old_probe_name.to_string()), + system_type_name: Some("int".to_string()), + user_type_schema: None, + user_type_name: None, + }, + ]; + restore_sqlserver_legacy_probe_output_names(&mut columns, &probe); + + assert_eq!(columns[0].name.as_deref(), Some(old_probe_name)); + assert_eq!(columns[1].name.as_deref(), Some(old_probe_name)); } #[test] @@ -3383,8 +3569,11 @@ mod tests { let setup = "\ IF OBJECT_ID('tempdb..#dbx_issue_4002') IS NOT NULL DROP TABLE #dbx_issue_4002; \ - CREATE TABLE #dbx_issue_4002 (id int NOT NULL, payload sql_variant NULL); \ - INSERT INTO #dbx_issue_4002 (id, payload) VALUES (1, CAST(N'legacy' AS nvarchar(20)))"; + CREATE TABLE #dbx_issue_4002 (\ + id int NOT NULL, ybbz int NULL, cytzrq date NULL, jzrq datetime NULL, payload sql_variant NULL\ + ); \ + INSERT INTO #dbx_issue_4002 (id, ybbz, cytzrq, jzrq, payload) \ + VALUES (1, 2, '2026-07-28', '2026-07-28T12:34:56', CAST(N'legacy' AS nvarchar(20)))"; client.simple_query(setup).await.unwrap().into_results().await.unwrap(); let sql = "SELECT id, payload FROM #dbx_issue_4002"; @@ -3429,6 +3618,39 @@ mod tests { assert_eq!(duplicate_rows[0].get::(0), Some(1)); assert_eq!(duplicate_rows[0].get::<&str, _>(1), Some("legacy")); + let wildcard_sql = "SELECT ybbz, cytzrq, jzrq, * FROM #dbx_issue_4002"; + let previous_wildcard_probe = super::SqlServerLegacyProbe { + source_sql: format!("({wildcard_sql}) AS [dbx_probe_source]"), + output_names: None, + output_name_overrides: Vec::new(), + }; + let previous_error = + super::describe_sqlserver_result_set_with_mode(&mut client, wildcard_sql, &previous_wildcard_probe, true) + .await + .unwrap_err(); + assert!(previous_error.contains("dbx_probe_source")); + + let wildcard_probe = super::sqlserver_legacy_probe(wildcard_sql).unwrap(); + let wildcard_columns = + super::describe_sqlserver_result_set_with_mode(&mut client, wildcard_sql, &wildcard_probe, true) + .await + .unwrap(); + assert_eq!(wildcard_columns.len(), 8); + assert_eq!(wildcard_columns[0].name.as_deref(), Some("ybbz")); + assert_eq!(wildcard_columns[1].name.as_deref(), Some("cytzrq")); + assert_eq!(wildcard_columns[2].name.as_deref(), Some("jzrq")); + assert_eq!(wildcard_columns[4].name.as_deref(), Some("ybbz")); + assert_eq!(wildcard_columns[5].name.as_deref(), Some("cytzrq")); + assert_eq!(wildcard_columns[6].name.as_deref(), Some("jzrq")); + assert!(is_sqlserver_variant_column(&wildcard_columns[7])); + + let wildcard_rewritten = build_sqlserver_unsafe_type_query(wildcard_sql, &wildcard_columns).unwrap(); + let wildcard_rows = client.query(wildcard_rewritten, &[]).await.unwrap().into_first_result().await.unwrap(); + assert_eq!(wildcard_rows.len(), 1); + assert_eq!(wildcard_rows[0].columns()[0].name(), "ybbz"); + assert_eq!(wildcard_rows[0].columns()[4].name(), "ybbz"); + assert_eq!(wildcard_rows[0].get::<&str, _>(7), 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)]]);