diff --git a/crates/dbx-core/src/connection.rs b/crates/dbx-core/src/connection.rs index 214e1dea1..af9577c03 100644 --- a/crates/dbx-core/src/connection.rs +++ b/crates/dbx-core/src/connection.rs @@ -59,7 +59,10 @@ impl AppState { configs.get(connection_id).map(|c| c.db_type.clone()) }; - let is_embedded = matches!(db_type, Some(DatabaseType::Sqlite) | Some(DatabaseType::DuckDb)); + let is_embedded = matches!( + db_type, + Some(DatabaseType::Sqlite) | Some(DatabaseType::DuckDb) + ); if is_embedded { return Ok(connection_id.to_string()); } @@ -95,28 +98,47 @@ impl AppState { } let (host, port) = self.connection_host_port(connection_id, &db_config).await?; + probe_connection_endpoint(&db_config, &host, port).await?; let url = connection_url_for_endpoint(&db_config, &host, port); let pool = match db_config.db_type { - DatabaseType::Mysql if db_config.needs_bare_mysql() => PoolKind::Mysql(db::mysql::connect_bare(&url).await?, true), + DatabaseType::Mysql if db_config.needs_bare_mysql() => { + PoolKind::Mysql(db::mysql::connect_bare(&url).await?, true) + } DatabaseType::Mysql => PoolKind::Mysql(db::mysql::connect(&url).await?, false), - DatabaseType::Doris | DatabaseType::StarRocks => PoolKind::Mysql(db::mysql::connect_bare(&url).await?, true), - DatabaseType::Postgres | DatabaseType::Redshift => PoolKind::Postgres(db::postgres::connect(&url).await?), - DatabaseType::Sqlite => PoolKind::Sqlite(db::sqlite::connect_path(&expand_tilde(&db_config.host)).await?), + DatabaseType::Doris | DatabaseType::StarRocks => { + PoolKind::Mysql(db::mysql::connect_bare(&url).await?, true) + } + DatabaseType::Postgres | DatabaseType::Redshift => { + PoolKind::Postgres(db::postgres::connect(&url).await?) + } + DatabaseType::Sqlite => { + PoolKind::Sqlite(db::sqlite::connect_path(&expand_tilde(&db_config.host)).await?) + } DatabaseType::Redis => { let con = db::redis_driver::connect(&url).await?; PoolKind::Redis(tokio::sync::Mutex::new(con)) } DatabaseType::DuckDb => { - let con = duckdb::Connection::open(&expand_tilde(&db_config.host)).map_err(|e| e.to_string())?; + let con = duckdb::Connection::open(&expand_tilde(&db_config.host)) + .map_err(|e| e.to_string())?; PoolKind::DuckDb(Arc::new(std::sync::Mutex::new(con))) } DatabaseType::MongoDb => { - let client = mongodb::Client::with_uri_str(&url).await.map_err(|e| e.to_string())?; + let client = db::mongo_driver::connect(&url).await?; + db::mongo_driver::test_connection(&client).await?; PoolKind::MongoDb(client) } DatabaseType::ClickHouse => { - let username = if db_config.username.is_empty() { None } else { Some(db_config.username.clone()) }; - let password = if db_config.password.is_empty() { None } else { Some(db_config.password.clone()) }; + let username = if db_config.username.is_empty() { + None + } else { + Some(db_config.username.clone()) + }; + let password = if db_config.password.is_empty() { + None + } else { + Some(db_config.password.clone()) + }; let client = db::clickhouse_driver::ChClient::new(&url, username, password); db::clickhouse_driver::test_connection(&client).await?; PoolKind::ClickHouse(client) @@ -137,7 +159,8 @@ impl AppState { &host, port, db_config.database.as_deref().unwrap_or("ORCL"), - &db_config.username, &db_config.password, + &db_config.username, + &db_config.password, ) .await?; PoolKind::Oracle(Arc::new(tokio::sync::Mutex::new(client))) @@ -196,8 +219,11 @@ impl AppState { ) -> Result { let is_single_conn = { let configs = self.configs.lock().await; - configs.get(connection_id) - .map(|c| c.db_type == DatabaseType::Oracle || c.db_type == DatabaseType::Elasticsearch) + configs + .get(connection_id) + .map(|c| { + c.db_type == DatabaseType::Oracle || c.db_type == DatabaseType::Elasticsearch + }) .unwrap_or(false) }; let pool_key = if is_single_conn { @@ -232,3 +258,22 @@ pub fn redacted_connection_url_for_endpoint( config.redacted_connection_url_with_host(host, port) } } + +pub async fn probe_connection_endpoint( + config: &ConnectionConfig, + host: &str, + port: u16, +) -> Result<(), String> { + match config.db_type { + DatabaseType::Sqlite | DatabaseType::DuckDb => Ok(()), + DatabaseType::MongoDb + if config + .connection_string + .as_deref() + .is_some_and(|value| !value.is_empty()) => + { + Ok(()) + } + _ => db::probe_tcp_endpoint(&format!("{:?}", config.db_type), host, port).await, + } +} diff --git a/crates/dbx-core/src/db/clickhouse_driver.rs b/crates/dbx-core/src/db/clickhouse_driver.rs index 5b906901d..72440ad1f 100644 --- a/crates/dbx-core/src/db/clickhouse_driver.rs +++ b/crates/dbx-core/src/db/clickhouse_driver.rs @@ -2,6 +2,7 @@ use reqwest::Client as HttpClient; use serde::Deserialize; use std::time::Instant; +use super::{connection_timeout, with_connection_timeout}; use crate::types::{ColumnInfo, DatabaseInfo, QueryResult, TableInfo}; pub struct ChClient { @@ -13,8 +14,12 @@ pub struct ChClient { impl ChClient { pub fn new(url: &str, username: Option, password: Option) -> Self { + let http = HttpClient::builder() + .connect_timeout(connection_timeout()) + .build() + .unwrap_or_else(|_| HttpClient::new()); Self { - http: HttpClient::new(), + http, base_url: url.trim_end_matches('/').to_string(), username, password, @@ -57,7 +62,11 @@ fn build_request(client: &ChClient, req: reqwest::RequestBuilder) -> reqwest::Re } } -async fn ch_query(client: &ChClient, sql: &str, database: Option<&str>) -> Result { +async fn ch_query( + client: &ChClient, + sql: &str, + database: Option<&str>, +) -> Result { let mut url = format!("{}/?default_format=JSONCompact", client.base_url); if let Some(db) = database { url.push_str(&format!("&database={}", db)); @@ -71,22 +80,37 @@ async fn ch_query(client: &ChClient, sql: &str, database: Option<&str>) -> Resul let body = resp.text().await.unwrap_or_default(); return Err(format!("ClickHouse error: {body}")); } - resp.json::().await.map_err(|e| format!("ClickHouse parse error: {e}")) + resp.json::() + .await + .map_err(|e| format!("ClickHouse parse error: {e}")) } pub async fn test_connection(client: &ChClient) -> Result<(), String> { let url = format!("{}/ping", client.base_url); let req = build_request(client, client.http.get(&url)); - req.send().await - .map_err(|e| format!("ClickHouse connection failed: {e}"))?; + with_connection_timeout("ClickHouse", async { + req.send() + .await + .map_err(|e| format!("ClickHouse connection failed: {e}")) + }) + .await?; Ok(()) } pub async fn list_databases(client: &ChClient) -> Result, String> { - let result = ch_query(client, "SELECT name FROM system.databases ORDER BY name", None).await?; - Ok(result.data.iter().map(|row| { - DatabaseInfo { name: row[0].as_str().unwrap_or("").to_string() } - }).collect()) + let result = ch_query( + client, + "SELECT name FROM system.databases ORDER BY name", + None, + ) + .await?; + Ok(result + .data + .iter() + .map(|row| DatabaseInfo { + name: row[0].as_str().unwrap_or("").to_string(), + }) + .collect()) } pub async fn list_tables(client: &ChClient, database: &str) -> Result, String> { @@ -95,17 +119,29 @@ pub async fn list_tables(client: &ChClient, database: &str) -> Result Result, String> { +pub async fn get_columns( + client: &ChClient, + database: &str, + table: &str, +) -> Result, String> { let sql = format!( "SELECT name, type, default_kind, default_expression, is_in_primary_key \ FROM system.columns WHERE database = '{}' AND table = '{}' ORDER BY position", @@ -113,28 +149,45 @@ pub async fn get_columns(client: &ChClient, database: &str, table: &str) -> Resu table.replace('\'', "\\'") ); let result = ch_query(client, &sql, Some(database)).await?; - Ok(result.data.iter().map(|row| { - let data_type = row.get(1).and_then(|v| v.as_str()).unwrap_or("").to_string(); - let is_nullable = data_type.starts_with("Nullable"); - let is_pk = row.get(4).and_then(|v| v.as_u64()).unwrap_or(0) == 1; - let default_kind = row.get(2).and_then(|v| v.as_str()).unwrap_or(""); - let default_expr = row.get(3).and_then(|v| v.as_str()).unwrap_or(""); - let column_default = if default_kind.is_empty() { None } else { Some(default_expr.to_string()) }; - ColumnInfo { - name: row[0].as_str().unwrap_or("").to_string(), - data_type, - is_nullable, - column_default, - is_primary_key: is_pk, - extra: None, comment: None, - numeric_precision: None, - numeric_scale: None, - character_maximum_length: None, - } - }).collect()) + Ok(result + .data + .iter() + .map(|row| { + let data_type = row + .get(1) + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + let is_nullable = data_type.starts_with("Nullable"); + let is_pk = row.get(4).and_then(|v| v.as_u64()).unwrap_or(0) == 1; + let default_kind = row.get(2).and_then(|v| v.as_str()).unwrap_or(""); + let default_expr = row.get(3).and_then(|v| v.as_str()).unwrap_or(""); + let column_default = if default_kind.is_empty() { + None + } else { + Some(default_expr.to_string()) + }; + ColumnInfo { + name: row[0].as_str().unwrap_or("").to_string(), + data_type, + is_nullable, + column_default, + is_primary_key: is_pk, + extra: None, + comment: None, + numeric_precision: None, + numeric_scale: None, + character_maximum_length: None, + } + }) + .collect()) } -pub async fn execute_query(client: &ChClient, database: &str, sql: &str) -> Result { +pub async fn execute_query( + client: &ChClient, + database: &str, + sql: &str, +) -> Result { let start = Instant::now(); let trimmed = sql.trim().to_uppercase(); @@ -154,7 +207,10 @@ pub async fn execute_query(client: &ChClient, database: &str, sql: &str) -> Resu truncated: false, }) } else { - let url = format!("{}/?default_format=JSONCompact&database={}", client.base_url, database); + let url = format!( + "{}/?default_format=JSONCompact&database={}", + client.base_url, database + ); let req = build_request(client, client.http.post(&url).body(sql.to_string())); let resp = req .send() diff --git a/crates/dbx-core/src/db/elasticsearch_driver.rs b/crates/dbx-core/src/db/elasticsearch_driver.rs index d1ea6d1e2..b026c2aab 100644 --- a/crates/dbx-core/src/db/elasticsearch_driver.rs +++ b/crates/dbx-core/src/db/elasticsearch_driver.rs @@ -1,6 +1,7 @@ use reqwest::Client as HttpClient; use serde::Deserialize; +use super::{connection_timeout, with_connection_timeout}; use crate::db::mongo_driver::MongoDocumentResult; pub struct EsClient { @@ -15,8 +16,12 @@ impl EsClient { (Some(u), Some(p)) if !u.is_empty() => Some((u.to_string(), p.to_string())), _ => None, }; + let http = HttpClient::builder() + .connect_timeout(connection_timeout()) + .build() + .unwrap_or_else(|_| HttpClient::new()); Self { - http: HttpClient::new(), + http, base_url: url.trim_end_matches('/').to_string(), auth, } @@ -62,10 +67,14 @@ impl Clone for EsClient { } pub async fn test_connection(client: &EsClient) -> Result<(), String> { - let resp = client.get("/") - .send() - .await - .map_err(|e| format!("Elasticsearch connection failed: {e}"))?; + let resp = with_connection_timeout("Elasticsearch", async { + client + .get("/") + .send() + .await + .map_err(|e| format!("Elasticsearch connection failed: {e}")) + }) + .await?; if !resp.status().is_success() { let body = resp.text().await.unwrap_or_default(); return Err(format!("Elasticsearch error: {body}")); @@ -79,7 +88,8 @@ struct CatIndex { } pub async fn list_indices(client: &EsClient) -> Result, String> { - let resp = client.get("/_cat/indices?format=json&h=index") + let resp = client + .get("/_cat/indices?format=json&h=index") .send() .await .map_err(|e| format!("Elasticsearch request failed: {e}"))?; @@ -87,7 +97,10 @@ pub async fn list_indices(client: &EsClient) -> Result, String> { let body = resp.text().await.unwrap_or_default(); return Err(format!("Elasticsearch error: {body}")); } - let indices: Vec = resp.json().await.map_err(|e| format!("Elasticsearch parse error: {e}"))?; + let indices: Vec = resp + .json() + .await + .map_err(|e| format!("Elasticsearch parse error: {e}"))?; let mut names: Vec = indices .into_iter() .filter(|i| !i.index.starts_with('.')) @@ -134,7 +147,8 @@ pub async fn find_documents( }); let path = format!("/{}/_search", index); - let resp = client.post(&path) + let resp = client + .post(&path) .json(&body) .send() .await @@ -145,17 +159,24 @@ pub async fn find_documents( return Err(format!("Elasticsearch error: {body}")); } - let result: SearchResponse = resp.json().await + let result: SearchResponse = resp + .json() + .await .map_err(|e| format!("Elasticsearch parse error: {e}"))?; - let documents: Vec = result.hits.hits.into_iter().map(|hit| { - let mut doc = match hit.source { - serde_json::Value::Object(map) => map, - _ => serde_json::Map::new(), - }; - doc.insert("_id".to_string(), serde_json::Value::String(hit.id)); - serde_json::Value::Object(doc) - }).collect(); + let documents: Vec = result + .hits + .hits + .into_iter() + .map(|hit| { + let mut doc = match hit.source { + serde_json::Value::Object(map) => map, + _ => serde_json::Map::new(), + }; + doc.insert("_id".to_string(), serde_json::Value::String(hit.id)); + serde_json::Value::Object(doc) + }) + .collect(); Ok(MongoDocumentResult { documents, @@ -168,11 +189,12 @@ pub async fn insert_document( index: &str, doc_json: &str, ) -> Result { - let doc: serde_json::Value = serde_json::from_str(doc_json) - .map_err(|e| format!("Invalid JSON: {e}"))?; + let doc: serde_json::Value = + serde_json::from_str(doc_json).map_err(|e| format!("Invalid JSON: {e}"))?; let path = format!("/{}/_doc?refresh=true", index); - let resp = client.post(&path) + let resp = client + .post(&path) .json(&doc) .send() .await @@ -183,7 +205,9 @@ pub async fn insert_document( return Err(format!("Elasticsearch error: {body}")); } - let result: serde_json::Value = resp.json().await + let result: serde_json::Value = resp + .json() + .await .map_err(|e| format!("Elasticsearch parse error: {e}"))?; Ok(result["_id"].as_str().unwrap_or("").to_string()) } @@ -194,11 +218,12 @@ pub async fn update_document( id: &str, doc_json: &str, ) -> Result { - let doc: serde_json::Value = serde_json::from_str(doc_json) - .map_err(|e| format!("Invalid JSON: {e}"))?; + let doc: serde_json::Value = + serde_json::from_str(doc_json).map_err(|e| format!("Invalid JSON: {e}"))?; let path = format!("/{}/_doc/{}?refresh=true", index, id); - let resp = client.put(&path) + let resp = client + .put(&path) .json(&doc) .send() .await @@ -212,13 +237,10 @@ pub async fn update_document( Ok(1) } -pub async fn delete_document( - client: &EsClient, - index: &str, - id: &str, -) -> Result { +pub async fn delete_document(client: &EsClient, index: &str, id: &str) -> Result { let path = format!("/{}/_doc/{}?refresh=true", index, id); - let resp = client.delete(&path) + let resp = client + .delete(&path) .send() .await .map_err(|e| format!("Elasticsearch request failed: {e}"))?; diff --git a/crates/dbx-core/src/db/mod.rs b/crates/dbx-core/src/db/mod.rs index 1219d4262..341e247c8 100644 --- a/crates/dbx-core/src/db/mod.rs +++ b/crates/dbx-core/src/db/mod.rs @@ -9,5 +9,39 @@ pub mod sqlite; pub mod sqlserver; pub mod ssh_tunnel; +use std::future::Future; +use std::time::Duration; + // Re-export types so that `db::QueryResult` etc. work within dbx-core pub use crate::types::*; + +pub const CONNECTION_TIMEOUT_SECS: u64 = 5; +pub const TCP_PROBE_TIMEOUT_SECS: u64 = 3; + +pub fn connection_timeout() -> Duration { + Duration::from_secs(CONNECTION_TIMEOUT_SECS) +} + +pub fn tcp_probe_timeout() -> Duration { + Duration::from_secs(TCP_PROBE_TIMEOUT_SECS) +} + +pub async fn with_connection_timeout(label: &str, future: F) -> Result +where + F: Future>, +{ + tokio::time::timeout(connection_timeout(), future) + .await + .map_err(|_| format!("{label} connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))? +} + +pub async fn probe_tcp_endpoint(label: &str, host: &str, port: u16) -> Result<(), String> { + tokio::time::timeout( + tcp_probe_timeout(), + tokio::net::TcpStream::connect((host, port)), + ) + .await + .map_err(|_| format!("{label} TCP connection timed out ({TCP_PROBE_TIMEOUT_SECS}s)"))? + .map(|_| ()) + .map_err(|e| format!("{label} TCP connection failed: {e}")) +} diff --git a/crates/dbx-core/src/db/mongo_driver.rs b/crates/dbx-core/src/db/mongo_driver.rs index d27a0cf9b..3d78b3dda 100644 --- a/crates/dbx-core/src/db/mongo_driver.rs +++ b/crates/dbx-core/src/db/mongo_driver.rs @@ -1,12 +1,34 @@ -use mongodb::{bson::{doc, Document, Bson}, Client}; +use mongodb::{ + bson::{doc, Bson, Document}, + Client, +}; use serde::{Deserialize, Serialize}; +use super::{connection_timeout, with_connection_timeout, CONNECTION_TIMEOUT_SECS}; + #[derive(Debug, Clone, Serialize, Deserialize)] pub struct MongoDocumentResult { pub documents: Vec, pub total: u64, } +pub async fn connect(url: &str) -> Result { + with_connection_timeout("MongoDB", async { + Client::with_uri_str(url) + .await + .map_err(|e| format!("MongoDB connection failed: {e}")) + }) + .await +} + +pub async fn test_connection(client: &Client) -> Result<(), String> { + tokio::time::timeout(connection_timeout(), client.list_database_names()) + .await + .map_err(|_| format!("MongoDB connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))? + .map(|_| ()) + .map_err(|e| format!("MongoDB connection failed: {e}")) +} + pub async fn list_databases(client: &Client) -> Result, String> { client .list_database_names() @@ -31,7 +53,10 @@ pub async fn find_documents( ) -> Result { let col = client.database(database).collection::(collection); - let total = col.count_documents(doc! {}).await.map_err(|e| e.to_string())?; + let total = col + .count_documents(doc! {}) + .await + .map_err(|e| e.to_string())?; let mut cursor = col .find(doc! {}) @@ -56,8 +81,7 @@ pub async fn insert_document( collection: &str, doc_json: &str, ) -> Result { - let doc: Document = serde_json::from_str(doc_json) - .map_err(|e| format!("Invalid JSON: {e}"))?; + let doc: Document = serde_json::from_str(doc_json).map_err(|e| format!("Invalid JSON: {e}"))?; let col = client.database(database).collection::(collection); let result = col.insert_one(doc).await.map_err(|e| e.to_string())?; Ok(format!("{}", result.inserted_id)) @@ -72,8 +96,8 @@ pub async fn update_document( ) -> Result { let oid = mongodb::bson::oid::ObjectId::parse_str(id) .map_err(|e| format!("Invalid ObjectId: {e}"))?; - let new_doc: Document = serde_json::from_str(doc_json) - .map_err(|e| format!("Invalid JSON: {e}"))?; + let new_doc: Document = + serde_json::from_str(doc_json).map_err(|e| format!("Invalid JSON: {e}"))?; let col = client.database(database).collection::(collection); let result = col .replace_one(doc! { "_id": oid }, new_doc) diff --git a/crates/dbx-core/src/db/mysql.rs b/crates/dbx-core/src/db/mysql.rs index d11860d12..a9bfbfb8e 100644 --- a/crates/dbx-core/src/db/mysql.rs +++ b/crates/dbx-core/src/db/mysql.rs @@ -4,7 +4,10 @@ use sqlx::mysql::{MySqlPool, MySqlPoolOptions, MySqlRow}; use sqlx::{Column, Executor, Row, TypeInfo, ValueRef}; use std::time::{Duration, Instant}; -use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo}; +use super::{connection_timeout, with_connection_timeout}; +use crate::types::{ + ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo, +}; fn quote_value(s: &str) -> String { format!("'{}'", s.replace('\\', "\\\\").replace('\'', "\\'")) @@ -12,13 +15,19 @@ fn quote_value(s: &str) -> String { fn get_str(row: &MySqlRow, idx: usize) -> String { row.try_get::(idx) - .or_else(|_| row.try_get::, _>(idx).map(|b| String::from_utf8_lossy(&b).to_string())) + .or_else(|_| { + row.try_get::, _>(idx) + .map(|b| String::from_utf8_lossy(&b).to_string()) + }) .unwrap_or_default() } fn get_str_by_name(row: &MySqlRow, name: &str) -> String { row.try_get::(name) - .or_else(|_| row.try_get::, _>(name).map(|b| String::from_utf8_lossy(&b).to_string())) + .or_else(|_| { + row.try_get::, _>(name) + .map(|b| String::from_utf8_lossy(&b).to_string()) + }) .unwrap_or_default() } @@ -43,7 +52,8 @@ fn numeric_metadata_i64_to_i32(value: Option) -> Option { } fn numeric_metadata_str_to_i32(value: Option) -> Option { - value.and_then(|v| v.parse::().ok()) + value + .and_then(|v| v.parse::().ok()) .and_then(|v| i32::try_from(v).ok()) } @@ -57,7 +67,9 @@ fn get_opt_i32(row: &MySqlRow, name: &str) -> Option { .flatten() .or_else(|| numeric_metadata_i64_to_i32(row.try_get::, _>(name).ok().flatten())) .or_else(|| numeric_metadata_u64_to_i32(row.try_get::, _>(name).ok().flatten())) - .or_else(|| numeric_metadata_str_to_i32(row.try_get::, _>(name).ok().flatten())) + .or_else(|| { + numeric_metadata_str_to_i32(row.try_get::, _>(name).ok().flatten()) + }) .or_else(|| { row.try_get::>, _>(name) .ok() @@ -95,7 +107,8 @@ fn mysql_value_to_json(row: &MySqlRow, idx: usize, type_name: &str) -> serde_jso return v; } if let Ok(v) = row.try_get::(idx) { - return serde_json::from_str::(&v).unwrap_or(serde_json::Value::String(v)); + return serde_json::from_str::(&v) + .unwrap_or(serde_json::Value::String(v)); } return serde_json::Value::Null; } @@ -138,13 +151,21 @@ fn mysql_value_to_json(row: &MySqlRow, idx: usize, type_name: &str) -> serde_jso row.try_get::(idx) .map(serde_json::Value::String) - .or_else(|_| row.try_get::(idx).map(|v| serde_json::Value::Number(v.into()))) - .or_else(|_| row.try_get::(idx).map(|v| serde_json::Value::Number(v.into()))) - .or_else(|_| row.try_get::(idx).map(|v| { - serde_json::Number::from_f64(v) - .map(serde_json::Value::Number) - .unwrap_or(serde_json::Value::Null) - })) + .or_else(|_| { + row.try_get::(idx) + .map(|v| serde_json::Value::Number(v.into())) + }) + .or_else(|_| { + row.try_get::(idx) + .map(|v| serde_json::Value::Number(v.into())) + }) + .or_else(|_| { + row.try_get::(idx).map(|v| { + serde_json::Number::from_f64(v) + .map(serde_json::Value::Number) + .unwrap_or(serde_json::Value::Null) + }) + }) .or_else(|_| row.try_get::(idx).map(serde_json::Value::Bool)) .or_else(|_| { row.try_get::, _>(idx) @@ -155,39 +176,52 @@ fn mysql_value_to_json(row: &MySqlRow, idx: usize, type_name: &str) -> serde_jso } pub async fn connect(url: &str) -> Result { - MySqlPoolOptions::new() - .max_connections(5) - .acquire_timeout(Duration::from_secs(10)) - .idle_timeout(Duration::from_secs(300)) - .connect(url) - .await - .map_err(|e| format!("MySQL connection failed: {e}")) + with_connection_timeout("MySQL", async { + MySqlPoolOptions::new() + .max_connections(5) + .acquire_timeout(connection_timeout()) + .idle_timeout(Duration::from_secs(300)) + .connect(url) + .await + .map_err(|e| format!("MySQL connection failed: {e}")) + }) + .await } pub async fn connect_bare(url: &str) -> Result { - let options: sqlx::mysql::MySqlConnectOptions = url.parse() + let options: sqlx::mysql::MySqlConnectOptions = url + .parse() .map_err(|e: sqlx::Error| format!("Invalid MySQL URL: {e}"))?; let options = options .no_engine_substitution(false) .set_names(false) .pipes_as_concat(false) .timezone(None); - MySqlPoolOptions::new() - .max_connections(5) - .acquire_timeout(Duration::from_secs(10)) - .idle_timeout(Duration::from_secs(300)) - .connect_with(options) - .await - .map_err(|e| format!("MySQL connection failed: {e}")) + with_connection_timeout("MySQL", async { + MySqlPoolOptions::new() + .max_connections(5) + .acquire_timeout(connection_timeout()) + .idle_timeout(Duration::from_secs(300)) + .connect_with(options) + .await + .map_err(|e| format!("MySQL connection failed: {e}")) + }) + .await } pub async fn list_databases(pool: &MySqlPool) -> Result, String> { - let rows: Vec = sqlx::raw_sql("SELECT SCHEMA_NAME FROM information_schema.SCHEMATA ORDER BY SCHEMA_NAME") - .fetch_all(pool) - .await - .map_err(|e| e.to_string())?; + let rows: Vec = + sqlx::raw_sql("SELECT SCHEMA_NAME FROM information_schema.SCHEMATA ORDER BY SCHEMA_NAME") + .fetch_all(pool) + .await + .map_err(|e| e.to_string())?; - Ok(rows.iter().map(|row| DatabaseInfo { name: get_str(row, 0) }).collect()) + Ok(rows + .iter() + .map(|row| DatabaseInfo { + name: get_str(row, 0), + }) + .collect()) } pub async fn list_tables(pool: &MySqlPool, database: &str) -> Result, String> { @@ -255,7 +289,11 @@ pub async fn execute_query(pool: &MySqlPool, sql: &str, bare: bool) -> Result = sqlx::raw_sql(sql) .fetch_all(pool) @@ -263,8 +301,16 @@ pub async fn execute_query(pool: &MySqlPool, sql: &str, bare: bool) -> Result = first.columns().iter().map(|c| c.name().to_string()).collect(); - let types: Vec = first.columns().iter().map(|c| c.type_info().name().to_string()).collect(); + let cols: Vec = first + .columns() + .iter() + .map(|c| c.name().to_string()) + .collect(); + let types: Vec = first + .columns() + .iter() + .map(|c| c.type_info().name().to_string()) + .collect(); (cols, types) } else { (vec![], vec![]) @@ -274,7 +320,13 @@ pub async fn execute_query(pool: &MySqlPool, sql: &str, bare: bool) -> Result Result = desc.columns().iter().map(|c| c.name().to_string()).collect(); - let column_types: Vec = desc.columns().iter().map(|c| c.type_info().name().to_string()).collect(); + let columns: Vec = desc + .columns() + .iter() + .map(|c| c.name().to_string()) + .collect(); + let column_types: Vec = desc + .columns() + .iter() + .map(|c| c.type_info().name().to_string()) + .collect(); let rows: Vec = sqlx::query(sql) .fetch_all(pool) @@ -300,7 +360,13 @@ pub async fn execute_query(pool: &MySqlPool, sql: &str, bare: bool) -> Result Result Result, String> { +pub async fn list_indexes( + pool: &MySqlPool, + database: &str, + table: &str, +) -> Result, String> { let sql = format!( "SELECT INDEX_NAME, GROUP_CONCAT(COLUMN_NAME ORDER BY SEQ_IN_INDEX) AS columns, \ MIN(NON_UNIQUE) = 0 AS is_unique, INDEX_NAME = 'PRIMARY' AS is_primary, \ @@ -352,7 +422,11 @@ pub async fn list_indexes(pool: &MySqlPool, database: &str, table: &str) -> Resu let cols_str = get_str_by_name(row, "columns"); IndexInfo { name: get_str_by_name(row, "INDEX_NAME"), - columns: cols_str.split(',').filter(|s| !s.is_empty()).map(|s| s.to_string()).collect(), + columns: cols_str + .split(',') + .filter(|s| !s.is_empty()) + .map(|s| s.to_string()) + .collect(), is_unique: row.get::("is_unique"), is_primary: row.get::("is_primary"), filter: None, @@ -364,7 +438,11 @@ pub async fn list_indexes(pool: &MySqlPool, database: &str, table: &str) -> Resu .collect()) } -pub async fn list_foreign_keys(pool: &MySqlPool, database: &str, table: &str) -> Result, String> { +pub async fn list_foreign_keys( + pool: &MySqlPool, + database: &str, + table: &str, +) -> Result, String> { let sql = format!( "SELECT kcu.CONSTRAINT_NAME, kcu.COLUMN_NAME, \ kcu.REFERENCED_TABLE_NAME, kcu.REFERENCED_COLUMN_NAME \ @@ -391,7 +469,11 @@ pub async fn list_foreign_keys(pool: &MySqlPool, database: &str, table: &str) -> .collect()) } -pub async fn list_triggers(pool: &MySqlPool, database: &str, table: &str) -> Result, String> { +pub async fn list_triggers( + pool: &MySqlPool, + database: &str, + table: &str, +) -> Result, String> { let sql = format!( "SELECT TRIGGER_NAME, EVENT_MANIPULATION, ACTION_TIMING \ FROM information_schema.TRIGGERS \ diff --git a/crates/dbx-core/src/db/oracle_driver.rs b/crates/dbx-core/src/db/oracle_driver.rs index 9e34ad8db..b60c6d295 100644 --- a/crates/dbx-core/src/db/oracle_driver.rs +++ b/crates/dbx-core/src/db/oracle_driver.rs @@ -1,15 +1,28 @@ use oracle_rs::{Config, Connection}; use std::time::Instant; -use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo}; +use super::{connection_timeout, CONNECTION_TIMEOUT_SECS}; +use crate::types::{ + ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo, +}; pub type OracleClient = Connection; -pub async fn connect(host: &str, port: u16, service: &str, user: &str, pass: &str) -> Result { +pub async fn connect( + host: &str, + port: u16, + service: &str, + user: &str, + pass: &str, +) -> Result { let config = Config::new(host, port, service, user, pass); - Connection::connect_with_config(config) - .await - .map_err(|e| format!("Oracle connection failed: {e}")) + tokio::time::timeout( + connection_timeout(), + Connection::connect_with_config(config), + ) + .await + .map_err(|_| format!("Oracle connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))? + .map_err(|e| format!("Oracle connection failed: {e}")) } fn value_to_json(val: &oracle_rs::Value) -> serde_json::Value { @@ -27,23 +40,29 @@ fn value_to_json(val: &oracle_rs::Value) -> serde_json::Value { } pub async fn list_databases(conn: &OracleClient) -> Result, String> { - let result = conn.query( - "SELECT username FROM all_users ORDER BY username", - &[], - ).await.map_err(|e| e.to_string())?; - Ok(result.rows.iter().map(|row| { - DatabaseInfo { name: row.get_string(0).unwrap_or("").to_string() } - }).collect()) + let result = conn + .query("SELECT username FROM all_users ORDER BY username", &[]) + .await + .map_err(|e| e.to_string())?; + Ok(result + .rows + .iter() + .map(|row| DatabaseInfo { + name: row.get_string(0).unwrap_or("").to_string(), + }) + .collect()) } pub async fn list_schemas(conn: &OracleClient) -> Result, String> { - let result = conn.query( - "SELECT username FROM all_users ORDER BY username", - &[], - ).await.map_err(|e| e.to_string())?; - Ok(result.rows.iter().map(|row| { - row.get_string(0).unwrap_or("").to_string() - }).collect()) + let result = conn + .query("SELECT username FROM all_users ORDER BY username", &[]) + .await + .map_err(|e| e.to_string())?; + Ok(result + .rows + .iter() + .map(|row| row.get_string(0).unwrap_or("").to_string()) + .collect()) } pub async fn list_tables(conn: &OracleClient, schema: &str) -> Result, String> { @@ -55,15 +74,21 @@ pub async fn list_tables(conn: &OracleClient, schema: &str) -> Result Result, String> { +pub async fn get_columns( + conn: &OracleClient, + schema: &str, + table: &str, +) -> Result, String> { let s = schema.replace('\'', "''"); let t = table.replace('\'', "''"); @@ -75,7 +100,9 @@ pub async fn get_columns(conn: &OracleClient, schema: &str, table: &str) -> Resu ), &[], ).await.map_err(|e| e.to_string())?; - let pk_names: std::collections::HashSet = pk_result.rows.iter() + let pk_names: std::collections::HashSet = pk_result + .rows + .iter() .filter_map(|row| row.get_string(0).map(|s| s.to_string())) .collect(); @@ -89,47 +116,56 @@ pub async fn get_columns(conn: &OracleClient, schema: &str, table: &str) -> Resu &[], ).await.map_err(|e| e.to_string())?; - Ok(col_result.rows.iter().map(|row| { - let name = row.get_string(0).unwrap_or("").to_string(); - let base = row.get_string(1).unwrap_or("").to_string(); - let data_len = row.get_i64(5).map(|v| v as i32); - let char_len = row.get_i64(6).map(|v| v as i32); - let num_prec = row.get_i64(3).map(|v| v as i32); - let num_scale = row.get_i64(4).map(|v| v as i32); - let data_type = match base.to_uppercase().as_str() { - "VARCHAR2" | "NVARCHAR2" | "CHAR" | "NCHAR" => { - let len = char_len.or(data_len); - match len { - Some(n) => format!("{base}({n})"), - None => base, + Ok(col_result + .rows + .iter() + .map(|row| { + let name = row.get_string(0).unwrap_or("").to_string(); + let base = row.get_string(1).unwrap_or("").to_string(); + let data_len = row.get_i64(5).map(|v| v as i32); + let char_len = row.get_i64(6).map(|v| v as i32); + let num_prec = row.get_i64(3).map(|v| v as i32); + let num_scale = row.get_i64(4).map(|v| v as i32); + let data_type = match base.to_uppercase().as_str() { + "VARCHAR2" | "NVARCHAR2" | "CHAR" | "NCHAR" => { + let len = char_len.or(data_len); + match len { + Some(n) => format!("{base}({n})"), + None => base, + } } + "NUMBER" => match (num_prec, num_scale) { + (Some(p), Some(s)) if s > 0 => format!("NUMBER({p},{s})"), + (Some(p), _) if p > 0 => format!("NUMBER({p})"), + _ => "NUMBER".to_string(), + }, + "RAW" => match data_len { + Some(n) => format!("RAW({n})"), + None => "RAW".to_string(), + }, + _ => base, + }; + ColumnInfo { + is_primary_key: pk_names.contains(&name), + name, + data_type, + is_nullable: row.get_string(2).unwrap_or("N") == "Y", + column_default: None, + extra: None, + comment: None, + numeric_precision: num_prec, + numeric_scale: num_scale, + character_maximum_length: char_len, } - "NUMBER" => match (num_prec, num_scale) { - (Some(p), Some(s)) if s > 0 => format!("NUMBER({p},{s})"), - (Some(p), _) if p > 0 => format!("NUMBER({p})"), - _ => "NUMBER".to_string(), - }, - "RAW" => match data_len { - Some(n) => format!("RAW({n})"), - None => "RAW".to_string(), - }, - _ => base, - }; - ColumnInfo { - is_primary_key: pk_names.contains(&name), - name, - data_type, - is_nullable: row.get_string(2).unwrap_or("N") == "Y", - column_default: None, - extra: None, comment: None, - numeric_precision: num_prec, - numeric_scale: num_scale, - character_maximum_length: char_len, - } - }).collect()) + }) + .collect()) } -pub async fn list_indexes(conn: &OracleClient, schema: &str, table: &str) -> Result, String> { +pub async fn list_indexes( + conn: &OracleClient, + schema: &str, + table: &str, +) -> Result, String> { let sql = format!( "SELECT i.INDEX_NAME, \ LISTAGG(ic.COLUMN_NAME, ',') WITHIN GROUP (ORDER BY ic.COLUMN_POSITION) AS columns, \ @@ -146,22 +182,34 @@ pub async fn list_indexes(conn: &OracleClient, schema: &str, table: &str) -> Res s = schema.replace('\'', "''"), t = table.replace('\'', "''") ); let result = conn.query(&sql, &[]).await.map_err(|e| e.to_string())?; - Ok(result.rows.iter().map(|row| { - let cols_str = row.get_string(1).unwrap_or(""); - IndexInfo { - name: row.get_string(0).unwrap_or("").to_string(), - columns: cols_str.split(',').filter(|s| !s.is_empty()).map(|s| s.to_string()).collect(), - is_unique: row.get_string(2).unwrap_or("") == "UNIQUE", - is_primary: row.get_i64(3).unwrap_or(0) == 1, - filter: None, - index_type: row.get_string(4).map(|s| s.to_string()), - included_columns: None, - comment: None, - } - }).collect()) + Ok(result + .rows + .iter() + .map(|row| { + let cols_str = row.get_string(1).unwrap_or(""); + IndexInfo { + name: row.get_string(0).unwrap_or("").to_string(), + columns: cols_str + .split(',') + .filter(|s| !s.is_empty()) + .map(|s| s.to_string()) + .collect(), + is_unique: row.get_string(2).unwrap_or("") == "UNIQUE", + is_primary: row.get_i64(3).unwrap_or(0) == 1, + filter: None, + index_type: row.get_string(4).map(|s| s.to_string()), + included_columns: None, + comment: None, + } + }) + .collect()) } -pub async fn list_foreign_keys(conn: &OracleClient, schema: &str, table: &str) -> Result, String> { +pub async fn list_foreign_keys( + conn: &OracleClient, + schema: &str, + table: &str, +) -> Result, String> { let sql = format!( "SELECT c.CONSTRAINT_NAME, cc.COLUMN_NAME, rc.TABLE_NAME, rcc.COLUMN_NAME \ FROM ALL_CONSTRAINTS c \ @@ -173,32 +221,41 @@ pub async fn list_foreign_keys(conn: &OracleClient, schema: &str, table: &str) - s = schema.replace('\'', "''"), t = table.replace('\'', "''") ); let result = conn.query(&sql, &[]).await.map_err(|e| e.to_string())?; - Ok(result.rows.iter().map(|row| { - ForeignKeyInfo { + Ok(result + .rows + .iter() + .map(|row| ForeignKeyInfo { name: row.get_string(0).unwrap_or("").to_string(), column: row.get_string(1).unwrap_or("").to_string(), ref_table: row.get_string(2).unwrap_or("").to_string(), ref_column: row.get_string(3).unwrap_or("").to_string(), - } - }).collect()) + }) + .collect()) } -pub async fn list_triggers(conn: &OracleClient, schema: &str, table: &str) -> Result, String> { +pub async fn list_triggers( + conn: &OracleClient, + schema: &str, + table: &str, +) -> Result, String> { let sql = format!( "SELECT TRIGGER_NAME, TRIGGERING_EVENT, TRIGGER_TYPE \ FROM ALL_TRIGGERS \ WHERE OWNER = '{s}' AND TABLE_NAME = '{t}' \ ORDER BY TRIGGER_NAME", - s = schema.replace('\'', "''"), t = table.replace('\'', "''") + s = schema.replace('\'', "''"), + t = table.replace('\'', "''") ); let result = conn.query(&sql, &[]).await.map_err(|e| e.to_string())?; - Ok(result.rows.iter().map(|row| { - TriggerInfo { + Ok(result + .rows + .iter() + .map(|row| TriggerInfo { name: row.get_string(0).unwrap_or("").to_string(), event: row.get_string(1).unwrap_or("").to_string(), timing: row.get_string(2).unwrap_or("").to_string(), - } - }).collect()) + }) + .collect()) } pub async fn execute_query(conn: &OracleClient, sql: &str) -> Result { @@ -214,13 +271,19 @@ pub async fn execute_query(conn: &OracleClient, sql: &str) -> Result = result.columns.iter().map(|c| c.name.clone()).collect(); - let rows: Vec> = result.rows.iter().map(|row| { - (0..columns.len()).map(|i| { - row.get(i) - .map(|v| value_to_json(v)) - .unwrap_or(serde_json::Value::Null) - }).collect() - }).collect(); + let rows: Vec> = result + .rows + .iter() + .map(|row| { + (0..columns.len()) + .map(|i| { + row.get(i) + .map(|v| value_to_json(v)) + .unwrap_or(serde_json::Value::Null) + }) + .collect() + }) + .collect(); Ok(QueryResult { columns, diff --git a/crates/dbx-core/src/db/postgres.rs b/crates/dbx-core/src/db/postgres.rs index 6c1bd7040..e779cd201 100644 --- a/crates/dbx-core/src/db/postgres.rs +++ b/crates/dbx-core/src/db/postgres.rs @@ -4,7 +4,10 @@ use sqlx::postgres::{PgPool, PgPoolOptions, PgRow}; use sqlx::{Column, Executor, Row, TypeInfo, ValueRef}; use std::time::{Duration, Instant}; -use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo}; +use super::{connection_timeout, with_connection_timeout}; +use crate::types::{ + ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo, +}; fn pg_temporal_to_json_value(row: &PgRow, idx: usize) -> Option { if let Ok(v) = row.try_get::, _>(idx) { @@ -34,7 +37,8 @@ fn pg_value_to_json(row: &PgRow, idx: usize, type_name: &str) -> serde_json::Val return v; } if let Ok(v) = row.try_get::(idx) { - return serde_json::from_str::(&v).unwrap_or(serde_json::Value::String(v)); + return serde_json::from_str::(&v) + .unwrap_or(serde_json::Value::String(v)); } return serde_json::Value::Null; } @@ -87,13 +91,16 @@ fn pg_value_to_json(row: &PgRow, idx: usize, type_name: &str) -> serde_json::Val } pub async fn connect(url: &str) -> Result { - PgPoolOptions::new() - .max_connections(5) - .acquire_timeout(Duration::from_secs(10)) - .idle_timeout(Duration::from_secs(300)) - .connect(url) - .await - .map_err(|e| format!("PostgreSQL connection failed: {e}")) + with_connection_timeout("PostgreSQL", async { + PgPoolOptions::new() + .max_connections(5) + .acquire_timeout(connection_timeout()) + .idle_timeout(Duration::from_secs(300)) + .connect(url) + .await + .map_err(|e| format!("PostgreSQL connection failed: {e}")) + }) + .await } pub async fn list_databases(pool: &PgPool) -> Result, String> { @@ -187,7 +194,9 @@ pub async fn get_columns( Ok(rows .iter() .map(|row| { - let full_type = row.get::, _>("full_type").unwrap_or_default(); + let full_type = row + .get::, _>("full_type") + .unwrap_or_default(); ColumnInfo { name: row.get::("column_name"), data_type: full_type, @@ -220,17 +229,26 @@ pub async fn execute_query(pool: &PgPool, sql: &str) -> Result, Vec) = if let Some(first) = rows.first() { + let (columns, column_types): (Vec, Vec) = if let Some(first) = rows.first() + { let cols = first.columns(); ( cols.iter().map(|c| c.name().to_string()).collect(), - cols.iter().map(|c| c.type_info().name().to_string()).collect(), + cols.iter() + .map(|c| c.type_info().name().to_string()) + .collect(), ) } else { let desc = pool.describe(sql).await.map_err(|e| e.to_string())?; ( - desc.columns().iter().map(|c| c.name().to_string()).collect(), - desc.columns().iter().map(|c| c.type_info().name().to_string()).collect(), + desc.columns() + .iter() + .map(|c| c.name().to_string()) + .collect(), + desc.columns() + .iter() + .map(|c| c.type_info().name().to_string()) + .collect(), ) }; @@ -272,7 +290,11 @@ pub async fn execute_query(pool: &PgPool, sql: &str) -> Result Result, String> { +pub async fn list_indexes( + pool: &PgPool, + schema: &str, + table: &str, +) -> Result, String> { let rows: Vec = sqlx::query( "SELECT i.relname AS index_name, \ array_agg(COALESCE(a.attname, pg_get_indexdef(ix.indexrelid, k.n::int, true)) ORDER BY k.n) AS columns, \ @@ -304,9 +326,15 @@ pub async fn list_indexes(pool: &PgPool, schema: &str, table: &str) -> Result = row.get::, _>("columns"); - let nkeyatts = row.get::, _>("nkeyatts").unwrap_or(all_cols.len() as i16) as usize; + let nkeyatts = row + .get::, _>("nkeyatts") + .unwrap_or(all_cols.len() as i16) as usize; let key_cols = all_cols[..nkeyatts].to_vec(); - let included = if nkeyatts < all_cols.len() { all_cols[nkeyatts..].to_vec() } else { vec![] }; + let included = if nkeyatts < all_cols.len() { + all_cols[nkeyatts..].to_vec() + } else { + vec![] + }; IndexInfo { name: row.get::("index_name"), columns: key_cols, @@ -314,14 +342,22 @@ pub async fn list_indexes(pool: &PgPool, schema: &str, table: &str) -> Result("is_primary"), filter: row.get::, _>("filter_expr"), index_type: row.get::, _>("index_type"), - included_columns: if included.is_empty() { None } else { Some(included) }, + included_columns: if included.is_empty() { + None + } else { + Some(included) + }, comment: row.get::, _>("index_comment"), } }) .collect()) } -pub async fn list_foreign_keys(pool: &PgPool, schema: &str, table: &str) -> Result, String> { +pub async fn list_foreign_keys( + pool: &PgPool, + schema: &str, + table: &str, +) -> Result, String> { let rows: Vec = sqlx::query( "SELECT kcu.constraint_name, kcu.column_name, \ ccu.table_name AS ref_table, ccu.column_name AS ref_column \ @@ -352,7 +388,11 @@ pub async fn list_foreign_keys(pool: &PgPool, schema: &str, table: &str) -> Resu .collect()) } -pub async fn list_triggers(pool: &PgPool, schema: &str, table: &str) -> Result, String> { +pub async fn list_triggers( + pool: &PgPool, + schema: &str, + table: &str, +) -> Result, String> { let rows: Vec = sqlx::query( "SELECT trigger_name, event_manipulation, action_timing \ FROM information_schema.triggers \ diff --git a/crates/dbx-core/src/db/redis_driver.rs b/crates/dbx-core/src/db/redis_driver.rs index f0dd1b03b..25dd9ebe0 100644 --- a/crates/dbx-core/src/db/redis_driver.rs +++ b/crates/dbx-core/src/db/redis_driver.rs @@ -1,6 +1,8 @@ use redis::{AsyncCommands, FromRedisValue, Value as RedisRawValue}; use serde::{Deserialize, Serialize}; +use super::{connection_timeout, CONNECTION_TIMEOUT_SECS}; + const STREAM_ENTRY_LIMIT: usize = 100; const DEFAULT_REDIS_DATABASES: u32 = 16; @@ -28,17 +30,20 @@ pub struct RedisValue { pub async fn connect(url: &str) -> Result { let client = redis::Client::open(url).map_err(|e| format!("Redis connection failed: {e}"))?; let mut con = tokio::time::timeout( - std::time::Duration::from_secs(10), + connection_timeout(), client.get_multiplexed_async_connection(), ) .await - .map_err(|_| "Redis connection timed out (10s)".to_string())? + .map_err(|_| format!("Redis connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))? .map_err(|e| format!("Redis connection failed: {e}"))?; - redis::cmd("PING") - .query_async::(&mut con) - .await - .map_err(|e| format!("Redis authentication failed or command rejected: {e}"))?; + tokio::time::timeout( + connection_timeout(), + redis::cmd("PING").query_async::(&mut con), + ) + .await + .map_err(|_| format!("Redis ping timed out ({CONNECTION_TIMEOUT_SECS}s)"))? + .map_err(|e| format!("Redis authentication failed or command rejected: {e}"))?; Ok(con) } diff --git a/crates/dbx-core/src/db/sqlite.rs b/crates/dbx-core/src/db/sqlite.rs index dd6a50d00..e0a55cb9f 100644 --- a/crates/dbx-core/src/db/sqlite.rs +++ b/crates/dbx-core/src/db/sqlite.rs @@ -2,7 +2,9 @@ use sqlx::sqlite::{SqliteConnectOptions, SqlitePool, SqlitePoolOptions, SqliteRo use sqlx::{Column, Executor, Row}; use std::time::{Duration, Instant}; -use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo}; +use crate::types::{ + ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo, +}; pub async fn connect_path(path: &str) -> Result { let mut options = SqliteConnectOptions::new() @@ -23,11 +25,16 @@ pub async fn connect_path(path: &str) -> Result { } fn is_network_path(path: &str) -> bool { - path.starts_with("\\\\") || path.starts_with("//") || path.contains("wsl.localhost") || path.contains("wsl$") + path.starts_with("\\\\") + || path.starts_with("//") + || path.contains("wsl.localhost") + || path.contains("wsl$") } pub async fn list_databases(_pool: &SqlitePool) -> Result, String> { - Ok(vec![DatabaseInfo { name: "main".to_string() }]) + Ok(vec![DatabaseInfo { + name: "main".to_string(), + }]) } pub async fn list_tables(pool: &SqlitePool, _schema: &str) -> Result, String> { @@ -44,13 +51,21 @@ pub async fn list_tables(pool: &SqlitePool, _schema: &str) -> Result("name"), - table_type: if t == "view" { "VIEW".to_string() } else { "BASE TABLE".to_string() }, + table_type: if t == "view" { + "VIEW".to_string() + } else { + "BASE TABLE".to_string() + }, } }) .collect()) } -pub async fn get_columns(pool: &SqlitePool, _schema: &str, table: &str) -> Result, String> { +pub async fn get_columns( + pool: &SqlitePool, + _schema: &str, + table: &str, +) -> Result, String> { let rows: Vec = sqlx::query(&format!("PRAGMA table_info(\"{}\")", table)) .fetch_all(pool) .await @@ -64,7 +79,8 @@ pub async fn get_columns(pool: &SqlitePool, _schema: &str, table: &str) -> Resul is_nullable: row.get::("notnull") == 0, column_default: row.get::, _>("dflt_value"), is_primary_key: row.get::("pk") > 0, - extra: None, comment: None, + extra: None, + comment: None, numeric_precision: None, numeric_scale: None, character_maximum_length: None, @@ -72,7 +88,11 @@ pub async fn get_columns(pool: &SqlitePool, _schema: &str, table: &str) -> Resul .collect()) } -pub async fn list_indexes(pool: &SqlitePool, _schema: &str, table: &str) -> Result, String> { +pub async fn list_indexes( + pool: &SqlitePool, + _schema: &str, + table: &str, +) -> Result, String> { let safe_table = table.replace('"', "\"\""); let idx_rows: Vec = sqlx::query(&format!("PRAGMA index_list(\"{safe_table}\")")) .fetch_all(pool) @@ -92,7 +112,10 @@ pub async fn list_indexes(pool: &SqlitePool, _schema: &str, table: &str) -> Resu .await .map_err(|e| e.to_string())?; - let columns: Vec = col_rows.iter().map(|r| r.get::("name")).collect(); + let columns: Vec = col_rows + .iter() + .map(|r| r.get::("name")) + .collect(); indexes.push(IndexInfo { name, @@ -108,7 +131,11 @@ pub async fn list_indexes(pool: &SqlitePool, _schema: &str, table: &str) -> Resu Ok(indexes) } -pub async fn list_foreign_keys(pool: &SqlitePool, _schema: &str, table: &str) -> Result, String> { +pub async fn list_foreign_keys( + pool: &SqlitePool, + _schema: &str, + table: &str, +) -> Result, String> { let rows: Vec = sqlx::query(&format!("PRAGMA foreign_key_list(\"{}\")", table)) .fetch_all(pool) .await @@ -125,7 +152,11 @@ pub async fn list_foreign_keys(pool: &SqlitePool, _schema: &str, table: &str) -> .collect()) } -pub async fn list_triggers(pool: &SqlitePool, _schema: &str, table: &str) -> Result, String> { +pub async fn list_triggers( + pool: &SqlitePool, + _schema: &str, + table: &str, +) -> Result, String> { let rows: Vec = sqlx::query( "SELECT name, sql FROM sqlite_master WHERE type = 'trigger' AND tbl_name = ? ORDER BY name", ) @@ -139,8 +170,20 @@ pub async fn list_triggers(pool: &SqlitePool, _schema: &str, table: &str) -> Res .map(|row| { let sql_text: String = row.get::, _>("sql").unwrap_or_default(); let upper = sql_text.to_uppercase(); - let timing = if upper.contains("BEFORE") { "BEFORE" } else if upper.contains("AFTER") { "AFTER" } else { "INSTEAD OF" }; - let event = if upper.contains("INSERT") { "INSERT" } else if upper.contains("UPDATE") { "UPDATE" } else { "DELETE" }; + let timing = if upper.contains("BEFORE") { + "BEFORE" + } else if upper.contains("AFTER") { + "AFTER" + } else { + "INSTEAD OF" + }; + let event = if upper.contains("INSERT") { + "INSERT" + } else if upper.contains("UPDATE") { + "UPDATE" + } else { + "DELETE" + }; TriggerInfo { name: row.get::("name"), event: event.to_string(), @@ -160,7 +203,11 @@ pub async fn execute_query(pool: &SqlitePool, sql: &str) -> Result = desc.columns().iter().map(|c| c.name().to_string()).collect(); + let columns: Vec = desc + .columns() + .iter() + .map(|c| c.name().to_string()) + .collect(); let rows: Vec = sqlx::query(sql) .fetch_all(pool) @@ -174,12 +221,17 @@ pub async fn execute_query(pool: &SqlitePool, sql: &str) -> Result(i) .map(serde_json::Value::String) - .or_else(|_| row.try_get::(i).map(|v| serde_json::Value::Number(v.into()))) - .or_else(|_| row.try_get::(i).map(|v| { - serde_json::Number::from_f64(v) - .map(serde_json::Value::Number) - .unwrap_or(serde_json::Value::Null) - })) + .or_else(|_| { + row.try_get::(i) + .map(|v| serde_json::Value::Number(v.into())) + }) + .or_else(|_| { + row.try_get::(i).map(|v| { + serde_json::Number::from_f64(v) + .map(serde_json::Value::Number) + .unwrap_or(serde_json::Value::Null) + }) + }) .or_else(|_| row.try_get::(i).map(serde_json::Value::Bool)) .unwrap_or(serde_json::Value::Null) }) diff --git a/crates/dbx-core/src/db/sqlserver.rs b/crates/dbx-core/src/db/sqlserver.rs index 90ee7936c..075ad3c38 100644 --- a/crates/dbx-core/src/db/sqlserver.rs +++ b/crates/dbx-core/src/db/sqlserver.rs @@ -1,21 +1,37 @@ use rust_decimal::Decimal; +use std::time::Instant; use tiberius::{AuthMethod, Client, Config}; use tokio::net::TcpStream; use tokio_util::compat::{Compat, TokioAsyncWriteCompatExt}; -use std::time::Instant; -use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo}; +use super::{connection_timeout, CONNECTION_TIMEOUT_SECS}; +use crate::types::{ + ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo, +}; pub type SqlServerClient = Client>; -pub async fn connect(host: &str, port: u16, user: &str, pass: &str, database: Option<&str>) -> Result { +pub async fn connect( + host: &str, + port: u16, + user: &str, + pass: &str, + database: Option<&str>, +) -> Result { match try_connect(host, port, user, pass, database, true).await { Ok(client) => Ok(client), Err(_) => try_connect(host, port, user, pass, database, false).await, } } -async fn try_connect(host: &str, port: u16, user: &str, pass: &str, database: Option<&str>, use_encryption: bool) -> Result { +async fn try_connect( + host: &str, + port: u16, + user: &str, + pass: &str, + database: Option<&str>, + use_encryption: bool, +) -> Result { let mut config = Config::new(); config.host(host); config.port(port); @@ -28,72 +44,107 @@ async fn try_connect(host: &str, port: u16, user: &str, pass: &str, database: Op config.encryption(tiberius::EncryptionLevel::NotSupported); } - let tcp = TcpStream::connect(config.get_addr()) + 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}"))?; - Client::connect(config, tcp.compat_write()) - .await - .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)"))? + .map_err(|e| format!("SQL Server connection failed: {e}")) } fn row_to_json(row: &tiberius::Row) -> Vec { - (0..row.len()).map(|i| { - if let Some(v) = row.try_get::<&str, _>(i).ok().flatten() { - serde_json::Value::String(v.to_string()) - } else if let Some(v) = row.try_get::(i).ok().flatten() { - serde_json::Value::String(v.to_string()) - } else if let Some(v) = row.try_get::(i).ok().flatten() { - serde_json::Value::Number(v.into()) - } else if let Some(v) = row.try_get::(i).ok().flatten() { - serde_json::Value::Number(v.into()) - } else if let Some(v) = row.try_get::(i).ok().flatten() { - serde_json::Number::from_f64(v).map(serde_json::Value::Number).unwrap_or(serde_json::Value::Null) - } else if let Some(v) = row.try_get::(i).ok().flatten() { - serde_json::Value::Bool(v) - } else { - serde_json::Value::Null - } - }).collect() + (0..row.len()) + .map(|i| { + if let Some(v) = row.try_get::<&str, _>(i).ok().flatten() { + serde_json::Value::String(v.to_string()) + } else if let Some(v) = row.try_get::(i).ok().flatten() { + serde_json::Value::String(v.to_string()) + } else if let Some(v) = row.try_get::(i).ok().flatten() { + serde_json::Value::Number(v.into()) + } else if let Some(v) = row.try_get::(i).ok().flatten() { + serde_json::Value::Number(v.into()) + } else if let Some(v) = row.try_get::(i).ok().flatten() { + serde_json::Number::from_f64(v) + .map(serde_json::Value::Number) + .unwrap_or(serde_json::Value::Null) + } else if let Some(v) = row.try_get::(i).ok().flatten() { + serde_json::Value::Bool(v) + } else { + serde_json::Value::Null + } + }) + .collect() } pub async fn list_databases(client: &mut SqlServerClient) -> Result, String> { - let stream = client.query("SELECT name FROM sys.databases ORDER BY name", &[]) - .await.map_err(|e| e.to_string())?; - let rows = stream.into_first_result().await.map_err(|e| e.to_string())?; - Ok(rows.iter().map(|row| { - DatabaseInfo { name: row.get::<&str, _>(0).unwrap_or("").to_string() } - }).collect()) + let stream = client + .query("SELECT name FROM sys.databases ORDER BY name", &[]) + .await + .map_err(|e| e.to_string())?; + let rows = stream + .into_first_result() + .await + .map_err(|e| e.to_string())?; + Ok(rows + .iter() + .map(|row| DatabaseInfo { + name: row.get::<&str, _>(0).unwrap_or("").to_string(), + }) + .collect()) } pub async fn list_schemas(client: &mut SqlServerClient) -> Result, String> { - let stream = client.query( - "SELECT SCHEMA_NAME FROM INFORMATION_SCHEMA.SCHEMATA \ + let stream = client + .query( + "SELECT SCHEMA_NAME FROM INFORMATION_SCHEMA.SCHEMATA \ WHERE SCHEMA_NAME NOT IN ('guest','INFORMATION_SCHEMA','sys') \ ORDER BY SCHEMA_NAME", - &[], - ).await.map_err(|e| e.to_string())?; - let rows = stream.into_first_result().await.map_err(|e| e.to_string())?; - Ok(rows.iter().map(|row| { - row.get::<&str, _>(0).unwrap_or("").to_string() - }).collect()) + &[], + ) + .await + .map_err(|e| e.to_string())?; + let rows = stream + .into_first_result() + .await + .map_err(|e| e.to_string())?; + Ok(rows + .iter() + .map(|row| row.get::<&str, _>(0).unwrap_or("").to_string()) + .collect()) } -pub async fn list_tables(client: &mut SqlServerClient, schema: &str) -> Result, String> { +pub async fn list_tables( + client: &mut SqlServerClient, + schema: &str, +) -> Result, String> { let sql = format!( "SELECT TABLE_NAME, TABLE_TYPE FROM INFORMATION_SCHEMA.TABLES WHERE TABLE_SCHEMA = '{}' ORDER BY TABLE_NAME", schema.replace('\'', "''") ); let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?; - let rows = stream.into_first_result().await.map_err(|e| e.to_string())?; - Ok(rows.iter().map(|row| { - TableInfo { + let rows = stream + .into_first_result() + .await + .map_err(|e| e.to_string())?; + Ok(rows + .iter() + .map(|row| TableInfo { name: row.get::<&str, _>(0).unwrap_or("").to_string(), table_type: row.get::<&str, _>(1).unwrap_or("BASE TABLE").to_string(), - } - }).collect()) + }) + .collect()) } -pub async fn get_columns(client: &mut SqlServerClient, schema: &str, table: &str) -> Result, String> { +pub async fn get_columns( + client: &mut SqlServerClient, + schema: &str, + table: &str, +) -> Result, String> { let sql = format!( "SELECT c.COLUMN_NAME, c.DATA_TYPE, c.IS_NULLABLE, c.COLUMN_DEFAULT, \ CASE WHEN kcu.COLUMN_NAME IS NOT NULL THEN 1 ELSE 0 END AS IS_PK, \ @@ -107,58 +158,69 @@ pub async fn get_columns(client: &mut SqlServerClient, schema: &str, table: &str s = schema.replace('\'', "''"), t = table.replace('\'', "''") ); let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?; - let rows = stream.into_first_result().await.map_err(|e| e.to_string())?; - Ok(rows.iter().map(|row| { - let base = row.get::<&str, _>(1).unwrap_or("").to_string(); - let max_len = row.get::(7); - let dt_prec = row.get::(8); - let num_prec = row.get::(5); - let num_scale = row.get::(6); - let data_type = match base.to_lowercase().as_str() { - "varchar" => match max_len { - Some(-1) => "varchar(max)".to_string(), - Some(n) => format!("varchar({n})"), - None => "varchar".to_string(), - }, - "nvarchar" => match max_len { - Some(-1) => "nvarchar(max)".to_string(), - Some(n) => format!("nvarchar({n})"), - None => "nvarchar".to_string(), - }, - "varbinary" => match max_len { - Some(-1) => "varbinary(max)".to_string(), - Some(n) if n > 0 => format!("varbinary({n})"), - _ => "varbinary".to_string(), - }, - "char" | "nchar" | "binary" => match max_len { - Some(n) if n > 0 => format!("{base}({n})"), + let rows = stream + .into_first_result() + .await + .map_err(|e| e.to_string())?; + Ok(rows + .iter() + .map(|row| { + let base = row.get::<&str, _>(1).unwrap_or("").to_string(); + let max_len = row.get::(7); + let dt_prec = row.get::(8); + let num_prec = row.get::(5); + let num_scale = row.get::(6); + let data_type = match base.to_lowercase().as_str() { + "varchar" => match max_len { + Some(-1) => "varchar(max)".to_string(), + Some(n) => format!("varchar({n})"), + None => "varchar".to_string(), + }, + "nvarchar" => match max_len { + Some(-1) => "nvarchar(max)".to_string(), + Some(n) => format!("nvarchar({n})"), + None => "nvarchar".to_string(), + }, + "varbinary" => match max_len { + Some(-1) => "varbinary(max)".to_string(), + Some(n) if n > 0 => format!("varbinary({n})"), + _ => "varbinary".to_string(), + }, + "char" | "nchar" | "binary" => match max_len { + Some(n) if n > 0 => format!("{base}({n})"), + _ => base, + }, + "decimal" | "numeric" => match (num_prec, num_scale) { + (Some(p), Some(s)) => format!("{base}({p},{s})"), + _ => base, + }, + "datetime2" | "datetimeoffset" | "time" => match dt_prec { + Some(p) => format!("{base}({p})"), + _ => base, + }, _ => base, + }; + ColumnInfo { + name: row.get::<&str, _>(0).unwrap_or("").to_string(), + data_type, + is_nullable: row.get::<&str, _>(2).unwrap_or("NO") == "YES", + column_default: row.get::<&str, _>(3).map(|s| s.to_string()), + is_primary_key: row.get::(4).unwrap_or(0) == 1, + extra: None, + comment: None, + numeric_precision: num_prec, + numeric_scale: num_scale, + character_maximum_length: max_len, } - "decimal" | "numeric" => match (num_prec, num_scale) { - (Some(p), Some(s)) => format!("{base}({p},{s})"), - _ => base, - }, - "datetime2" | "datetimeoffset" | "time" => match dt_prec { - Some(p) => format!("{base}({p})"), - _ => base, - }, - _ => base, - }; - ColumnInfo { - name: row.get::<&str, _>(0).unwrap_or("").to_string(), - data_type, - is_nullable: row.get::<&str, _>(2).unwrap_or("NO") == "YES", - column_default: row.get::<&str, _>(3).map(|s| s.to_string()), - is_primary_key: row.get::(4).unwrap_or(0) == 1, - extra: None, comment: None, - numeric_precision: num_prec, - numeric_scale: num_scale, - character_maximum_length: max_len, - } - }).collect()) + }) + .collect()) } -pub async fn list_indexes(client: &mut SqlServerClient, schema: &str, table: &str) -> Result, String> { +pub async fn list_indexes( + client: &mut SqlServerClient, + schema: &str, + table: &str, +) -> Result, String> { let sql = format!( "SELECT i.name, \ STRING_AGG(CASE WHEN ic.is_included_column = 0 THEN c.name END, ',') WITHIN GROUP (ORDER BY ic.key_ordinal) AS columns, \ @@ -174,24 +236,42 @@ pub async fn list_indexes(client: &mut SqlServerClient, schema: &str, table: &st s = schema.replace('\'', "''"), t = table.replace('\'', "''") ); let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?; - let rows = stream.into_first_result().await.map_err(|e| e.to_string())?; - Ok(rows.iter().map(|row| { - let cols_str = row.get::<&str, _>(1).unwrap_or(""); - let inc_str = row.get::<&str, _>(5).unwrap_or(""); - IndexInfo { - name: row.get::<&str, _>(0).unwrap_or("").to_string(), - columns: cols_str.split(',').filter(|s| !s.is_empty()).map(|s| s.to_string()).collect(), - is_unique: row.get::(2).unwrap_or(false), - is_primary: row.get::(3).unwrap_or(false), - filter: row.get::<&str, _>(6).map(|s| s.to_string()), - index_type: row.get::<&str, _>(4).map(|s| s.to_string()), - included_columns: if inc_str.is_empty() { None } else { Some(inc_str.split(',').map(|s| s.to_string()).collect()) }, - comment: None, - } - }).collect()) + let rows = stream + .into_first_result() + .await + .map_err(|e| e.to_string())?; + Ok(rows + .iter() + .map(|row| { + let cols_str = row.get::<&str, _>(1).unwrap_or(""); + let inc_str = row.get::<&str, _>(5).unwrap_or(""); + IndexInfo { + name: row.get::<&str, _>(0).unwrap_or("").to_string(), + columns: cols_str + .split(',') + .filter(|s| !s.is_empty()) + .map(|s| s.to_string()) + .collect(), + is_unique: row.get::(2).unwrap_or(false), + is_primary: row.get::(3).unwrap_or(false), + filter: row.get::<&str, _>(6).map(|s| s.to_string()), + index_type: row.get::<&str, _>(4).map(|s| s.to_string()), + included_columns: if inc_str.is_empty() { + None + } else { + Some(inc_str.split(',').map(|s| s.to_string()).collect()) + }, + comment: None, + } + }) + .collect()) } -pub async fn list_foreign_keys(client: &mut SqlServerClient, schema: &str, table: &str) -> Result, String> { +pub async fn list_foreign_keys( + client: &mut SqlServerClient, + schema: &str, + table: &str, +) -> Result, String> { let sql = format!( "SELECT fk.name, c.name, rt.name, rc.name \ FROM sys.foreign_keys fk \ @@ -204,18 +284,26 @@ pub async fn list_foreign_keys(client: &mut SqlServerClient, schema: &str, table s = schema.replace('\'', "''"), t = table.replace('\'', "''") ); let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?; - let rows = stream.into_first_result().await.map_err(|e| e.to_string())?; - Ok(rows.iter().map(|row| { - ForeignKeyInfo { + let rows = stream + .into_first_result() + .await + .map_err(|e| e.to_string())?; + Ok(rows + .iter() + .map(|row| ForeignKeyInfo { name: row.get::<&str, _>(0).unwrap_or("").to_string(), column: row.get::<&str, _>(1).unwrap_or("").to_string(), ref_table: row.get::<&str, _>(2).unwrap_or("").to_string(), ref_column: row.get::<&str, _>(3).unwrap_or("").to_string(), - } - }).collect()) + }) + .collect()) } -pub async fn list_triggers(client: &mut SqlServerClient, schema: &str, table: &str) -> Result, String> { +pub async fn list_triggers( + client: &mut SqlServerClient, + schema: &str, + table: &str, +) -> Result, String> { let sql = format!( "SELECT t.name, te.type_desc, CASE WHEN t.is_instead_of_trigger = 1 THEN 'INSTEAD OF' ELSE 'AFTER' END \ FROM sys.triggers t \ @@ -225,14 +313,18 @@ pub async fn list_triggers(client: &mut SqlServerClient, schema: &str, table: &s s = schema.replace('\'', "''"), t = table.replace('\'', "''") ); let stream = client.query(&*sql, &[]).await.map_err(|e| e.to_string())?; - let rows = stream.into_first_result().await.map_err(|e| e.to_string())?; - Ok(rows.iter().map(|row| { - TriggerInfo { + let rows = stream + .into_first_result() + .await + .map_err(|e| e.to_string())?; + Ok(rows + .iter() + .map(|row| TriggerInfo { name: row.get::<&str, _>(0).unwrap_or("").to_string(), event: row.get::<&str, _>(1).unwrap_or("").to_string(), timing: row.get::<&str, _>(2).unwrap_or("AFTER").to_string(), - } - }).collect()) + }) + .collect()) } pub async fn execute_query(client: &mut SqlServerClient, sql: &str) -> Result { @@ -245,12 +337,23 @@ pub async fn execute_query(client: &mut SqlServerClient, sql: &str) -> Result>()) + let columns_meta = stream + .columns() + .await + .map_err(|e| e.to_string())? + .map(|cols| { + cols.iter() + .map(|c| c.name().to_string()) + .collect::>() + }) .unwrap_or_default(); - let rows = stream.into_first_result().await.map_err(|e| e.to_string())?; - let result_rows: Vec> = rows.iter().map(|row| row_to_json(row)).collect(); + let rows = stream + .into_first_result() + .await + .map_err(|e| e.to_string())?; + let result_rows: Vec> = + rows.iter().map(|row| row_to_json(row)).collect(); Ok(QueryResult { columns: columns_meta, diff --git a/crates/dbx-core/src/db/ssh_tunnel.rs b/crates/dbx-core/src/db/ssh_tunnel.rs index 2bd9f55fd..b5b9897f7 100644 --- a/crates/dbx-core/src/db/ssh_tunnel.rs +++ b/crates/dbx-core/src/db/ssh_tunnel.rs @@ -9,6 +9,8 @@ use tokio::net::TcpListener; use tokio::sync::Mutex; use tokio::task::JoinHandle; +use super::{connection_timeout, CONNECTION_TIMEOUT_SECS}; + struct SshClient; impl client::Handler for SshClient { @@ -35,16 +37,25 @@ async fn connect_and_authenticate( ..Default::default() }); - let mut session = client::connect(config, (ssh_host, ssh_port), SshClient {}) - .await - .map_err(|e| format!("SSH connection failed: {e}"))?; + let mut session = tokio::time::timeout( + connection_timeout(), + client::connect(config, (ssh_host, ssh_port), SshClient {}), + ) + .await + .map_err(|_| format!("SSH connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))? + .map_err(|e| format!("SSH connection failed: {e}"))?; if !ssh_key_path.is_empty() { - let passphrase = if ssh_key_passphrase.is_empty() { None } else { Some(ssh_key_passphrase) }; + let passphrase = if ssh_key_passphrase.is_empty() { + None + } else { + Some(ssh_key_passphrase) + }; let key_pair = load_secret_key(ssh_key_path, passphrase) .map_err(|e| format!("Failed to load SSH key: {e}"))?; - let auth_res = session - .authenticate_publickey( + let auth_res = tokio::time::timeout( + connection_timeout(), + session.authenticate_publickey( ssh_user, PrivateKeyWithHashAlg::new( Arc::new(key_pair), @@ -55,17 +66,22 @@ async fn connect_and_authenticate( .flatten() .flatten(), ), - ) - .await - .map_err(|e| format!("SSH key auth failed: {e}"))?; + ), + ) + .await + .map_err(|_| format!("SSH key auth timed out ({CONNECTION_TIMEOUT_SECS}s)"))? + .map_err(|e| format!("SSH key auth failed: {e}"))?; if !auth_res.success() { return Err("SSH public key authentication failed".to_string()); } } else if !ssh_password.is_empty() { - let auth_res = session - .authenticate_password(ssh_user, ssh_password) - .await - .map_err(|e| format!("SSH password auth failed: {e}"))?; + let auth_res = tokio::time::timeout( + connection_timeout(), + session.authenticate_password(ssh_user, ssh_password), + ) + .await + .map_err(|_| format!("SSH password auth timed out ({CONNECTION_TIMEOUT_SECS}s)"))? + .map_err(|e| format!("SSH password auth failed: {e}"))?; if !auth_res.success() { return Err("SSH password authentication failed".to_string()); } @@ -167,11 +183,21 @@ impl TunnelManager { ) -> Result { let local_port = portpicker::pick_unused_port().ok_or("No available port")?; - let session = - connect_and_authenticate(ssh_host, ssh_port, ssh_user, ssh_password, ssh_key_path, ssh_key_passphrase) - .await?; + let session = connect_and_authenticate( + ssh_host, + ssh_port, + ssh_user, + ssh_password, + ssh_key_path, + ssh_key_passphrase, + ) + .await?; - let bind_addr = if expose_to_lan { "0.0.0.0" } else { "127.0.0.1" }; + let bind_addr = if expose_to_lan { + "0.0.0.0" + } else { + "127.0.0.1" + }; let listener = TcpListener::bind((bind_addr, local_port)) .await .map_err(|e| format!("Failed to bind local port: {e}"))?; diff --git a/src-tauri/src/commands/connection.rs b/src-tauri/src/commands/connection.rs index 555aa1e79..ee6e997ff 100644 --- a/src-tauri/src/commands/connection.rs +++ b/src-tauri/src/commands/connection.rs @@ -2,8 +2,8 @@ use std::sync::Arc; use tauri::State; pub use dbx_core::connection::{ - connection_url_for_endpoint, expand_tilde, redacted_connection_url_for_endpoint, AppState, - PoolKind, + connection_url_for_endpoint, expand_tilde, probe_connection_endpoint, + redacted_connection_url_for_endpoint, AppState, PoolKind, }; use dbx_core::db; use dbx_core::models::connection::{ConnectionConfig, DatabaseType}; @@ -50,6 +50,7 @@ pub async fn test_connection( config.id.as_str() }; let (host, port) = state.connection_host_port(connection_id, &config).await?; + let probe_result = probe_connection_endpoint(&config, &host, port).await; let url = connection_url_for_endpoint(&config, &host, port); let target = redacted_connection_url_for_endpoint(&config, &host, port); log::info!( @@ -57,102 +58,109 @@ pub async fn test_connection( config.db_type, target ); - let result = match config.db_type { - DatabaseType::Mysql if config.needs_bare_mysql() => { - match db::mysql::connect_bare(&url).await { + let result = match probe_result { + Err(e) => Err(e), + Ok(()) => match config.db_type { + DatabaseType::Mysql if config.needs_bare_mysql() => { + match db::mysql::connect_bare(&url).await { + Ok(pool) => { + pool.close().await; + Ok("Connection successful".to_string()) + } + Err(e) => Err(e), + } + } + DatabaseType::Mysql => match db::mysql::connect(&url).await { Ok(pool) => { pool.close().await; Ok("Connection successful".to_string()) } Err(e) => Err(e), - } - }, - DatabaseType::Mysql => { - match db::mysql::connect(&url).await { - Ok(pool) => { - pool.close().await; - Ok("Connection successful".to_string()) + }, + DatabaseType::Doris | DatabaseType::StarRocks => { + match db::mysql::connect_bare(&url).await { + Ok(pool) => { + pool.close().await; + Ok("Connection successful".to_string()) + } + Err(e) => Err(e), } - Err(e) => Err(e), } - }, - DatabaseType::Doris | DatabaseType::StarRocks => { - match db::mysql::connect_bare(&url).await { - Ok(pool) => { - pool.close().await; - Ok("Connection successful".to_string()) + DatabaseType::Postgres | DatabaseType::Redshift => { + match db::postgres::connect(&url).await { + Ok(pool) => { + pool.close().await; + Ok("Connection successful".to_string()) + } + Err(e) => Err(e), } - Err(e) => Err(e), } - }, - DatabaseType::Postgres | DatabaseType::Redshift => match db::postgres::connect(&url).await { - Ok(pool) => { - pool.close().await; - Ok("Connection successful".to_string()) + DatabaseType::Sqlite => { + match db::sqlite::connect_path(&expand_tilde(&config.host)).await { + Ok(pool) => { + pool.close().await; + Ok("Connection successful".to_string()) + } + Err(e) => Err(e), + } } - Err(e) => Err(e), - }, - DatabaseType::Sqlite => match db::sqlite::connect_path(&expand_tilde(&config.host)).await { - Ok(pool) => { - pool.close().await; - Ok("Connection successful".to_string()) - } - Err(e) => Err(e), - }, - DatabaseType::Redis => { - db::redis_driver::connect(&url) - .await - .map(|_| "Connection successful".to_string()) - } - DatabaseType::DuckDb => { - duckdb::Connection::open(&expand_tilde(&config.host)) - .map(|_| "Connection successful".to_string()) - .map_err(|e| e.to_string()) - } - DatabaseType::MongoDb => match mongodb::Client::with_uri_str(&url).await { - Ok(client) => client - .list_database_names() + DatabaseType::Redis => db::redis_driver::connect(&url) .await + .map(|_| "Connection successful".to_string()), + DatabaseType::DuckDb => duckdb::Connection::open(&expand_tilde(&config.host)) .map(|_| "Connection successful".to_string()) .map_err(|e| e.to_string()), - Err(e) => Err(e.to_string()), + DatabaseType::MongoDb => match db::mongo_driver::connect(&url).await { + Ok(client) => db::mongo_driver::test_connection(&client) + .await + .map(|_| "Connection successful".to_string()), + Err(e) => Err(e.to_string()), + }, + DatabaseType::ClickHouse => { + let username = if config.username.is_empty() { + None + } else { + Some(config.username.clone()) + }; + let password = if config.password.is_empty() { + None + } else { + Some(config.password.clone()) + }; + let client = db::clickhouse_driver::ChClient::new(&url, username, password); + db::clickhouse_driver::test_connection(&client) + .await + .map(|_| "Connection successful".to_string()) + } + DatabaseType::SqlServer => db::sqlserver::connect( + &host, + port, + &config.username, + &config.password, + config.database.as_deref(), + ) + .await + .map(|_| "Connection successful".to_string()), + DatabaseType::Oracle => db::oracle_driver::connect( + &host, + port, + config.database.as_deref().unwrap_or("ORCL"), + &config.username, + &config.password, + ) + .await + .map(|_| "Connection successful".to_string()), + DatabaseType::Elasticsearch => { + let client = db::elasticsearch_driver::EsClient::new( + &url, + Some(&config.username), + Some(&config.password), + ); + db::elasticsearch_driver::test_connection(&client) + .await + .map(|_| "Connection successful".to_string()) + } }, - DatabaseType::ClickHouse => { - let username = if config.username.is_empty() { None } else { Some(config.username.clone()) }; - let password = if config.password.is_empty() { None } else { Some(config.password.clone()) }; - let client = db::clickhouse_driver::ChClient::new(&url, username, password); - db::clickhouse_driver::test_connection(&client) - .await - .map(|_| "Connection successful".to_string()) - } - DatabaseType::SqlServer => db::sqlserver::connect( - &host, - port, - &config.username, - &config.password, - config.database.as_deref(), - ) - .await - .map(|_| "Connection successful".to_string()), - DatabaseType::Oracle => db::oracle_driver::connect( - &host, - port, - config.database.as_deref().unwrap_or("ORCL"), - &config.username, - &config.password, - ) - .await - .map(|_| "Connection successful".to_string()), - DatabaseType::Elasticsearch => { - let client = db::elasticsearch_driver::EsClient::new( - &url, - Some(&config.username), - Some(&config.password), - ); - db::elasticsearch_driver::test_connection(&client) - .await - .map(|_| "Connection successful".to_string()) - } }; if config.ssh_enabled && !config.ssh_host.is_empty() { @@ -170,29 +178,48 @@ pub async fn connect_db( let id = config.id.clone(); let (host, port) = state.connection_host_port(&id, &config).await?; + probe_connection_endpoint(&config, &host, port).await?; let url = connection_url_for_endpoint(&config, &host, port); let pool = match config.db_type { - DatabaseType::Mysql if config.needs_bare_mysql() => PoolKind::Mysql(db::mysql::connect_bare(&url).await?, true), + DatabaseType::Mysql if config.needs_bare_mysql() => { + PoolKind::Mysql(db::mysql::connect_bare(&url).await?, true) + } DatabaseType::Mysql => PoolKind::Mysql(db::mysql::connect(&url).await?, false), - DatabaseType::Doris | DatabaseType::StarRocks => PoolKind::Mysql(db::mysql::connect_bare(&url).await?, true), - DatabaseType::Postgres | DatabaseType::Redshift => PoolKind::Postgres(db::postgres::connect(&url).await?), - DatabaseType::Sqlite => PoolKind::Sqlite(db::sqlite::connect_path(&expand_tilde(&config.host)).await?), + DatabaseType::Doris | DatabaseType::StarRocks => { + PoolKind::Mysql(db::mysql::connect_bare(&url).await?, true) + } + DatabaseType::Postgres | DatabaseType::Redshift => { + PoolKind::Postgres(db::postgres::connect(&url).await?) + } + DatabaseType::Sqlite => { + PoolKind::Sqlite(db::sqlite::connect_path(&expand_tilde(&config.host)).await?) + } DatabaseType::Redis => { let con = db::redis_driver::connect(&url).await?; PoolKind::Redis(tokio::sync::Mutex::new(con)) } DatabaseType::DuckDb => { - let con = duckdb::Connection::open(&expand_tilde(&config.host)).map_err(|e| e.to_string())?; + let con = + duckdb::Connection::open(&expand_tilde(&config.host)).map_err(|e| e.to_string())?; PoolKind::DuckDb(std::sync::Arc::new(std::sync::Mutex::new(con))) } DatabaseType::MongoDb => { - let client = mongodb::Client::with_uri_str(&url).await.map_err(|e| e.to_string())?; + let client = db::mongo_driver::connect(&url).await?; + db::mongo_driver::test_connection(&client).await?; PoolKind::MongoDb(client) } DatabaseType::ClickHouse => { - let username = if config.username.is_empty() { None } else { Some(config.username.clone()) }; - let password = if config.password.is_empty() { None } else { Some(config.password.clone()) }; + let username = if config.username.is_empty() { + None + } else { + Some(config.username.clone()) + }; + let password = if config.password.is_empty() { + None + } else { + Some(config.password.clone()) + }; let client = db::clickhouse_driver::ChClient::new(&url, username, password); db::clickhouse_driver::test_connection(&client).await?; PoolKind::ClickHouse(client) @@ -201,7 +228,8 @@ pub async fn connect_db( let client = db::sqlserver::connect( &host, port, - &config.username, &config.password, + &config.username, + &config.password, config.database.as_deref(), ) .await?; @@ -252,13 +280,13 @@ pub async fn disconnect_db( PoolKind::Mysql(p, _) => p.close().await, PoolKind::Postgres(p) => p.close().await, PoolKind::Sqlite(p) => p.close().await, - PoolKind::Redis(_) => {}, - PoolKind::DuckDb(_) => {}, - PoolKind::MongoDb(_) => {}, - PoolKind::ClickHouse(_) => {}, - PoolKind::SqlServer(_) => {}, - PoolKind::Oracle(_) => {}, - PoolKind::Elasticsearch(_) => {}, + PoolKind::Redis(_) => {} + PoolKind::DuckDb(_) => {} + PoolKind::MongoDb(_) => {} + PoolKind::ClickHouse(_) => {} + PoolKind::SqlServer(_) => {} + PoolKind::Oracle(_) => {} + PoolKind::Elasticsearch(_) => {} } } } diff --git a/src-web/src/routes/connection.rs b/src-web/src/routes/connection.rs index 86ae4f691..4a57a50d7 100644 --- a/src-web/src/routes/connection.rs +++ b/src-web/src/routes/connection.rs @@ -41,10 +41,22 @@ pub async fn test_connection( .insert(temp_id.clone(), config.clone()); // Try to connect - let result = app.get_or_create_pool(&temp_id, config.database.as_deref()).await; + let result = app + .get_or_create_pool(&temp_id, config.database.as_deref()) + .await; - // Clean up - app.connections.lock().await.remove(&temp_id); + // Clean up any pool keys created for the temporary connection, including + // database-scoped keys like "__test_uuid:database". + let mut connections = app.connections.lock().await; + let temp_keys: Vec = connections + .keys() + .filter(|key| key.starts_with(&temp_id)) + .cloned() + .collect(); + for key in temp_keys { + connections.remove(&key); + } + drop(connections); app.configs.lock().await.remove(&temp_id); match result {