diff --git a/Cargo.lock b/Cargo.lock index 4b11079aa..59c01149d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -8021,6 +8021,7 @@ dependencies = [ "rustls-native-certs 0.6.3", "rustls-pemfile", "thiserror 1.0.69", + "tokio", "tokio-rustls 0.24.1", "tokio-util", "tracing", diff --git a/crates/dbx-core/Cargo.toml b/crates/dbx-core/Cargo.toml index d2ee2ffe3..000adbbff 100644 --- a/crates/dbx-core/Cargo.toml +++ b/crates/dbx-core/Cargo.toml @@ -23,7 +23,7 @@ sqlparser = "0.62.0" redis = { version = "0.32.2", features = ["tokio-comp", "tls-rustls", "tokio-rustls-comp"] } rustls = { version = "0.23", features = ["aws-lc-rs"] } duckdb = "1.3.2" -tiberius = { version = "0.12.3", default-features = false, features = ["tds73", "chrono", "rust_decimal", "rustls"] } +tiberius = { version = "0.12.3", default-features = false, features = ["tds73", "chrono", "rust_decimal", "rustls", "sql-browser-tokio"] } reqwest = { version = "0.12", default-features = false, features = ["json", "stream", "rustls-tls", "socks"] } futures = "0.3" iana-time-zone = "0.1" diff --git a/crates/dbx-core/src/db/sqlserver.rs b/crates/dbx-core/src/db/sqlserver.rs index 567d6d7fe..bc668a11e 100644 --- a/crates/dbx-core/src/db/sqlserver.rs +++ b/crates/dbx-core/src/db/sqlserver.rs @@ -1,7 +1,7 @@ use futures::TryStreamExt; use rust_decimal::Decimal; use std::time::Instant; -use tiberius::{AuthMethod, Client, ColumnData, Config, FromSql, QueryItem, QueryStream}; +use tiberius::{AuthMethod, Client, ColumnData, Config, FromSql, QueryItem, QueryStream, SqlBrowser}; use tokio::net::TcpStream; use tokio_util::compat::{Compat, TokioAsyncWriteCompatExt}; @@ -13,6 +13,22 @@ use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryRes pub type SqlServerClient = Client>; const SIMPLE_QUERY_MODULE_KEYWORDS: &[&str] = &["FUNCTION", "PROC", "PROCEDURE", "TRIGGER", "VIEW"]; +#[derive(Debug, PartialEq, Eq)] +struct SqlServerEndpoint<'a> { + host: &'a str, + instance_name: Option<&'a str>, +} + +fn sqlserver_endpoint(host: &str) -> SqlServerEndpoint<'_> { + if let Some((server, instance)) = host.split_once('\\') { + if !server.trim().is_empty() && !instance.trim().is_empty() { + return SqlServerEndpoint { host: server.trim(), instance_name: Some(instance.trim()) }; + } + } + + SqlServerEndpoint { host: host.trim(), instance_name: None } +} + fn query_result_row_limit(max_rows: Option) -> usize { max_rows.unwrap_or(MAX_ROWS).max(1) } @@ -39,8 +55,13 @@ async fn try_connect( use_encryption: bool, ) -> Result { let mut config = Config::new(); - config.host(host); - config.port(port); + let endpoint = sqlserver_endpoint(host); + config.host(endpoint.host); + if let Some(instance_name) = endpoint.instance_name { + config.instance_name(instance_name); + } else { + config.port(port); + } config.authentication(AuthMethod::sql_server(user, pass)); if let Some(db) = database { config.database(db); @@ -50,10 +71,17 @@ async fn try_connect( config.encryption(tiberius::EncryptionLevel::NotSupported); } - let tcp = tokio::time::timeout(connection_timeout(), TcpStream::connect(config.get_addr())) - .await - .map_err(|_| format!("SQL Server connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))? - .map_err(|e| format!("SQL Server connection failed: {e}"))?; + let tcp = if endpoint.instance_name.is_some() { + tokio::time::timeout(connection_timeout(), TcpStream::connect_named(&config)) + .await + .map_err(|_| format!("SQL Server connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))? + .map_err(|e| format!("SQL Server connection failed: {e}"))? + } else { + tokio::time::timeout(connection_timeout(), TcpStream::connect(config.get_addr())) + .await + .map_err(|_| format!("SQL Server connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))? + .map_err(|e| format!("SQL Server connection failed: {e}"))? + }; tokio::time::timeout(connection_timeout(), Client::connect(config, tcp.compat_write())) .await .map_err(|_| format!("SQL Server handshake timed out ({CONNECTION_TIMEOUT_SECS}s)"))? @@ -702,6 +730,38 @@ mod tests { use std::time::Instant; use tiberius::{ColumnData, IntoSql}; + #[test] + fn sqlserver_endpoint_splits_named_instance_hosts() { + assert_eq!( + super::sqlserver_endpoint(r"192.168.1.10\SQL2022"), + super::SqlServerEndpoint { host: "192.168.1.10", instance_name: Some("SQL2022") } + ); + assert_eq!( + super::sqlserver_endpoint(r" db.example.com\SQLEXPRESS "), + super::SqlServerEndpoint { host: "db.example.com", instance_name: Some("SQLEXPRESS") } + ); + } + + #[test] + fn sqlserver_endpoint_keeps_regular_hosts() { + assert_eq!( + super::sqlserver_endpoint("db.example.com"), + super::SqlServerEndpoint { host: "db.example.com", instance_name: None } + ); + assert_eq!( + super::sqlserver_endpoint(r"db.example.com\"), + super::SqlServerEndpoint { host: r"db.example.com\", instance_name: None } + ); + } + + #[test] + fn sqlserver_connect_uses_named_instance_resolution() { + let source = include_str!("sqlserver.rs"); + let try_connect = source.split("async fn try_connect").nth(1).unwrap(); + let try_connect = try_connect.split("fn row_to_json").next().unwrap(); + assert!(try_connect.contains("connect_named(&config)")); + } + #[test] fn sqlserver_module_definitions_require_simple_query_batch() { assert!(requires_simple_query_batch("CREATE FUNCTION dbo.fn_demo() RETURNS INT AS BEGIN RETURN 1; END;")); diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index d68bce401..3820ea461 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -39,7 +39,7 @@ redis = { version = "0.32.2", features = ["tokio-comp", "tls-rustls", "tokio-rus rustls = { version = "0.23", features = ["aws-lc-rs"] } portpicker = "0.1.1" duckdb = "1.3.2" -tiberius = { version = "0.12.3", default-features = false, features = ["tds73", "chrono", "rust_decimal", "rustls"] } +tiberius = { version = "0.12.3", default-features = false, features = ["tds73", "chrono", "rust_decimal", "rustls", "sql-browser-tokio"] } tokio-util = { version = "0.7", features = ["compat"] } reqwest = { version = "0.12", features = ["json", "stream"] } futures = "0.3"