use crate::db::db_types::{QueryResult, SqlValue};
use crate::db::storage::{
StorageEngine, StorageError, TableSchema, WhereClause,
};
use std::sync::Arc;
#[derive(Clone)]
pub struct AsyncStorageEngine {
inner: Arc<tokio::sync::Mutex<StorageEngine>>,
}
impl Default for AsyncStorageEngine {
fn default() -> Self {
Self::new()
}
}
impl AsyncStorageEngine {
pub fn new() -> Self {
Self {
inner: Arc::new(tokio::sync::Mutex::new(StorageEngine::new())),
}
}
pub fn from_sync(engine: StorageEngine) -> Self {
Self {
inner: Arc::new(tokio::sync::Mutex::new(engine)),
}
}
pub async fn save_to_file(&self, path: &str) -> Result<(), StorageError> {
let engine = self.inner.lock().await;
engine.save_to_file(path)
}
pub async fn load_from_file(path: &str) -> Result<Self, StorageError> {
let engine = StorageEngine::load_from_file(path)?;
Ok(Self::from_sync(engine))
}
pub async fn create_table(&self, schema: TableSchema) -> Result<(), StorageError> {
let mut engine = self.inner.lock().await;
engine.create_table(schema)
}
pub async fn drop_table(&self, table_name: &str) -> Result<(), StorageError> {
let mut engine = self.inner.lock().await;
engine.drop_table(table_name)
}
pub async fn insert(
&self,
table_name: &str,
values: Vec<SqlValue>,
) -> Result<u64, StorageError> {
let mut engine = self.inner.lock().await;
engine.insert(table_name, values)
}
pub async fn select(
&self,
table_name: &str,
columns: Option<Vec<String>>,
where_clause: Option<WhereClause>,
limit: Option<usize>,
offset: Option<usize>,
) -> Result<QueryResult, StorageError> {
let engine = self.inner.lock().await;
engine.select(table_name, columns, where_clause, limit, offset)
}
pub async fn update(
&self,
table_name: &str,
updates: std::collections::HashMap<String, SqlValue>,
where_clause: Option<WhereClause>,
) -> Result<u64, StorageError> {
let mut engine = self.inner.lock().await;
engine.update(table_name, updates, where_clause)
}
pub async fn delete(
&self,
table_name: &str,
where_clause: Option<WhereClause>,
) -> Result<u64, StorageError> {
let mut engine = self.inner.lock().await;
engine.delete(table_name, where_clause)
}
pub async fn get_table_schema(&self, table_name: &str) -> Option<TableSchema> {
let engine = self.inner.lock().await;
engine.get_table_schema(table_name).cloned()
}
pub async fn get_table_names(&self) -> Vec<String> {
let engine = self.inner.lock().await;
engine.get_table_names()
}
pub async fn clear(&self) {
let mut engine = self.inner.lock().await;
*engine = StorageEngine::new();
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::db::storage::{ColumnDefinition, ColumnType, TableSchema};
fn users_schema() -> TableSchema {
TableSchema {
name: "users".to_string(),
columns: vec![
ColumnDefinition {
name: "id".to_string(),
column_type: ColumnType::Integer,
nullable: false,
default: None,
unique: true,
},
ColumnDefinition {
name: "name".to_string(),
column_type: ColumnType::Text,
nullable: false,
default: None,
unique: false,
},
ColumnDefinition {
name: "age".to_string(),
column_type: ColumnType::Integer,
nullable: true,
default: None,
unique: false,
},
],
primary_key: Some("id".to_string()),
}
}
#[tokio::test]
async fn test_async_crud() {
let engine = AsyncStorageEngine::new();
engine.create_table(users_schema()).await.unwrap();
let id = engine
.insert(
"users",
vec![
SqlValue::I32(1),
SqlValue::String("Alice".to_string()),
SqlValue::I32(25),
],
)
.await
.unwrap();
assert_eq!(id, 1);
let result = engine.select("users", None, None, None, None).await.unwrap();
assert_eq!(result.rows.len(), 1);
assert_eq!(
result.rows[0].get("name"),
Some(&SqlValue::String("Alice".to_string()))
);
let mut updates = std::collections::HashMap::new();
updates.insert("age".to_string(), SqlValue::I32(26));
let affected = engine.update("users", updates, None).await.unwrap();
assert_eq!(affected, 1);
let affected = engine.delete("users", None).await.unwrap();
assert_eq!(affected, 1);
let result = engine.select("users", None, None, None, None).await.unwrap();
assert!(result.rows.is_empty());
}
#[tokio::test]
async fn test_async_concurrent_inserts() {
let engine = AsyncStorageEngine::new();
engine.create_table(users_schema()).await.unwrap();
let mut handles = Vec::new();
for i in 0..50u32 {
let engine = engine.clone();
handles.push(tokio::spawn(async move {
engine
.insert(
"users",
vec![
SqlValue::I32(i as i32),
SqlValue::String(format!("user_{}", i)),
SqlValue::I32(i as i32 * 2),
],
)
.await
.unwrap();
}));
}
for h in handles {
h.await.unwrap();
}
let result = engine.select("users", None, None, None, None).await.unwrap();
assert_eq!(result.rows.len(), 50);
let mut ids: Vec<i32> = result
.rows
.iter()
.filter_map(|r| r.get("id").and_then(|v| v.as_i32()))
.collect();
ids.sort();
let expected: Vec<i32> = (0..50).collect();
assert_eq!(ids, expected);
}
#[tokio::test]
async fn test_async_concurrent_reads_writes() {
let engine = AsyncStorageEngine::new();
engine.create_table(users_schema()).await.unwrap();
for i in 0..10u32 {
engine
.insert(
"users",
vec![
SqlValue::I32(i as i32),
SqlValue::String(format!("user_{}", i)),
SqlValue::I32(i as i32),
],
)
.await
.unwrap();
}
let reader = engine.clone();
let read_handle = tokio::spawn(async move {
let result = reader.select("users", None, None, None, None).await.unwrap();
result.rows.len()
});
let mut writers = Vec::new();
for i in 0..10u32 {
let engine = engine.clone();
writers.push(tokio::spawn(async move {
let mut updates = std::collections::HashMap::new();
updates.insert("age".to_string(), SqlValue::I32(i as i32 + 100));
engine
.update(
"users",
updates,
Some(WhereClause::Eq("id".to_string(), SqlValue::I32(i as i32))),
)
.await
.unwrap();
}));
}
let read_count = read_handle.await.unwrap();
for w in writers {
w.await.unwrap();
}
assert_eq!(read_count, 10);
}
#[tokio::test]
async fn test_async_persistence() {
let path = format!("/tmp/torm_async_storage_{}.tormdb", uuid::Uuid::new_v4());
let _ = std::fs::remove_file(&path);
{
let engine = AsyncStorageEngine::new();
engine.create_table(users_schema()).await.unwrap();
engine
.insert(
"users",
vec![
SqlValue::I32(1),
SqlValue::String("Alice".to_string()),
SqlValue::I32(25),
],
)
.await
.unwrap();
engine.save_to_file(&path).await.unwrap();
}
{
let engine = AsyncStorageEngine::load_from_file(&path).await.unwrap();
let result = engine.select("users", None, None, None, None).await.unwrap();
assert_eq!(result.rows.len(), 1);
assert_eq!(
result.rows[0].get("name"),
Some(&SqlValue::String("Alice".to_string()))
);
}
let _ = std::fs::remove_file(&path);
}
}