diff --git a/src/cli/ast/common.rs b/src/cli/ast/common.rs index 48f1739..ab9356d 100644 --- a/src/cli/ast/common.rs +++ b/src/cli/ast/common.rs @@ -1,4 +1,5 @@ -use crate::cli::{ast::parser::Parser, table::Value, tokenizer::token::TokenTypes}; +use crate::cli::{ast::parser::Parser, tokenizer::token::TokenTypes}; +use crate::db::table::Value; use hex::decode; // Returns an error if the current token does not match the given token type diff --git a/src/cli/ast/create_statement.rs b/src/cli/ast/create_statement.rs index 2514abf..983ae16 100644 --- a/src/cli/ast/create_statement.rs +++ b/src/cli/ast/create_statement.rs @@ -1,4 +1,5 @@ -use crate::cli::{ast::{parser::Parser, CreateTableStatement, SqlStatement::{self, CreateTable}, common::expect_token_type}, table::{ColumnDefinition, DataType}, tokenizer::token::TokenTypes}; +use crate::cli::{ast::{parser::Parser, CreateTableStatement, SqlStatement::{self, CreateTable}, common::expect_token_type}, tokenizer::token::TokenTypes}; +use crate::db::table::{ColumnDefinition, DataType}; pub fn build(parser: &mut Parser) -> Result { parser.advance()?; @@ -36,7 +37,7 @@ fn table_statement(parser: &mut Parser) -> Result { } fn column_definitions(parser: &mut Parser) -> Result, String> { - let mut columns: Vec = vec![]; + let mut columns: Vec = vec![]; expect_token_type(parser, TokenTypes::LeftParen)?; parser.advance()?; diff --git a/src/cli/ast/insert_statement.rs b/src/cli/ast/insert_statement.rs index 6f0adb6..e4edebd 100644 --- a/src/cli/ast/insert_statement.rs +++ b/src/cli/ast/insert_statement.rs @@ -1,4 +1,5 @@ -use crate::cli::{ast::{parser::Parser, common::token_to_value, common::expect_token_type, InsertIntoStatement, SqlStatement::{self, InsertInto}}, table::Value, tokenizer::token::TokenTypes}; +use crate::cli::{ast::{parser::Parser, common::token_to_value, common::expect_token_type, InsertIntoStatement, SqlStatement::{self, InsertInto}}, tokenizer::token::TokenTypes}; +use crate::db::table::Value; pub fn build(parser: &mut Parser) -> Result { parser.advance()?; diff --git a/src/cli/ast/mod.rs b/src/cli/ast/mod.rs index 6e76401..bfacf90 100644 --- a/src/cli/ast/mod.rs +++ b/src/cli/ast/mod.rs @@ -1,4 +1,5 @@ -use crate::cli::{self, table::{ColumnDefinition, Value}, tokenizer::token::TokenTypes}; +use crate::cli::tokenizer::{scanner::Token, token::TokenTypes}; +use crate::db::table::{ColumnDefinition, Value}; mod common; mod create_statement; @@ -98,7 +99,7 @@ impl StatementBuilder for DefaultStatementBuilder { } } -pub fn generate(tokens: Vec) -> Vec> { +pub fn generate(tokens: Vec) -> Vec> { let mut results: Vec> = vec![]; let mut parser = parser::Parser::new(tokens); let builder : &dyn StatementBuilder = &DefaultStatementBuilder; diff --git a/src/cli/ast/select_statement.rs b/src/cli/ast/select_statement.rs index e992222..548e065 100644 --- a/src/cli/ast/select_statement.rs +++ b/src/cli/ast/select_statement.rs @@ -41,7 +41,6 @@ fn get_table_name(parser: &mut Parser) -> Result { Ok(result) } - fn get_where_clause(parser: &mut Parser) -> Result, String> { if expect_token_type(parser, TokenTypes::Where).is_err() { return Ok(None); @@ -153,7 +152,7 @@ fn get_limit(parser: &mut Parser) -> Result, String> { mod tests { use super::*; use crate::cli::tokenizer::scanner::Token; - use crate::cli::table::Value; + use crate::db::table::Value; fn token(tt: TokenTypes, val: &'static str) -> Token<'static> { Token { diff --git a/src/cli/mod.rs b/src/cli/mod.rs index d168cd1..a5ef2e4 100644 --- a/src/cli/mod.rs +++ b/src/cli/mod.rs @@ -1,12 +1,13 @@ use std::io; -mod ast; -mod table; +use crate::db; +pub mod ast; mod tokenizer; pub fn cli() { clear_screen(); println!("Welcome to the MollyDB CLI"); let mut line_count = 1; + let mut database = db::database::Database::new(); loop { print!("({:03}) > ", line_count); @@ -31,8 +32,29 @@ pub fn cli() { let tokens = tokenizer::tokenize(input); println!("{:?}", tokens); let ast = ast::generate(tokens); - for result in ast { - println!("{:?}", result); + for sql_statement in ast { + println!("{:?}", sql_statement); + match sql_statement { + Ok(statement) => { + let result = database.execute(statement); + if let Ok(values) = result { + if let Some(rows) = values { + for row in rows { + println!("{:?}", row); + } + } + else { + println!("Executed Successfully"); + } + } + else { + println!("Error: {}", result.unwrap_err()); + } + }, + Err(error) => { + println!("Error: {}", error); + }, + } } } } diff --git a/src/cli/table.rs b/src/cli/table.rs deleted file mode 100644 index c91d190..0000000 --- a/src/cli/table.rs +++ /dev/null @@ -1,42 +0,0 @@ - -#[derive(Debug, PartialEq)] -pub enum DataType { - Integer, - Real, - Text, - Blob, - Null, -} - -pub struct _Table { - name: String, - columns: Vec, - rows: Vec<_Row>, - length: usize, -} - -#[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, -} - -struct _Row { - primary_key: usize, - values: Vec, -} - -#[derive(Debug, PartialEq)] -pub enum Value { - Integer(i64), - Real(f64), - Text(String), - Blob(Vec), - Null -} diff --git a/src/db/database.rs b/src/db/database.rs new file mode 100644 index 0000000..f13e221 --- /dev/null +++ b/src/db/database.rs @@ -0,0 +1,71 @@ +use crate::db::table::{Table, Value}; +use crate::cli::ast::{SqlStatement, CreateTableStatement, InsertIntoStatement, SelectStatement}; +use std::collections::HashMap; + +pub struct Database { + tables: HashMap, +} + +impl Database { + pub fn new() -> Self { + Self { + tables: HashMap::new(), + } + } + + pub fn execute(&mut self, sql_statement: SqlStatement) -> Result>>, String> { + return match sql_statement { + SqlStatement::CreateTable(statement) => { + self.create_table(statement)?; + Ok(None) + }, + SqlStatement::InsertInto(statement) => { + self.insert_into_table(statement)?; + Ok(None) + }, + SqlStatement::Select(statement) => { + let rows = self.select_from_table(statement)?; + Ok(Some(rows)) + }, + } + } + + 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_name = statement.table_name; + self.tables.insert(table_name.clone(), Table::new(table_name, statement.columns)); + Ok(()) + } + + fn insert_into_table(&mut self, statement: InsertIntoStatement) -> Result<(), String> { + let table = self.get_table_mut(&statement.table_name)?; + table.insert(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)?; + Ok(rows) + } + + fn has_table(&self, table_name: &str) -> bool { + self.tables.contains_key(table_name) + } + + fn get_table(&self, table_name: &str) -> Result<&Table, String> { + if !self.has_table(table_name) { + return Err(format!("Table {} does not exist", table_name)); + } + Ok(self.tables.get(table_name).unwrap()) + } + + fn get_table_mut(&mut self, table_name: &str) -> Result<&mut Table, String> { + if !self.has_table(table_name) { + return Err(format!("Table {} does not exist", table_name)); + } + Ok(self.tables.get_mut(table_name).unwrap()) + } +} \ No newline at end of file diff --git a/src/db/mod.rs b/src/db/mod.rs new file mode 100644 index 0000000..b66a9b6 --- /dev/null +++ b/src/db/mod.rs @@ -0,0 +1,2 @@ +pub mod database; +pub mod table; diff --git a/src/db/table.rs b/src/db/table.rs new file mode 100644 index 0000000..f7250df --- /dev/null +++ b/src/db/table.rs @@ -0,0 +1,202 @@ +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/main.rs b/src/main.rs index dacf017..674b10d 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,4 +1,5 @@ mod cli; +mod db; fn main() { cli::cli();