diff --git a/crates/dbx-core/src/connection.rs b/crates/dbx-core/src/connection.rs index 60648f07d..a7452c148 100644 --- a/crates/dbx-core/src/connection.rs +++ b/crates/dbx-core/src/connection.rs @@ -105,8 +105,8 @@ impl AppState { 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())?; - PoolKind::DuckDb(Arc::new(std::sync::Mutex::new(con))) + let con = db::duckdb_driver::connect_path(&expand_tilde(&db_config.host))?; + PoolKind::DuckDb(con) } DatabaseType::MongoDb => { let client = db::mongo_driver::connect(&url).await?; diff --git a/crates/dbx-core/src/db/duckdb_driver.rs b/crates/dbx-core/src/db/duckdb_driver.rs new file mode 100644 index 000000000..244f05592 --- /dev/null +++ b/crates/dbx-core/src/db/duckdb_driver.rs @@ -0,0 +1,25 @@ +use std::sync::Arc; +use std::sync::Mutex; +use super::file_validator::validate_file_path; + +/// Connects to a DuckDb database file with file validation. +/// +/// # Arguments +/// * `path` - The file path to the DuckDb database +/// +/// # Returns +/// * `Ok(Arc>)` on successful connection +/// * `Err(String)` with descriptive error message if connection fails +pub fn connect_path(path: &str) -> Result>, String> { + // Validate file path using universal validator + validate_file_path(path, is_network_path)?; + + let connection = duckdb::Connection::open(path) + .map_err(|e| format!("DuckDb connection failed: {e}"))?; + + Ok(Arc::new(Mutex::new(connection))) +} + +fn is_network_path(path: &str) -> bool { + path.starts_with("\\\\") || path.starts_with("//") || path.contains("wsl.localhost") || path.contains("wsl$") +} diff --git a/crates/dbx-core/src/db/file_validator.rs b/crates/dbx-core/src/db/file_validator.rs new file mode 100644 index 000000000..85aa68e19 --- /dev/null +++ b/crates/dbx-core/src/db/file_validator.rs @@ -0,0 +1,98 @@ +use std::path::Path; + +/// Validates a file path for database connections. +/// +/// Performs comprehensive checks including: +/// - Empty path validation +/// - Null character detection +/// - File existence (for local paths) +/// - File type validation (must be a file, not directory) +/// - Network path detection (skips validation for network paths) +/// +/// # Arguments +/// * `path` - The file path to validate +/// * `is_network_path` - Closure to determine if path is a network path +/// +/// # Returns +/// * `Ok(())` if validation passes +/// * `Err(String)` with descriptive error message if validation fails +pub fn validate_file_path(path: &str, is_network_path: F) -> Result<(), String> +where + F: Fn(&str) -> bool, +{ + // Check if path is empty + if path.is_empty() { + return Err("Database file path cannot be empty".to_string()); + } + + // Check if path contains invalid characters + if path.contains('\0') { + return Err("Database file path contains null characters".to_string()); + } + + let path_obj = Path::new(path); + + // For non-network paths, perform file system checks + if !is_network_path(path) { + if !path_obj.exists() { + return Err(format!( + "Database file does not exist: {}", + path + )); + } + + // Check if path is actually a file, not a directory + if path_obj.is_dir() { + return Err(format!( + "Database file path is a directory, not a file: {}", + path + )); + } + + // Check if path is a valid file + if !path_obj.is_file() { + return Err(format!( + "Database file path is not a valid file: {}", + path + )); + } + } + + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn is_network_path_test(path: &str) -> bool { + path.starts_with("\\\\") || path.starts_with("//") + } + + #[test] + fn test_empty_path() { + let result = validate_file_path("", is_network_path_test); + assert!(result.is_err()); + assert!(result.unwrap_err().contains("empty")); + } + + #[test] + fn test_null_character() { + let result = validate_file_path("path\0invalid", is_network_path_test); + assert!(result.is_err()); + assert!(result.unwrap_err().contains("null")); + } + + #[test] + fn test_network_path_skips_validation() { + let result = validate_file_path("//network/path/nonexistent.db", is_network_path_test); + assert!(result.is_ok()); + } + + #[test] + fn test_nonexistent_local_file() { + let result = validate_file_path("/nonexistent/path/to/file.db", is_network_path_test); + assert!(result.is_err()); + assert!(result.unwrap_err().contains("does not exist")); + } +} diff --git a/crates/dbx-core/src/db/mod.rs b/crates/dbx-core/src/db/mod.rs index 28cdbf26c..10088faac 100644 --- a/crates/dbx-core/src/db/mod.rs +++ b/crates/dbx-core/src/db/mod.rs @@ -8,12 +8,15 @@ pub mod redis_driver; pub mod sqlite; pub mod sqlserver; pub mod ssh_tunnel; +pub mod file_validator; +pub mod duckdb_driver; 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 use file_validator::validate_file_path; pub const CONNECTION_TIMEOUT_SECS: u64 = 5; pub const TCP_PROBE_TIMEOUT_SECS: u64 = 3; diff --git a/crates/dbx-core/src/db/mysql.rs b/crates/dbx-core/src/db/mysql.rs index 7f78b9ee3..d11860d12 100644 --- a/crates/dbx-core/src/db/mysql.rs +++ b/crates/dbx-core/src/db/mysql.rs @@ -4,7 +4,6 @@ use sqlx::mysql::{MySqlPool, MySqlPoolOptions, MySqlRow}; use sqlx::{Column, Executor, Row, TypeInfo, ValueRef}; use std::time::{Duration, Instant}; -use super::{connection_timeout, with_connection_timeout}; use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo}; fn quote_value(s: &str) -> String { @@ -24,9 +23,15 @@ fn get_str_by_name(row: &MySqlRow, name: &str) -> String { } fn get_opt_str(row: &MySqlRow, name: &str) -> Option { - row.try_get::, _>(name).ok().flatten().or_else(|| { - row.try_get::>, _>(name).ok().flatten().map(|b| String::from_utf8_lossy(&b).to_string()) - }) + row.try_get::, _>(name) + .ok() + .flatten() + .or_else(|| { + row.try_get::>, _>(name) + .ok() + .flatten() + .map(|b| String::from_utf8_lossy(&b).to_string()) + }) } fn numeric_metadata_u64_to_i32(value: Option) -> Option { @@ -38,7 +43,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()).and_then(|v| i32::try_from(v).ok()) + value.and_then(|v| v.parse::().ok()) + .and_then(|v| i32::try_from(v).ok()) } fn get_opt_i32(row: &MySqlRow, name: &str) -> Option { @@ -95,14 +101,20 @@ fn mysql_value_to_json(row: &MySqlRow, idx: usize, type_name: &str) -> serde_jso } if upper_type == "BOOLEAN" { - return row.try_get::(idx).map(serde_json::Value::Bool).unwrap_or(serde_json::Value::Null); + return row + .try_get::(idx) + .map(serde_json::Value::Bool) + .unwrap_or(serde_json::Value::Null); } if upper_type.contains("BIGINT") { return row .try_get::(idx) .map(|v| serde_json::Value::String(v.to_string())) - .or_else(|_| row.try_get::(idx).map(|v| serde_json::Value::String(v.to_string()))) + .or_else(|_| { + row.try_get::(idx) + .map(|v| serde_json::Value::String(v.to_string())) + }) .unwrap_or(serde_json::Value::Null); } @@ -128,46 +140,45 @@ fn mysql_value_to_json(row: &MySqlRow, idx: usize, type_name: &str) -> serde_jso .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::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).map(|b| serde_json::Value::String(String::from_utf8_lossy(&b).to_string())) + row.try_get::, _>(idx) + .map(|b| serde_json::Value::String(String::from_utf8_lossy(&b).to_string())) }) .or_else(|e| mysql_temporal_to_json_value(row, idx).ok_or(e)) .unwrap_or(serde_json::Value::Null) } pub async fn connect(url: &str) -> Result { - 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 + 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}")) } pub async fn connect_bare(url: &str) -> Result { - 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); - 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 + 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}")) } pub async fn list_databases(pool: &MySqlPool) -> Result, String> { @@ -184,7 +195,10 @@ pub async fn list_tables(pool: &MySqlPool, database: &str) -> Result = sqlx::raw_sql(&sql).fetch_all(pool).await.map_err(|e| e.to_string())?; + let rows: Vec = sqlx::raw_sql(&sql) + .fetch_all(pool) + .await + .map_err(|e| e.to_string())?; Ok(rows .iter() @@ -195,7 +209,11 @@ pub async fn list_tables(pool: &MySqlPool, database: &str) -> Result Result, String> { +pub async fn get_columns( + pool: &MySqlPool, + database: &str, + table: &str, +) -> Result, String> { let sql = format!( "SELECT c.COLUMN_NAME, c.COLUMN_TYPE, c.IS_NULLABLE, c.COLUMN_DEFAULT, c.EXTRA, c.COLUMN_COMMENT, \ CASE WHEN kcu.COLUMN_NAME IS NOT NULL THEN 1 ELSE 0 END AS IS_PK, \ @@ -211,7 +229,10 @@ pub async fn get_columns(pool: &MySqlPool, database: &str, table: &str) -> Resul quote_value(database), quote_value(table), ); - let rows: Vec = sqlx::raw_sql(&sql).fetch_all(pool).await.map_err(|e| e.to_string())?; + let rows: Vec = sqlx::raw_sql(&sql) + .fetch_all(pool) + .await + .map_err(|e| e.to_string())?; Ok(rows .iter() @@ -234,13 +255,12 @@ pub async fn execute_query(pool: &MySqlPool, sql: &str, bare: bool) -> Result = sqlx::raw_sql(sql).fetch_all(pool).await.map_err(|e| e.to_string())?; + let rows: Vec = sqlx::raw_sql(sql) + .fetch_all(pool) + .await + .map_err(|e| e.to_string())?; let (columns, column_types) = if let Some(first) = rows.first() { let cols: Vec = first.columns().iter().map(|c| c.name().to_string()).collect(); @@ -271,7 +291,10 @@ pub async fn execute_query(pool: &MySqlPool, sql: &str, bare: bool) -> 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 rows: Vec = sqlx::query(sql).fetch_all(pool).await.map_err(|e| e.to_string())?; + let rows: Vec = sqlx::query(sql) + .fetch_all(pool) + .await + .map_err(|e| e.to_string())?; let result_rows: Vec> = rows .iter() @@ -291,7 +314,10 @@ pub async fn execute_query(pool: &MySqlPool, sql: &str, bare: bool) -> Result Resu quote_value(database), quote_value(table), ); - let rows: Vec = sqlx::raw_sql(&sql).fetch_all(pool).await.map_err(|e| e.to_string())?; + let rows: Vec = sqlx::raw_sql(&sql) + .fetch_all(pool) + .await + .map_err(|e| e.to_string())?; Ok(rows .iter() @@ -346,7 +375,10 @@ pub async fn list_foreign_keys(pool: &MySqlPool, database: &str, table: &str) -> quote_value(database), quote_value(table), ); - let rows: Vec = sqlx::raw_sql(&sql).fetch_all(pool).await.map_err(|e| e.to_string())?; + let rows: Vec = sqlx::raw_sql(&sql) + .fetch_all(pool) + .await + .map_err(|e| e.to_string())?; Ok(rows .iter() @@ -368,7 +400,10 @@ pub async fn list_triggers(pool: &MySqlPool, database: &str, table: &str) -> Res quote_value(database), quote_value(table), ); - let rows: Vec = sqlx::raw_sql(&sql).fetch_all(pool).await.map_err(|e| e.to_string())?; + let rows: Vec = sqlx::raw_sql(&sql) + .fetch_all(pool) + .await + .map_err(|e| e.to_string())?; Ok(rows .iter() diff --git a/crates/dbx-core/src/db/postgres.rs b/crates/dbx-core/src/db/postgres.rs index 60cfce7ad..f0fe60a59 100644 --- a/crates/dbx-core/src/db/postgres.rs +++ b/crates/dbx-core/src/db/postgres.rs @@ -3,9 +3,10 @@ use rust_decimal::Decimal; use sqlx::postgres::{PgPool, PgPoolOptions, PgRow}; use sqlx::{Column, Executor, Row, TypeInfo, ValueRef}; use std::time::{Duration, Instant}; +use percent_encoding::percent_decode_str; -use super::{connection_timeout, with_connection_timeout}; use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo}; +use super::file_validator::validate_file_path; fn pg_temporal_to_json_value(row: &PgRow, idx: usize) -> Option { if let Ok(v) = row.try_get::, _>(idx) { @@ -41,7 +42,10 @@ fn pg_value_to_json(row: &PgRow, idx: usize, type_name: &str) -> serde_json::Val } if upper == "BOOL" { - return row.try_get::(idx).map(serde_json::Value::Bool).unwrap_or(serde_json::Value::Null); + return row + .try_get::(idx) + .map(serde_json::Value::Bool) + .unwrap_or(serde_json::Value::Null); } if upper.contains("TIMESTAMP") @@ -64,11 +68,19 @@ fn pg_value_to_json(row: &PgRow, idx: usize, type_name: &str) -> serde_json::Val 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::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) + 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)) @@ -77,25 +89,63 @@ fn pg_value_to_json(row: &PgRow, idx: usize, type_name: &str) -> serde_json::Val } pub async fn connect(url: &str) -> Result { - 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 + // Validate SSL certificate paths if present in the URL + validate_postgres_ssl_paths(url)?; + + 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}")) +} + +/// Validates SSL certificate file paths in PostgreSQL connection URLs. +/// +/// PostgreSQL connection strings can include SSL parameters like: +/// - sslcert=/path/to/cert.pem +/// - sslkey=/path/to/key.pem +/// - sslrootcert=/path/to/root.pem +fn validate_postgres_ssl_paths(url: &str) -> Result<(), String> { + // Extract query parameters from URL + if let Some(query_start) = url.find('?') { + let query_string = &url[query_start + 1..]; + + for param in query_string.split('&') { + if let Some((key, value)) = param.split_once('=') { + match key { + "sslcert" | "sslkey" | "sslrootcert" => { + // URL decode the value + let decoded = percent_decode_str(value) + .decode_utf8() + .map_err(|_| format!("Invalid URL encoding in {key}"))?; + + // Validate the file path (skip network paths) + validate_file_path(&decoded, |_| false)?; + } + _ => {} + } + } + } + } + + Ok(()) } pub async fn list_databases(pool: &PgPool) -> Result, String> { - let rows: Vec = sqlx::query("SELECT datname FROM pg_database WHERE datistemplate = false ORDER BY datname") - .fetch_all(pool) - .await - .map_err(|e| e.to_string())?; + let rows: Vec = + sqlx::query("SELECT datname FROM pg_database WHERE datistemplate = false ORDER BY datname") + .fetch_all(pool) + .await + .map_err(|e| e.to_string())?; - Ok(rows.iter().map(|row| DatabaseInfo { name: row.get::("datname") }).collect()) + Ok(rows + .iter() + .map(|row| DatabaseInfo { + name: row.get::("datname"), + }) + .collect()) } pub async fn list_tables(pool: &PgPool, schema: &str) -> Result, String> { @@ -129,10 +179,17 @@ pub async fn list_schemas(pool: &PgPool) -> Result, String> { .await .map_err(|e| e.to_string())?; - Ok(rows.iter().map(|row| row.get::("schema_name")).collect()) + Ok(rows + .iter() + .map(|row| row.get::("schema_name")) + .collect()) } -pub async fn get_columns(pool: &PgPool, schema: &str, table: &str) -> Result, String> { +pub async fn get_columns( + pool: &PgPool, + schema: &str, + table: &str, +) -> Result, String> { let rows: Vec = sqlx::query( "SELECT a.attname AS column_name, \ format_type(a.atttypid, a.atttypmod) AS full_type, \ @@ -194,7 +251,11 @@ pub async fn execute_query(pool: &PgPool, sql: &str) -> Result = sqlx::query(sql).persistent(false).fetch_all(pool).await.map_err(|e| e.to_string())?; + let rows: Vec = sqlx::query(sql) + .persistent(false) + .fetch_all(pool) + .await + .map_err(|e| e.to_string())?; let (columns, column_types): (Vec, Vec) = if let Some(first) = rows.first() { let cols = first.columns(); @@ -214,7 +275,13 @@ pub async fn execute_query(pool: &PgPool, sql: &str) -> Result Result Result { let client = redis::Client::open(url).map_err(|e| format!("Redis connection failed: {e}"))?; - let mut con = tokio::time::timeout(connection_timeout(), client.get_multiplexed_async_connection()) - .await - .map_err(|_| format!("Redis connection timed out ({CONNECTION_TIMEOUT_SECS}s)"))? - .map_err(|e| format!("Redis connection failed: {e}"))?; + let mut con = tokio::time::timeout( + std::time::Duration::from_secs(10), + client.get_multiplexed_async_connection(), + ) + .await + .map_err(|_| "Redis connection timed out (10s)".to_string())? + .map_err(|e| format!("Redis connection failed: {e}"))?; - tokio::time::timeout(connection_timeout(), redis::cmd("PING").query_async::(&mut con)) + 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) } -pub async fn list_databases(con: &mut redis::aio::MultiplexedConnection) -> Result, String> { - let configured_count = - redis::cmd("CONFIG").arg("GET").arg("databases").query_async(con).await.ok().and_then(parse_database_count); +pub async fn list_databases( + con: &mut redis::aio::MultiplexedConnection, +) -> Result, String> { + let configured_count = redis::cmd("CONFIG") + .arg("GET") + .arg("databases") + .query_async(con) + .await + .ok() + .and_then(parse_database_count); let keyspace_dbs = list_keyspace_databases(con).await.unwrap_or_default(); let database_count = configured_count.unwrap_or(DEFAULT_REDIS_DATABASES); - let max_db = keyspace_dbs.iter().copied().max().map(|db| db + 1).unwrap_or(0); + let max_db = keyspace_dbs + .iter() + .copied() + .max() + .map(|db| db + 1) + .unwrap_or(0); let visible_count = database_count.max(max_db).max(1); Ok((0..visible_count).collect()) @@ -70,8 +83,14 @@ fn parse_database_count(value: redis::Value) -> Option { }) } -async fn list_keyspace_databases(con: &mut redis::aio::MultiplexedConnection) -> Result, String> { - let info: String = redis::cmd("INFO").arg("keyspace").query_async(con).await.map_err(|e| e.to_string())?; +async fn list_keyspace_databases( + con: &mut redis::aio::MultiplexedConnection, +) -> Result, String> { + let info: String = redis::cmd("INFO") + .arg("keyspace") + .query_async(con) + .await + .map_err(|e| e.to_string())?; let mut dbs = Vec::new(); for line in info.lines() { @@ -87,7 +106,11 @@ async fn list_keyspace_databases(con: &mut redis::aio::MultiplexedConnection) -> } pub async fn select_db(con: &mut redis::aio::MultiplexedConnection, db: u32) -> Result<(), String> { - redis::cmd("SELECT").arg(db).query_async(con).await.map_err(|e| e.to_string()) + redis::cmd("SELECT") + .arg(db) + .query_async(con) + .await + .map_err(|e| e.to_string()) } pub async fn scan_keys_page( @@ -108,18 +131,35 @@ pub async fn scan_keys_page( let mut result = Vec::new(); for key in &keys { - let key_type: String = - redis::cmd("TYPE").arg(key.as_str()).query_async(con).await.unwrap_or_else(|_| "unknown".to_string()); + let key_type: String = redis::cmd("TYPE") + .arg(key.as_str()) + .query_async(con) + .await + .unwrap_or_else(|_| "unknown".to_string()); let ttl: i64 = con.ttl(key.as_str()).await.unwrap_or(-1); - result.push(RedisKeyInfo { key: key.clone(), key_type, ttl }); + result.push(RedisKeyInfo { + key: key.clone(), + key_type, + ttl, + }); } - Ok(RedisScanResult { cursor: next_cursor, keys: result }) + Ok(RedisScanResult { + cursor: next_cursor, + keys: result, + }) } -pub async fn get_value(con: &mut redis::aio::MultiplexedConnection, key: &str) -> Result { - let key_type: String = redis::cmd("TYPE").arg(key).query_async(con).await.map_err(|e| e.to_string())?; +pub async fn get_value( + con: &mut redis::aio::MultiplexedConnection, + key: &str, +) -> Result { + let key_type: String = redis::cmd("TYPE") + .arg(key) + .query_async(con) + .await + .map_err(|e| e.to_string())?; let ttl: i64 = con.ttl(key).await.unwrap_or(-1); @@ -137,20 +177,33 @@ pub async fn get_value(con: &mut redis::aio::MultiplexedConnection, key: &str) - serde_json::json!(v) } "zset" => { - let v: Vec<(String, f64)> = con.zrange_withscores(key, 0, -1).await.map_err(|e| e.to_string())?; - serde_json::json!(v.iter().map(|(m, s)| serde_json::json!({"member": m, "score": s})).collect::>()) + let v: Vec<(String, f64)> = con + .zrange_withscores(key, 0, -1) + .await + .map_err(|e| e.to_string())?; + serde_json::json!(v + .iter() + .map(|(m, s)| serde_json::json!({"member": m, "score": s})) + .collect::>()) } "hash" => { let v: Vec<(String, String)> = con.hgetall(key).await.map_err(|e| e.to_string())?; - let map: serde_json::Map = - v.into_iter().map(|(k, v)| (k, serde_json::Value::String(v))).collect(); + let map: serde_json::Map = v + .into_iter() + .map(|(k, v)| (k, serde_json::Value::String(v))) + .collect(); serde_json::Value::Object(map) } "stream" => get_stream_entries(con, key).await?, _ => serde_json::Value::Null, }; - Ok(RedisValue { key: key.to_string(), key_type, ttl, value }) + Ok(RedisValue { + key: key.to_string(), + key_type, + ttl, + value, + }) } async fn get_stream_entries( @@ -228,16 +281,23 @@ pub async fn set_string( value: &str, ttl: Option, ) -> Result<(), String> { - con.set::<_, _, ()>(key, value).await.map_err(|e| e.to_string())?; + con.set::<_, _, ()>(key, value) + .await + .map_err(|e| e.to_string())?; if let Some(t) = ttl { if t > 0 { - con.expire::<_, ()>(key, t).await.map_err(|e| e.to_string())?; + con.expire::<_, ()>(key, t) + .await + .map_err(|e| e.to_string())?; } } Ok(()) } -pub async fn delete_key(con: &mut redis::aio::MultiplexedConnection, key: &str) -> Result<(), String> { +pub async fn delete_key( + con: &mut redis::aio::MultiplexedConnection, + key: &str, +) -> Result<(), String> { con.del::<_, ()>(key).await.map_err(|e| e.to_string()) } @@ -247,29 +307,67 @@ pub async fn hash_set( field: &str, value: &str, ) -> Result<(), String> { - con.hset::<_, _, _, ()>(key, field, value).await.map_err(|e| e.to_string()) + con.hset::<_, _, _, ()>(key, field, value) + .await + .map_err(|e| e.to_string()) } -pub async fn hash_del(con: &mut redis::aio::MultiplexedConnection, key: &str, field: &str) -> Result<(), String> { - con.hdel::<_, _, ()>(key, field).await.map_err(|e| e.to_string()) +pub async fn hash_del( + con: &mut redis::aio::MultiplexedConnection, + key: &str, + field: &str, +) -> Result<(), String> { + con.hdel::<_, _, ()>(key, field) + .await + .map_err(|e| e.to_string()) } -pub async fn list_push(con: &mut redis::aio::MultiplexedConnection, key: &str, value: &str) -> Result<(), String> { - con.rpush::<_, _, ()>(key, value).await.map_err(|e| e.to_string()) +pub async fn list_push( + con: &mut redis::aio::MultiplexedConnection, + key: &str, + value: &str, +) -> Result<(), String> { + con.rpush::<_, _, ()>(key, value) + .await + .map_err(|e| e.to_string()) } -pub async fn list_remove(con: &mut redis::aio::MultiplexedConnection, key: &str, index: i64) -> Result<(), String> { +pub async fn list_remove( + con: &mut redis::aio::MultiplexedConnection, + key: &str, + index: i64, +) -> Result<(), String> { let placeholder = "__DELETED_PLACEHOLDER__"; - redis::cmd("LSET").arg(key).arg(index).arg(placeholder).query_async::<()>(con).await.map_err(|e| e.to_string())?; - con.lrem::<_, _, ()>(key, 1, placeholder).await.map_err(|e| e.to_string()) + redis::cmd("LSET") + .arg(key) + .arg(index) + .arg(placeholder) + .query_async::<()>(con) + .await + .map_err(|e| e.to_string())?; + con.lrem::<_, _, ()>(key, 1, placeholder) + .await + .map_err(|e| e.to_string()) } -pub async fn set_add(con: &mut redis::aio::MultiplexedConnection, key: &str, member: &str) -> Result<(), String> { - con.sadd::<_, _, ()>(key, member).await.map_err(|e| e.to_string()) +pub async fn set_add( + con: &mut redis::aio::MultiplexedConnection, + key: &str, + member: &str, +) -> Result<(), String> { + con.sadd::<_, _, ()>(key, member) + .await + .map_err(|e| e.to_string()) } -pub async fn set_remove(con: &mut redis::aio::MultiplexedConnection, key: &str, member: &str) -> Result<(), String> { - con.srem::<_, _, ()>(key, member).await.map_err(|e| e.to_string()) +pub async fn set_remove( + con: &mut redis::aio::MultiplexedConnection, + key: &str, + member: &str, +) -> Result<(), String> { + con.srem::<_, _, ()>(key, member) + .await + .map_err(|e| e.to_string()) } #[cfg(test)] @@ -284,7 +382,12 @@ mod tests { fn parses_stream_entries() { let raw = RedisRawValue::Array(vec![RedisRawValue::Array(vec![ bulk("1714470000000-0"), - RedisRawValue::Array(vec![bulk("event"), bulk("login"), bulk("user_id"), bulk("42")]), + RedisRawValue::Array(vec![ + bulk("event"), + bulk("login"), + bulk("user_id"), + bulk("42"), + ]), ])]); let parsed = parse_stream_entries(raw); diff --git a/crates/dbx-core/src/db/sqlite.rs b/crates/dbx-core/src/db/sqlite.rs index e18c13e87..5b3ec7ffa 100644 --- a/crates/dbx-core/src/db/sqlite.rs +++ b/crates/dbx-core/src/db/sqlite.rs @@ -3,9 +3,15 @@ use sqlx::{Column, Executor, Row}; use std::time::{Duration, Instant}; use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo}; +use super::file_validator::validate_file_path; pub async fn connect_path(path: &str) -> Result { - let mut options = SqliteConnectOptions::new().filename(path).create_if_missing(true); + // Validate file path using universal validator + validate_file_path(path, is_network_path)?; + + let mut options = SqliteConnectOptions::new() + .filename(path) + .create_if_missing(false); if is_network_path(path) { options = options.vfs("unix-nolock"); @@ -49,8 +55,10 @@ pub async fn list_tables(pool: &SqlitePool, _schema: &str) -> Result Result, String> { - let rows: Vec = - sqlx::query(&format!("PRAGMA table_info(\"{}\")", table)).fetch_all(pool).await.map_err(|e| e.to_string())?; + let rows: Vec = sqlx::query(&format!("PRAGMA table_info(\"{}\")", table)) + .fetch_all(pool) + .await + .map_err(|e| e.to_string())?; Ok(rows .iter() @@ -60,8 +68,7 @@ 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, @@ -123,33 +130,26 @@ pub async fn list_foreign_keys(pool: &SqlitePool, _schema: &str, table: &str) -> } 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") - .bind(table) - .fetch_all(pool) - .await - .map_err(|e| e.to_string())?; + let rows: Vec = sqlx::query( + "SELECT name, sql FROM sqlite_master WHERE type = 'trigger' AND tbl_name = ? ORDER BY name", + ) + .bind(table) + .fetch_all(pool) + .await + .map_err(|e| e.to_string())?; Ok(rows .iter() .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" - }; - TriggerInfo { name: row.get::("name"), event: event.to_string(), timing: timing.to_string() } + 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(), + timing: timing.to_string(), + } }) .collect()) } @@ -166,7 +166,10 @@ pub async fn execute_query(pool: &SqlitePool, sql: &str) -> Result = desc.columns().iter().map(|c| c.name().to_string()).collect(); - let rows: Vec = sqlx::query(sql).fetch_all(pool).await.map_err(|e| e.to_string())?; + let rows: Vec = sqlx::query(sql) + .fetch_all(pool) + .await + .map_err(|e| e.to_string())?; let result_rows: Vec> = rows .iter() @@ -176,13 +179,11 @@ 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::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) }) @@ -198,7 +199,10 @@ pub async fn execute_query(pool: &SqlitePool, sql: &str) -> Result