Formatted

This commit is contained in:
Agent-43 2026-05-05 13:57:04 +05:30
parent 0ee433763a
commit 80e1cf927e
9 changed files with 143 additions and 339 deletions

View File

@ -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 .

View File

@ -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<Mutex<duckdb::Connection>>)` on successful connection
/// * `Err(String)` with descriptive error message if connection fails
@ -14,8 +14,7 @@ pub fn connect_path(path: &str) -> Result<Arc<Mutex<duckdb::Connection>>, 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)))
}

View File

@ -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));
}
}

View File

@ -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;

View File

@ -23,15 +23,9 @@ fn get_str_by_name(row: &MySqlRow, name: &str) -> String {
}
fn get_opt_str(row: &MySqlRow, name: &str) -> Option<String> {
row.try_get::<Option<String>, _>(name)
.ok()
.flatten()
.or_else(|| {
row.try_get::<Option<Vec<u8>>, _>(name)
.ok()
.flatten()
.map(|b| String::from_utf8_lossy(&b).to_string())
})
row.try_get::<Option<String>, _>(name).ok().flatten().or_else(|| {
row.try_get::<Option<Vec<u8>>, _>(name).ok().flatten().map(|b| String::from_utf8_lossy(&b).to_string())
})
}
fn numeric_metadata_u64_to_i32(value: Option<u64>) -> Option<i32> {
@ -43,8 +37,7 @@ fn numeric_metadata_i64_to_i32(value: Option<i64>) -> Option<i32> {
}
fn numeric_metadata_str_to_i32(value: Option<String>) -> Option<i32> {
value.and_then(|v| v.parse::<i64>().ok())
.and_then(|v| i32::try_from(v).ok())
value.and_then(|v| v.parse::<i64>().ok()).and_then(|v| i32::try_from(v).ok())
}
fn get_opt_i32(row: &MySqlRow, name: &str) -> Option<i32> {
@ -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::<bool, _>(idx)
.map(serde_json::Value::Bool)
.unwrap_or(serde_json::Value::Null);
return row.try_get::<bool, _>(idx).map(serde_json::Value::Bool).unwrap_or(serde_json::Value::Null);
}
if upper_type.contains("BIGINT") {
return row
.try_get::<i64, _>(idx)
.map(|v| serde_json::Value::String(v.to_string()))
.or_else(|_| {
row.try_get::<u64, _>(idx)
.map(|v| serde_json::Value::String(v.to_string()))
})
.or_else(|_| row.try_get::<u64, _>(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::<i64, _>(idx).map(|v| serde_json::Value::Number(v.into())))
.or_else(|_| row.try_get::<u64, _>(idx).map(|v| serde_json::Value::Number(v.into())))
.or_else(|_| row.try_get::<f64, _>(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::<f64, _>(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::<bool, _>(idx).map(serde_json::Value::Bool))
.or_else(|_| {
row.try_get::<Vec<u8>, _>(idx)
.map(|b| serde_json::Value::String(String::from_utf8_lossy(&b).to_string()))
row.try_get::<Vec<u8>, _>(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<MySqlPool, String> {
}
pub async fn connect_bare(url: &str) -> Result<MySqlPool, String> {
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<Vec<TableIn
"SELECT TABLE_NAME, TABLE_TYPE FROM information_schema.TABLES WHERE TABLE_SCHEMA = {} ORDER BY TABLE_NAME",
quote_value(database),
);
let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql)
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
let rows: Vec<MySqlRow> = 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<Vec<TableIn
.collect())
}
pub async fn get_columns(
pool: &MySqlPool,
database: &str,
table: &str,
) -> Result<Vec<ColumnInfo>, String> {
pub async fn get_columns(pool: &MySqlPool, database: &str, table: &str) -> Result<Vec<ColumnInfo>, 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<MySqlRow> = sqlx::raw_sql(&sql)
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
let rows: Vec<MySqlRow> = 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<Qu
let start = Instant::now();
let trimmed = sql.trim().to_uppercase();
if trimmed.starts_with("SELECT") || trimmed.starts_with("SHOW") || trimmed.starts_with("DESCRIBE") || trimmed.starts_with("EXPLAIN") {
if trimmed.starts_with("SELECT")
|| trimmed.starts_with("SHOW")
|| trimmed.starts_with("DESCRIBE")
|| trimmed.starts_with("EXPLAIN")
{
if bare {
let rows: Vec<MySqlRow> = sqlx::raw_sql(sql)
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
let rows: Vec<MySqlRow> = 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<String> = 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<Qu
let columns: Vec<String> = desc.columns().iter().map(|c| c.name().to_string()).collect();
let column_types: Vec<String> = desc.columns().iter().map(|c| c.type_info().name().to_string()).collect();
let rows: Vec<MySqlRow> = sqlx::query(sql)
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
let rows: Vec<MySqlRow> = sqlx::query(sql).fetch_all(pool).await.map_err(|e| e.to_string())?;
let result_rows: Vec<Vec<serde_json::Value>> = rows
.iter()
@ -320,10 +290,7 @@ pub async fn execute_query(pool: &MySqlPool, sql: &str, bare: bool) -> Result<Qu
})
}
} else {
let result = sqlx::raw_sql(sql)
.execute(pool)
.await
.map_err(|e| e.to_string())?;
let result = sqlx::raw_sql(sql).execute(pool).await.map_err(|e| e.to_string())?;
Ok(QueryResult {
columns: vec![],
@ -347,10 +314,7 @@ pub async fn list_indexes(pool: &MySqlPool, database: &str, table: &str) -> Resu
quote_value(database),
quote_value(table),
);
let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql)
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
let rows: Vec<MySqlRow> = 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<MySqlRow> = sqlx::raw_sql(&sql)
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
let rows: Vec<MySqlRow> = 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<MySqlRow> = sqlx::raw_sql(&sql)
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
let rows: Vec<MySqlRow> = sqlx::raw_sql(&sql).fetch_all(pool).await.map_err(|e| e.to_string())?;
Ok(rows
.iter()

View File

@ -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<serde_json::Value> {
if let Ok(v) = row.try_get::<DateTime<Utc>, _>(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::<bool, _>(idx)
.map(serde_json::Value::Bool)
.unwrap_or(serde_json::Value::Null);
return row.try_get::<bool, _>(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::<String, _>(idx)
.map(serde_json::Value::String)
.or_else(|_| {
row.try_get::<i64, _>(idx)
.map(|v| serde_json::Value::Number(v.into()))
})
.or_else(|_| {
row.try_get::<i32, _>(idx)
.map(|v| serde_json::Value::Number(v.into()))
})
.or_else(|_| row.try_get::<i64, _>(idx).map(|v| serde_json::Value::Number(v.into())))
.or_else(|_| row.try_get::<i32, _>(idx).map(|v| serde_json::Value::Number(v.into())))
.or_else(|_| {
row.try_get::<f64, _>(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::<bool, _>(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<PgPool, String> {
// 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<PgPool, String> {
}
/// 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<Vec<DatabaseInfo>, String> {
let rows: Vec<PgRow> =
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<PgRow> = 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::<String, _>("datname"),
})
.collect())
Ok(rows.iter().map(|row| DatabaseInfo { name: row.get::<String, _>("datname") }).collect())
}
pub async fn list_tables(pool: &PgPool, schema: &str) -> Result<Vec<TableInfo>, String> {
@ -182,17 +165,10 @@ pub async fn list_schemas(pool: &PgPool) -> Result<Vec<String>, String> {
.await
.map_err(|e| e.to_string())?;
Ok(rows
.iter()
.map(|row| row.get::<String, _>("schema_name"))
.collect())
Ok(rows.iter().map(|row| row.get::<String, _>("schema_name")).collect())
}
pub async fn get_columns(
pool: &PgPool,
schema: &str,
table: &str,
) -> Result<Vec<ColumnInfo>, String> {
pub async fn get_columns(pool: &PgPool, schema: &str, table: &str) -> Result<Vec<ColumnInfo>, String> {
let rows: Vec<PgRow> = 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<QueryResult, Stri
|| trimmed.starts_with("WITH")
|| trimmed.starts_with("TABLE")
{
let rows: Vec<PgRow> = sqlx::query(sql)
.persistent(false)
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
let rows: Vec<PgRow> = sqlx::query(sql).persistent(false).fetch_all(pool).await.map_err(|e| e.to_string())?;
let (columns, column_types): (Vec<String>, Vec<String>) = if let Some(first) = rows.first() {
let cols = first.columns();
@ -278,13 +250,7 @@ pub async fn execute_query(pool: &PgPool, sql: &str) -> Result<QueryResult, Stri
.iter()
.map(|row| {
(0..row.len())
.map(|i| {
pg_value_to_json(
row,
i,
column_types.get(i).map(String::as_str).unwrap_or(""),
)
})
.map(|i| pg_value_to_json(row, i, column_types.get(i).map(String::as_str).unwrap_or("")))
.collect()
})
.collect();
@ -297,10 +263,7 @@ pub async fn execute_query(pool: &PgPool, sql: &str) -> Result<QueryResult, Stri
truncated: false,
})
} else {
let result = sqlx::query(sql)
.execute(pool)
.await
.map_err(|e| e.to_string())?;
let result = sqlx::query(sql).execute(pool).await.map_err(|e| e.to_string())?;
Ok(QueryResult {
columns: vec![],

View File

@ -27,44 +27,26 @@ pub struct RedisValue {
pub async fn connect(url: &str) -> Result<redis::aio::MultiplexedConnection, String> {
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::<String>(&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::<String>(&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<Vec<u32>, 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<Vec<u32>, 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<u32> {
})
}
async fn list_keyspace_databases(
con: &mut redis::aio::MultiplexedConnection,
) -> Result<Vec<u32>, 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<Vec<u32>, 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<RedisValue, String> {
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<RedisValue, String> {
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::<Vec<_>>())
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::<Vec<_>>())
}
"hash" => {
let v: Vec<(String, String)> = con.hgetall(key).await.map_err(|e| e.to_string())?;
let map: serde_json::Map<String, serde_json::Value> = v
.into_iter()
.map(|(k, v)| (k, serde_json::Value::String(v)))
.collect();
let map: serde_json::Map<String, serde_json::Value> =
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<i64>,
) -> 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);

View File

@ -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<SqlitePool, String> {
// 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<Vec<TableIn
}
pub async fn get_columns(pool: &SqlitePool, _schema: &str, table: &str) -> Result<Vec<ColumnInfo>, String> {
let rows: Vec<SqliteRow> = sqlx::query(&format!("PRAGMA table_info(\"{}\")", table))
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
let rows: Vec<SqliteRow> =
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::<i32, _>("notnull") == 0,
column_default: row.get::<Option<String>, _>("dflt_value"),
is_primary_key: row.get::<i32, _>("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<Vec<TriggerInfo>, String> {
let rows: Vec<SqliteRow> = 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<SqliteRow> =
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::<Option<String>, _>("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::<String, _>("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::<String, _>("name"), event: event.to_string(), timing: timing.to_string() }
})
.collect())
}
@ -166,10 +170,7 @@ pub async fn execute_query(pool: &SqlitePool, sql: &str) -> Result<QueryResult,
let desc = pool.describe(sql).await.map_err(|e| e.to_string())?;
let columns: Vec<String> = desc.columns().iter().map(|c| c.name().to_string()).collect();
let rows: Vec<SqliteRow> = sqlx::query(sql)
.fetch_all(pool)
.await
.map_err(|e| e.to_string())?;
let rows: Vec<SqliteRow> = sqlx::query(sql).fetch_all(pool).await.map_err(|e| e.to_string())?;
let result_rows: Vec<Vec<serde_json::Value>> = rows
.iter()
@ -179,11 +180,13 @@ pub async fn execute_query(pool: &SqlitePool, sql: &str) -> Result<QueryResult,
row.try_get::<String, _>(i)
.map(serde_json::Value::String)
.or_else(|_| row.try_get::<i64, _>(i).map(|v| serde_json::Value::Number(v.into())))
.or_else(|_| row.try_get::<f64, _>(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::<f64, _>(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::<bool, _>(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<QueryResult,
truncated: false,
})
} else {
let result = sqlx::query(sql)
.execute(pool)
.await
.map_err(|e| e.to_string())?;
let result = sqlx::query(sql).execute(pool).await.map_err(|e| e.to_string())?;
Ok(QueryResult {
columns: vec![],

View File

@ -9,7 +9,7 @@ use tokio::net::TcpListener;
use tokio::sync::Mutex;
use tokio::task::JoinHandle;
use super::{connection_timeout, CONNECTION_TIMEOUT_SECS, file_validator::validate_file_path};
use super::{connection_timeout, file_validator::validate_file_path, CONNECTION_TIMEOUT_SECS};
struct SshClient;
@ -43,7 +43,7 @@ async fn connect_and_authenticate(
if !ssh_key_path.is_empty() {
// Validate SSH key file path
validate_file_path(ssh_key_path, |_| false)?;
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 = tokio::time::timeout(