diff --git a/src/db/table/helpers/common.rs b/src/db/table/helpers/common.rs index d96006a..8486735 100644 --- a/src/db/table/helpers/common.rs +++ b/src/db/table/helpers/common.rs @@ -1,7 +1,7 @@ use crate::db::table::{Table, Value, DataType}; use crate::interpreter::ast::{SelectStatementColumns, WhereStackElement, OrderByClause, LimitClause}; use crate::db::table::helpers::where_stack::matches_where_stack; -use crate::db::table::helpers::{order_by_clause::get_ordered_row_indicies, limit_clause::get_limited_row_indicies}; +use crate::db::table::helpers::{order_by_clause::get_ordered_row_indicies, limit_clause::get_limited_rows}; pub fn validate_and_clone_row(table: &Table, row: &Vec) -> Result, String> { if row.len() != table.width() { @@ -54,7 +54,7 @@ pub fn get_columns_from_row(table: &Table, row: &Vec, selected_columns: & } else { let specific_selected_columns = selected_columns.columns()?; for (i, column) in table.columns.iter().enumerate() { - if (*specific_selected_columns).contains(&column.name) { + if (*specific_selected_columns).contains(&&column.name) { row_values.push(row[i].clone()); } } @@ -70,7 +70,8 @@ pub fn get_row_indicies_matching_clauses(table: &Table, where_clause: &Option, limit_clause: &LimitClause) -> Result, String> { +pub fn get_limited_rows(mut rows: Vec, limit_clause: &LimitClause) -> Result, String> { let mut index: usize = 0; if let Some(offset) = &limit_clause.offset && let Value::Integer(offset) = offset { index = *offset as usize; } - if index >= rows.len() { - return Ok(vec![]); + if index >= rows.len() { + rows.truncate(0); + return Ok(rows); } - + rows.drain(0..index); + let limit = match limit_clause.limit { Value::Integer(limit) => { if limit < 0 { rows.len() } else { - min((limit as usize)+index, rows.len()) + min(limit as usize, rows.len()) } }, - _ => return Err("Limit must be an integer".to_string()), // The parser should have already validated this + _ => unreachable!() // validated by parser }; + rows.truncate(limit); - let mut limited_rows: Vec = vec![]; - for i in index..limit { - limited_rows.push(rows[i]); - } - return Ok(limited_rows); + return Ok(rows); } @@ -50,7 +49,8 @@ mod tests { #[test] fn no_offset_and_limit_is_equal_to_rows_length() { let limit_clause = generate_limit_clause(10, None); - let result = get_limited_row_indicies(default_rows(), &limit_clause); + let table = default_rows(); + let result = get_limited_rows(table, &limit_clause); assert!(result.is_ok()); assert_eq!(default_rows(), result.unwrap()); } @@ -58,7 +58,8 @@ mod tests { #[test] fn no_offset_and_limit_is_greater_than_rows_length() { let limit_clause = generate_limit_clause(15, None); - let result = get_limited_row_indicies(default_rows(), &limit_clause); + let table = default_rows(); + let result = get_limited_rows(table, &limit_clause); assert!(result.is_ok()); assert_eq!(default_rows(), result.unwrap()); } @@ -66,7 +67,8 @@ mod tests { #[test] fn no_offset_and_limit_is_less_than_rows_length() { let limit_clause = generate_limit_clause(5, None); - let result = get_limited_row_indicies(default_rows(), &limit_clause); + let table = default_rows(); + let result = get_limited_rows(table, &limit_clause); assert!(result.is_ok()); let expected = vec![0, 1, 2, 3, 4]; assert_eq!(expected, result.unwrap()); @@ -75,7 +77,8 @@ mod tests { #[test] fn no_offset_and_negative_limit_returns_all_rows() { let limit_clause = generate_limit_clause(-1, None); - let result = get_limited_row_indicies(default_rows(), &limit_clause); + let table = default_rows(); + let result = get_limited_rows(table, &limit_clause); assert!(result.is_ok()); assert_eq!(default_rows(), result.unwrap()); } @@ -83,16 +86,18 @@ mod tests { #[test] fn offset_and_limit_is_generated_correctly() { let limit_clause = generate_limit_clause(5, Some(1)); - let result = get_limited_row_indicies(default_rows(), &limit_clause); + let table = default_rows(); + let result = get_limited_rows(table, &limit_clause); assert!(result.is_ok()); - let expected = vec![1, 2, 3, 4, 5]; + let expected: Vec = vec![1, 2, 3, 4, 5]; assert_eq!(expected, result.unwrap()); } #[test] fn offset_is_greater_than_rows_length_returns_empty_rows() { let limit_clause = generate_limit_clause(5, Some(10)); - let result = get_limited_row_indicies(default_rows(), &limit_clause); + let table = default_rows(); + let result = get_limited_rows(table, &limit_clause); assert!(result.is_ok()); let expected: Vec = vec![]; assert_eq!(expected, result.unwrap()); diff --git a/src/db/table/helpers/order_by_clause.rs b/src/db/table/helpers/order_by_clause.rs index a857d93..9a61cb6 100644 --- a/src/db/table/helpers/order_by_clause.rs +++ b/src/db/table/helpers/order_by_clause.rs @@ -9,19 +9,20 @@ use crate::db::table::Value; // This sorting algorithm will always return a stable sort, this is given by all of the order columns // then the input order of the rows is maintained with any required tie breaking. pub fn get_ordered_row_indicies(table: &Table, mut row_indicies: Vec, order_by_clauses: &Vec) -> Result, String> { + let columns: Vec<&String> = table.columns.iter().map(|column| &column.name).collect(); row_indicies.sort_by(|a, b| { - perform_comparions(table, &table.rows[*a], &table.rows[*b], order_by_clauses) + perform_comparions(&columns, &table.rows[*a], &table.rows[*b], order_by_clauses) }); return Ok(row_indicies); } -fn perform_comparions(table: &Table, row1: &Vec, row2: &Vec, order_by_clauses: &Vec) -> Ordering { +pub fn perform_comparions(columns: &Vec<&String>, row1: &Vec, row2: &Vec, order_by_clauses: &Vec) -> Ordering { let mut result = Ordering::Equal; for comparison in order_by_clauses { - let index = table.get_index_of_column(&comparison.column); + let index = get_index_of_column(columns, &comparison.column); let index = match index { Ok(index) => index, - Err(_) => return Ordering::Equal, // Bad but should never happen because we've validated the columns in the parser + Err(_) => unreachable!(), }; let ordering = row1[index].compare(&row2[index], &comparison.direction); if ordering != Ordering::Equal { @@ -32,6 +33,16 @@ fn perform_comparions(table: &Table, row1: &Vec, row2: &Vec, order return result; } +fn get_index_of_column(columns: &Vec<&String>, column_name: &String) -> Result { + let result = columns.iter().position(|column| column_name == *column); + if let Some(index) = result { + return Ok(index); + } + else { + return Err(format!("Column {} does not exist in table", column_name)); + } +} + #[cfg(test)] mod tests { use super::*; diff --git a/src/db/table/mod.rs b/src/db/table/mod.rs index 066a317..d8f365f 100644 --- a/src/db/table/mod.rs +++ b/src/db/table/mod.rs @@ -175,4 +175,8 @@ impl Table { } return Err(format!("Column {} does not exist in table {}", column, self.name)); } + + pub fn get_columns(&self) -> Vec<&String> { + self.columns.iter().map(|column| &column.name).collect() + } } \ No newline at end of file diff --git a/src/db/table/select/mod.rs b/src/db/table/select/mod.rs index bceb871..8f60009 100644 --- a/src/db/table/select/mod.rs +++ b/src/db/table/select/mod.rs @@ -1,14 +1,37 @@ mod select_statement; mod set_operator_evaluator; use crate::db::{database::Database, table::Value}; -use crate::interpreter::ast::{SelectStatementStack, SetOperator, SelectStatementStackElement}; +use crate::interpreter::ast::{SelectStatementStack, SetOperator, SelectStatementStackElement, SelectStatementColumns}; +use crate::db::table::helpers::{order_by_clause::{perform_comparions}, limit_clause::get_limited_rows}; + pub fn select_statement_stack(database: &Database, statement: SelectStatementStack) -> Result>, String> { let mut evaluator = set_operator_evaluator::SetOperatorEvaluator::new(); + let statement_columns = statement.columns.columns(); + let mut columns: Option> = match statement_columns { + Err(_) => None, + Ok(columns_list) => Some(columns_list), + }; for element in statement.elements { match element { SelectStatementStackElement::SelectStatement(select_statement) => { let table = database.get_table(&select_statement.table_name)?; + columns = match columns { + None => Some(table.get_columns()), + Some(columns) => { + if statement.columns == SelectStatementColumns::All { + if table.get_columns() != columns { + return Err(format!("Columns mismatch between SELECT statements in Union")); + } + } + else { + if statement.columns.columns()? != columns { + return Err(format!("Columns mismatch between SELECT statements in Union")); + } + } + Some(columns) + }, + }; let rows = select_statement::select_statement(table, &select_statement)?; evaluator.push(rows); } @@ -30,7 +53,20 @@ pub fn select_statement_stack(database: &Database, statement: SelectStatementSta } } } - let result = evaluator.result()?; + let mut result = evaluator.result()?; + if let Some(order_by_clause) = statement.order_by_clause { + result.sort_by(|a, b| { + if let Some(columns) = &columns { + perform_comparions(&columns, a, b, &order_by_clause) + } + else { + unreachable!() + } + }); + } + if let Some(limit_clause) = statement.limit_clause { + result = get_limited_rows(result, &limit_clause)?; + } Ok(result) } @@ -46,6 +82,7 @@ mod tests { fn select_statement_stack_with_multiple_set_operators_works_correctly() { let database = default_database(); let statement = SelectStatementStack { + columns: SelectStatementColumns::All, elements: vec![SelectStatementStackElement::SelectStatement(SelectStatement { table_name: "users".to_string(), columns: SelectStatementColumns::All, @@ -71,6 +108,7 @@ mod tests { fn select_statement_stack_with_set_operator_works_correctly() { let database = default_database(); let statement = SelectStatementStack { + columns: SelectStatementColumns::All, elements: vec![ SelectStatementStackElement::SelectStatement(SelectStatement { table_name: "users".to_string(), @@ -107,6 +145,7 @@ mod tests { fn select_statement_stack_works_correctly_with_multiple_set_operators() { let database = default_database(); let statement = SelectStatementStack { + columns: SelectStatementColumns::All, elements: vec![SelectStatementStackElement::SelectStatement(SelectStatement { table_name: "users".to_string(), columns: SelectStatementColumns::All, diff --git a/src/interpreter/ast/mod.rs b/src/interpreter/ast/mod.rs index 25fb204..e3314a7 100644 --- a/src/interpreter/ast/mod.rs +++ b/src/interpreter/ast/mod.rs @@ -42,6 +42,7 @@ pub struct InsertIntoStatement { #[derive(Debug, PartialEq)] pub struct SelectStatementStack { + pub columns: SelectStatementColumns, pub elements: Vec, pub order_by_clause: Option>, pub limit_clause: Option, @@ -109,17 +110,17 @@ pub struct ColumnValue { pub value: Value, } -#[derive(Debug, PartialEq)] +#[derive(Debug, PartialEq, Clone)] pub enum SelectStatementColumns { All, Specific(Vec), } impl SelectStatementColumns { - pub fn columns(&self) -> Result<&Vec, String> { + pub fn columns(&self) -> Result, String> { return match self { SelectStatementColumns::All => Err("Cannot get columns from all columns".to_string()), - SelectStatementColumns::Specific(columns) => Ok(columns), + SelectStatementColumns::Specific(columns) => Ok(columns.iter().map(|column| column).collect()), } } } @@ -355,6 +356,7 @@ mod tests { let expected = vec![ Ok(DatabaseSqlStatement { sql_statement: SqlStatement::Select(SelectStatementStack { + columns: SelectStatementColumns::All, elements: vec![SelectStatementStackElement::SelectStatement(SelectStatement { table_name: "users".to_string(), columns: SelectStatementColumns::All, @@ -402,7 +404,6 @@ mod tests { token(TokenTypes::EOF, ""), ]; let result = generate(tokens); - println!("{:?}", result); assert!(result[0].is_err()); assert!(result[1].is_ok()); let expected = vec![ @@ -448,6 +449,7 @@ mod tests { let expected = vec![ Ok(DatabaseSqlStatement { sql_statement: SqlStatement::Select(SelectStatementStack { + columns: SelectStatementColumns::All, elements: vec![SelectStatementStackElement::SelectStatement(SelectStatement { table_name: "users".to_string(), columns: SelectStatementColumns::All, diff --git a/src/interpreter/ast/parser.rs b/src/interpreter/ast/parser.rs index 4196181..1b8ab2e 100644 --- a/src/interpreter/ast/parser.rs +++ b/src/interpreter/ast/parser.rs @@ -148,6 +148,7 @@ mod tests { parser.advance()?; parser.advance_past_semicolon()?; return Ok(SqlStatement::Select(SelectStatementStack { + columns: SelectStatementColumns::All, elements: vec![SelectStatementStackElement::SelectStatement(SelectStatement { table_name: "users".to_string(), columns: SelectStatementColumns::All, @@ -202,6 +203,7 @@ mod tests { // Select let result = parser.next_statement(builder); let expected = Some(Ok(SqlStatement::Select(SelectStatementStack { + columns: SelectStatementColumns::All, elements: vec![SelectStatementStackElement::SelectStatement(SelectStatement { table_name: "users".to_string(), columns: SelectStatementColumns::All, diff --git a/src/interpreter/ast/select_statement_stack.rs b/src/interpreter/ast/select_statement_stack.rs index e47231f..70d83cc 100644 --- a/src/interpreter/ast/select_statement_stack.rs +++ b/src/interpreter/ast/select_statement_stack.rs @@ -1,6 +1,6 @@ use crate::interpreter::ast::helpers::order_by_clause::get_order_by; use crate::interpreter::ast::helpers::limit_clause::get_limit; -use crate::interpreter::ast::{parser::Parser, SqlStatement, SelectStatementStack, SelectStatementStackElement, SetOperator, SelectStackOperators}; +use crate::interpreter::ast::{parser::Parser, SqlStatement, SelectStatementStack, SelectStatementStackElement, SetOperator, SelectStackOperators, SelectStatementColumns}; use crate::interpreter::ast::helpers::select_statement; use crate::interpreter::ast::Parentheses; use crate::interpreter::tokenizer::token::TokenTypes; @@ -8,10 +8,12 @@ use crate::interpreter::tokenizer::token::TokenTypes; // Returns a SelectStatementStack which is an RPN representation of the SELECT statements and set operators. pub fn build(parser: &mut Parser) -> Result { let mut statement_stack = SelectStatementStack { + columns: SelectStatementColumns::All, elements: vec![], order_by_clause: None, limit_clause: None, }; + let mut columns = None; let mut set_operator_stack: Vec = vec![]; loop { @@ -19,6 +21,15 @@ pub fn build(parser: &mut Parser) -> Result { match token.token_type { TokenTypes::Select => { let mut statement = select_statement::get_statement(parser)?; + columns = match columns { + None => Some(statement.columns.clone()), + Some(columns) => { + if statement.columns != columns { + return Err("Columns mismatch between SELECT statements in Union".to_string()); + } + Some(columns) + }, + }; if parser.current_token()?.token_type != TokenTypes::SemiColon { if statement.order_by_clause.is_some() || statement.limit_clause.is_some() { return Err("ORDER BY, or LIMIT clause not allowed with UNION SELECT statements".to_string()); @@ -104,7 +115,10 @@ pub fn build(parser: &mut Parser) -> Result { return Err("Mismatched parentheses found.".to_string()); } } - + match columns { + Some(columns) => statement_stack.columns = columns, + None => return Err("Error parsing SELECT statement. Columns not found.".to_string()), + } return Ok(SqlStatement::Select(statement_stack)); } @@ -187,6 +201,7 @@ mod tests { assert!(result.is_ok()); let statement = result.unwrap(); let expected = SqlStatement::Select(SelectStatementStack { + columns: SelectStatementColumns::All, elements: vec![expected_simple_select_statement(1)], order_by_clause: None, limit_clause: None, @@ -203,10 +218,10 @@ mod tests { tokens.append(&mut vec![token(TokenTypes::SemiColon, ";")]); let mut parser = Parser::new(tokens); let result = build(&mut parser); - println!("{:?}", result); assert!(result.is_ok()); let statement = result.unwrap(); let expected = SqlStatement::Select(SelectStatementStack { + columns: SelectStatementColumns::All, elements: vec![ expected_simple_select_statement(1), expected_simple_select_statement(2), @@ -231,10 +246,10 @@ mod tests { tokens.append(&mut vec![token(TokenTypes::SemiColon, ";")]); let mut parser = Parser::new(tokens); let result = build(&mut parser); - println!("{:?}", result); assert!(result.is_ok()); let statement = result.unwrap(); let expected = SqlStatement::Select(SelectStatementStack { + columns: SelectStatementColumns::All, elements: vec![ expected_simple_select_statement(1), expected_simple_select_statement(2), @@ -268,10 +283,10 @@ mod tests { tokens.append(&mut vec![token(TokenTypes::SemiColon, ";")]); let mut parser = Parser::new(tokens); let result = build(&mut parser); - println!("{:?}", result); assert!(result.is_ok()); let statement = result.unwrap(); let expected = SqlStatement::Select(SelectStatementStack { + columns: SelectStatementColumns::All, elements: vec![ expected_simple_select_statement(1), expected_simple_select_statement(2), @@ -320,10 +335,10 @@ mod tests { ]; let mut parser = Parser::new(tokens); let result = build(&mut parser); - println!("{:?}", result); assert!(result.is_ok()); let statement = result.unwrap(); let expected = SqlStatement::Select(SelectStatementStack { + columns: SelectStatementColumns::Specific(vec!["name".to_string()]), elements: vec![ SelectStatementStackElement::SelectStatement(SelectStatement { table_name: "employees".to_string(), @@ -383,6 +398,7 @@ mod tests { assert!(result.is_ok()); let statement = result.unwrap(); let expected = SqlStatement::Select(SelectStatementStack { + columns: SelectStatementColumns::All, elements: vec![ expected_simple_select_statement(1), expected_simple_select_statement(2), @@ -418,6 +434,7 @@ mod tests { assert!(result.is_ok()); let statement = result.unwrap(); let expected = SqlStatement::Select(SelectStatementStack { + columns: SelectStatementColumns::All, elements: vec![ expected_simple_select_statement(1), expected_simple_select_statement(2), @@ -431,4 +448,27 @@ mod tests { }); assert_eq!(expected, statement); } + + #[test] + fn select_statement_with_columns_mismatch_is_generated_correctly() { + // SELECT id, name FROM users UNION SELECT name FROM users; + let tokens = vec![ + token(TokenTypes::Select, "SELECT"), + token(TokenTypes::Identifier, "id"), + token(TokenTypes::Comma, ","), + token(TokenTypes::Identifier, "name"), + token(TokenTypes::From, "FROM"), + token(TokenTypes::Identifier, "users"), + token(TokenTypes::Union, "UNION"), + token(TokenTypes::Select, "SELECT"), + token(TokenTypes::Identifier, "name"), + token(TokenTypes::From, "FROM"), + token(TokenTypes::Identifier, "users"), + token(TokenTypes::SemiColon, ";") + ]; + let mut parser = Parser::new(tokens); + let result = build(&mut parser); + assert!(result.is_err()); + assert_eq!(result.unwrap_err(), "Columns mismatch between SELECT statements in Union".to_string()); + } } \ No newline at end of file diff --git a/tests/set_operators.rs b/tests/set_operators.rs index 3ff8aa2..4c0d56e 100644 --- a/tests/set_operators.rs +++ b/tests/set_operators.rs @@ -26,24 +26,71 @@ fn test_set_operators() { } #[test] -fn test_set_operators_with_clauses_and_parentheses() { - // 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, 'Jane'), (4, 'Jack'); - // (SELECT id, name FROM users WHERE id > 1 INTERSECT SELECT id, name FROM users WHERE id < 4) - // ORDER BY name ASC, id DESC LIMIT 1; - // "; - // let mut result = run_sql(&mut database, sql); - // println!("{:?}", result); - // assert!(result.iter().all(|result| result.is_ok())); - // let expected = vec![ - // vec![Value::Integer(3), Value::Text("Jane".to_string())], - // ]; - // test_utils::assert_table_rows_eq_unordered(expected, result.pop().unwrap().unwrap().unwrap()); - // assert!(result.into_iter().all(|result| result.is_ok() && result.unwrap().is_none())); - // NOT CURRENTLY SUPPORTED NEED TO ADD. -} \ No newline at end of file +fn test_set_operators_order_by_clauses_and_parentheses() { + let mut database = Database::new(); + let sql = " + CREATE TABLE users ( + id INTEGER, + name TEXT + ); + INSERT INTO users (id, name) VALUES (1, 'John'), (2, 'zane'), (3, 'Jane'), (4, 'Jack'); + (SELECT id, name FROM users WHERE id > 1 INTERSECT SELECT id, name FROM users WHERE id < 4) + ORDER BY name ASC, id DESC; + (SELECT id, name FROM users WHERE id > 1 INTERSECT SELECT id, name FROM users WHERE id < 4) + ORDER BY name DESC LIMIT 1; + "; + let mut result = run_sql(&mut database, sql); + assert!(result.iter().all(|result| result.is_ok())); + let expected_second = vec![ + vec![Value::Integer(3), Value::Text("Jane".to_string())], + vec![Value::Integer(2), Value::Text("zane".to_string())], + ]; + let expected_first = vec![ + vec![Value::Integer(2), Value::Text("zane".to_string())], + ]; + assert_eq!(expected_first, result.pop().unwrap().unwrap().unwrap()); + assert_eq!(expected_second, result.pop().unwrap().unwrap().unwrap()); + assert!(result.into_iter().all(|result| result.is_ok() && result.unwrap().is_none())); +} + +#[test] +fn test_set_operators_with_different_tables_and_clause() { + let mut database = Database::new(); + let sql = " + CREATE TABLE users1 ( + id INTEGER, + name TEXT + ); + CREATE TABLE users2 ( + id INTEGER, + name TEXT + ); + CREATE TABLE users3 ( + employee_id INTEGER, + name TEXT + ); + INSERT INTO users1 (id, name) VALUES (1, 'John'), (2, 'Jane'), (3, 'Jim'), (4, 'Jack'); + INSERT INTO users2 (id, name) VALUES (1, 'Fletcher'), (2, 'Jane'), (3, 'Jim'), (4, 'Fletcher'); + SELECT name FROM users1 UNION SELECT name FROM users2; + SELECT * FROM users1 UNION SELECT * FROM users3; + "; + let mut result = run_sql(&mut database, sql); + let first_result = result.pop().unwrap(); + assert!(first_result.is_err()); + let expected_second = "Execution Error with statement starting on line 17 \n Error: Columns mismatch between SELECT statements in Union".to_string(); + assert_eq!(expected_second, first_result.unwrap_err()); + assert!(result.iter().all(|result| result.is_ok())); + let expected_first = vec![ + vec![Value::Text("John".to_string())], + vec![Value::Text("Fletcher".to_string())], + vec![Value::Text("Jane".to_string())], + vec![Value::Text("Jim".to_string())], + vec![Value::Text("Jack".to_string())], + ]; + test_utils::assert_table_rows_eq_unordered(expected_first, result.pop().unwrap().unwrap().unwrap()); + assert!(result.into_iter().all(|result| result.is_ok() && result.unwrap().is_none())); + +} + +// ADD TESTS with two seperate tables with different columns and using SELECT * +// two tables with same columns and using SELECT * \ No newline at end of file