diff --git a/src/db/database.rs b/src/db/database.rs index 5329690..e493dd9 100644 --- a/src/db/database.rs +++ b/src/db/database.rs @@ -8,7 +8,7 @@ use crate::interpreter::ast::SqlStatement; use std::collections::HashMap; pub struct Database { - pub tables: HashMap, + pub tables: HashMap>>, pub transaction: TransactionLog, } @@ -54,7 +54,7 @@ impl Database { Ok(None) } SqlStatement::DropTable(statement) => { - drop_table::drop_table(self, statement)?; + drop_table::drop_table(self, statement, self.transaction.in_transaction())?; self.transaction.append_entry(sql_statement_clone, vec![])?; Ok(None) } @@ -114,20 +114,65 @@ impl Database { pub fn has_table(&self, table_name: &str) -> bool { self.tables.contains_key(table_name) + && !self.tables.get(table_name).is_none() + && !self.tables.get(table_name).unwrap().is_empty() + && self + .tables + .get(table_name) + .unwrap() + .last() + .unwrap() + .is_some() } pub fn get_table(&self, table_name: &str) -> Result<&Table, String> { if !self.has_table(table_name) { return Err(format!("Table `{}` does not exist", table_name)); } - Ok(self.tables.get(table_name).unwrap()) + let table = self.tables.get(table_name).unwrap().last().unwrap(); + match table { + Some(table) => Ok(table), + _ => Err(format!("Table `{}` does not exist", table_name)), + } } pub fn get_table_mut(&mut self, table_name: &str) -> Result<&mut Table, String> { if !self.has_table(table_name) { return Err(format!("Table `{}` does not exist", table_name)); } - Ok(self.tables.get_mut(table_name).unwrap()) + let table = self.tables.get_mut(table_name).unwrap().last_mut().unwrap(); + match table { + Some(table) => Ok(table), + _ => Err(format!("Table `{}` does not exist", table_name)), + } + } + + pub fn push_table_change(&mut self, table_name: &str, table: Table) { + if !self.has_table(table_name) { + self.tables + .insert(table_name.to_string(), vec![Some(table)]); + } else { + self.tables.get_mut(table_name).unwrap().push(Some(table)); + } + } + + pub fn pop_table_change(&mut self, table_name: &str) -> Result { + if !self.has_table(table_name) { + return Err(format!("Table `{}` does not exist", table_name)); + } + + let table = self.tables.get_mut(table_name).unwrap().pop().unwrap(); + + // Check if vector is empty before removing key + let is_empty = self.tables.get(table_name).unwrap().is_empty(); + if is_empty { + self.tables.remove(table_name); + } + + match table { + Some(table) => Ok(table), + _ => Err(format!("Table `{}` does not exist", table_name)), + } } } @@ -140,7 +185,7 @@ mod tests { Database { tables: HashMap::from([( "users".to_string(), - Table::new( + vec![Some(Table::new( "users".to_string(), vec![ ColumnDefinition { @@ -154,7 +199,7 @@ mod tests { constraints: vec![], }, ], - ), + ))], )]), transaction: TransactionLog { entries: None }, } diff --git a/src/db/table/operations/alter_table/mod.rs b/src/db/table/operations/alter_table/mod.rs index ee9fa98..ddbd5a8 100644 --- a/src/db/table/operations/alter_table/mod.rs +++ b/src/db/table/operations/alter_table/mod.rs @@ -9,14 +9,9 @@ pub fn alter_table( ) -> Result<(), String> { return match statement.action { AlterTableAction::RenameTable { new_table_name } => { - let table = database.tables.remove(&statement.table_name); - match table { - Some(mut table) => { - table.change_name(new_table_name, is_transaction); - database.tables.insert(table.name()?.clone(), table); - } - None => return Err(format!("Table `{}` does not exist", statement.table_name)), - }; + let mut table = database.pop_table_change(statement.table_name.as_str())?; + table.change_name(new_table_name.clone(), is_transaction); + database.push_table_change(new_table_name.as_str(), table); Ok(()) } AlterTableAction::RenameColumn { @@ -113,9 +108,9 @@ mod tests { }; let result = alter_table(&mut database, statement, false); assert!(result.is_ok()); - assert!(!database.tables.contains_key("users")); - assert!(database.tables.contains_key("new_users")); - assert!(database.tables.get("new_users").unwrap().name().unwrap() == "new_users"); + assert!(!database.has_table("users")); + assert!(database.has_table("new_users")); + assert!(database.get_table("new_users").unwrap().name().unwrap() == "new_users"); } #[test] diff --git a/src/db/table/operations/create_table/mod.rs b/src/db/table/operations/create_table/mod.rs index 8336be9..0fe8542 100644 --- a/src/db/table/operations/create_table/mod.rs +++ b/src/db/table/operations/create_table/mod.rs @@ -17,7 +17,9 @@ pub fn create_table( } } let table = Table::new(statement.table_name, statement.columns); - database.tables.insert(table.name()?.clone(), table); + database + .tables + .insert(table.name()?.clone(), vec![Some(table)]); Ok(()) } diff --git a/src/db/table/operations/drop_table/mod.rs b/src/db/table/operations/drop_table/mod.rs index 38190c4..1217fa1 100644 --- a/src/db/table/operations/drop_table/mod.rs +++ b/src/db/table/operations/drop_table/mod.rs @@ -1,7 +1,11 @@ use crate::db::database::Database; use crate::interpreter::ast::{DropTableStatement, ExistenceCheck}; -pub fn drop_table(database: &mut Database, statement: DropTableStatement) -> Result<(), String> { +pub fn drop_table( + database: &mut Database, + statement: DropTableStatement, + is_transaction: bool, +) -> Result<(), String> { if !database.has_table(&statement.table_name) { match statement.existence_check { Some(ExistenceCheck::IfExists) => { @@ -12,7 +16,15 @@ pub fn drop_table(database: &mut Database, statement: DropTableStatement) -> Res } } } - database.tables.remove(&statement.table_name); + if is_transaction { + database + .tables + .get_mut(&statement.table_name) + .unwrap() + .push(None); + } else { + database.tables.remove(&statement.table_name); + } Ok(()) } @@ -28,7 +40,7 @@ mod tests { existence_check: None, }; let mut database = default_database(); - let result = drop_table(&mut database, statement); + let result = drop_table(&mut database, statement, false); assert!(result.is_ok()); assert!(!database.has_table("users")); } @@ -40,7 +52,7 @@ mod tests { existence_check: None, }; let mut database = Database::new(); - let result = drop_table(&mut database, statement); + let result = drop_table(&mut database, statement, false); assert!(result.is_err()); assert_eq!("Table `users` does not exist", result.err().unwrap()); } @@ -52,7 +64,7 @@ mod tests { existence_check: Some(ExistenceCheck::IfExists), }; let mut database = Database::new(); - let result = drop_table(&mut database, statement); + let result = drop_table(&mut database, statement, false); assert!(result.is_ok()); } } diff --git a/src/db/table/test_utils.rs b/src/db/table/test_utils.rs index bf06133..d093b0c 100644 --- a/src/db/table/test_utils.rs +++ b/src/db/table/test_utils.rs @@ -70,7 +70,9 @@ pub fn default_table() -> Table { #[cfg(test)] pub fn default_database() -> Database { let mut database = Database::new(); - database.tables.insert("users".to_string(), default_table()); + database + .tables + .insert("users".to_string(), vec![Some(default_table())]); database } diff --git a/src/db/transactions/rollback.rs b/src/db/transactions/rollback.rs index 7277b33..93876f2 100644 --- a/src/db/transactions/rollback.rs +++ b/src/db/transactions/rollback.rs @@ -24,15 +24,22 @@ pub fn rollback_transaction_entry( } AlterTableAction::RenameTable { ref new_table_name } => { // It is now under the new name - let mut table = database - .tables - .remove(new_table_name.as_str()) - .ok_or(format!("Table `{}` does not exist", new_table_name))?; + let mut table = database.pop_table_change(new_table_name.as_str())?; table.rollback_name(); - database.tables.insert(table.name()?.clone(), table); + database.push_table_change(statement.table_name.as_str(), table); } }, SqlStatement::Select(_) => {} // These should be kept in the log but obv do nothing. + SqlStatement::CreateTable(_) => { + database.tables.remove(statement.table_name.as_str()); + } + SqlStatement::DropTable(statement) => { + database + .tables + .get_mut(statement.table_name.as_str()) + .unwrap() + .pop(); + } _ => return Err("UNSUPPORTED".to_string()), } return Ok(()); diff --git a/tests/transaction_test.rs b/tests/transaction_test.rs index 2eef1e3..bf042cb 100644 --- a/tests/transaction_test.rs +++ b/tests/transaction_test.rs @@ -29,7 +29,6 @@ fn test_transaction() { SELECT * FROM new_users; "; let result = run_sql(&mut database, sql); - println!("{:?}", result); let expected = vec![ Ok(None), Ok(None), @@ -62,3 +61,61 @@ fn test_transaction() { assert_eq!(expected[i], *result); } } + +#[test] +fn test_transaction_create_table() { + let mut database = Database::new(); + let sql = " + BEGIN; + CREATE TABLE users ( + id INTEGER, + name TEXT + ); + SELECT * FROM users; + ROLLBACK; + SELECT * FROM users; + "; + let result = run_sql(&mut database, sql); + let expected = vec![ + Ok(None), + Ok(None), + Ok(Some(vec![])), + Ok(None), + Err("Execution Error with statement starting on line 9 \n Error: Table `users` does not exist".to_string()), + ]; + for (i, result) in result.iter().enumerate() { + assert_eq!(expected[i], *result); + } +} + +#[test] +fn test_transaction_drop_table() { + let mut database = Database::new(); + let sql = " + CREATE TABLE users ( + id INTEGER, + name TEXT + ); + INSERT INTO users (id, name) VALUES (1, 'John'); + BEGIN; + SELECT * FROM users; + DROP TABLE users; + SELECT * FROM users; + ROLLBACK; + SELECT * FROM users; + "; + let result = run_sql(&mut database, sql); + let expected = vec![ + Ok(None), + Ok(None), + Ok(None), + Ok(Some(vec![Row(vec![Value::Integer(1), Value::Text("John".to_string())])])), + Ok(None), + Err("Execution Error with statement starting on line 10 \n Error: Table `users` does not exist".to_string()), + Ok(None), + Ok(Some(vec![Row(vec![Value::Integer(1), Value::Text("John".to_string())])])), + ]; + for (i, result) in result.iter().enumerate() { + assert_eq!(expected[i], *result); + } +}