diff --git a/src/cli/ast/insert_statement.rs b/src/cli/ast/insert_statement.rs index e4edebd..76ebb6c 100644 --- a/src/cli/ast/insert_statement.rs +++ b/src/cli/ast/insert_statement.rs @@ -53,11 +53,28 @@ fn into_statement(parser: &mut Parser) -> Result { } } - return Ok(InsertInto(InsertIntoStatement { + let statement = InsertIntoStatement { table_name: table_name, columns: columns, values: values, - })); + }; + validate_insert_statement(&statement)?; + return Ok(InsertInto(statement)); +} + +fn validate_insert_statement(statement: &InsertIntoStatement) -> Result<(), String> { + for row in &statement.values { + if row.len() != statement.values[0].len() { + return Err(format!("Rows have different lengths")); + } + } + + if let Some(columns) = &statement.columns { + if columns.len() != statement.values[0].len() { + return Err(format!("Columns and values have different lengths")); + } + } + return Ok(()); } fn get_values(parser: &mut Parser) -> Result, String> { @@ -279,4 +296,63 @@ mod tests { let result = build(&mut parser); assert!(result.is_err()); } + + #[test] + fn insert_with_different_lengths_is_error() { + // INSERT INTO users VALUES (1, "Alice"), (2, "Bob", "Charlie"); + let tokens = vec![ + token(TokenTypes::Insert, "INSERT"), + token(TokenTypes::Into, "INTO"), + token(TokenTypes::Identifier, "users"), + token(TokenTypes::Values, "VALUES"), + token(TokenTypes::LeftParen, "("), + token(TokenTypes::IntLiteral, "1"), + token(TokenTypes::Comma, ","), + token(TokenTypes::String, "Alice"), + token(TokenTypes::RightParen, ")"), + token(TokenTypes::Comma, ","), + token(TokenTypes::LeftParen, "("), + token(TokenTypes::IntLiteral, "2"), + token(TokenTypes::Comma, ","), + token(TokenTypes::String, "Bob"), + token(TokenTypes::Comma, ","), + token(TokenTypes::String, "Charlie"), + token(TokenTypes::RightParen, ")"), + token(TokenTypes::SemiColon, ";"), + ]; + let mut parser = Parser::new(tokens); + let result = build(&mut parser); + assert!(result.is_err()); + let expected = Err("Rows have different lengths".to_string()); + assert_eq!(expected, result); + } + + #[test] + fn insert_with_different_column_and_value_lengths_is_error() { + // INSERT INTO users (id, name) VALUES (1, "Alice", "Bob"); + let tokens = vec![ + token(TokenTypes::Insert, "INSERT"), + token(TokenTypes::Into, "INTO"), + token(TokenTypes::Identifier, "users"), + token(TokenTypes::LeftParen, "("), + token(TokenTypes::Identifier, "id"), + token(TokenTypes::Comma, ","), + token(TokenTypes::Identifier, "name"), + token(TokenTypes::RightParen, ")"), + token(TokenTypes::Values, "VALUES"), + token(TokenTypes::LeftParen, "("), + token(TokenTypes::IntLiteral, "1"), + token(TokenTypes::Comma, ","), + token(TokenTypes::String, "Alice"), + token(TokenTypes::Comma, ","), + token(TokenTypes::String, "Bob"), + token(TokenTypes::RightParen, ")"), + token(TokenTypes::SemiColon, ";"), + ]; + let mut parser = Parser::new(tokens); + let result = build(&mut parser); + assert!(result.is_err()); + let expected = Err("Columns and values have different lengths".to_string()); + assert_eq!(expected, result); + } } \ No newline at end of file diff --git a/src/cli/ast/mod.rs b/src/cli/ast/mod.rs index bfacf90..340a386 100644 --- a/src/cli/ast/mod.rs +++ b/src/cli/ast/mod.rs @@ -42,6 +42,15 @@ pub enum SelectStatementColumns { Specific(Vec), } +impl SelectStatementColumns { + pub fn columns(&self) -> Result<&Vec, String> { + return match self { + SelectStatementColumns::All => Err("Cannot get columns from all columns".to_string()), + SelectStatementColumns::Specific(columns) => Ok(columns), + } + } +} + #[derive(Debug, PartialEq)] pub enum Operator { Equals, diff --git a/src/db/database.rs b/src/db/database.rs index f13e221..89064f3 100644 --- a/src/db/database.rs +++ b/src/db/database.rs @@ -1,5 +1,7 @@ use crate::db::table::{Table, Value}; use crate::cli::ast::{SqlStatement, CreateTableStatement, InsertIntoStatement, SelectStatement}; +use crate::db::table::select; +use crate::db::table::insert; use std::collections::HashMap; pub struct Database { @@ -34,20 +36,20 @@ impl Database { if self.has_table(&statement.table_name) { return Err(format!("Table {} already exists", statement.table_name)); } - let table_name = statement.table_name; - self.tables.insert(table_name.clone(), Table::new(table_name, statement.columns)); + let table = Table::new(statement.table_name, statement.columns) ; + self.tables.insert(table.name.clone(), table); Ok(()) } fn insert_into_table(&mut self, statement: InsertIntoStatement) -> Result<(), String> { let table = self.get_table_mut(&statement.table_name)?; - table.insert(statement)?; + insert::insert(table, statement)?; Ok(()) } fn select_from_table(&mut self, statement: SelectStatement) -> Result>, String> { let table = self.get_table(&statement.table_name)?; - let rows = table.select(statement)?; + let rows = select::select(table, statement)?; Ok(rows) } @@ -68,4 +70,72 @@ impl Database { } Ok(self.tables.get_mut(table_name).unwrap()) } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::cli::ast::CreateTableStatement; + use crate::db::table::{ColumnDefinition, DataType}; + + + fn default_database() -> Database { + Database { + tables: HashMap::from([ + ("users".to_string(), Table::new("users".to_string(), vec![ + ColumnDefinition { + name: "id".to_string(), + data_type: DataType::Integer, + constraints: vec![] + }, + ColumnDefinition { + name: "name".to_string(), + data_type: DataType::Text, + constraints: vec![] + }, + ])) + ]) + } + } + + #[test] + fn create_table_generates_proper_table() { + let statement = CreateTableStatement { + table_name: "users".to_string(), + columns: vec![ + ColumnDefinition { + name: "id".to_string(), + data_type: DataType::Integer, + constraints: vec![] + }, + ], + }; + let mut database = Database::new(); + assert!(database.create_table(statement).is_ok()); + assert!(database.has_table("users")); + } + + #[test] + fn has_table_returns_proper_response() { + let database = default_database(); + assert!(database.has_table("users")); + assert!(!database.has_table("not_users")); + } + + #[test] + fn get_table_funcs_returns_proper_table() { + let mut database = default_database(); + let table = database.get_table("users"); + assert!(table.is_ok()); + assert_eq!(table.unwrap().name, "users"); + let table = database.get_table("not_users"); + assert!(table.is_err()); + assert_eq!(table.unwrap_err(), "Table not_users does not exist"); + let table = database.get_table_mut("users"); + assert!(table.is_ok()); + assert_eq!(table.unwrap().name, "users"); + let table = database.get_table_mut("not_users"); + assert!(table.is_err()); + assert_eq!(table.unwrap_err(), "Table not_users does not exist"); + } } \ No newline at end of file diff --git a/src/db/mod.rs b/src/db/mod.rs index b66a9b6..b0f1354 100644 --- a/src/db/mod.rs +++ b/src/db/mod.rs @@ -1,2 +1,2 @@ pub mod database; -pub mod table; +pub mod table; \ No newline at end of file diff --git a/src/db/table.rs b/src/db/table.rs deleted file mode 100644 index f7250df..0000000 --- a/src/db/table.rs +++ /dev/null @@ -1,202 +0,0 @@ -use crate::cli::ast::{InsertIntoStatement, Operator, SelectStatement, SelectStatementColumns, WhereClause}; - -#[derive(Debug, PartialEq)] -pub enum DataType { - Integer, - Real, - Text, - Blob, - Null, -} - - - -#[derive(Debug, PartialEq)] -pub struct ColumnDefinition { - pub name: String, - pub data_type: DataType, - pub constraints: Vec, -} - -#[derive(Debug, PartialEq)] -pub struct ColumnConstraint { - pub constraint_type: String, -} - - -#[derive(Debug, PartialEq, PartialOrd)] -pub enum Value { - Integer(i64), - Real(f64), - Text(String), - Blob(Vec), - Null -} - -impl Value { - pub fn get_type(&self) -> DataType { - match self { - Value::Integer(_) => DataType::Integer, - Value::Real(_) => DataType::Real, - Value::Text(_) => DataType::Text, - Value::Blob(_) => DataType::Blob, - Value::Null => DataType::Null, - } - } - - pub fn clone(&self) -> Value { - match self { - Value::Integer(value) => Value::Integer(*value), - Value::Real(value) => Value::Real(*value), - Value::Text(value) => Value::Text(value.clone()), - Value::Blob(value) => Value::Blob(value.clone()), - Value::Null => Value::Null, - } - } -} - -pub struct Table { - _name: String, - columns: Vec, - rows: Vec>, -} - -impl Table { - pub fn new(_name: String, columns: Vec) -> Self { - Self { - _name, - columns, - rows: vec![], - } - } - - pub fn insert(&mut self, statement: InsertIntoStatement) -> Result<(), String> { - // Validate columns - if let Some(columns) = statement.columns { - if columns.len() != self.columns.len() { - return Err(format!("Columns have incorrect width")); - } - for (i, column) in columns.iter().enumerate() { - if column != &self.columns[i].name { - return Err(format!("Column mismatch")); - } - } - } - - let mut rows: Vec> = vec![]; - // Validate row inserts - for row in statement.values { - if row.len() != self.width() { - return Err(format!("Rows have incorrect width")); - } - let row_values = self.validate_and_clone_row(&row)?; - rows.push(row_values); - } - - // Insert rows - for row in rows { - self.rows.push(row); - } - return Ok(()); - } - - pub fn select(&self, statement: SelectStatement) -> Result>, String> { - let mut rows: Vec> = vec![]; - if let Some(where_clause) = statement.where_clause { - for row in self.rows.iter() { - if self.matches_where_clause(&row, &where_clause) { - rows.push(self.get_columns_from_row(&row, &statement.columns)?); - } - } - } else { - for row in self.rows.iter() { - rows.push(self.get_columns_from_row(&row, &statement.columns)?); - } - } - return Ok(rows); - } - - fn matches_where_clause(&self, row: &Vec, where_clause: &WhereClause) -> bool { - let column_value = self.get_column_from_row(row, &where_clause.column); - if column_value.get_type() != where_clause.value.get_type() { - return false; - } - - match where_clause.operator { - Operator::Equals => { - return *column_value == where_clause.value; - }, - Operator::NotEquals => { - return *column_value != where_clause.value; - }, - _ => { - match column_value.get_type() { - DataType::Integer | DataType::Real => { - match where_clause.operator { - Operator::LessThan => { - return *column_value < where_clause.value; - }, - Operator::GreaterThan => { - return *column_value > where_clause.value; - }, - Operator::LessEquals => { - return *column_value <= where_clause.value; - }, - Operator::GreaterEquals => { - return *column_value >= where_clause.value; - }, - _ => { - return false; - }, - } - }, - _ => { - return false; - }, - } - } - } - } - - fn get_column_from_row<'a>(&self, row: &'a Vec, column: &String) -> &'a Value { - for (i, value) in row.iter().enumerate() { - if self.columns[i].name == *column { - return &value; - } - } - return &Value::Null; - } - - fn get_columns_from_row(&self, row: &Vec, columns: &SelectStatementColumns) -> Result, String> { - let mut row_values: Vec = vec![]; - if *columns == SelectStatementColumns::All { - return Ok(self.validate_and_clone_row(row)?); - } else { - for (i, column) in self.columns.iter().enumerate() { - if self.columns.contains(column) { - row_values.push(row[i].clone()); - } - } - } - return Ok(row_values); - } - - fn width(&self) -> usize { - self.columns.len() - } - - fn validate_and_clone_row(&self, row: &Vec) -> Result, String> { - if row.len() != self.width() { - return Err(format!("Rows have incorrect width")); - } - - let mut row_values: Vec = vec![]; - for (i, value) in row.iter().enumerate() { - if value.get_type() != self.columns[i].data_type && value.get_type() != DataType::Null { - return Err(format!("Data type mismatch for column {}", self.columns[i].name)); - } - row_values.push(row[i].clone()); - } - return Ok(row_values); - } -} \ No newline at end of file diff --git a/src/db/table/common.rs b/src/db/table/common.rs new file mode 100644 index 0000000..b379a42 --- /dev/null +++ b/src/db/table/common.rs @@ -0,0 +1,16 @@ +use crate::db::table::{Table, Value, DataType}; + +pub fn validate_and_clone_row(table: &Table, row: &Vec) -> Result, String> { + if row.len() != table.width() { + return Err(format!("Rows have incorrect width")); + } + + let mut row_values: Vec = vec![]; + for (i, value) in row.iter().enumerate() { + if value.get_type() != table.columns[i].data_type && value.get_type() != DataType::Null { + return Err(format!("Data type mismatch for column {}", table.columns[i].name)); + } + row_values.push(row[i].clone()); + } + return Ok(row_values); +} \ No newline at end of file diff --git a/src/db/table/insert/mod.rs b/src/db/table/insert/mod.rs new file mode 100644 index 0000000..49b0c50 --- /dev/null +++ b/src/db/table/insert/mod.rs @@ -0,0 +1,110 @@ +use std::collections::{HashMap, VecDeque}; + +use crate::db::table::{Table, Value}; +use crate::cli::ast::InsertIntoStatement; +use crate::db::table::common::validate_and_clone_row; + + +pub fn insert(table: &mut Table, statement: InsertIntoStatement) -> Result<(), String> { + // Validate columns + if let Some(columns) = &statement.columns { + for column in columns { + if table.columns.iter().find(|c| c.name == *column).is_none() { + return Err(format!("Column '{}' does not exist in table", column)); + } + } + } + + let mut rows: Vec> = vec![]; + // Creates a hash map from the statement values with the columns as the keys + // The values are stored in a queue to match the order of the columns, we push back to the queue + // and then pop off the front when creating the rows. + // Todo: make this logic simpler. + if let Some(statement_columns) = &statement.columns { + let mut map: HashMap<&String, VecDeque> = HashMap::new(); + for (i, column) in statement_columns.iter().enumerate() { + map.insert(column, VecDeque::new()); + for row in statement.values.iter() { + map.get_mut(column).unwrap().push_back(row[i].clone()); + } + } + for _ in 0..statement.values.len() { + let mut row: Vec = vec![]; + for table_column in table.columns.iter() { + if map.contains_key(&table_column.name) { + let queue = map.get_mut(&table_column.name).unwrap(); + let value = queue.pop_front().unwrap(); + row.push(value); + } + else { + row.push(Value::Null); + } + } + rows.push(row); + } + } else { + // Inserts entire row in the order provided in the statement + for row in statement.values { + let row_values = validate_and_clone_row(table, &row)?; + rows.push(row_values); + } + } + + // Insert rows + for row in rows { + table.rows.push(row); + } + return Ok(()); +} + + + +#[cfg(test)] +mod tests { + use super::*; + use crate::db::table::{Table, Value, DataType, ColumnDefinition}; + + fn default_table() -> Table { + Table::new( + "users".to_string(), + vec![ + ColumnDefinition {name: "id".to_string(), data_type: DataType::Integer, constraints: vec![]}, + ColumnDefinition {name: "name".to_string(), data_type: DataType::Text, constraints: vec![]}, + ColumnDefinition {name: "age".to_string(), data_type: DataType::Integer, constraints: vec![]}, + ColumnDefinition {name: "money".to_string(), data_type: DataType::Real, constraints: vec![]}, + ] + ) + } + + #[test] + fn insert_into_table_is_generated_correctly() { + let mut table = default_table(); + let statement = InsertIntoStatement { + table_name: "users".to_string(), + columns: None, + values: vec![vec![Value::Integer(1), Value::Text("John".to_string()), Value::Integer(25), Value::Real(1000.0)]], + }; + assert!(insert(&mut table, statement).is_ok()); + let expected = vec![vec![Value::Integer(1), Value::Text("John".to_string()), Value::Integer(25), Value::Real(1000.0)]]; + assert_eq!(table.rows, expected); + } + + #[test] + fn insert_into_table_with_columns_is_generated_correctly() { + let mut table = default_table(); + let statement = InsertIntoStatement { + table_name: "users".to_string(), + columns: Some(vec!["id".to_string(), "name".to_string()]), + values: vec![ + vec![Value::Integer(1), Value::Text("John".to_string())], + vec![Value::Integer(2), Value::Text("Jane".to_string())], + ], + }; + assert!(insert(&mut table, statement).is_ok()); + let expected = vec![ + vec![Value::Integer(1), Value::Text("John".to_string()), Value::Null, Value::Null], + vec![Value::Integer(2), Value::Text("Jane".to_string()), Value::Null, Value::Null], + ]; + assert_eq!(table.rows, expected); + } +} \ No newline at end of file diff --git a/src/db/table/mod.rs b/src/db/table/mod.rs new file mode 100644 index 0000000..8f372d8 --- /dev/null +++ b/src/db/table/mod.rs @@ -0,0 +1,85 @@ +pub mod select; +pub mod insert; +pub mod common; + +#[derive(Debug, PartialEq)] +pub enum DataType { + Integer, + Real, + Text, + Blob, + Null, +} + +#[derive(Debug, PartialEq)] +pub struct ColumnDefinition { + pub name: String, + pub data_type: DataType, + pub constraints: Vec, +} + +#[derive(Debug, PartialEq)] +pub struct ColumnConstraint { + pub constraint_type: String, +} + +#[derive(Debug, PartialEq, PartialOrd)] +pub enum Value { + Integer(i64), + Real(f64), + Text(String), + Blob(Vec), + Null +} + +impl Value { + pub fn get_type(&self) -> DataType { + match self { + Value::Integer(_) => DataType::Integer, + Value::Real(_) => DataType::Real, + Value::Text(_) => DataType::Text, + Value::Blob(_) => DataType::Blob, + Value::Null => DataType::Null, + } + } + + pub fn clone(&self) -> Value { + match self { + Value::Integer(value) => Value::Integer(*value), + Value::Real(value) => Value::Real(*value), + Value::Text(value) => Value::Text(value.clone()), + Value::Blob(value) => Value::Blob(value.clone()), + Value::Null => Value::Null, + } + } +} + +#[derive(Debug)] +pub struct Table { + pub name: String, + pub columns: Vec, + pub rows: Vec>, +} + +impl Table { + pub fn new(name: String, columns: Vec) -> Self { + Self { + name, + columns, + rows: vec![], + } + } + + pub fn get_column_from_row<'a>(&self, row: &'a Vec, column: &String) -> &'a Value { + for (i, value) in row.iter().enumerate() { + if self.columns[i].name == *column { + return &value; + } + } + return &Value::Null; + } + + fn width(&self) -> usize { + self.columns.len() + } +} \ No newline at end of file diff --git a/src/db/table/select/mod.rs b/src/db/table/select/mod.rs new file mode 100644 index 0000000..7a88682 --- /dev/null +++ b/src/db/table/select/mod.rs @@ -0,0 +1,150 @@ +pub mod where_clause; +use crate::db::table::{Table, Value}; +use crate::cli::ast::SelectStatement; +use crate::cli::ast::SelectStatementColumns; +use crate::db::table::common::validate_and_clone_row; + + +pub fn select(table: &Table, statement: SelectStatement) -> Result>, String> { + let mut rows: Vec> = vec![]; + if let Some(where_clause) = statement.where_clause { + for row in table.rows.iter() { + if where_clause::matches_where_clause(table, &row, &where_clause) { + rows.push(get_columns_from_row(table, &row, &statement.columns)?); + } + } + } else { + for row in table.rows.iter() { + rows.push(get_columns_from_row(table, &row, &statement.columns)?); + } + } + return Ok(rows); +} + +pub fn get_columns_from_row(table: &Table, row: &Vec, selected_columns: &SelectStatementColumns) -> Result, String> { + let mut row_values: Vec = vec![]; + if *selected_columns == SelectStatementColumns::All { + return Ok(validate_and_clone_row(table, row)?); + } else { + let specific_selected_columns = selected_columns.columns()?; + for (i, column) in table.columns.iter().enumerate() { + if (*specific_selected_columns).contains(&column.name) { + row_values.push(row[i].clone()); + } + } + } + return Ok(row_values); +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::db::table::{Table, Value, DataType, ColumnDefinition}; + use crate::cli::ast::SelectStatementColumns; + use crate::cli::ast::Operator; + use crate::cli::ast::WhereClause; + + fn default_table() -> Table { + Table { + name: "users".to_string(), + columns: vec![ + ColumnDefinition {name: "id".to_string(), data_type: DataType::Integer, constraints: vec![]}, + ColumnDefinition {name: "name".to_string(), data_type: DataType::Text, constraints: vec![]}, + ColumnDefinition {name: "age".to_string(), data_type: DataType::Integer, constraints: vec![]}, + ColumnDefinition {name: "money".to_string(), data_type: DataType::Real, constraints: vec![]}, + ], + rows: vec![ + vec![Value::Integer(1), Value::Text("John".to_string()), Value::Integer(25), Value::Real(1000.0)], + vec![Value::Integer(2), Value::Text("Jane".to_string()), Value::Integer(30), Value::Real(2000.0)], + vec![Value::Integer(3), Value::Text("Jim".to_string()), Value::Integer(35), Value::Real(3000.0)], + vec![Value::Integer(4), Value::Null, Value::Integer(40), Value::Real(4000.0)], + ], + } + } + + #[test] + fn select_with_all_tokens_is_generated_correctly() { + let table = default_table(); + let statement = SelectStatement { + table_name: "users".to_string(), + columns: SelectStatementColumns::All, + where_clause: None, + order_by_clause: None, + limit_clause: None, + }; + let result = select(&table, statement); + assert!(result.is_ok()); + let expected = vec![ + vec![Value::Integer(1), Value::Text("John".to_string()), Value::Integer(25), Value::Real(1000.0)], + vec![Value::Integer(2), Value::Text("Jane".to_string()), Value::Integer(30), Value::Real(2000.0)], + vec![Value::Integer(3), Value::Text("Jim".to_string()), Value::Integer(35), Value::Real(3000.0)], + vec![Value::Integer(4), Value::Null, Value::Integer(40), Value::Real(4000.0)], + ]; + assert_eq!(expected, result.unwrap()); + } + + #[test] + fn select_specific_columns_is_generated_correctly() { + let table = default_table(); + let statement = SelectStatement { + table_name: "users".to_string(), + columns: SelectStatementColumns::Specific(vec!["name".to_string(), "age".to_string()]), + where_clause: None, + order_by_clause: None, + limit_clause: None, + }; + let result = select(&table, statement); + assert!(result.is_ok()); + let expected = vec![ + vec![Value::Text("John".to_string()), Value::Integer(25)], + vec![Value::Text("Jane".to_string()), Value::Integer(30)], + vec![Value::Text("Jim".to_string()), Value::Integer(35)], + vec![Value::Null, Value::Integer(40)], + ]; + assert_eq!(expected, result.unwrap()); + } + + #[test] + fn select_with_where_clause_is_generated_correctly() { + let table = default_table(); + let statement = SelectStatement { + table_name: "users".to_string(), + columns: SelectStatementColumns::All, + where_clause: Some(WhereClause { + column: "name".to_string(), + operator: Operator::Equals, + value: Value::Text("John".to_string()), + }), + order_by_clause: None, + limit_clause: None, + }; + let result = select(&table, statement); + assert!(result.is_ok()); + let expected = vec![ + vec![Value::Integer(1), Value::Text("John".to_string()), Value::Integer(25), Value::Real(1000.0)], + ]; + assert_eq!(expected, result.unwrap()); + } + + #[test] + fn select_with_where_clause_using_column_not_included_in_selected_columns() { + let table = default_table(); + let statement = SelectStatement { + table_name: "users".to_string(), + columns: SelectStatementColumns::Specific(vec!["name".to_string(), "age".to_string()]), + where_clause: Some(WhereClause { + column: "money".to_string(), + operator: Operator::Equals, + value: Value::Real(1000.0), + }), + order_by_clause: None, + limit_clause: None, + }; + let result = select(&table, statement); + assert!(result.is_ok()); + let expected = vec![ + vec![Value::Text("John".to_string()), Value::Integer(25)], + ]; + assert_eq!(expected, result.unwrap()); + } +} \ No newline at end of file diff --git a/src/db/table/select/where_clause.rs b/src/db/table/select/where_clause.rs new file mode 100644 index 0000000..2a7dfa9 --- /dev/null +++ b/src/db/table/select/where_clause.rs @@ -0,0 +1,146 @@ +use crate::cli::ast::{Operator, WhereClause}; +use crate::db::table::{Table, Value, DataType}; + +pub fn matches_where_clause(table: &Table, row: &Vec, where_clause: &WhereClause) -> bool { + let column_value = table.get_column_from_row(row, &where_clause.column); + if column_value.get_type() != where_clause.value.get_type() { + return false; + } + + match where_clause.operator { + Operator::Equals => { + return *column_value == where_clause.value; + }, + Operator::NotEquals => { + return *column_value != where_clause.value; + }, + _ => { + match column_value.get_type() { + DataType::Integer | DataType::Real | DataType::Text => { + match where_clause.operator { + Operator::LessThan => { + return *column_value < where_clause.value; + }, + Operator::GreaterThan => { + return *column_value > where_clause.value; + }, + Operator::LessEquals => { + return *column_value <= where_clause.value; + }, + Operator::GreaterEquals => { + return *column_value >= where_clause.value; + }, + _ => { + return false; + }, + } + }, + _ => { + return false; + }, + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::db::table::{Table, Value, DataType, ColumnDefinition}; + + #[test] + fn matches_where_clause_returns_true_if_row_matches_where_clause() { + let table = Table::new("users".to_string(), vec![ + ColumnDefinition { + name:"id".to_string(), + data_type:DataType::Integer, + constraints: vec![] + }, + ]); + let row = vec![Value::Integer(1)]; + let where_clause = WhereClause {column:"id".to_string(),operator:Operator::Equals,value:Value::Integer(1)}; + assert!(matches_where_clause(&table, &row, &where_clause)); + } + + #[test] + fn matches_where_clause_returns_false_if_row_does_not_match_where_clause() { + let table = Table::new("users".to_string(), vec![ + ColumnDefinition {name:"id".to_string(),data_type:DataType::Integer, constraints: vec![] }, + ]); + let row = vec![Value::Integer(2)]; + let where_clause = WhereClause {column:"id".to_string(),operator:Operator::Equals,value:Value::Integer(1)}; + assert!(!matches_where_clause(&table, &row, &where_clause)); + } + + #[test] + fn matches_where_clause_handles_different_data_types() { + let table = Table::new("users".to_string(), vec![ + ColumnDefinition { + name:"id".to_string(), + data_type:DataType::Integer, + constraints: vec![] + }, + ]); + let row = vec![Value::Integer(1)]; + let where_clause = WhereClause {column:"id".to_string(),operator:Operator::Equals,value:Value::Text("Fletcher".to_string())}; + assert!(!matches_where_clause(&table, &row, &where_clause)); + } + + #[test] + fn matches_where_clause_handles_different_operators() { + let table = Table::new("users".to_string(), vec![ + ColumnDefinition {name:"id".to_string(),data_type:DataType::Integer, constraints: vec![] }, + ]); + let row = vec![Value::Integer(10)]; + let where_clause = WhereClause {column:"id".to_string(),operator:Operator::GreaterThan,value:Value::Integer(0)}; + assert!(matches_where_clause(&table, &row, &where_clause)); + let where_clause = WhereClause {column:"id".to_string(),operator:Operator::GreaterEquals,value:Value::Integer(0)}; + assert!(matches_where_clause(&table, &row, &where_clause)); + let where_clause = WhereClause {column:"id".to_string(),operator:Operator::LessThan,value:Value::Integer(20)}; + assert!(matches_where_clause(&table, &row, &where_clause)); + let where_clause = WhereClause {column:"id".to_string(),operator:Operator::LessEquals,value:Value::Integer(20)}; + assert!(matches_where_clause(&table, &row, &where_clause)); + let where_clause = WhereClause {column:"id".to_string(),operator:Operator::NotEquals,value:Value::Integer(10)}; + assert!(!matches_where_clause(&table, &row, &where_clause)); + } + + #[test] + fn matches_where_clause_handles_string_comparison() { + let table = Table::new("users".to_string(), vec![ + ColumnDefinition {name:"name".to_string(),data_type:DataType::Text, constraints: vec![] }, + ]); + let row = vec![Value::Text("lop".to_string())]; + let where_clause = WhereClause {column:"name".to_string(),operator:Operator::GreaterEquals,value:Value::Text("abc".to_string())}; + assert!(matches_where_clause(&table, &row, &where_clause)); + let where_clause = WhereClause {column:"name".to_string(),operator:Operator::LessEquals,value:Value::Text("lop".to_string())}; + assert!(matches_where_clause(&table, &row, &where_clause)); + let where_clause = WhereClause {column:"name".to_string(),operator:Operator::GreaterThan,value:Value::Text("xyz".to_string())}; + assert!(!matches_where_clause(&table, &row, &where_clause)); + let where_clause = WhereClause {column:"name".to_string(),operator:Operator::LessThan,value:Value::Text("abc".to_string())}; + assert!(!matches_where_clause(&table, &row, &where_clause)); + let where_clause = WhereClause {column:"name".to_string(),operator:Operator::NotEquals,value:Value::Text("abc".to_string())}; + assert!(matches_where_clause(&table, &row, &where_clause)); + let where_clause = WhereClause {column:"name".to_string(),operator:Operator::Equals,value:Value::Text("lop".to_string())}; + assert!(matches_where_clause(&table, &row, &where_clause)); + } + + #[test] + fn matches_where_clause_handles_null() { + let table = Table::new("users".to_string(), vec![ + ColumnDefinition {name:"id".to_string(),data_type:DataType::Integer, constraints: vec![] }, + ]); + let row = vec![Value::Null]; + let where_clause = WhereClause {column:"id".to_string(),operator:Operator::GreaterEquals,value:Value::Integer(1)}; + assert!(!matches_where_clause(&table, &row, &where_clause)); + } + + #[test] + fn matches_where_clause_handles_invalid_operator_for_data_type() { + let table = Table::new("users".to_string(), vec![ + ColumnDefinition {name:"id".to_string(),data_type:DataType::Blob, constraints: vec![] }, + ]); + let row = vec![Value::Blob(vec![1, 2, 3])]; + let where_clause = WhereClause {column:"id".to_string(),operator:Operator::GreaterEquals,value:Value::Blob(vec![1, 2, 3])}; + assert!(!matches_where_clause(&table, &row, &where_clause)); + } +} \ No newline at end of file