use crate::config::{ConnectionConfig, DatabaseEngine, DatabasesConfig};
use crate::error::{AppError, AppResult};
use serde_json::Value as JsonValue;
use sqlx::mysql::{MySqlPool, MySqlPoolOptions, MySqlRow};
use sqlx::postgres::{PgPool, PgPoolOptions, PgRow};
use sqlx::sqlite::{SqlitePool, SqlitePoolOptions, SqliteRow};
use sqlx::{Column, Row, TypeInfo};
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use tracing::{debug, info, instrument};
pub enum DatabasePool {
MySQL(MySqlPool),
Postgres(PgPool),
SQLite(SqlitePool),
}
impl DatabasePool {
pub fn engine(&self) -> DatabaseEngine {
match self {
DatabasePool::MySQL(_) => DatabaseEngine::MySQL,
DatabasePool::Postgres(_) => DatabaseEngine::Postgres,
DatabasePool::SQLite(_) => DatabaseEngine::SQLite,
}
}
}
#[derive(Debug, Clone, serde::Serialize)]
pub struct ConnectionInfo {
pub name: String,
pub engine: String,
pub is_default: bool,
}
#[derive(Clone)]
pub struct ConnectionManager {
pools: Arc<HashMap<String, DatabasePool>>,
default_connection: Option<String>,
}
impl ConnectionManager {
#[instrument(skip(config), fields(databases = config.databases.len()))]
pub async fn new(config: &DatabasesConfig) -> AppResult<Self> {
info!("Initializing connection manager");
let mut pools = HashMap::new();
for conn_config in &config.databases {
let name = conn_config.name().to_string();
info!("Creating connection pool for '{}' ({})", name, conn_config.engine());
let pool = Self::create_pool(conn_config).await?;
pools.insert(name, pool);
}
info!(
"Connection manager initialized with {} database(s)",
pools.len()
);
Ok(Self {
pools: Arc::new(pools),
default_connection: config.default_connection.clone(),
})
}
async fn create_pool(config: &ConnectionConfig) -> AppResult<DatabasePool> {
let url = config.connection_url();
match config {
ConnectionConfig::MySQL(_) => {
let pool = MySqlPoolOptions::new()
.max_connections(config.max_connections())
.min_connections(config.min_connections())
.acquire_timeout(Duration::from_secs(config.connect_timeout_secs()))
.idle_timeout(Duration::from_secs(600))
.max_lifetime(Duration::from_secs(1800))
.connect(&url)
.await?;
Ok(DatabasePool::MySQL(pool))
}
ConnectionConfig::Postgres(_) => {
let pool = PgPoolOptions::new()
.max_connections(config.max_connections())
.min_connections(config.min_connections())
.acquire_timeout(Duration::from_secs(config.connect_timeout_secs()))
.idle_timeout(Duration::from_secs(600))
.max_lifetime(Duration::from_secs(1800))
.connect(&url)
.await?;
Ok(DatabasePool::Postgres(pool))
}
ConnectionConfig::SQLite(_) => {
let pool = SqlitePoolOptions::new()
.max_connections(config.max_connections())
.acquire_timeout(Duration::from_secs(config.connect_timeout_secs()))
.connect(&url)
.await?;
Ok(DatabasePool::SQLite(pool))
}
}
}
pub fn get_pool(&self, connection: Option<&str>) -> AppResult<&DatabasePool> {
let name = connection
.map(|s| s.to_string())
.or_else(|| self.default_connection.clone())
.ok_or(AppError::NoDefaultConnection)?;
self.pools
.get(&name)
.ok_or_else(|| AppError::ConnectionNotFound(name))
}
pub fn get_engine(&self, connection: Option<&str>) -> AppResult<DatabaseEngine> {
Ok(self.get_pool(connection)?.engine())
}
pub fn connection_count(&self) -> usize {
self.pools.len()
}
pub fn list_connections(&self) -> Vec<ConnectionInfo> {
self.pools
.iter()
.map(|(name, pool)| ConnectionInfo {
name: name.clone(),
engine: pool.engine().to_string(),
is_default: Some(name) == self.default_connection.as_ref(),
})
.collect()
}
#[instrument(skip(self, params), fields(query_len = query.len()))]
pub async fn query(
&self,
connection: Option<&str>,
query: &str,
params: Vec<JsonValue>,
) -> AppResult<Vec<serde_json::Map<String, JsonValue>>> {
let pool = self.get_pool(connection)?;
validate_query(query, QueryType::Select)?;
debug!("Executing SELECT query with {} parameters", params.len());
match pool {
DatabasePool::MySQL(p) => self.mysql_query(p, query, params).await,
DatabasePool::Postgres(p) => self.postgres_query(p, query, params).await,
DatabasePool::SQLite(p) => self.sqlite_query(p, query, params).await,
}
}
#[instrument(skip(self, params), fields(query_len = query.len()))]
pub async fn insert(
&self,
connection: Option<&str>,
query: &str,
params: Vec<JsonValue>,
) -> AppResult<i64> {
let pool = self.get_pool(connection)?;
validate_query(query, QueryType::Insert)?;
debug!("Executing INSERT query with {} parameters", params.len());
match pool {
DatabasePool::MySQL(p) => self.mysql_insert(p, query, params).await,
DatabasePool::Postgres(p) => self.postgres_insert(p, query, params).await,
DatabasePool::SQLite(p) => self.sqlite_insert(p, query, params).await,
}
}
#[instrument(skip(self, params), fields(query_len = query.len()))]
pub async fn update(
&self,
connection: Option<&str>,
query: &str,
params: Vec<JsonValue>,
) -> AppResult<u64> {
let pool = self.get_pool(connection)?;
validate_query(query, QueryType::Update)?;
debug!("Executing UPDATE query with {} parameters", params.len());
match pool {
DatabasePool::MySQL(p) => self.mysql_update(p, query, params).await,
DatabasePool::Postgres(p) => self.postgres_update(p, query, params).await,
DatabasePool::SQLite(p) => self.sqlite_update(p, query, params).await,
}
}
#[instrument(skip(self, params), fields(query_len = query.len()))]
pub async fn delete(
&self,
connection: Option<&str>,
query: &str,
params: Vec<JsonValue>,
) -> AppResult<u64> {
let pool = self.get_pool(connection)?;
validate_query(query, QueryType::Delete)?;
debug!("Executing DELETE query with {} parameters", params.len());
match pool {
DatabasePool::MySQL(p) => self.mysql_delete(p, query, params).await,
DatabasePool::Postgres(p) => self.postgres_delete(p, query, params).await,
DatabasePool::SQLite(p) => self.sqlite_delete(p, query, params).await,
}
}
#[instrument(skip(self))]
pub async fn list_tables(&self, connection: Option<&str>) -> AppResult<Vec<String>> {
let pool = self.get_pool(connection)?;
debug!("Listing database tables");
match pool {
DatabasePool::MySQL(p) => self.mysql_list_tables(p).await,
DatabasePool::Postgres(p) => self.postgres_list_tables(p).await,
DatabasePool::SQLite(p) => self.sqlite_list_tables(p).await,
}
}
#[instrument(skip(self))]
pub async fn describe_table(
&self,
connection: Option<&str>,
table_name: &str,
) -> AppResult<Vec<serde_json::Map<String, JsonValue>>> {
validate_identifier(table_name)?;
let pool = self.get_pool(connection)?;
debug!("Describing table: {}", table_name);
match pool {
DatabasePool::MySQL(p) => self.mysql_describe_table(p, table_name).await,
DatabasePool::Postgres(p) => self.postgres_describe_table(p, table_name).await,
DatabasePool::SQLite(p) => self.sqlite_describe_table(p, table_name).await,
}
}
pub async fn health_check(&self, connection: Option<&str>) -> AppResult<bool> {
let pool = self.get_pool(connection)?;
match pool {
DatabasePool::MySQL(p) => {
sqlx::query("SELECT 1").fetch_one(p).await?;
}
DatabasePool::Postgres(p) => {
sqlx::query("SELECT 1").fetch_one(p).await?;
}
DatabasePool::SQLite(p) => {
sqlx::query("SELECT 1").fetch_one(p).await?;
}
}
Ok(true)
}
async fn mysql_query(
&self,
pool: &MySqlPool,
query: &str,
params: Vec<JsonValue>,
) -> AppResult<Vec<serde_json::Map<String, JsonValue>>> {
let mut q = sqlx::query(query);
for param in params {
q = bind_mysql_value(q, param);
}
let rows: Vec<MySqlRow> = q.fetch_all(pool).await?;
Ok(rows.iter().map(mysql_row_to_json).collect())
}
async fn mysql_insert(
&self,
pool: &MySqlPool,
query: &str,
params: Vec<JsonValue>,
) -> AppResult<i64> {
let mut q = sqlx::query(query);
for param in params {
q = bind_mysql_value(q, param);
}
let result = q.execute(pool).await?;
Ok(result.last_insert_id() as i64)
}
async fn mysql_update(
&self,
pool: &MySqlPool,
query: &str,
params: Vec<JsonValue>,
) -> AppResult<u64> {
let mut q = sqlx::query(query);
for param in params {
q = bind_mysql_value(q, param);
}
let result = q.execute(pool).await?;
Ok(result.rows_affected())
}
async fn mysql_delete(
&self,
pool: &MySqlPool,
query: &str,
params: Vec<JsonValue>,
) -> AppResult<u64> {
let mut q = sqlx::query(query);
for param in params {
q = bind_mysql_value(q, param);
}
let result = q.execute(pool).await?;
Ok(result.rows_affected())
}
async fn mysql_list_tables(&self, pool: &MySqlPool) -> AppResult<Vec<String>> {
let rows: Vec<MySqlRow> = sqlx::query("SHOW TABLES").fetch_all(pool).await?;
Ok(rows
.iter()
.filter_map(|row| row.try_get::<String, _>(0).ok())
.collect())
}
async fn mysql_describe_table(
&self,
pool: &MySqlPool,
table_name: &str,
) -> AppResult<Vec<serde_json::Map<String, JsonValue>>> {
let query = format!("DESCRIBE `{}`", escape_identifier(table_name));
let rows: Vec<MySqlRow> = sqlx::query(&query).fetch_all(pool).await?;
Ok(rows.iter().map(mysql_row_to_json).collect())
}
async fn postgres_query(
&self,
pool: &PgPool,
query: &str,
params: Vec<JsonValue>,
) -> AppResult<Vec<serde_json::Map<String, JsonValue>>> {
let converted_query = convert_placeholders_to_postgres(query);
let mut q = sqlx::query(&converted_query);
for param in params {
q = bind_postgres_value(q, param);
}
let rows: Vec<PgRow> = q.fetch_all(pool).await?;
Ok(rows.iter().map(postgres_row_to_json).collect())
}
async fn postgres_insert(
&self,
pool: &PgPool,
query: &str,
params: Vec<JsonValue>,
) -> AppResult<i64> {
let converted_query = convert_placeholders_to_postgres(query);
let query_with_returning = if !converted_query.to_uppercase().contains("RETURNING") {
format!("{} RETURNING id", converted_query.trim_end_matches(';'))
} else {
converted_query
};
let mut q = sqlx::query_scalar::<_, i64>(&query_with_returning);
for param in params {
q = bind_postgres_scalar_value(q, param);
}
match q.fetch_optional(pool).await? {
Some(id) => Ok(id),
None => Ok(0), }
}
async fn postgres_update(
&self,
pool: &PgPool,
query: &str,
params: Vec<JsonValue>,
) -> AppResult<u64> {
let converted_query = convert_placeholders_to_postgres(query);
let mut q = sqlx::query(&converted_query);
for param in params {
q = bind_postgres_value(q, param);
}
let result = q.execute(pool).await?;
Ok(result.rows_affected())
}
async fn postgres_delete(
&self,
pool: &PgPool,
query: &str,
params: Vec<JsonValue>,
) -> AppResult<u64> {
let converted_query = convert_placeholders_to_postgres(query);
let mut q = sqlx::query(&converted_query);
for param in params {
q = bind_postgres_value(q, param);
}
let result = q.execute(pool).await?;
Ok(result.rows_affected())
}
async fn postgres_list_tables(&self, pool: &PgPool) -> AppResult<Vec<String>> {
let rows: Vec<PgRow> = sqlx::query(
"SELECT tablename FROM pg_tables WHERE schemaname = 'public' ORDER BY tablename",
)
.fetch_all(pool)
.await?;
Ok(rows
.iter()
.filter_map(|row| row.try_get::<String, _>(0).ok())
.collect())
}
async fn postgres_describe_table(
&self,
pool: &PgPool,
table_name: &str,
) -> AppResult<Vec<serde_json::Map<String, JsonValue>>> {
let rows: Vec<PgRow> = sqlx::query(
r#"
SELECT
column_name as "Field",
data_type as "Type",
is_nullable as "Null",
column_default as "Default",
CASE WHEN pk.column_name IS NOT NULL THEN 'PRI' ELSE '' END as "Key"
FROM information_schema.columns c
LEFT JOIN (
SELECT ku.column_name
FROM information_schema.table_constraints tc
JOIN information_schema.key_column_usage ku
ON tc.constraint_name = ku.constraint_name
WHERE tc.table_name = $1 AND tc.constraint_type = 'PRIMARY KEY'
) pk ON c.column_name = pk.column_name
WHERE c.table_name = $1
ORDER BY c.ordinal_position
"#,
)
.bind(table_name)
.fetch_all(pool)
.await?;
Ok(rows.iter().map(postgres_row_to_json).collect())
}
async fn sqlite_query(
&self,
pool: &SqlitePool,
query: &str,
params: Vec<JsonValue>,
) -> AppResult<Vec<serde_json::Map<String, JsonValue>>> {
let mut q = sqlx::query(query);
for param in params {
q = bind_sqlite_value(q, param);
}
let rows: Vec<SqliteRow> = q.fetch_all(pool).await?;
Ok(rows.iter().map(sqlite_row_to_json).collect())
}
async fn sqlite_insert(
&self,
pool: &SqlitePool,
query: &str,
params: Vec<JsonValue>,
) -> AppResult<i64> {
let mut q = sqlx::query(query);
for param in params {
q = bind_sqlite_value(q, param);
}
let result = q.execute(pool).await?;
Ok(result.last_insert_rowid())
}
async fn sqlite_update(
&self,
pool: &SqlitePool,
query: &str,
params: Vec<JsonValue>,
) -> AppResult<u64> {
let mut q = sqlx::query(query);
for param in params {
q = bind_sqlite_value(q, param);
}
let result = q.execute(pool).await?;
Ok(result.rows_affected())
}
async fn sqlite_delete(
&self,
pool: &SqlitePool,
query: &str,
params: Vec<JsonValue>,
) -> AppResult<u64> {
let mut q = sqlx::query(query);
for param in params {
q = bind_sqlite_value(q, param);
}
let result = q.execute(pool).await?;
Ok(result.rows_affected())
}
async fn sqlite_list_tables(&self, pool: &SqlitePool) -> AppResult<Vec<String>> {
let rows: Vec<SqliteRow> = sqlx::query(
"SELECT name FROM sqlite_master WHERE type='table' AND name NOT LIKE 'sqlite_%' ORDER BY name",
)
.fetch_all(pool)
.await?;
Ok(rows
.iter()
.filter_map(|row| row.try_get::<String, _>(0).ok())
.collect())
}
async fn sqlite_describe_table(
&self,
pool: &SqlitePool,
table_name: &str,
) -> AppResult<Vec<serde_json::Map<String, JsonValue>>> {
let query = format!("PRAGMA table_info('{}')", escape_identifier(table_name));
let rows: Vec<SqliteRow> = sqlx::query(&query).fetch_all(pool).await?;
Ok(rows
.iter()
.map(|row| {
let mut map = serde_json::Map::new();
map.insert(
"Field".to_string(),
JsonValue::String(row.try_get::<String, _>("name").unwrap_or_default()),
);
map.insert(
"Type".to_string(),
JsonValue::String(row.try_get::<String, _>("type").unwrap_or_default()),
);
let notnull: i32 = row.try_get("notnull").unwrap_or(0);
map.insert(
"Null".to_string(),
JsonValue::String(if notnull == 1 { "NO" } else { "YES" }.to_string()),
);
map.insert(
"Default".to_string(),
row.try_get::<String, _>("dflt_value")
.map(JsonValue::String)
.unwrap_or(JsonValue::Null),
);
let pk: i32 = row.try_get("pk").unwrap_or(0);
map.insert(
"Key".to_string(),
JsonValue::String(if pk == 1 { "PRI" } else { "" }.to_string()),
);
map
})
.collect())
}
}
#[derive(Debug, Clone, Copy)]
enum QueryType {
Select,
Insert,
Update,
Delete,
}
fn validate_query(query: &str, expected_type: QueryType) -> AppResult<()> {
let normalized = query.trim().to_uppercase();
if query.contains(';') && query.trim().ends_with(';') {
let statements: Vec<&str> = query.split(';').filter(|s| !s.trim().is_empty()).collect();
if statements.len() > 1 {
return Err(AppError::Security(
"Multiple SQL statements are not allowed".to_string(),
));
}
}
let starts_with_expected = match expected_type {
QueryType::Select => normalized.starts_with("SELECT"),
QueryType::Insert => normalized.starts_with("INSERT"),
QueryType::Update => normalized.starts_with("UPDATE"),
QueryType::Delete => normalized.starts_with("DELETE"),
};
if !starts_with_expected {
return Err(AppError::Validation(format!(
"Expected {:?} query, but received different query type",
expected_type
)));
}
let dangerous_patterns = [
"DROP ", "TRUNCATE ", "ALTER ", "CREATE ", "GRANT ", "REVOKE ",
"LOAD_FILE", "INTO OUTFILE", "INTO DUMPFILE", "BENCHMARK(", "SLEEP(",
"PG_SLEEP(", ];
for pattern in dangerous_patterns {
if normalized.contains(pattern) {
return Err(AppError::Security(format!(
"Dangerous SQL operation detected: {}",
pattern.trim()
)));
}
}
Ok(())
}
fn validate_identifier(identifier: &str) -> AppResult<()> {
if identifier.is_empty() || identifier.len() > 64 {
return Err(AppError::Validation(
"Identifier must be 1-64 characters".to_string(),
));
}
if !identifier
.chars()
.all(|c| c.is_alphanumeric() || c == '_' || c == '$')
{
return Err(AppError::Security(
"Invalid characters in identifier".to_string(),
));
}
if identifier.chars().next().map(|c| c.is_ascii_digit()) == Some(true) {
return Err(AppError::Validation(
"Identifier cannot start with a digit".to_string(),
));
}
Ok(())
}
fn escape_identifier(identifier: &str) -> String {
identifier.replace('`', "``").replace('\'', "''")
}
fn convert_placeholders_to_postgres(query: &str) -> String {
let mut result = String::with_capacity(query.len());
let mut param_num = 1;
let mut chars = query.chars().peekable();
while let Some(c) = chars.next() {
if c == '?' {
result.push('$');
result.push_str(¶m_num.to_string());
param_num += 1;
} else if c == '\'' {
result.push(c);
while let Some(sc) = chars.next() {
result.push(sc);
if sc == '\'' {
if chars.peek() == Some(&'\'') {
result.push(chars.next().unwrap());
} else {
break;
}
}
}
} else {
result.push(c);
}
}
result
}
fn bind_mysql_value<'q>(
query: sqlx::query::Query<'q, sqlx::MySql, sqlx::mysql::MySqlArguments>,
value: JsonValue,
) -> sqlx::query::Query<'q, sqlx::MySql, sqlx::mysql::MySqlArguments> {
match value {
JsonValue::Null => query.bind(None::<String>),
JsonValue::Bool(b) => query.bind(b),
JsonValue::Number(n) => {
if let Some(i) = n.as_i64() {
query.bind(i)
} else if let Some(f) = n.as_f64() {
query.bind(f)
} else {
query.bind(n.to_string())
}
}
JsonValue::String(s) => query.bind(s),
JsonValue::Array(arr) => query.bind(serde_json::to_string(&arr).unwrap_or_default()),
JsonValue::Object(obj) => query.bind(serde_json::to_string(&obj).unwrap_or_default()),
}
}
fn bind_postgres_value<'q>(
query: sqlx::query::Query<'q, sqlx::Postgres, sqlx::postgres::PgArguments>,
value: JsonValue,
) -> sqlx::query::Query<'q, sqlx::Postgres, sqlx::postgres::PgArguments> {
match value {
JsonValue::Null => query.bind(None::<String>),
JsonValue::Bool(b) => query.bind(b),
JsonValue::Number(n) => {
if let Some(i) = n.as_i64() {
query.bind(i)
} else if let Some(f) = n.as_f64() {
query.bind(f)
} else {
query.bind(n.to_string())
}
}
JsonValue::String(s) => query.bind(s),
JsonValue::Array(arr) => query.bind(serde_json::to_string(&arr).unwrap_or_default()),
JsonValue::Object(obj) => query.bind(serde_json::to_string(&obj).unwrap_or_default()),
}
}
fn bind_postgres_scalar_value<'q, T>(
query: sqlx::query::QueryScalar<'q, sqlx::Postgres, T, sqlx::postgres::PgArguments>,
value: JsonValue,
) -> sqlx::query::QueryScalar<'q, sqlx::Postgres, T, sqlx::postgres::PgArguments> {
match value {
JsonValue::Null => query.bind(None::<String>),
JsonValue::Bool(b) => query.bind(b),
JsonValue::Number(n) => {
if let Some(i) = n.as_i64() {
query.bind(i)
} else if let Some(f) = n.as_f64() {
query.bind(f)
} else {
query.bind(n.to_string())
}
}
JsonValue::String(s) => query.bind(s),
JsonValue::Array(arr) => query.bind(serde_json::to_string(&arr).unwrap_or_default()),
JsonValue::Object(obj) => query.bind(serde_json::to_string(&obj).unwrap_or_default()),
}
}
fn bind_sqlite_value<'q>(
query: sqlx::query::Query<'q, sqlx::Sqlite, sqlx::sqlite::SqliteArguments<'q>>,
value: JsonValue,
) -> sqlx::query::Query<'q, sqlx::Sqlite, sqlx::sqlite::SqliteArguments<'q>> {
match value {
JsonValue::Null => query.bind(None::<String>),
JsonValue::Bool(b) => query.bind(b),
JsonValue::Number(n) => {
if let Some(i) = n.as_i64() {
query.bind(i)
} else if let Some(f) = n.as_f64() {
query.bind(f)
} else {
query.bind(n.to_string())
}
}
JsonValue::String(s) => query.bind(s),
JsonValue::Array(arr) => query.bind(serde_json::to_string(&arr).unwrap_or_default()),
JsonValue::Object(obj) => query.bind(serde_json::to_string(&obj).unwrap_or_default()),
}
}
fn mysql_row_to_json(row: &MySqlRow) -> serde_json::Map<String, JsonValue> {
let mut map = serde_json::Map::new();
for column in row.columns() {
let name = column.name().to_string();
let type_name = column.type_info().name();
let value: JsonValue = match type_name {
"BOOLEAN" | "BOOL" | "TINYINT(1)" => row
.try_get::<bool, _>(column.ordinal())
.map(JsonValue::Bool)
.unwrap_or(JsonValue::Null),
"TINYINT" | "SMALLINT" | "MEDIUMINT" | "INT" | "INTEGER" | "BIGINT" => row
.try_get::<i64, _>(column.ordinal())
.map(|v| JsonValue::Number(v.into()))
.unwrap_or(JsonValue::Null),
"TINYINT UNSIGNED" | "SMALLINT UNSIGNED" | "MEDIUMINT UNSIGNED" | "INT UNSIGNED"
| "BIGINT UNSIGNED" => row
.try_get::<u64, _>(column.ordinal())
.map(|v| JsonValue::Number(v.into()))
.unwrap_or(JsonValue::Null),
"FLOAT" | "DOUBLE" | "DECIMAL" => row
.try_get::<f64, _>(column.ordinal())
.ok()
.and_then(|v| serde_json::Number::from_f64(v))
.map(JsonValue::Number)
.unwrap_or(JsonValue::Null),
"JSON" => row
.try_get::<JsonValue, _>(column.ordinal())
.unwrap_or(JsonValue::Null),
_ => row
.try_get::<String, _>(column.ordinal())
.map(JsonValue::String)
.unwrap_or(JsonValue::Null),
};
map.insert(name, value);
}
map
}
fn postgres_row_to_json(row: &PgRow) -> serde_json::Map<String, JsonValue> {
let mut map = serde_json::Map::new();
for column in row.columns() {
let name = column.name().to_string();
let type_name = column.type_info().name();
let value: JsonValue = match type_name {
"BOOL" => row
.try_get::<bool, _>(column.ordinal())
.map(JsonValue::Bool)
.unwrap_or(JsonValue::Null),
"INT2" | "INT4" | "INT8" => row
.try_get::<i64, _>(column.ordinal())
.map(|v| JsonValue::Number(v.into()))
.unwrap_or(JsonValue::Null),
"FLOAT4" | "FLOAT8" | "NUMERIC" => row
.try_get::<f64, _>(column.ordinal())
.ok()
.and_then(|v| serde_json::Number::from_f64(v))
.map(JsonValue::Number)
.unwrap_or(JsonValue::Null),
"JSON" | "JSONB" => row
.try_get::<JsonValue, _>(column.ordinal())
.unwrap_or(JsonValue::Null),
_ => row
.try_get::<String, _>(column.ordinal())
.map(JsonValue::String)
.unwrap_or(JsonValue::Null),
};
map.insert(name, value);
}
map
}
fn sqlite_row_to_json(row: &SqliteRow) -> serde_json::Map<String, JsonValue> {
let mut map = serde_json::Map::new();
for column in row.columns() {
let name = column.name().to_string();
let type_name = column.type_info().name().to_uppercase();
let value: JsonValue = match type_name.as_str() {
"BOOLEAN" | "BOOL" => row
.try_get::<bool, _>(column.ordinal())
.map(JsonValue::Bool)
.unwrap_or(JsonValue::Null),
"INTEGER" | "INT" | "BIGINT" | "SMALLINT" | "TINYINT" => row
.try_get::<i64, _>(column.ordinal())
.map(|v| JsonValue::Number(v.into()))
.unwrap_or(JsonValue::Null),
"REAL" | "FLOAT" | "DOUBLE" => row
.try_get::<f64, _>(column.ordinal())
.ok()
.and_then(|v| serde_json::Number::from_f64(v))
.map(JsonValue::Number)
.unwrap_or(JsonValue::Null),
_ => row
.try_get::<String, _>(column.ordinal())
.map(JsonValue::String)
.unwrap_or(JsonValue::Null),
};
map.insert(name, value);
}
map
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_convert_placeholders() {
assert_eq!(
convert_placeholders_to_postgres("SELECT * FROM users WHERE id = ?"),
"SELECT * FROM users WHERE id = $1"
);
assert_eq!(
convert_placeholders_to_postgres("INSERT INTO users (name, age) VALUES (?, ?)"),
"INSERT INTO users (name, age) VALUES ($1, $2)"
);
assert_eq!(
convert_placeholders_to_postgres("SELECT * FROM users WHERE name = '?' AND id = ?"),
"SELECT * FROM users WHERE name = '?' AND id = $1"
);
}
#[test]
fn test_validate_identifier() {
assert!(validate_identifier("users").is_ok());
assert!(validate_identifier("user_table").is_ok());
assert!(validate_identifier("User123").is_ok());
assert!(validate_identifier("").is_err());
assert!(validate_identifier("123table").is_err());
assert!(validate_identifier("table;drop").is_err());
}
#[test]
fn test_validate_query() {
assert!(validate_query("SELECT * FROM users", QueryType::Select).is_ok());
assert!(validate_query("INSERT INTO users VALUES (1)", QueryType::Insert).is_ok());
assert!(validate_query("INSERT INTO users VALUES (1)", QueryType::Select).is_err());
assert!(validate_query("SELECT SLEEP(10)", QueryType::Select).is_err());
assert!(validate_query("SELECT pg_sleep(10)", QueryType::Select).is_err());
}
}