diff --git a/.husky/pre-commit b/.husky/pre-commit index 291c534c0..870148c97 100644 --- a/.husky/pre-commit +++ b/.husky/pre-commit @@ -1,2 +1,3 @@ npx lint-staged cd src-tauri && cargo fmt && git add -u . +cd ../crates/dbx-core && cargo fmt --all && git add -u . diff --git a/crates/dbx-core/src/db/duckdb_driver.rs b/crates/dbx-core/src/db/duckdb_driver.rs index 244f05592..3b36bd82f 100644 --- a/crates/dbx-core/src/db/duckdb_driver.rs +++ b/crates/dbx-core/src/db/duckdb_driver.rs @@ -1,12 +1,12 @@ +use super::file_validator::validate_file_path; 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 @@ -14,8 +14,7 @@ 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}"))?; + let connection = duckdb::Connection::open(path).map_err(|e| format!("DuckDb connection failed: {e}"))?; Ok(Arc::new(Mutex::new(connection))) } diff --git a/crates/dbx-core/src/db/file_validator.rs b/crates/dbx-core/src/db/file_validator.rs index 85aa68e19..48f4ff073 100644 --- a/crates/dbx-core/src/db/file_validator.rs +++ b/crates/dbx-core/src/db/file_validator.rs @@ -1,18 +1,18 @@ 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 @@ -35,26 +35,17 @@ where // 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 - )); + 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 - )); + 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 - )); + return Err(format!("Database file path is not a valid file: {}", path)); } } diff --git a/crates/dbx-core/src/db/mod.rs b/crates/dbx-core/src/db/mod.rs index 10088faac..a77f56fb1 100644 --- a/crates/dbx-core/src/db/mod.rs +++ b/crates/dbx-core/src/db/mod.rs @@ -1,5 +1,7 @@ pub mod clickhouse_driver; +pub mod duckdb_driver; pub mod elasticsearch_driver; +pub mod file_validator; pub mod mongo_driver; pub mod mysql; pub mod oracle_driver; @@ -8,8 +10,6 @@ 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; diff --git a/crates/dbx-core/src/db/mysql.rs b/crates/dbx-core/src/db/mysql.rs index 2411af113..f5d7da774 100644 --- a/crates/dbx-core/src/db/mysql.rs +++ b/crates/dbx-core/src/db/mysql.rs @@ -23,15 +23,9 @@ 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 { @@ -43,8 +37,7 @@ 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 { @@ -101,20 +94,14 @@ 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); } @@ -140,15 +127,14 @@ 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) @@ -168,13 +154,9 @@ pub async fn connect(url: &str) -> Result { } 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); + 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); super::with_connection_timeout("MySQL", async { MySqlPoolOptions::new() .max_connections(5) @@ -201,10 +183,7 @@ 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() @@ -215,11 +194,7 @@ 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, \ @@ -235,10 +210,7 @@ pub async fn get_columns( 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() @@ -261,12 +233,13 @@ 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(); @@ -297,10 +270,7 @@ 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() @@ -320,10 +290,7 @@ 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() @@ -381,10 +345,7 @@ 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() @@ -406,10 +367,7 @@ 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 f0a779a11..2074597d9 100644 --- a/crates/dbx-core/src/db/postgres.rs +++ b/crates/dbx-core/src/db/postgres.rs @@ -1,12 +1,12 @@ use chrono::{DateTime, NaiveDate, NaiveDateTime, NaiveTime, Utc}; +use percent_encoding::percent_decode_str; 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 crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo}; use super::file_validator::validate_file_path; +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) { @@ -42,10 +42,7 @@ 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") @@ -68,19 +65,11 @@ 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)) @@ -91,7 +80,7 @@ fn pg_value_to_json(row: &PgRow, idx: usize, type_name: &str) -> serde_json::Val pub async fn connect(url: &str) -> Result { // Validate SSL certificate paths if present in the URL validate_postgres_ssl_paths(url)?; - + super::with_connection_timeout("PostgreSQL", async { PgPoolOptions::new() .max_connections(5) @@ -105,7 +94,7 @@ pub async fn connect(url: &str) -> Result { } /// 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 @@ -114,7 +103,7 @@ 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 { @@ -123,7 +112,7 @@ fn validate_postgres_ssl_paths(url: &str) -> Result<(), String> { 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)?; } @@ -132,23 +121,17 @@ fn validate_postgres_ssl_paths(url: &str) -> Result<(), String> { } } } - + 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> { @@ -182,17 +165,10 @@ 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, \ @@ -254,11 +230,7 @@ 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(); @@ -278,13 +250,7 @@ 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( - super::connection_timeout(), - client.get_multiplexed_async_connection(), - ) - .await - .map_err(|_| format!("Redis connection timed out ({}s)", super::CONNECTION_TIMEOUT_SECS))? - .map_err(|e| format!("Redis connection failed: {e}"))?; + let mut con = tokio::time::timeout(super::connection_timeout(), client.get_multiplexed_async_connection()) + .await + .map_err(|_| format!("Redis connection timed out ({}s)", super::CONNECTION_TIMEOUT_SECS))? + .map_err(|e| format!("Redis connection failed: {e}"))?; - tokio::time::timeout( - super::connection_timeout(), - redis::cmd("PING").query_async::(&mut con), - ) - .await - .map_err(|_| format!("Redis ping timed out ({}s)", super::CONNECTION_TIMEOUT_SECS))? - .map_err(|e| format!("Redis authentication failed or command rejected: {e}"))?; + tokio::time::timeout(super::connection_timeout(), redis::cmd("PING").query_async::(&mut con)) + .await + .map_err(|_| format!("Redis ping timed out ({}s)", super::CONNECTION_TIMEOUT_SECS))? + .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()) @@ -86,14 +68,8 @@ 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() { @@ -109,11 +85,7 @@ async fn list_keyspace_databases( } 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( @@ -134,35 +106,18 @@ 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); @@ -180,33 +135,20 @@ pub async fn get_value( 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( @@ -284,23 +226,16 @@ 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()) } @@ -310,67 +245,29 @@ 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)] @@ -385,12 +282,7 @@ 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 5b3ec7ffa..b83672f0e 100644 --- a/crates/dbx-core/src/db/sqlite.rs +++ b/crates/dbx-core/src/db/sqlite.rs @@ -2,16 +2,14 @@ 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 super::file_validator::validate_file_path; +use crate::types::{ColumnInfo, DatabaseInfo, ForeignKeyInfo, IndexInfo, QueryResult, TableInfo, TriggerInfo}; pub async fn connect_path(path: &str) -> Result { // Validate file path using universal validator validate_file_path(path, is_network_path)?; - let mut options = SqliteConnectOptions::new() - .filename(path) - .create_if_missing(false); + let mut options = SqliteConnectOptions::new().filename(path).create_if_missing(false); if is_network_path(path) { options = options.vfs("unix-nolock"); @@ -55,10 +53,8 @@ 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() @@ -68,7 +64,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, @@ -130,26 +127,33 @@ 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,10 +170,7 @@ 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() @@ -179,11 +180,13 @@ 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) }) @@ -199,10 +202,7 @@ pub async fn execute_query(pool: &SqlitePool, sql: &str) -> Result