diff --git a/src/db/database.rs b/src/db/database.rs index 858688b..a3da50f 100644 --- a/src/db/database.rs +++ b/src/db/database.rs @@ -2,8 +2,8 @@ use crate::db::table::core::{row::Row, table::Table}; use crate::db::table::operations::{ alter_table, create_table, delete, drop_table, insert, select, update, }; -use crate::db::transactions::rollback::rollback_transaction_entry; -use crate::db::transactions::{TransactionEntry, TransactionLog}; +use crate::db::transactions::TransactionLog; +use crate::db::transactions::{commit::commit_transaction, rollback::rollback_statement}; use crate::interpreter::ast::SqlStatement; use std::collections::HashMap; @@ -24,7 +24,7 @@ impl Database { let sql_statement_clone = sql_statement.clone(); return match sql_statement { SqlStatement::CreateTable(statement) => { - create_table::create_table(self, statement)?; + create_table::create_table(self, statement, self.transaction.in_transaction())?; self.transaction.append_entry(sql_statement_clone, vec![])?; Ok(None) } @@ -40,15 +40,17 @@ impl Database { Ok(Some(result)) } SqlStatement::UpdateStatement(statement) => { + let is_transaction = self.transaction.in_transaction(); let table = self.get_table_mut(&statement.table_name)?; - let rows_updated = update::update(table, statement)?; + let rows_updated = update::update(table, statement, is_transaction)?; self.transaction .append_entry(sql_statement_clone, rows_updated)?; Ok(None) } SqlStatement::DeleteStatement(statement) => { + let is_transaction = self.transaction.in_transaction(); let table = self.get_table_mut(&statement.table_name)?; - let rows_deleted = delete::delete(table, statement)?; + let rows_deleted = delete::delete(table, statement, is_transaction)?; self.transaction .append_entry(sql_statement_clone, rows_deleted)?; Ok(None) @@ -68,37 +70,11 @@ impl Database { Ok(None) } SqlStatement::Commit => { - let transaction_log = self.transaction.commit_transaction()?; - for transaction_entry in transaction_log.get_entries()?.iter() { - match transaction_entry { - TransactionEntry::Statement(statement) => { - let table = self.get_table_mut(&statement.table_name)?; - table.commit_transaction(&statement.affected_rows)?; - } - TransactionEntry::Savepoint(_) => {} - } - } - + commit_transaction(self)?; Ok(None) } - SqlStatement::Rollback(_) => { - if let Some(transaction_log) = self.transaction.commit_transaction()?.entries { - // We roll back in reverse order because of dependencies. - for transaction_entry in transaction_log.iter().rev() { - match transaction_entry { - TransactionEntry::Statement(statement) => { - // TODO: Some matching needs to be here for table based operations. - // CURRENTLY SUPPORTED STATEMENTS ARE: - // - ALTER TABLE RENAME COLUMN, ALTER TABLE ADD COLUMN, ALTER TABLE DROP COLUMN, ALTER TABLE RENAME TABLE - // - CREATE TABLE, DROP TABLE - rollback_transaction_entry(self, &statement)?; - } - TransactionEntry::Savepoint(_) => {} - } - } - } else { - return Err("No transaction is currently active".to_string()); - } + SqlStatement::Rollback(statement) => { + rollback_statement(self, &statement)?; Ok(None) } SqlStatement::Savepoint(_) => { diff --git a/src/db/table/core/table.rs b/src/db/table/core/table.rs index 2eca446..1b06f65 100644 --- a/src/db/table/core/table.rs +++ b/src/db/table/core/table.rs @@ -9,7 +9,8 @@ use std::ops::{Index, IndexMut}; pub struct Table { pub name: NameStack, pub columns: ColumnStack, - rows: Vec, + pub rows: Vec, + length: usize, } #[derive(Debug)] @@ -37,6 +38,7 @@ impl Table { name: NameStack { stack: vec![name] }, columns: ColumnStack::new(columns), rows: vec![], + length: 0, } } @@ -56,19 +58,33 @@ impl Table { } pub fn get(&self, i: usize) -> Option<&Row> { - self.rows.get(i)?.stack.last() + if i < self.length { + self.rows.get(i)?.stack.last() + } else { + None + } } pub fn iter(&self) -> impl Iterator { - self.rows.iter().map(|s| s.stack.last().unwrap()) + self.rows + .iter() + .take(self.length) + .map(|s| s.stack.last().unwrap()) } pub fn iter_mut(&mut self) -> impl Iterator { - self.rows.iter_mut().map(|s| s.stack.last_mut().unwrap()) + self.rows + .iter_mut() + .take(self.length) + .map(|s| s.stack.last_mut().unwrap()) } pub fn len(&self) -> usize { - self.rows.len() + self.length + } + + pub fn set_length(&mut self, length: usize) { + self.length = length; } pub fn swap(&mut self, a: usize, b: usize) { @@ -79,17 +95,23 @@ impl Table { pub fn get_rows_clone(&self) -> Vec { self.rows .iter() + .take(self.length) .map(|s| s.stack.last().unwrap().clone()) .collect() } pub fn get_rows(&self) -> Vec<&Row> { - self.rows.iter().map(|s| s.stack.last().unwrap()).collect() + self.rows + .iter() + .take(self.length) + .map(|s| s.stack.last().unwrap()) + .collect() } pub fn get_rows_mut(&mut self) -> Vec<&mut Row> { self.rows .iter_mut() + .take(self.length) .map(|s| s.stack.last_mut().unwrap()) .collect() } @@ -104,14 +126,17 @@ impl Table { } pub fn set_rows(&mut self, rows: Vec) { + self.length = rows.len(); self.rows = rows.into_iter().map(|r| RowStack::new(r)).collect(); } pub fn push(&mut self, row: Row) { + self.length += 1; self.rows.push(RowStack::new(row)); } pub fn pop(&mut self) -> Option { + self.length -= 1; self.rows.pop().and_then(|mut value| value.stack.pop()) } @@ -123,7 +148,16 @@ impl Table { } else { return Err("Error committing transaction. Row stack is empty".to_string()); } - // TODO: Add commit for column stack and name stack. + } + if self.columns.stack.len() > 1 { + let last_column_stack = self.columns.stack.pop().unwrap(); + self.columns = ColumnStack::new(last_column_stack); + } + if self.name.stack.len() > 1 { + let last_name = self.name.stack.pop().unwrap(); + self.name = NameStack { + stack: vec![last_name], + }; } Ok(()) } diff --git a/src/db/table/operations/create_table/mod.rs b/src/db/table/operations/create_table/mod.rs index 0fe8542..70d49b0 100644 --- a/src/db/table/operations/create_table/mod.rs +++ b/src/db/table/operations/create_table/mod.rs @@ -5,6 +5,7 @@ use crate::interpreter::ast::{CreateTableStatement, ExistenceCheck}; pub fn create_table( database: &mut Database, statement: CreateTableStatement, + is_transaction: bool, ) -> Result<(), String> { if database.has_table(&statement.table_name) { match statement.existence_check { @@ -16,10 +17,18 @@ pub fn create_table( } } } - let table = Table::new(statement.table_name, statement.columns); - database - .tables - .insert(table.name()?.clone(), vec![Some(table)]); + let table = Table::new(statement.table_name.clone(), statement.columns); + if is_transaction && database.tables.contains_key(&statement.table_name) { + database + .tables + .get_mut(&statement.table_name) + .unwrap() + .push(Some(table)); + } else { + database + .tables + .insert(table.name()?.clone(), vec![Some(table)]); + } Ok(()) } @@ -42,7 +51,7 @@ mod tests { }], }; let mut database = Database::new(); - assert!(create_table(&mut database, statement).is_ok()); + assert!(create_table(&mut database, statement, false).is_ok()); assert!(database.has_table("users")); } @@ -58,7 +67,7 @@ mod tests { }], }; let mut database = default_database(); - let result = create_table(&mut database, statement); + let result = create_table(&mut database, statement, false); assert!(result.is_err()); assert_eq!("Table users already exists", result.err().unwrap()); } @@ -75,7 +84,27 @@ mod tests { }], }; let mut database = default_database(); - let result = create_table(&mut database, statement); + let result = create_table(&mut database, statement, false); + assert!(result.is_ok()); + } + + #[test] + fn create_table_with_transaction_clause_works_correctly() { + let statement = CreateTableStatement { + table_name: "users".to_string(), + existence_check: None, + columns: vec![ColumnDefinition { + name: "id".to_string(), + data_type: DataType::Integer, + constraints: vec![], + }], + }; + let mut database = Database::new(); + let result = create_table(&mut database, statement, true); assert!(result.is_ok()); + assert!(database.has_table("users")); + let table = database.tables.get("users").unwrap(); + assert!(table.len() == 1); + assert!(table.first().unwrap().is_some()); } } diff --git a/src/db/table/operations/delete/mod.rs b/src/db/table/operations/delete/mod.rs index da00074..345bfb3 100644 --- a/src/db/table/operations/delete/mod.rs +++ b/src/db/table/operations/delete/mod.rs @@ -4,18 +4,30 @@ use crate::db::table::core::table::Table; use crate::db::table::operations::helpers::common::get_row_indicies_matching_clauses; use crate::interpreter::ast::DeleteStatement; -pub fn delete(table: &mut Table, statement: DeleteStatement) -> Result, String> { - let row_indicies_to_delete = get_row_indicies_matching_clauses( +pub fn delete( + table: &mut Table, + statement: DeleteStatement, + is_transaction: bool, +) -> Result, String> { + let mut row_indicies_to_delete = get_row_indicies_matching_clauses( table, &statement.where_clause, &statement.order_by_clause, &statement.limit_clause, )?; - swap_remove_bulk(table, &row_indicies_to_delete)?; + // We get omega saved here by the fact that we don't need to guarentee the order of the rows after rollbacks. + // This means we can swap the semi-deleted rows to the end of the table and then set the length of the table + // to the length of the table minus the number of semi-deleted rows. Then on rollback we can just extend the length of the table. + // to then include the deleted rows. if we commit, we pop off the end of the table until at the desired length. + swap_remove_bulk(table, &mut row_indicies_to_delete, is_transaction)?; Ok(row_indicies_to_delete) } -fn swap_remove_bulk(table: &mut Table, row_indicies: &Vec) -> Result<(), String> { +fn swap_remove_bulk( + table: &mut Table, + row_indicies: &mut Vec, + is_transaction: bool, +) -> Result<(), String> { if table.len() == 0 { if row_indicies.len() != 0 { unreachable!(); @@ -25,7 +37,8 @@ fn swap_remove_bulk(table: &mut Table, row_indicies: &Vec) -> Result<(), let table_len = table.len() - 1; let mut row_indicies_set = row_indicies.iter().collect::>(); let mut right_pointer = 0; - let mut iter = row_indicies.iter(); + let mut iter = row_indicies.iter().rev(); // We recieve the indexes in ascending order, + // We reverse them to get rid of the furtherst indexes first. while let Some(to_swap) = iter.next() { if *to_swap == (table_len - right_pointer) { @@ -37,8 +50,12 @@ fn swap_remove_bulk(table: &mut Table, row_indicies: &Vec) -> Result<(), right_pointer += 1; } } - for _ in 0..right_pointer { - table.pop(); + if is_transaction { + table.set_length(table.len() - right_pointer); + } else { + for _ in 0..right_pointer { + table.pop(); + } } Ok(()) } @@ -67,7 +84,7 @@ mod tests { order_by_clause: None, limit_clause: None, }; - let result = delete(&mut table, statement); + let result = delete(&mut table, statement, false); assert!(result.is_ok()); let expected = vec![ Row(vec![ @@ -158,7 +175,7 @@ mod tests { offset: Some(2), }), }; - let result = delete(&mut table, statement); + let result = delete(&mut table, statement, false); assert!(result.is_ok()); let expected = vec![ Row(vec![ @@ -214,7 +231,7 @@ mod tests { order_by_clause: None, limit_clause: None, }; - let result = delete(&mut table, statement); + let result = delete(&mut table, statement, false); assert!(result.is_ok()); let row_indicies = result.unwrap(); assert_eq!(vec![1, 2, 3], row_indicies); @@ -236,7 +253,7 @@ mod tests { order_by_clause: None, limit_clause: None, }; - let result = delete(&mut table, statement); + let result = delete(&mut table, statement, false); assert!(result.is_ok()); let expected = vec![]; assert_table_rows_eq_unordered(expected, table.get_rows_clone()); @@ -252,7 +269,7 @@ mod tests { order_by_clause: None, limit_clause: None, }; - let result = delete(&mut table, statement); + let result = delete(&mut table, statement, false); assert!(result.is_ok()); } @@ -310,7 +327,7 @@ mod tests { offset: Some(1), }), }; - let result = delete(&mut table, statement); + let result = delete(&mut table, statement, false); assert!(result.is_ok()); let deleted_indices = result.unwrap(); assert_eq!(deleted_indices.len(), 2); @@ -356,7 +373,7 @@ mod tests { order_by_clause: None, limit_clause: None, }; - let result = delete(&mut table, statement); + let result = delete(&mut table, statement, false); assert!(result.is_ok()); let deleted_indices = result.unwrap(); assert_eq!(deleted_indices, vec![0]); diff --git a/src/db/table/operations/drop_table/mod.rs b/src/db/table/operations/drop_table/mod.rs index 1217fa1..351c1af 100644 --- a/src/db/table/operations/drop_table/mod.rs +++ b/src/db/table/operations/drop_table/mod.rs @@ -67,4 +67,19 @@ mod tests { let result = drop_table(&mut database, statement, false); assert!(result.is_ok()); } + + #[test] + fn drop_table_with_transaction_clause_works_correctly() { + let statement = DropTableStatement { + table_name: "users".to_string(), + existence_check: None, + }; + let mut database = default_database(); + let result = drop_table(&mut database, statement, true); + assert!(result.is_ok()); + assert!(!database.has_table("users")); + let table = database.tables.get("users").unwrap(); + assert!(table.first().unwrap().is_some()); + assert!(table.last().unwrap().is_none()); + } } diff --git a/src/db/table/operations/update/mod.rs b/src/db/table/operations/update/mod.rs index 48e3398..6b137c3 100644 --- a/src/db/table/operations/update/mod.rs +++ b/src/db/table/operations/update/mod.rs @@ -2,14 +2,23 @@ use crate::db::table::core::{table::Table, value::DataType}; use crate::db::table::operations::helpers::common::get_row_indicies_matching_clauses; use crate::interpreter::ast::{ColumnValue, UpdateStatement}; -pub fn update(table: &mut Table, statement: UpdateStatement) -> Result, String> { +pub fn update( + table: &mut Table, + statement: UpdateStatement, + is_transaction: bool, +) -> Result, String> { let row_indicies = get_row_indicies_matching_clauses( table, &statement.where_clause, &statement.order_by_clause, &statement.limit_clause, )?; - update_rows_from_indicies(table, &row_indicies, statement.update_values)?; + update_rows_from_indicies( + table, + &row_indicies, + statement.update_values, + is_transaction, + )?; Ok(row_indicies) } @@ -17,6 +26,7 @@ fn update_rows_from_indicies( table: &mut Table, row_indicies: &Vec, update_values: Vec, + is_transaction: bool, ) -> Result<(), String> { for row_index in row_indicies { for update_value in &update_values { @@ -30,6 +40,9 @@ fn update_rows_from_indicies( update_value.value.get_type() )); } + if is_transaction { + table.get_row_stacks_mut()[*row_index].append_clone(); + } table[*row_index][column_index] = update_value.value.clone(); } } @@ -62,7 +75,7 @@ mod tests { order_by_clause: None, limit_clause: None, }; - let result = update(&mut table, statement); + let result = update(&mut table, statement, false); assert!(result.is_ok()); let expected = vec![ Row(vec![ @@ -163,7 +176,7 @@ mod tests { offset: Some(2), }), }; - let result = update(&mut table, statement); + let result = update(&mut table, statement, false); assert!(result.is_ok()); let expected = vec![ Row(vec![ @@ -235,7 +248,7 @@ mod tests { order_by_clause: None, limit_clause: None, }; - let result = update(&mut table, statement); + let result = update(&mut table, statement, false); assert!(result.is_ok()); let row_indicies = result.unwrap(); assert_eq!(vec![1, 2, 3], row_indicies); @@ -289,7 +302,7 @@ mod tests { order_by_clause: None, limit_clause: None, }; - let result = update(&mut table, statement); + let result = update(&mut table, statement, false); assert!(result.is_ok()); assert_eq!(result.unwrap(), vec![]); let expected = vec![]; @@ -309,7 +322,7 @@ mod tests { order_by_clause: None, limit_clause: None, }; - let result = update(&mut table, statement); + let result = update(&mut table, statement, false); assert!(result.is_err()); assert_eq!( result.err().unwrap(), @@ -330,7 +343,7 @@ mod tests { order_by_clause: None, limit_clause: None, }; - let result = update(&mut table, statement); + let result = update(&mut table, statement, false); assert!(result.is_err()); assert_eq!( result.err().unwrap(), @@ -351,7 +364,7 @@ mod tests { order_by_clause: None, limit_clause: None, }; - let result = update(&mut table, statement); + let result = update(&mut table, statement, false); assert!(result.is_ok()); assert_eq!(result.unwrap(), vec![0, 1, 2, 3]); let expected = vec![ @@ -382,4 +395,55 @@ mod tests { ]; assert_table_rows_eq_unordered(expected, table.get_rows_clone()); } + + #[test] + fn update_with_transaction_works_correctly() { + let mut table = default_table(); + let statement = UpdateStatement { + table_name: "users".to_string(), + update_values: vec![ColumnValue { + column: "name".to_string(), + value: Value::Text("Fletcher".to_string()), + }], + where_clause: None, + order_by_clause: None, + limit_clause: None, + }; + let result = update(&mut table, statement, true); + assert!(result.is_ok()); + assert_eq!(result.unwrap(), vec![0, 1, 2, 3]); + let expected = vec![ + Row(vec![ + Value::Integer(1), + Value::Text("Fletcher".to_string()), + Value::Integer(25), + Value::Real(1000.0), + ]), + Row(vec![ + Value::Integer(2), + Value::Text("Fletcher".to_string()), + Value::Integer(30), + Value::Real(2000.0), + ]), + Row(vec![ + Value::Integer(3), + Value::Text("Fletcher".to_string()), + Value::Integer(35), + Value::Real(3000.0), + ]), + Row(vec![ + Value::Integer(4), + Value::Text("Fletcher".to_string()), + Value::Integer(40), + Value::Real(4000.0), + ]), + ]; + assert_table_rows_eq_unordered(expected, table.get_rows_clone()); + assert!( + table + .get_row_stacks_mut() + .iter() + .all(|row_stack| row_stack.stack.len() == 2) + ); + } } diff --git a/src/db/transactions/commit.rs b/src/db/transactions/commit.rs new file mode 100644 index 0000000..6754e36 --- /dev/null +++ b/src/db/transactions/commit.rs @@ -0,0 +1,16 @@ +use crate::db::database::Database; +use crate::db::transactions::TransactionEntry; + +pub fn commit_transaction(database: &mut Database) -> Result<(), String> { + let transaction_log = database.transaction.commit_transaction()?; + for transaction_entry in transaction_log.get_entries()?.iter() { + match transaction_entry { + TransactionEntry::Statement(statement) => { + let table = database.get_table_mut(&statement.table_name)?; + table.commit_transaction(&statement.affected_rows)?; + } + TransactionEntry::Savepoint(_) => {} + } + } + Ok(()) +} diff --git a/src/db/transactions/mod.rs b/src/db/transactions/mod.rs index c610d2a..d41e9c8 100644 --- a/src/db/transactions/mod.rs +++ b/src/db/transactions/mod.rs @@ -1,4 +1,5 @@ use crate::interpreter::ast::SqlStatement; +pub mod commit; pub mod rollback; #[derive(Debug, PartialEq, Clone)] @@ -75,6 +76,13 @@ impl TransactionLog { Ok(()) } + pub fn savepoint_exists(&self, savepoint_name: &String) -> Result { + Ok(self.get_entries()?.iter().any(|entry| match entry { + TransactionEntry::Savepoint(savepoint) => savepoint.name == *savepoint_name, + _ => false, + })) + } + pub fn begin_transaction(&mut self) { self.entries = Some(vec![]); } @@ -87,6 +95,10 @@ impl TransactionLog { Ok(transaction_log) } + pub fn pop_entry(&mut self) -> Result, String> { + Ok(self.get_entries_mut()?.pop()) + } + pub fn get_entries(&self) -> Result<&Vec, String> { self.entries .as_ref() diff --git a/src/db/transactions/rollback.rs b/src/db/transactions/rollback.rs index 6090ab3..92ba79a 100644 --- a/src/db/transactions/rollback.rs +++ b/src/db/transactions/rollback.rs @@ -1,24 +1,74 @@ use crate::db::database::Database; -use crate::db::transactions::StatementEntry; -use crate::interpreter::ast::{AlterTableAction, SqlStatement}; +use crate::db::transactions::{StatementEntry, TransactionEntry}; +use crate::interpreter::ast::{AlterTableAction, RollbackStatement, SqlStatement}; + +pub fn rollback_statement( + database: &mut Database, + statement: &RollbackStatement, +) -> Result<(), String> { + if !database.transaction.in_transaction() { + return Err("No transaction is currently active".to_string()); + } + + if let Some(savepoint_name) = &statement.savepoint_name { + // First make sure the savepoint exists + if !database.transaction.savepoint_exists(savepoint_name)? { + return Err(format!("Savepoint `{}` does not exist", savepoint_name)); + } + // Rollback to savepoint - keep transaction active + let mut current_entry = database.transaction.pop_entry()?; + while current_entry.is_some() { + match current_entry.unwrap() { + TransactionEntry::Statement(transaction_statement) => { + rollback_transaction_entry(database, &transaction_statement)?; + } + TransactionEntry::Savepoint(savepoint_statement) => { + if savepoint_statement.name == *savepoint_name { + break; + } + } + } + current_entry = database.transaction.pop_entry()?; + } + } else { + // Full rollback - commit transaction to get entries and clear state + if let Some(transaction_log) = database.transaction.commit_transaction()?.entries { + // COMMIT TRANSACTIONS CLEARS THIS WITH TAKE + for transaction_entry in transaction_log.iter().rev() { + match transaction_entry { + TransactionEntry::Statement(statement) => { + // TODO: Some matching needs to be here for table based operations. + // CURRENTLY SUPPORTED STATEMENTS ARE: + // - ALTER TABLE RENAME COLUMN, ALTER TABLE ADD COLUMN, ALTER TABLE DROP COLUMN, ALTER TABLE RENAME TABLE + // - CREATE TABLE, DROP TABLE + // - INSERT INTO, UPDATE, DELETE + rollback_transaction_entry(database, &statement)?; + } + TransactionEntry::Savepoint(_) => {} + } + } + } + } + Ok(()) +} pub fn rollback_transaction_entry( database: &mut Database, - statement: &StatementEntry, + statement_entry: &StatementEntry, ) -> Result<(), String> { - match &statement.statement { + match &statement_entry.statement { SqlStatement::AlterTable(alter_table) => match alter_table.action { AlterTableAction::RenameColumn { .. } => { - let table = database.get_table_mut(&statement.table_name)?; + let table = database.get_table_mut(&statement_entry.table_name)?; table.rollback_columns(); } AlterTableAction::AddColumn { .. } => { - let table = database.get_table_mut(&statement.table_name)?; + let table = database.get_table_mut(&statement_entry.table_name)?; table.rollback_columns(); table.rollback_all_rows(); } AlterTableAction::DropColumn { .. } => { - let table = database.get_table_mut(&statement.table_name)?; + let table = database.get_table_mut(&statement_entry.table_name)?; table.rollback_columns(); table.rollback_all_rows(); } @@ -26,20 +76,40 @@ pub fn rollback_transaction_entry( // It is now under the new name let mut table = database.pop_table_change(&new_table_name)?; table.rollback_name(); - database.push_table_change(&statement.table_name, table); + database.push_table_change(&statement_entry.table_name, table); } }, SqlStatement::Select(_) => {} // These should be kept in the log but obv do nothing. SqlStatement::CreateTable(_) => { - database.tables.remove(&statement.table_name); - } - SqlStatement::DropTable(statement) => { database .tables - .get_mut(&statement.table_name) + .get_mut(&statement_entry.table_name) .unwrap() .pop(); } + SqlStatement::DropTable(_) => { + // For drop table rollback, we need to pop the None that was pushed during the drop + if let Some(table_versions) = database.tables.get_mut(&statement_entry.table_name) { + table_versions.pop(); + } + } + SqlStatement::InsertInto(_) => { + let table = database.get_table_mut(&statement_entry.table_name)?; + for _ in &statement_entry.affected_rows { + table.get_row_stacks_mut().pop(); // We can pop all the rows off because they always get pushed to the end + } + table.set_length(table.len() - statement_entry.affected_rows.len()); + } + SqlStatement::UpdateStatement(_) => { + let table = database.get_table_mut(&statement_entry.table_name)?; + for index in &statement_entry.affected_rows { + table.get_row_stacks_mut()[*index].stack.pop(); + } + } + SqlStatement::DeleteStatement(_) => { + let table = database.get_table_mut(&statement_entry.table_name)?; + table.set_length(table.len() + statement_entry.affected_rows.len()); + } _ => return Err("UNSUPPORTED".to_string()), } return Ok(()); diff --git a/tests/transaction_test.rs b/tests/transaction_test.rs index bf042cb..7be1a9b 100644 --- a/tests/transaction_test.rs +++ b/tests/transaction_test.rs @@ -101,6 +101,13 @@ fn test_transaction_drop_table() { SELECT * FROM users; DROP TABLE users; SELECT * FROM users; + CREATE TABLE users ( + id INTEGER, + name TEXT + ); + SELECT * FROM users; + DROP TABLE users; + SELECT * FROM users; ROLLBACK; SELECT * FROM users; "; @@ -113,9 +120,360 @@ fn test_transaction_drop_table() { 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![])), + Ok(None), + Err("Execution Error with statement starting on line 17 \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); } } + +#[test] +fn test_transaction_insert_into() { + let mut database = Database::new(); + let sql = " + CREATE TABLE users ( + id INTEGER, + name TEXT + ); + INSERT INTO users (id, name) VALUES (1, 'John'); + SELECT * FROM users; + BEGIN; + INSERT INTO users (id, name) VALUES (2, 'Jane'); + SELECT * FROM users; + ROLLBACK; + SELECT * FROM users; + "; + let result = run_sql(&mut database, sql); + let expected = vec![ + Ok(None), + Ok(None), + Ok(Some(vec![Row(vec![ + Value::Integer(1), + Value::Text("John".to_string()), + ])])), + Ok(None), + Ok(None), + Ok(Some(vec![ + Row(vec![Value::Integer(1), Value::Text("John".to_string())]), + Row(vec![Value::Integer(2), Value::Text("Jane".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); + } +} + +#[test] +fn test_transaction_update() { + let mut database = Database::new(); + let sql = " + CREATE TABLE users ( + id INTEGER, + name TEXT + ); + INSERT INTO users (id, name) VALUES (1, 'John'); + SELECT * FROM users; + BEGIN; + UPDATE users SET name = 'Jane' WHERE id = 1; + SELECT * FROM users; + ROLLBACK; + SELECT * FROM users; + "; + let result = run_sql(&mut database, sql); + let expected = vec![ + Ok(None), + Ok(None), + Ok(Some(vec![Row(vec![ + Value::Integer(1), + Value::Text("John".to_string()), + ])])), + Ok(None), + Ok(None), + Ok(Some(vec![Row(vec![ + Value::Integer(1), + Value::Text("Jane".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); + } +} + +#[test] +fn test_transaction_delete() { + let mut database = Database::new(); + let sql = " + CREATE TABLE users ( + id INTEGER, + name TEXT + ); + INSERT INTO users (id, name) VALUES (1, 'John'), (2, 'Jane'), (3, 'Jim'), (4, 'Jill'); + SELECT * FROM users; + BEGIN; + DELETE FROM users WHERE id = 1; + SELECT * FROM users; + DELETE FROM users WHERE id >= 3; + SELECT * FROM users; + ROLLBACK; + SELECT * FROM users; + "; + let result = run_sql(&mut database, sql); + let expected = vec![ + Ok(None), + Ok(None), + Ok(Some(vec![ + Row(vec![Value::Integer(1), Value::Text("John".to_string())]), + Row(vec![Value::Integer(2), Value::Text("Jane".to_string())]), + Row(vec![Value::Integer(3), Value::Text("Jim".to_string())]), + Row(vec![Value::Integer(4), Value::Text("Jill".to_string())]), + ])), + Ok(None), + Ok(None), + Ok(Some(vec![ + Row(vec![Value::Integer(2), Value::Text("Jane".to_string())]), + Row(vec![Value::Integer(3), Value::Text("Jim".to_string())]), + Row(vec![Value::Integer(4), Value::Text("Jill".to_string())]), + ])), + Ok(None), + Ok(Some(vec![Row(vec![ + Value::Integer(2), + Value::Text("Jane".to_string()), + ])])), + Ok(None), + Ok(Some(vec![ + Row(vec![Value::Integer(1), Value::Text("John".to_string())]), + Row(vec![Value::Integer(2), Value::Text("Jane".to_string())]), + Row(vec![Value::Integer(3), Value::Text("Jim".to_string())]), + Row(vec![Value::Integer(4), Value::Text("Jill".to_string())]), + ])), + ]; + for (i, result) in result.iter().enumerate() { + match (&expected[i], result) { + (Ok(None), Ok(None)) => { + assert_eq!(expected[i], *result); + } + (Ok(Some(expected_rows)), Ok(Some(actual_rows))) => { + test_utils::assert_eq_table_rows_unordered( + expected_rows.clone(), + actual_rows.clone(), + ); + } + (_, _) => { + assert_eq!(expected[i], *result); + } + } + } +} + +#[test] +fn test_transaction_savepoint() { + let mut database = Database::new(); + let sql = " + CREATE TABLE users ( + id INTEGER, + name TEXT + ); + BEGIN; + SAVEPOINT savepoint_name; + INSERT INTO users (id, name) VALUES (1, 'John'); + SELECT * FROM users; + ROLLBACK TO SAVEPOINT savepoint_name; + SELECT * FROM users; + ROLLBACK; + SELECT * FROM users; + "; + let result = run_sql(&mut database, sql); + let expected = vec![ + Ok(None), + Ok(None), + Ok(None), + Ok(None), + Ok(Some(vec![Row(vec![ + Value::Integer(1), + Value::Text("John".to_string()), + ])])), + Ok(None), + Ok(Some(vec![])), + Ok(None), + Ok(Some(vec![])), + ]; + for (i, result) in result.iter().enumerate() { + assert_eq!(expected[i], *result); + } +} + +#[test] +fn test_transaction_rollback_to_savepoint_that_does_not_exist() { + let mut database = Database::new(); + let sql = " + BEGIN; + SAVEPOINT savepoint_name; + RELEASE SAVEPOINT savepoint_name; + ROLLBACK TO SAVEPOINT savepoint_name; + ROLLBACK; + "; + let result = run_sql(&mut database, sql); + let expected = vec![ + Ok(None), + Ok(None), + Ok(None), + Err("Execution Error with statement starting on line 5 \n Error: Savepoint `savepoint_name` does not exist".to_string()), + Ok(None), + ]; + for (i, result) in result.iter().enumerate() { + assert_eq!(expected[i], *result); + } +} + +#[test] +fn test_transaction_commit_with_savepoint() { + let mut database = Database::new(); + let sql = " + CREATE TABLE users ( + id INTEGER, + name TEXT + ); + BEGIN; + SAVEPOINT savepoint_name; + INSERT INTO users (id, name) VALUES (1, 'John'); + SELECT * FROM users; + ROLLBACK TO SAVEPOINT savepoint_name; + INSERT INTO users (id, name) VALUES (2, 'Jane'); + SELECT * FROM users; + COMMIT; + SELECT * FROM users; + "; + let result = run_sql(&mut database, sql); + let expected = vec![ + Ok(None), + Ok(None), + Ok(None), + Ok(None), + Ok(Some(vec![Row(vec![ + Value::Integer(1), + Value::Text("John".to_string()), + ])])), + Ok(None), + Ok(None), + Ok(Some(vec![Row(vec![ + Value::Integer(2), + Value::Text("Jane".to_string()), + ])])), + Ok(None), + Ok(Some(vec![Row(vec![ + Value::Integer(2), + Value::Text("Jane".to_string()), + ])])), + ]; + for (i, result) in result.iter().enumerate() { + assert_eq!(expected[i], *result); + } +} + +#[test] +fn test_transaction_commit_with_many_changes() { + let mut database = Database::new(); + let sql = " + BEGIN; + CREATE TABLE users ( + id INTEGER, + name TEXT + ); + INSERT INTO users (id, name) VALUES (1, 'John'); + SAVEPOINT savepoint_name; + DROP TABLE users; + SELECT * FROM users; + ROLLBACK TO SAVEPOINT savepoint_name; + SELECT * FROM users; + ALTER TABLE users ADD COLUMN age INTEGER; + ALTER TABLE users RENAME COLUMN age TO new_age; + SELECT * FROM users; + ALTER TABLE users RENAME TO new_users; + COMMIT; + SELECT * FROM new_users; + "; + let result = run_sql(&mut database, sql); + + let expected = vec![ + Ok(None), + Ok(None), + Ok(None), + Ok(None), + 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())])])), + Ok(None), + Ok(None), + Ok(Some(vec![Row(vec![Value::Integer(1), Value::Text("John".to_string()), Value::Null])])), + Ok(None), + Err("Execution Error with statement starting on line 17 \n Error: Table `users` does not exist".to_string()), + Ok(Some(vec![Row(vec![Value::Integer(1), Value::Text("John".to_string()), Value::Null])])), + ]; + + for (i, result_item) in result.iter().enumerate() { + assert_eq!(expected[i], *result_item); + } +} + +#[test] +fn test_transaction_with_multiple_savepoints() { + let mut database = Database::new(); + let sql = " + CREATE TABLE users ( + id INTEGER, + name TEXT + ); + BEGIN; + SAVEPOINT savepoint1; + INSERT INTO users (id, name) VALUES (1, 'John'); + SAVEPOINT savepoint2; + INSERT INTO users (id, name) VALUES (2, 'Jane'); + SELECT * FROM users; + ROLLBACK TO SAVEPOINT savepoint2; + SELECT * FROM users; + ROLLBACK TO SAVEPOINT savepoint1; + SELECT * FROM users; + COMMIT; + SELECT * FROM users; + "; + let result = run_sql(&mut database, sql); + let expected = vec![ + Ok(None), + Ok(None), + Ok(None), + Ok(None), + Ok(None), + Ok(None), + Ok(Some(vec![ + Row(vec![Value::Integer(1), Value::Text("John".to_string())]), + Row(vec![Value::Integer(2), Value::Text("Jane".to_string())]), + ])), + Ok(None), + Ok(Some(vec![Row(vec![ + Value::Integer(1), + Value::Text("John".to_string()), + ])])), + Ok(None), + Ok(Some(vec![])), + Ok(None), + Ok(Some(vec![])), + ]; + for (i, result) in result.iter().enumerate() { + assert_eq!(expected[i], *result); + } +}