diff --git a/src-tauri/src/commands/connection.rs b/src-tauri/src/commands/connection.rs index 8e31e1c38..1b61d2276 100644 --- a/src-tauri/src/commands/connection.rs +++ b/src-tauri/src/commands/connection.rs @@ -79,7 +79,8 @@ impl AppState { } } - let url = db_config.connection_url(); + let (host, port) = self.connection_host_port(connection_id, &db_config).await?; + let url = connection_url_for_endpoint(&db_config, &host, port); let pool = match db_config.db_type { DatabaseType::Mysql => PoolKind::Mysql(db::mysql::connect(&url).await?), DatabaseType::Postgres => PoolKind::Postgres(db::postgres::connect(&url).await?), @@ -103,18 +104,23 @@ impl AppState { } DatabaseType::SqlServer => { let client = db::sqlserver::connect( - &db_config.host, db_config.port, - &db_config.username, &db_config.password, + &host, + port, + &db_config.username, + &db_config.password, db_config.database.as_deref(), - ).await?; + ) + .await?; PoolKind::SqlServer(std::sync::Arc::new(tokio::sync::Mutex::new(client))) } DatabaseType::Oracle => { let client = db::oracle_driver::connect( - &db_config.host, db_config.port, + &host, + port, db_config.database.as_deref().unwrap_or("ORCL"), &db_config.username, &db_config.password, - ).await?; + ) + .await?; PoolKind::Oracle(std::sync::Arc::new(tokio::sync::Mutex::new(client))) } }; @@ -123,6 +129,36 @@ impl AppState { Ok(pool_key) } + async fn connection_host_port( + &self, + connection_id: &str, + config: &ConnectionConfig, + ) -> Result<(String, u16), String> { + if !config.ssh_enabled || config.ssh_host.is_empty() { + return Ok((config.host.clone(), config.port)); + } + + if let Some(local_port) = self.tunnels.local_port(connection_id).await { + return Ok(("127.0.0.1".to_string(), local_port)); + } + + let local_port = self + .tunnels + .start_tunnel( + connection_id, + &config.ssh_host, + config.ssh_port, + &config.ssh_user, + &config.ssh_password, + &config.ssh_key_path, + &config.host, + config.port, + ) + .await?; + + Ok(("127.0.0.1".to_string(), local_port)) + } + pub async fn reconnect_pool( &self, connection_id: &str, @@ -153,6 +189,26 @@ fn connections_file(app: &AppHandle) -> Result { Ok(dir.join("connections.json")) } +fn connection_url_for_endpoint(config: &ConnectionConfig, host: &str, port: u16) -> String { + if host == config.host && port == config.port { + config.connection_url() + } else { + config.connection_url_with_host(host, port) + } +} + +fn redacted_connection_url_for_endpoint( + config: &ConnectionConfig, + host: &str, + port: u16, +) -> String { + if host == config.host && port == config.port { + config.redacted_connection_url() + } else { + config.redacted_connection_url_with_host(host, port) + } +} + #[tauri::command] pub async fn save_connections( app: AppHandle, @@ -177,64 +233,95 @@ pub async fn load_connections(app: AppHandle) -> Result, S } #[tauri::command] -pub async fn test_connection(config: ConnectionConfig) -> Result { - let url = config.connection_url(); +pub async fn test_connection( + state: State<'_, Arc>, + config: ConnectionConfig, +) -> Result { + let tunnel_id = format!("{}:test", config.id); + let connection_id = if config.ssh_enabled && !config.ssh_host.is_empty() { + tunnel_id.as_str() + } else { + config.id.as_str() + }; + let (host, port) = state.connection_host_port(connection_id, &config).await?; + let url = connection_url_for_endpoint(&config, &host, port); + let target = redacted_connection_url_for_endpoint(&config, &host, port); log::info!( "[test_connection] db_type={:?} target={}", config.db_type, - config.redacted_connection_url() + target ); - match config.db_type { - DatabaseType::Mysql => { - let pool = db::mysql::connect(&url).await?; - pool.close().await; - Ok("Connection successful".to_string()) - } - DatabaseType::Postgres => { - let pool = db::postgres::connect(&url).await?; - pool.close().await; - Ok("Connection successful".to_string()) - } - DatabaseType::Sqlite => { - let pool = db::sqlite::connect(&url).await?; - pool.close().await; - Ok("Connection successful".to_string()) - } + let result = match config.db_type { + DatabaseType::Mysql => match db::mysql::connect(&url).await { + Ok(pool) => { + pool.close().await; + Ok("Connection successful".to_string()) + } + Err(e) => Err(e), + }, + DatabaseType::Postgres => match db::postgres::connect(&url).await { + Ok(pool) => { + pool.close().await; + Ok("Connection successful".to_string()) + } + Err(e) => Err(e), + }, + DatabaseType::Sqlite => match db::sqlite::connect(&url).await { + Ok(pool) => { + pool.close().await; + Ok("Connection successful".to_string()) + } + Err(e) => Err(e), + }, DatabaseType::Redis => { - let _con = db::redis_driver::connect(&url).await?; - Ok("Connection successful".to_string()) + db::redis_driver::connect(&url) + .await + .map(|_| "Connection successful".to_string()) } DatabaseType::DuckDb => { - let _con = duckdb::Connection::open(&config.host).map_err(|e| e.to_string())?; - Ok("Connection successful".to_string()) - } - DatabaseType::MongoDb => { - let client = mongodb::Client::with_uri_str(&url).await.map_err(|e| e.to_string())?; - client.list_database_names().await.map_err(|e| e.to_string())?; - Ok("Connection successful".to_string()) + duckdb::Connection::open(&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() + .await + .map(|_| "Connection successful".to_string()) + .map_err(|e| e.to_string()), + Err(e) => Err(e.to_string()), + }, DatabaseType::ClickHouse => { let client = db::clickhouse_driver::ChClient::new(&url); - db::clickhouse_driver::test_connection(&client).await?; - Ok("Connection successful".to_string()) - } - DatabaseType::SqlServer => { - let _client = db::sqlserver::connect( - &config.host, config.port, - &config.username, &config.password, - config.database.as_deref(), - ).await?; - Ok("Connection successful".to_string()) - } - DatabaseType::Oracle => { - let _client = db::oracle_driver::connect( - &config.host, config.port, - config.database.as_deref().unwrap_or("ORCL"), - &config.username, &config.password, - ).await?; - Ok("Connection successful".to_string()) + 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()), + }; + + if config.ssh_enabled && !config.ssh_host.is_empty() { + state.tunnels.stop_tunnel(&tunnel_id).await; } + + result } #[tauri::command] @@ -244,16 +331,8 @@ pub async fn connect_db( ) -> Result { let id = config.id.clone(); - let url = if config.ssh_enabled && !config.ssh_host.is_empty() { - let local_port = state.tunnels.start_tunnel( - &id, &config.ssh_host, config.ssh_port, - &config.ssh_user, &config.ssh_password, &config.ssh_key_path, - &config.host, config.port, - ).await?; - config.connection_url_with_host("127.0.0.1", local_port) - } else { - config.connection_url() - }; + let (host, port) = state.connection_host_port(&id, &config).await?; + let url = connection_url_for_endpoint(&config, &host, port); let pool = match config.db_type { DatabaseType::Mysql => PoolKind::Mysql(db::mysql::connect(&url).await?), @@ -278,17 +357,23 @@ pub async fn connect_db( } DatabaseType::SqlServer => { let client = db::sqlserver::connect( - &config.host, config.port, + &host, + port, &config.username, &config.password, config.database.as_deref(), - ).await?; - PoolKind::SqlServer(std::sync::Arc::new(tokio::sync::Mutex::new(client))) } + ) + .await?; + PoolKind::SqlServer(std::sync::Arc::new(tokio::sync::Mutex::new(client))) + } DatabaseType::Oracle => { let client = db::oracle_driver::connect( - &config.host, config.port, + &host, + port, config.database.as_deref().unwrap_or("ORCL"), - &config.username, &config.password, - ).await?; + &config.username, + &config.password, + ) + .await?; PoolKind::Oracle(std::sync::Arc::new(tokio::sync::Mutex::new(client))) } }; diff --git a/src-tauri/src/commands/mod.rs b/src-tauri/src/commands/mod.rs index 82aa9206d..4c62e8d2b 100644 --- a/src-tauri/src/commands/mod.rs +++ b/src-tauri/src/commands/mod.rs @@ -5,3 +5,4 @@ pub mod mongo_cmd; pub mod query; pub mod redis_cmd; pub mod schema; +pub mod update; diff --git a/src-tauri/src/commands/update.rs b/src-tauri/src/commands/update.rs new file mode 100644 index 000000000..fdf9ef7c2 --- /dev/null +++ b/src-tauri/src/commands/update.rs @@ -0,0 +1,98 @@ +use serde::{Deserialize, Serialize}; + +const LATEST_RELEASE_URL: &str = "https://api.github.com/repos/t8y2/dbx/releases/latest"; + +#[derive(Debug, Deserialize)] +struct GithubRelease { + tag_name: String, + name: Option, + html_url: String, + body: Option, +} + +#[derive(Debug, Serialize)] +pub struct UpdateInfo { + pub current_version: String, + pub latest_version: String, + pub update_available: bool, + pub release_name: String, + pub release_url: String, + pub release_notes: String, +} + +#[tauri::command] +pub async fn check_for_updates() -> Result { + let client = reqwest::Client::new(); + let release = client + .get(LATEST_RELEASE_URL) + .header(reqwest::header::USER_AGENT, "dbx-update-checker") + .send() + .await + .map_err(|e| format!("Failed to check updates: {e}"))? + .error_for_status() + .map_err(|e| format!("Failed to check updates: {e}"))? + .json::() + .await + .map_err(|e| format!("Failed to parse update response: {e}"))?; + + let current_version = env!("CARGO_PKG_VERSION").to_string(); + let latest_version = normalize_version(&release.tag_name); + + Ok(UpdateInfo { + update_available: is_newer_version(&latest_version, ¤t_version), + current_version, + latest_version, + release_name: release.name.unwrap_or_else(|| release.tag_name.clone()), + release_url: release.html_url, + release_notes: release.body.unwrap_or_default(), + }) +} + +fn normalize_version(version: &str) -> String { + version.trim().trim_start_matches('v').to_string() +} + +fn parse_version(version: &str) -> Vec { + normalize_version(version) + .split(['.', '-', '+']) + .map(|part| part.parse::().unwrap_or(0)) + .collect() +} + +fn is_newer_version(latest: &str, current: &str) -> bool { + let latest_parts = parse_version(latest); + let current_parts = parse_version(current); + let max_len = latest_parts.len().max(current_parts.len()); + + for i in 0..max_len { + let latest_part = *latest_parts.get(i).unwrap_or(&0); + let current_part = *current_parts.get(i).unwrap_or(&0); + if latest_part > current_part { + return true; + } + if latest_part < current_part { + return false; + } + } + + false +} + +#[cfg(test)] +mod tests { + use super::{is_newer_version, normalize_version}; + + #[test] + fn normalizes_tag_versions() { + assert_eq!(normalize_version("v1.2.3"), "1.2.3"); + assert_eq!(normalize_version(" 0.2.0 "), "0.2.0"); + } + + #[test] + fn compares_semver_like_versions() { + assert!(is_newer_version("0.2.1", "0.2.0")); + assert!(is_newer_version("1.0.0", "0.9.9")); + assert!(!is_newer_version("0.2.0", "0.2.0")); + assert!(!is_newer_version("0.1.9", "0.2.0")); + } +} diff --git a/src-tauri/src/db/ssh_tunnel.rs b/src-tauri/src/db/ssh_tunnel.rs index c4537d847..6509a0bce 100644 --- a/src-tauri/src/db/ssh_tunnel.rs +++ b/src-tauri/src/db/ssh_tunnel.rs @@ -79,6 +79,14 @@ impl TunnelManager { Ok(local_port) } + pub async fn local_port(&self, connection_id: &str) -> Option { + self.tunnels + .lock() + .await + .get(connection_id) + .map(|(_, port)| *port) + } + pub async fn stop_tunnel(&self, connection_id: &str) { if let Some((mut child, _)) = self.tunnels.lock().await.remove(connection_id) { let _ = child.kill().await; diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 880637524..ca7298422 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -75,6 +75,7 @@ pub fn run() { commands::history::load_history, commands::history::clear_history, commands::history::delete_history_entry, + commands::update::check_for_updates, ]) .run(tauri::generate_context!()) .expect("error while running tauri application"); diff --git a/src/App.vue b/src/App.vue index a1c0da021..a638c288f 100644 --- a/src/App.vue +++ b/src/App.vue @@ -1,7 +1,7 @@