diff --git a/src/db/database.rs b/src/db/database.rs index 5e027ab..f376198 100644 --- a/src/db/database.rs +++ b/src/db/database.rs @@ -4,6 +4,7 @@ use crate::db::table::select; use crate::db::table::insert; use crate::db::table::delete; use crate::db::table::update; +use crate::db::table::create_table; use std::collections::HashMap; pub struct Database { @@ -43,38 +44,29 @@ impl Database { } fn create_table(&mut self, statement: CreateTableStatement) -> Result<(), String> { - if self.has_table(&statement.table_name) { - return Err(format!("Table {} already exists", statement.table_name)); - } - let table = Table::new(statement.table_name, statement.columns) ; - self.tables.insert(table.name.clone(), table); - Ok(()) + create_table::create_table(self, statement) } fn insert_into_table(&mut self, statement: InsertIntoStatement) -> Result<(), String> { let table = self.get_table_mut(&statement.table_name)?; - insert::insert(table, statement)?; - Ok(()) + insert::insert(table, statement) } fn select_statement_stack(&mut self, statement: SelectStatementStack) -> Result>, String> { - let rows = select::select_statement_stack(self, statement)?; - Ok(rows) + select::select_statement_stack(self, statement) } fn delete_from_table(&mut self, statement: DeleteStatement) -> Result<(), String> { let table = self.get_table_mut(&statement.table_name)?; - delete::delete(table, statement)?; - Ok(()) + delete::delete(table, statement) } fn update_table(&mut self, statement: UpdateStatement) -> Result<(), String> { let table = self.get_table_mut(&statement.table_name)?; - update::update(table, statement)?; - Ok(()) + update::update(table, statement) } - fn has_table(&self, table_name: &str) -> bool { + pub fn has_table(&self, table_name: &str) -> bool { self.tables.contains_key(table_name) } @@ -96,7 +88,6 @@ impl Database { #[cfg(test)] mod tests { use super::*; - use crate::interpreter::ast::CreateTableStatement; use crate::db::table::{ColumnDefinition, DataType}; @@ -119,23 +110,6 @@ mod tests { } } - #[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(); diff --git a/src/db/table/create_table/mod.rs b/src/db/table/create_table/mod.rs new file mode 100644 index 0000000..5908863 --- /dev/null +++ b/src/db/table/create_table/mod.rs @@ -0,0 +1,71 @@ +use crate::db::database::Database; +use crate::interpreter::ast::{CreateTableStatement, ExistenceCheck}; +use crate::db::table::Table; + + +pub fn create_table(database: &mut Database, statement: CreateTableStatement) -> Result<(), String> { + if database.has_table(&statement.table_name) { + match statement.existence_check { + Some(ExistenceCheck::IfExists) => { + return Ok(()); + } + _ => { + return Err(format!("Table {} already exists", statement.table_name)); + } + } + } + let table = Table::new(statement.table_name, statement.columns) ; + database.tables.insert(table.name.clone(), table); + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::interpreter::ast::CreateTableStatement; + use crate::db::table::{ColumnDefinition, DataType}; + use crate::db::table::test_utils::default_database; + + #[test] + fn create_table_generates_proper_table() { + 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(); + assert!(create_table(&mut database, statement).is_ok()); + assert!(database.has_table("users")); + } + + #[test] + fn create_table_errors_when_table_already_exists() { + 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 = default_database(); + let result = create_table(&mut database, statement); + assert!(result.is_err()); + assert_eq!("Table users already exists", result.err().unwrap()); + } + + #[test] + fn create_table_with_if_exists_clause_does_not_error_when_table_already_exists() { + let statement = CreateTableStatement { + table_name: "users".to_string(), + existence_check: Some(ExistenceCheck::IfExists), + columns: vec![ColumnDefinition { name: "id".to_string(), data_type: DataType::Integer, constraints: vec![] }], + }; + let mut database = default_database(); + let result = create_table(&mut database, statement); + assert!(result.is_ok()); + } +} \ No newline at end of file diff --git a/src/db/table/mod.rs b/src/db/table/mod.rs index d8f365f..041fc85 100644 --- a/src/db/table/mod.rs +++ b/src/db/table/mod.rs @@ -7,6 +7,7 @@ pub mod select; pub mod insert; pub mod delete; pub mod update; +pub mod create_table; pub mod helpers; #[cfg(test)] pub mod test_utils; diff --git a/src/interpreter/ast/create_statement.rs b/src/interpreter/ast/create_statement.rs index 947877d..c8ac41b 100644 --- a/src/interpreter/ast/create_statement.rs +++ b/src/interpreter/ast/create_statement.rs @@ -1,6 +1,6 @@ use crate::interpreter::{ ast::{ - parser::Parser, CreateTableStatement, SqlStatement::{self, CreateTable}, + parser::Parser, CreateTableStatement, SqlStatement::{self, CreateTable}, ExistenceCheck, helpers::common::{expect_token_type, get_table_name} }, tokenizer::token::TokenTypes @@ -16,10 +16,9 @@ pub fn build(parser: &mut Parser) -> Result { TokenTypes::Table => { statement = table_statement(parser); }, - TokenTypes::Index => { - statement = index_statement(parser); + _ => { + return Err(parser.format_error()) }, - _ => return Err(parser.format_error()), } // Ensure SemiColon @@ -28,13 +27,25 @@ pub fn build(parser: &mut Parser) -> Result { } fn table_statement(parser: &mut Parser) -> Result { - // Get the table name - let table_name = get_table_name(parser)?; parser.advance()?; + let existence_check = match parser.current_token()?.token_type { + TokenTypes::If => { + parser.advance()?; + expect_token_type(parser, TokenTypes::Not)?; + parser.advance()?; + expect_token_type(parser, TokenTypes::Exists)?; + parser.advance()?; + Some(ExistenceCheck::IfExists) + }, + _ => None, + }; + + let table_name = get_table_name(parser)?; let column_definitions = column_definitions(parser)?; return Ok(CreateTable(CreateTableStatement { table_name, + existence_check, columns: column_definitions, })); } @@ -95,10 +106,6 @@ fn token_to_data_type(parser: &mut Parser) -> Result { }; } -fn index_statement(_parser: &mut Parser) -> Result { - return Err("Index statements not yet implemented".to_string()); -} - #[cfg(test)] mod tests { @@ -126,6 +133,7 @@ mod tests { let result = build(&mut parser); let expected = SqlStatement::CreateTable(CreateTableStatement { table_name: "users".to_string(), + existence_check: None, columns: vec![ ColumnDefinition { name: "id".to_string(), @@ -211,17 +219,29 @@ mod tests { } #[test] - fn index_statement_not_implemented() { - // CREATE INDEX my_index; + fn create_table_with_if_exists_clause() { + // CREATE TABLE IF NOT EXISTS users (id INTEGER); let tokens = vec![ token(TokenTypes::Create, "CREATE"), - token(TokenTypes::Index, "INDEX"), - token(TokenTypes::Identifier, "my_index"), + token(TokenTypes::Table, "TABLE"), + token(TokenTypes::If, "IF"), + token(TokenTypes::Not, "NOT"), + token(TokenTypes::Exists, "EXISTS"), + token(TokenTypes::Identifier, "users"), + token(TokenTypes::LeftParen, "("), + token(TokenTypes::Identifier, "id"), + token(TokenTypes::Integer, "INTEGER"), + token(TokenTypes::RightParen, ")"), token(TokenTypes::SemiColon, ";"), token(TokenTypes::EOF, ""), ]; let mut parser = Parser::new(tokens); let result = build(&mut parser); - assert!(result.is_err()); + let expected = SqlStatement::CreateTable(CreateTableStatement { + table_name: "users".to_string(), + existence_check: Some(ExistenceCheck::IfExists), + columns: vec![ColumnDefinition { name: "id".to_string(), data_type: DataType::Integer, constraints: vec![] }], + }); + assert_eq!(result.unwrap(), expected); } } \ No newline at end of file diff --git a/src/interpreter/ast/delete_statement.rs b/src/interpreter/ast/delete_statement.rs index a04648f..030df8b 100644 --- a/src/interpreter/ast/delete_statement.rs +++ b/src/interpreter/ast/delete_statement.rs @@ -12,8 +12,8 @@ use crate::interpreter::{ pub fn build(parser: &mut Parser) -> Result { parser.advance()?; expect_token_type(parser, TokenTypes::From)?; - let table_name = get_table_name(parser)?; parser.advance()?; + let table_name = get_table_name(parser)?; let where_clause = get_where_clause(parser)?; let order_by_clause = get_order_by(parser)?; let limit_clause = get_limit(parser)?; diff --git a/src/interpreter/ast/helpers/common.rs b/src/interpreter/ast/helpers/common.rs index d16d816..b90ff7c 100644 --- a/src/interpreter/ast/helpers/common.rs +++ b/src/interpreter/ast/helpers/common.rs @@ -70,10 +70,10 @@ pub fn tokens_to_identifier_list(parser: &mut Parser) -> Result, Str } pub fn get_table_name(parser: &mut Parser) -> Result { - parser.advance()?; let token = parser.current_token()?; expect_token_type(parser, TokenTypes::Identifier)?; let result = token.value.to_string(); + parser.advance()?; Ok(result) } diff --git a/src/interpreter/ast/helpers/select_statement.rs b/src/interpreter/ast/helpers/select_statement.rs index 8684483..37f656e 100644 --- a/src/interpreter/ast/helpers/select_statement.rs +++ b/src/interpreter/ast/helpers/select_statement.rs @@ -13,8 +13,8 @@ pub fn get_statement(parser: &mut Parser) -> Result { parser.advance()?; let columns = get_columns(parser)?; expect_token_type(parser, TokenTypes::From)?; - let table_name = get_table_name(parser)?; parser.advance()?; + let table_name = get_table_name(parser)?; let where_clause: Option> = get_where_clause(parser)?; let order_by_clause = get_order_by(parser)?; let limit_clause = get_limit(parser)?; diff --git a/src/interpreter/ast/insert_statement.rs b/src/interpreter/ast/insert_statement.rs index aee6135..1cf3a31 100644 --- a/src/interpreter/ast/insert_statement.rs +++ b/src/interpreter/ast/insert_statement.rs @@ -28,8 +28,8 @@ pub fn build(parser: &mut Parser) -> Result { } fn into_statement(parser: &mut Parser) -> Result { - let table_name = get_table_name(parser)?; parser.advance()?; + let table_name = get_table_name(parser)?; let token = parser.current_token()?; let columns = match token.token_type { diff --git a/src/interpreter/ast/mod.rs b/src/interpreter/ast/mod.rs index e3314a7..cbb736f 100644 --- a/src/interpreter/ast/mod.rs +++ b/src/interpreter/ast/mod.rs @@ -30,9 +30,16 @@ pub enum SqlStatement { #[derive(Debug, PartialEq)] pub struct CreateTableStatement { pub table_name: String, + pub existence_check: Option, pub columns: Vec, } +#[derive(Debug, PartialEq)] +pub enum ExistenceCheck { // Eventually expand to temp tables + IfNotExists, + IfExists, +} + #[derive(Debug, PartialEq)] pub struct InsertIntoStatement { pub table_name: String, diff --git a/src/interpreter/ast/parser.rs b/src/interpreter/ast/parser.rs index 1b8ab2e..2608d56 100644 --- a/src/interpreter/ast/parser.rs +++ b/src/interpreter/ast/parser.rs @@ -130,6 +130,7 @@ mod tests { parser.advance_past_semicolon()?; return Ok(SqlStatement::CreateTable(CreateTableStatement { table_name: "users".to_string(), + existence_check: None, columns: vec![], })); } @@ -187,6 +188,7 @@ mod tests { let result = parser.next_statement(builder); let expected = Some(Ok(SqlStatement::CreateTable(CreateTableStatement { table_name: "users".to_string(), + existence_check: None, columns: vec![], }))); assert_eq!(result, expected); diff --git a/src/interpreter/ast/update_statement.rs b/src/interpreter/ast/update_statement.rs index 4e8b201..1f1b3e6 100644 --- a/src/interpreter/ast/update_statement.rs +++ b/src/interpreter/ast/update_statement.rs @@ -7,10 +7,9 @@ use crate::interpreter::ast::helpers::where_stack::get_where_clause; use crate::interpreter::tokenizer::token::TokenTypes; pub fn build(parser: &mut Parser) -> Result { - + parser.advance()?; let table_name = get_table_name(parser)?; // Ensure Set - parser.advance()?; expect_token_type(parser, TokenTypes::Set)?; let update_values = get_update_values(parser)?; let where_clause = get_where_clause(parser)?; diff --git a/src/interpreter/tokenizer/mod.rs b/src/interpreter/tokenizer/mod.rs index e829dc6..34cb02c 100644 --- a/src/interpreter/tokenizer/mod.rs +++ b/src/interpreter/tokenizer/mod.rs @@ -83,7 +83,7 @@ mod tests { ORDER BY GROUP HAVING DISTINCT ALL AS ASC DESC INNER LEFT RIGHT FULL OUTER JOIN ON UNION LIMIT OFFSET - AND OR IN EXISTS + AND OR IN EXISTS IF CASE WHEN THEN ELSE END = != < <= > >= COUNT SUM AVG MIN MAX @@ -141,6 +141,7 @@ mod tests { token(TokenTypes::Or, "OR", 12, 9), token(TokenTypes::In, "IN", 15, 9), token(TokenTypes::Exists, "EXISTS", 18, 9), + token(TokenTypes::If, "IF", 25, 9), token(TokenTypes::Case, "CASE", 8, 10), token(TokenTypes::When, "WHEN", 13, 10), token(TokenTypes::Then, "THEN", 18, 10), diff --git a/src/interpreter/tokenizer/scanner.rs b/src/interpreter/tokenizer/scanner.rs index b83f3a7..beb4f78 100644 --- a/src/interpreter/tokenizer/scanner.rs +++ b/src/interpreter/tokenizer/scanner.rs @@ -166,6 +166,7 @@ impl<'a> Scanner<'a> { slice if slice.eq_ignore_ascii_case("OR") => TokenTypes::Or, slice if slice.eq_ignore_ascii_case("IN") => TokenTypes::In, slice if slice.eq_ignore_ascii_case("EXISTS") => TokenTypes::Exists, + slice if slice.eq_ignore_ascii_case("IF") => TokenTypes::If, slice if slice.eq_ignore_ascii_case("CASE") => TokenTypes::Case, slice if slice.eq_ignore_ascii_case("WHEN") => TokenTypes::When, slice if slice.eq_ignore_ascii_case("THEN") => TokenTypes::Then, diff --git a/src/interpreter/tokenizer/token.rs b/src/interpreter/tokenizer/token.rs index 1706f7b..b59d72e 100644 --- a/src/interpreter/tokenizer/token.rs +++ b/src/interpreter/tokenizer/token.rs @@ -12,7 +12,7 @@ pub enum TokenTypes { Inner, Left, Right, Full, Outer, Join, On, Limit, Offset, Union, Intersect, Except, // Logical Operators - And, Or, In, Exists, + And, Or, In, Exists, If, Case, When, Then, Else, End, Is, Equals, NotEquals, LessThan, LessEquals, GreaterThan, GreaterEquals, // Aggregate Functions