From 451ef6d684f9bed2e58f7b9ec3ff460d9e986dbe Mon Sep 17 00:00:00 2001 From: Fletcher555 Date: Sun, 21 Sep 2025 11:20:51 -0400 Subject: [PATCH 1/6] Refactor column name to column name object to allow alias' --- src/db/table/operations/delete/mod.rs | 9 ++--- src/db/table/operations/helpers/common.rs | 4 +-- .../operations/helpers/order_by_clause.rs | 7 ++-- src/db/table/operations/select/mod.rs | 28 +++++++-------- .../operations/select/select_statement.rs | 31 +++++++++-------- src/db/table/operations/update/mod.rs | 5 +-- src/interpreter/ast/delete_statement.rs | 5 +-- src/interpreter/ast/helpers/common.rs | 28 +++++++-------- .../ast/helpers/order_by_clause.rs | 11 +++--- .../ast/helpers/select_statement.rs | 34 +++++++++---------- src/interpreter/ast/mod.rs | 23 ++++++++++--- src/interpreter/ast/parser.rs | 4 +-- src/interpreter/ast/select_statement_stack.rs | 19 ++++++----- src/interpreter/ast/statement_builder.rs | 4 +-- src/interpreter/ast/update_statement.rs | 5 +-- 15 files changed, 119 insertions(+), 98 deletions(-) diff --git a/src/db/table/operations/delete/mod.rs b/src/db/table/operations/delete/mod.rs index 345bfb3..84e8a82 100644 --- a/src/db/table/operations/delete/mod.rs +++ b/src/db/table/operations/delete/mod.rs @@ -66,6 +66,7 @@ mod tests { use crate::db::table::core::{row::Row, value::Value}; use crate::db::table::test_utils::{assert_table_rows_eq_unordered, default_table}; use crate::interpreter::ast::LimitClause; + use crate::interpreter::ast::SelectStatementColumn; use crate::interpreter::ast::{ Operand, Operator, OrderByClause, OrderByDirection, SelectableStack, SelectableStackElement, WhereCondition, WhereStackElement, @@ -165,9 +166,9 @@ mod tests { })]), order_by_clause: Some(OrderByClause { columns: SelectableStack { - selectables: vec![SelectableStackElement::Column("id".to_string())], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("id".to_string()))], }, - column_names: vec!["id".to_string()], + column_names: vec![SelectStatementColumn::new("id".to_string())], directions: vec![OrderByDirection::Desc], }), limit_clause: Some(LimitClause { @@ -317,9 +318,9 @@ mod tests { })]), order_by_clause: Some(OrderByClause { columns: SelectableStack { - selectables: vec![SelectableStackElement::Column("id".to_string())], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("id".to_string()))], }, - column_names: vec!["id".to_string()], + column_names: vec![SelectStatementColumn::new("id".to_string())], directions: vec![OrderByDirection::Desc], }), limit_clause: Some(LimitClause { diff --git a/src/db/table/operations/helpers/common.rs b/src/db/table/operations/helpers/common.rs index 8c6f892..833bd9f 100644 --- a/src/db/table/operations/helpers/common.rs +++ b/src/db/table/operations/helpers/common.rs @@ -60,10 +60,10 @@ pub fn get_columns_from_row( } } SelectableStackElement::Column(value) => { - if let Some(value) = column_values.get(value) { + if let Some(value) = column_values.get(&value.column_name) { row_values.push((*value).clone()); } else { - return Err(format!("Invalid column name: {}", value)); + return Err(format!("Invalid column name: {}", value.column_name)); } } SelectableStackElement::Value(value) => { diff --git a/src/db/table/operations/helpers/order_by_clause.rs b/src/db/table/operations/helpers/order_by_clause.rs index 7ba31cb..902dffe 100644 --- a/src/db/table/operations/helpers/order_by_clause.rs +++ b/src/db/table/operations/helpers/order_by_clause.rs @@ -42,7 +42,8 @@ mod tests { use crate::interpreter::ast::OrderByDirection; use crate::interpreter::ast::SelectableStack; use crate::interpreter::ast::SelectableStackElement; - + use crate::interpreter::ast::SelectStatementColumn; + #[test] fn apply_order_by_from_precomputed_single_column_asc() { let mut to_order = vec!["second", "fourth", "third", "first"]; @@ -56,9 +57,9 @@ mod tests { let order_by_clause = OrderByClause { columns: SelectableStack { - selectables: vec![SelectableStackElement::Column("age".to_string())], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("age".to_string()))], }, - column_names: vec!["age".to_string()], + column_names: vec![SelectStatementColumn::new("age".to_string())], directions: vec![OrderByDirection::Asc], }; diff --git a/src/db/table/operations/select/mod.rs b/src/db/table/operations/select/mod.rs index 85261a7..1046285 100644 --- a/src/db/table/operations/select/mod.rs +++ b/src/db/table/operations/select/mod.rs @@ -22,7 +22,7 @@ pub fn select_statement_stack( SelectStatementStackElement::SelectStatement(select_statement) => { let table = database.get_table(&select_statement.table_name)?; let expanded_column_names = - expand_all_column_names(table, &select_statement.column_names)?; + expand_all_column_names(table, select_statement.column_names.iter().map(|column| &column.column_name).collect::>())?; match &column_names { Some(column_names) => { if expanded_column_names.len() != column_names.len() { @@ -82,7 +82,7 @@ pub fn select_statement_stack( .as_ref() .ok_or_else(|| "No column names found".to_string())? .iter() - .position(|column_name| column_name == order_by_column_name) + .position(|column_name| *column_name == order_by_column_name.column_name) .ok_or_else(|| { "Ordering column name not found in selected columns".to_string() })?, @@ -122,18 +122,18 @@ pub fn select_statement_stack( // TODO: add this logic in evaluation too fn expand_all_column_names( table: &Table, - column_names: &Vec, + column_names: Vec<&String>, ) -> Result, String> { let mut new = vec![]; - for column in column_names { - if *column == "*".to_string() { + for column in &column_names { + if **column == "*".to_string() { for name in table.get_column_names()? { - if !column_names.contains(name) { + if !column_names.contains(&name) { new.push(name.clone()); } } } else { - new.push(column.clone()); + new.push((*column).clone()); } } Ok(new) @@ -144,7 +144,7 @@ mod tests { use crate::db::table::core::value::Value; use crate::db::table::test_utils::default_database; use crate::interpreter::ast::{ - LogicalOperator, Operand, Operator, SelectMode, SelectStatement, SelectableStack, + LogicalOperator, Operand, Operator, SelectMode, SelectStatement, SelectableStack, SelectStatementColumn, SelectableStackElement, WhereCondition, WhereStackElement, }; @@ -159,7 +159,7 @@ mod tests { columns: SelectableStack { selectables: vec![SelectableStackElement::All], }, - column_names: vec!["*".to_string()], + column_names: vec![SelectStatementColumn::new("*".to_string())], where_clause: None, order_by_clause: None, limit_clause: None, @@ -210,7 +210,7 @@ mod tests { columns: SelectableStack { selectables: vec![SelectableStackElement::All], }, - column_names: vec!["*".to_string()], + column_names: vec![SelectStatementColumn::new("*".to_string())], where_clause: Some(vec![WhereStackElement::Condition(WhereCondition { l_side: Operand::Identifier("id".to_string()), operator: Operator::Equals, @@ -225,7 +225,7 @@ mod tests { columns: SelectableStack { selectables: vec![SelectableStackElement::All], }, - column_names: vec!["*".to_string()], + column_names: vec![SelectStatementColumn::new("*".to_string())], where_clause: None, order_by_clause: None, limit_clause: None, @@ -257,7 +257,7 @@ mod tests { columns: SelectableStack { selectables: vec![SelectableStackElement::All], }, - column_names: vec!["*".to_string()], + column_names: vec![SelectStatementColumn::new("*".to_string())], where_clause: None, order_by_clause: None, limit_clause: None, @@ -268,7 +268,7 @@ mod tests { columns: SelectableStack { selectables: vec![SelectableStackElement::All], }, - column_names: vec!["*".to_string()], + column_names: vec![SelectStatementColumn::new("*".to_string())], where_clause: Some(vec![ WhereStackElement::Condition(WhereCondition { l_side: Operand::Identifier("id".to_string()), @@ -292,7 +292,7 @@ mod tests { columns: SelectableStack { selectables: vec![SelectableStackElement::All], }, - column_names: vec!["*".to_string()], + column_names: vec![SelectStatementColumn::new("*".to_string())], where_clause: Some(vec![WhereStackElement::Condition(WhereCondition { l_side: Operand::Identifier("id".to_string()), operator: Operator::Equals, diff --git a/src/db/table/operations/select/select_statement.rs b/src/db/table/operations/select/select_statement.rs index e50ab35..c47cfd1 100644 --- a/src/db/table/operations/select/select_statement.rs +++ b/src/db/table/operations/select/select_statement.rs @@ -78,6 +78,7 @@ mod tests { use crate::db::table::test_utils::{assert_table_rows_eq_unordered, default_table}; use crate::interpreter::ast::Operand; use crate::interpreter::ast::SelectMode; + use crate::interpreter::ast::SelectStatementColumn; use crate::interpreter::ast::WhereCondition; use crate::interpreter::ast::WhereStackElement; use crate::interpreter::ast::{ @@ -94,7 +95,7 @@ mod tests { columns: SelectableStack { selectables: vec![SelectableStackElement::All], }, - column_names: vec!["*".to_string()], + column_names: vec![SelectStatementColumn::new("*".to_string())], where_clause: None, order_by_clause: None, limit_clause: None, @@ -138,11 +139,11 @@ mod tests { mode: SelectMode::All, columns: SelectableStack { selectables: vec![ - SelectableStackElement::Column("name".to_string()), - SelectableStackElement::Column("age".to_string()), + SelectableStackElement::Column(SelectStatementColumn::new("name".to_string())), + SelectableStackElement::Column(SelectStatementColumn::new("age".to_string())), ], }, - column_names: vec!["name".to_string(), "age".to_string()], + column_names: vec![SelectStatementColumn::new("name".to_string()), SelectStatementColumn::new("age".to_string())], where_clause: None, order_by_clause: None, limit_clause: None, @@ -167,7 +168,7 @@ mod tests { columns: SelectableStack { selectables: vec![SelectableStackElement::All], }, - column_names: vec!["*".to_string()], + column_names: vec![SelectStatementColumn::new("*".to_string())], where_clause: Some(vec![WhereStackElement::Condition(WhereCondition { l_side: Operand::Identifier("name".to_string()), operator: Operator::Equals, @@ -195,11 +196,11 @@ mod tests { mode: SelectMode::All, columns: SelectableStack { selectables: vec![ - SelectableStackElement::Column("name".to_string()), - SelectableStackElement::Column("age".to_string()), + SelectableStackElement::Column(SelectStatementColumn::new("name".to_string())), + SelectableStackElement::Column(SelectStatementColumn::new("age".to_string())), ], }, - column_names: vec!["name".to_string(), "age".to_string()], + column_names: vec![SelectStatementColumn::new("name".to_string()), SelectStatementColumn::new("age".to_string())], where_clause: Some(vec![WhereStackElement::Condition(WhereCondition { l_side: Operand::Identifier("money".to_string()), operator: Operator::Equals, @@ -226,7 +227,7 @@ mod tests { columns: SelectableStack { selectables: vec![SelectableStackElement::All], }, - column_names: vec!["*".to_string()], + column_names: vec![SelectStatementColumn::new("*".to_string())], where_clause: None, order_by_clause: None, limit_clause: Some(LimitClause { @@ -254,7 +255,7 @@ mod tests { columns: SelectableStack { selectables: vec![SelectableStackElement::All], }, - column_names: vec!["*".to_string()], + column_names: vec![SelectStatementColumn::new("*".to_string())], where_clause: Some(vec![WhereStackElement::Condition(WhereCondition { l_side: Operand::Identifier("column_not_included".to_string()), operator: Operator::Equals, @@ -280,13 +281,13 @@ mod tests { columns: SelectableStack { selectables: vec![SelectableStackElement::All], }, - column_names: vec!["*".to_string()], + column_names: vec![SelectStatementColumn::new("*".to_string())], where_clause: None, order_by_clause: Some(OrderByClause { columns: SelectableStack { - selectables: vec![SelectableStackElement::Column("money".to_string())], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("money".to_string()))], }, - column_names: vec!["money".to_string()], + column_names: vec![SelectStatementColumn::new("money".to_string())], directions: vec![OrderByDirection::Desc], }), limit_clause: None, @@ -347,10 +348,10 @@ mod tests { ]); let statement = SelectStatement { table_name: "users".to_string(), - column_names: vec!["name".to_string()], + column_names: vec![SelectStatementColumn::new("name".to_string())], mode: SelectMode::Distinct, columns: SelectableStack { - selectables: vec![SelectableStackElement::Column("name".to_string())], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("name".to_string()))], }, where_clause: None, order_by_clause: None, diff --git a/src/db/table/operations/update/mod.rs b/src/db/table/operations/update/mod.rs index 6b137c3..c18ac6d 100644 --- a/src/db/table/operations/update/mod.rs +++ b/src/db/table/operations/update/mod.rs @@ -57,6 +57,7 @@ mod tests { }; use crate::db::table::test_utils::{assert_table_rows_eq_unordered, default_table}; use crate::interpreter::ast::ColumnValue; + use crate::interpreter::ast::SelectStatementColumn; use crate::interpreter::ast::{ LimitClause, Operand, Operator, OrderByClause, OrderByDirection, SelectableStack, SelectableStackElement, WhereCondition, WhereStackElement, @@ -166,9 +167,9 @@ mod tests { })]), order_by_clause: Some(OrderByClause { columns: SelectableStack { - selectables: vec![SelectableStackElement::Column("id".to_string())], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("id".to_string()))], }, - column_names: vec!["id".to_string()], + column_names: vec![SelectStatementColumn::new("id".to_string())], directions: vec![OrderByDirection::Desc], }), limit_clause: Some(LimitClause { diff --git a/src/interpreter/ast/delete_statement.rs b/src/interpreter/ast/delete_statement.rs index 5cf0e5b..541cd66 100644 --- a/src/interpreter/ast/delete_statement.rs +++ b/src/interpreter/ast/delete_statement.rs @@ -38,6 +38,7 @@ mod tests { use crate::interpreter::ast::OrderByDirection; use crate::interpreter::ast::SelectableStack; use crate::interpreter::ast::SelectableStackElement; + use crate::interpreter::ast::SelectStatementColumn; use crate::interpreter::ast::WhereCondition; use crate::interpreter::ast::WhereStackElement; use crate::interpreter::ast::test_utils::token; @@ -98,9 +99,9 @@ mod tests { })]), order_by_clause: Some(OrderByClause { columns: SelectableStack { - selectables: vec![SelectableStackElement::Column("id".to_string())], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("id".to_string()))], }, - column_names: vec!["id".to_string()], + column_names: vec![SelectStatementColumn::new("id".to_string())], directions: vec![OrderByDirection::Asc], }), limit_clause: Some(LimitClause { diff --git a/src/interpreter/ast/helpers/common.rs b/src/interpreter/ast/helpers/common.rs index 79cd50f..fd4ea86 100644 --- a/src/interpreter/ast/helpers/common.rs +++ b/src/interpreter/ast/helpers/common.rs @@ -1,6 +1,6 @@ use crate::interpreter::{ ast::{ - ExistenceCheck, LogicalOperator, MathOperator, Operator, OrderByDirection, SelectableStack, + ExistenceCheck, LogicalOperator, MathOperator, Operator, OrderByDirection, SelectableStack, SelectStatementColumn, SelectableStackElement, helpers::token::token_to_value, parser::Parser, }, tokenizer::token::TokenTypes, @@ -24,11 +24,19 @@ pub fn get_table_name(parser: &mut Parser) -> Result { Ok(result) } +pub const TOKENS_NEEDING_SPECIAL_HANDLING: [TokenTypes; 5] = [ + TokenTypes::From, + TokenTypes::SemiColon, + TokenTypes::Where, + TokenTypes::Order, + TokenTypes::Limit, +]; + pub fn get_selectables( parser: &mut Parser, allow_multiple: bool, order_by_directions: &mut Option<&mut Vec>, - selectable_names: &mut Option<&mut Vec>, + selectable_names: &mut Option<&mut Vec>, ) -> Result { #[derive(PartialEq)] enum ExtendedSelectableStackElement { @@ -55,15 +63,7 @@ pub fn get_selectables( // Tokens needing special handling // TODO: more tokens should be added here (e.g. Group for GROUP BY) - if [ - TokenTypes::From, - TokenTypes::SemiColon, - TokenTypes::Where, - TokenTypes::Order, - TokenTypes::Limit, - ] - .contains(&token.token_type) - { + if TOKENS_NEEDING_SPECIAL_HANDLING.contains(&token.token_type) { // Default ordering is ASC if !expect_new_value && let Some(order_by_directions_vector) = order_by_directions { order_by_directions_vector.push(OrderByDirection::Asc); @@ -85,7 +85,7 @@ pub fn get_selectables( if !allow_multiple { return Err("Unexpected token: COMMA".to_string()); } else if let Some(selectable_names_vector) = selectable_names { - selectable_names_vector.push(current_name); + selectable_names_vector.push(SelectStatementColumn::new(current_name)); } // Default ordering is ASC if !expect_new_value && let Some(order_by_directions_vector) = order_by_directions { @@ -226,7 +226,7 @@ pub fn get_selectables( TokenTypes::HexLiteral => SelectableStackElement::Value(token_to_value(parser)?), TokenTypes::Null => SelectableStackElement::Value(token_to_value(parser)?), // TODO: handle ValueList (arrays) - TokenTypes::Identifier => SelectableStackElement::Column(token.value.to_string()), // TODO: verify it's a column, AND handle multi-tokens columns with AS (table_name.column_name) + TokenTypes::Identifier => SelectableStackElement::Column(SelectStatementColumn::new(token.value.to_string())), // TODO: verify it's a column, AND handle multi-tokens columns with AS (table_name.column_name) _ => return Err(parser.format_error()), // TODO: better error handling }; output.push(element); @@ -245,7 +245,7 @@ pub fn get_selectables( } if let Some(selectable_names_vector) = selectable_names { - selectable_names_vector.push(current_name); + selectable_names_vector.push(SelectStatementColumn::new(current_name)); } Ok(SelectableStack { diff --git a/src/interpreter/ast/helpers/order_by_clause.rs b/src/interpreter/ast/helpers/order_by_clause.rs index 49f1f3c..fd20703 100644 --- a/src/interpreter/ast/helpers/order_by_clause.rs +++ b/src/interpreter/ast/helpers/order_by_clause.rs @@ -32,6 +32,7 @@ pub fn get_order_by(parser: &mut Parser) -> Result, String #[cfg(test)] mod tests { use super::*; + use crate::interpreter::ast::SelectStatementColumn; use crate::interpreter::ast::test_utils::token; use crate::interpreter::ast::{OrderByDirection, SelectableStack, SelectableStackElement}; @@ -51,9 +52,9 @@ mod tests { let order_by_clause = result.unwrap(); let expected = Some(OrderByClause { columns: SelectableStack { - selectables: vec![SelectableStackElement::Column("id".to_string())], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("id".to_string()))], }, - column_names: vec!["id".to_string()], + column_names: vec![SelectStatementColumn::new("id".to_string())], directions: vec![OrderByDirection::Asc], }); assert_eq!(expected, order_by_clause); @@ -96,11 +97,11 @@ mod tests { let expected = Some(OrderByClause { columns: SelectableStack { selectables: vec![ - SelectableStackElement::Column("id".to_string()), - SelectableStackElement::Column("name".to_string()), + SelectableStackElement::Column(SelectStatementColumn::new("id".to_string())), + SelectableStackElement::Column(SelectStatementColumn::new("name".to_string())), ], }, - column_names: vec!["id".to_string(), "name".to_string()], + column_names: vec![SelectStatementColumn::new("id".to_string()), SelectStatementColumn::new("name".to_string())], directions: vec![OrderByDirection::Asc, OrderByDirection::Desc], }); assert_eq!(expected, order_by_clause); diff --git a/src/interpreter/ast/helpers/select_statement.rs b/src/interpreter/ast/helpers/select_statement.rs index 24db22f..3739800 100644 --- a/src/interpreter/ast/helpers/select_statement.rs +++ b/src/interpreter/ast/helpers/select_statement.rs @@ -1,6 +1,6 @@ use crate::interpreter::{ ast::{ - SelectMode, SelectStatement, SelectableStack, WhereStackElement, + SelectMode, SelectStatement, SelectableStack, SelectStatementColumn, WhereStackElement, helpers::{ common::{get_selectables, get_table_name}, limit_clause::get_limit, @@ -41,8 +41,8 @@ pub fn get_statement(parser: &mut Parser) -> Result { }); } -fn get_columns_and_names(parser: &mut Parser) -> Result<(SelectableStack, Vec), String> { - let mut column_names: Vec = vec![]; +fn get_columns_and_names(parser: &mut Parser) -> Result<(SelectableStack, Vec), String> { + let mut column_names: Vec = vec![]; Ok(( get_selectables(parser, true, &mut None, &mut Some(&mut column_names))?, column_names, @@ -85,7 +85,7 @@ mod tests { columns: SelectableStack { selectables: vec![SelectableStackElement::All], }, - column_names: vec!["*".to_string()], + column_names: vec![SelectStatementColumn::new("*".to_string())], where_clause: None, order_by_clause: None, limit_clause: None, @@ -113,9 +113,9 @@ mod tests { table_name: "guests".to_string(), mode: SelectMode::All, columns: SelectableStack { - selectables: vec![SelectableStackElement::Column("id".to_string())], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("id".to_string()))], }, - column_names: vec!["id".to_string()], + column_names: vec![SelectStatementColumn::new("id".to_string())], where_clause: None, order_by_clause: None, limit_clause: None, @@ -146,11 +146,11 @@ mod tests { mode: SelectMode::All, columns: SelectableStack { selectables: vec![ - SelectableStackElement::Column("id".to_string()), - SelectableStackElement::Column("name".to_string()), + SelectableStackElement::Column(SelectStatementColumn::new("id".to_string())), + SelectableStackElement::Column(SelectStatementColumn::new("name".to_string())), ], }, - column_names: vec!["id".to_string(), "name".to_string()], + column_names: vec![SelectStatementColumn::new("id".to_string()), SelectStatementColumn::new("name".to_string())], where_clause: None, order_by_clause: None, limit_clause: None, @@ -193,9 +193,9 @@ mod tests { table_name: "guests".to_string(), mode: SelectMode::All, columns: SelectableStack { - selectables: vec![SelectableStackElement::Column("id".to_string())], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("id".to_string()))], }, - column_names: vec!["id".to_string()], + column_names: vec![SelectStatementColumn::new("id".to_string())], where_clause: Some(vec![WhereStackElement::Condition(WhereCondition { l_side: Operand::Identifier("id".to_string()), operator: Operator::Equals, @@ -204,12 +204,12 @@ mod tests { order_by_clause: Some(OrderByClause { columns: SelectableStack { selectables: vec![ - SelectableStackElement::Column("id".to_string()), - SelectableStackElement::Column("name".to_string()), - SelectableStackElement::Column("age".to_string()), + SelectableStackElement::Column(SelectStatementColumn::new("id".to_string())), + SelectableStackElement::Column(SelectStatementColumn::new("name".to_string())), + SelectableStackElement::Column(SelectStatementColumn::new("age".to_string())), ], }, - column_names: vec!["id".to_string(), "name".to_string(), "age".to_string()], + column_names: vec![SelectStatementColumn::new("id".to_string()), SelectStatementColumn::new("name".to_string()), SelectStatementColumn::new("age".to_string())], directions: vec![ OrderByDirection::Asc, OrderByDirection::Desc, @@ -243,10 +243,10 @@ mod tests { statement, SelectStatement { table_name: "guests".to_string(), - column_names: vec!["id".to_string()], + column_names: vec![SelectStatementColumn::new("id".to_string())], mode: SelectMode::Distinct, columns: SelectableStack { - selectables: vec![SelectableStackElement::Column("id".to_string())], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("id".to_string()))], }, where_clause: None, order_by_clause: None, diff --git a/src/interpreter/ast/mod.rs b/src/interpreter/ast/mod.rs index f606215..724fee6 100644 --- a/src/interpreter/ast/mod.rs +++ b/src/interpreter/ast/mod.rs @@ -105,7 +105,7 @@ impl SetOperator { #[derive(Debug, PartialEq, Clone)] pub struct SelectStatement { pub table_name: String, - pub column_names: Vec, + pub column_names: Vec, pub mode: SelectMode, pub columns: SelectableStack, pub where_clause: Option>, @@ -113,6 +113,19 @@ pub struct SelectStatement { pub limit_clause: Option, } +#[derive(Debug, PartialEq, Clone)] +pub struct SelectStatementColumn { + pub column_name: String, + pub alias: Option, + pub table_name: Option, +} + +impl SelectStatementColumn { + pub fn new(column_name: String) -> Self { + Self { column_name, alias: None, table_name: None } + } +} + #[derive(Debug, PartialEq, Clone)] pub struct DeleteStatement { pub table_name: String, @@ -216,7 +229,7 @@ pub struct SelectableStack { #[derive(Debug, PartialEq, Clone)] pub enum SelectableStackElement { All, - Column(String), + Column(SelectStatementColumn), Value(Value), ValueList(Vec), // TODO: add column as data type in Value Function(FunctionSignature), @@ -316,7 +329,7 @@ pub enum OrderByDirection { #[derive(Debug, PartialEq, Clone)] pub struct OrderByClause { pub columns: SelectableStack, - pub column_names: Vec, + pub column_names: Vec, pub directions: Vec, } @@ -483,7 +496,7 @@ mod tests { columns: SelectableStack { selectables: vec![SelectableStackElement::All], }, - column_names: vec!["*".to_string()], + column_names: vec![SelectStatementColumn::new("*".to_string())], where_clause: None, order_by_clause: None, limit_clause: None, @@ -577,7 +590,7 @@ mod tests { columns: SelectableStack { selectables: vec![SelectableStackElement::All], }, - column_names: vec!["*".to_string()], + column_names: vec![SelectStatementColumn::new("*".to_string())], where_clause: None, order_by_clause: None, limit_clause: None, diff --git a/src/interpreter/ast/parser.rs b/src/interpreter/ast/parser.rs index c26759c..94e455d 100644 --- a/src/interpreter/ast/parser.rs +++ b/src/interpreter/ast/parser.rs @@ -130,7 +130,7 @@ mod tests { use crate::interpreter::ast::statement_builder::MockStatementBuilder; use crate::interpreter::ast::test_utils::{token, token_with_location}; use crate::interpreter::ast::{ - CreateTableStatement, InsertIntoStatement, SelectMode, SelectStatement, + CreateTableStatement, InsertIntoStatement, SelectMode, SelectStatement, SelectStatementColumn, SelectStatementStack, SelectStatementStackElement, SelectableStack, SelectableStackElement, }; @@ -198,7 +198,7 @@ mod tests { columns: SelectableStack { selectables: vec![SelectableStackElement::All], }, - column_names: vec!["*".to_string()], + column_names: vec![SelectStatementColumn::new("*".to_string())], where_clause: None, order_by_clause: None, limit_clause: None, diff --git a/src/interpreter/ast/select_statement_stack.rs b/src/interpreter/ast/select_statement_stack.rs index add3f12..cab348e 100644 --- a/src/interpreter/ast/select_statement_stack.rs +++ b/src/interpreter/ast/select_statement_stack.rs @@ -171,6 +171,7 @@ mod tests { use crate::interpreter::ast::SelectableStack; use crate::interpreter::ast::SelectableStackElement; use crate::interpreter::ast::SetOperator; + use crate::interpreter::ast::SelectStatementColumn; use crate::interpreter::ast::WhereCondition; use crate::interpreter::ast::WhereStackElement; use crate::interpreter::ast::test_utils::token; @@ -197,7 +198,7 @@ mod tests { columns: SelectableStack { selectables: vec![SelectableStackElement::All], }, - column_names: vec!["*".to_string()], + column_names: vec![SelectStatementColumn::new("*".to_string())], where_clause: Some(vec![WhereStackElement::Condition(WhereCondition { l_side: Operand::Identifier("id".to_string()), operator: Operator::Equals, @@ -359,9 +360,9 @@ mod tests { table_name: "employees".to_string(), mode: SelectMode::All, columns: SelectableStack { - selectables: vec![SelectableStackElement::Column("name".to_string())], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("name".to_string()))], }, - column_names: vec!["name".to_string()], + column_names: vec![SelectStatementColumn::new("name".to_string())], where_clause: Some(vec![WhereStackElement::Condition(WhereCondition { l_side: Operand::Identifier("name".to_string()), operator: Operator::Equals, @@ -374,9 +375,9 @@ mod tests { table_name: "employees".to_string(), mode: SelectMode::All, columns: SelectableStack { - selectables: vec![SelectableStackElement::Column("name".to_string())], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("name".to_string()))], }, - column_names: vec!["name".to_string()], + column_names: vec![SelectStatementColumn::new("name".to_string())], where_clause: Some(vec![WhereStackElement::Condition(WhereCondition { l_side: Operand::Identifier("name".to_string()), operator: Operator::Equals, @@ -389,9 +390,9 @@ mod tests { ], order_by_clause: Some(OrderByClause { columns: SelectableStack { - selectables: vec![SelectableStackElement::Column("name".to_string())], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("name".to_string()))], }, - column_names: vec!["name".to_string()], + column_names: vec![SelectStatementColumn::new("name".to_string())], directions: vec![OrderByDirection::Asc], }), limit_clause: Some(LimitClause { @@ -431,9 +432,9 @@ mod tests { ], order_by_clause: Some(OrderByClause { columns: SelectableStack { - selectables: vec![SelectableStackElement::Column("name".to_string())], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("name".to_string()))], }, - column_names: vec!["name".to_string()], + column_names: vec![SelectStatementColumn::new("name".to_string())], directions: vec![OrderByDirection::Asc], }), limit_clause: Some(LimitClause { diff --git a/src/interpreter/ast/statement_builder.rs b/src/interpreter/ast/statement_builder.rs index 039430e..27adec1 100644 --- a/src/interpreter/ast/statement_builder.rs +++ b/src/interpreter/ast/statement_builder.rs @@ -76,7 +76,7 @@ impl StatementBuilder for DefaultStatementBuilder { pub struct MockStatementBuilder; #[cfg(test)] use crate::interpreter::ast::{ - CreateTableStatement, InsertIntoStatement, SelectMode, SelectStatement, SelectStatementStack, + CreateTableStatement, InsertIntoStatement, SelectMode, SelectStatement, SelectStatementColumn, SelectStatementStack, SelectStatementStackElement, SelectableStack, SelectableStackElement, }; @@ -113,7 +113,7 @@ impl StatementBuilder for MockStatementBuilder { columns: SelectableStack { selectables: vec![SelectableStackElement::All], }, - column_names: vec!["*".to_string()], + column_names: vec![SelectStatementColumn::new("*".to_string())], where_clause: None, order_by_clause: None, limit_clause: None, diff --git a/src/interpreter/ast/update_statement.rs b/src/interpreter/ast/update_statement.rs index 7899cfd..5b6d0ff 100644 --- a/src/interpreter/ast/update_statement.rs +++ b/src/interpreter/ast/update_statement.rs @@ -72,6 +72,7 @@ mod tests { use crate::interpreter::ast::OrderByDirection; use crate::interpreter::ast::SelectableStack; use crate::interpreter::ast::SelectableStackElement; + use crate::interpreter::ast::SelectStatementColumn; use crate::interpreter::ast::WhereCondition; use crate::interpreter::ast::WhereStackElement; use crate::interpreter::ast::test_utils::token; @@ -230,9 +231,9 @@ mod tests { })]), order_by_clause: Some(OrderByClause { columns: SelectableStack { - selectables: vec![SelectableStackElement::Column("id".to_string())], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("id".to_string()))], }, - column_names: vec!["id".to_string()], + column_names: vec![SelectStatementColumn::new("id".to_string())], directions: vec![OrderByDirection::Asc], }), limit_clause: Some(LimitClause { From e6438b4aa16ab9d48b6d8aa785da7a4110640b11 Mon Sep 17 00:00:00 2001 From: Fletcher555 Date: Sun, 21 Sep 2025 12:16:05 -0400 Subject: [PATCH 2/6] Simple alias working within AST --- src/interpreter/ast/helpers/common.rs | 25 ++++++++- .../ast/helpers/order_by_clause.rs | 1 + .../ast/helpers/select_statement.rs | 56 +++++++++++++++---- tatus | 50 +++++++++++++++++ 4 files changed, 117 insertions(+), 15 deletions(-) create mode 100644 tatus diff --git a/src/interpreter/ast/helpers/common.rs b/src/interpreter/ast/helpers/common.rs index fd4ea86..b63ed46 100644 --- a/src/interpreter/ast/helpers/common.rs +++ b/src/interpreter/ast/helpers/common.rs @@ -47,6 +47,7 @@ pub fn get_selectables( let mut operators: Vec = vec![]; let mut depth = 0; let mut current_name = "".to_string(); + let mut current_alias: Option = None; let mut first = true; let mut expect_new_value = false; // Will be set after a valid ASC or DESC to ensure proper syntax @@ -85,7 +86,9 @@ pub fn get_selectables( if !allow_multiple { return Err("Unexpected token: COMMA".to_string()); } else if let Some(selectable_names_vector) = selectable_names { - selectable_names_vector.push(SelectStatementColumn::new(current_name)); + let mut column = SelectStatementColumn::new(current_name.clone()); + column.alias = current_alias.clone(); + selectable_names_vector.push(column); } // Default ordering is ASC if !expect_new_value && let Some(order_by_directions_vector) = order_by_directions { @@ -93,6 +96,7 @@ pub fn get_selectables( } expect_new_value = false; current_name = "".to_string(); + current_alias = None; } else { current_name += token.value; } @@ -226,7 +230,20 @@ pub fn get_selectables( TokenTypes::HexLiteral => SelectableStackElement::Value(token_to_value(parser)?), TokenTypes::Null => SelectableStackElement::Value(token_to_value(parser)?), // TODO: handle ValueList (arrays) - TokenTypes::Identifier => SelectableStackElement::Column(SelectStatementColumn::new(token.value.to_string())), // TODO: verify it's a column, AND handle multi-tokens columns with AS (table_name.column_name) + TokenTypes::Identifier => { + let mut column = SelectStatementColumn::new(token.value.to_string()); + if let Ok(peek_token) = parser.peek_token() { + if peek_token.token_type == TokenTypes::As { + parser.advance()?; + parser.advance()?; + expect_token_type(parser, TokenTypes::Identifier)?; + let alias = parser.current_token()?.value.to_string(); + column.alias = Some(alias.clone()); + current_alias = Some(alias); + } + } + SelectableStackElement::Column(column) // TODO: verify it's a column, AND handle multi-tokens columns with AS (table_name.column_name) + } _ => return Err(parser.format_error()), // TODO: better error handling }; output.push(element); @@ -245,7 +262,9 @@ pub fn get_selectables( } if let Some(selectable_names_vector) = selectable_names { - selectable_names_vector.push(SelectStatementColumn::new(current_name)); + let mut column = SelectStatementColumn::new(current_name); + column.alias = current_alias; + selectable_names_vector.push(column); } Ok(SelectableStack { diff --git a/src/interpreter/ast/helpers/order_by_clause.rs b/src/interpreter/ast/helpers/order_by_clause.rs index fd20703..05c988c 100644 --- a/src/interpreter/ast/helpers/order_by_clause.rs +++ b/src/interpreter/ast/helpers/order_by_clause.rs @@ -48,6 +48,7 @@ mod tests { ]; let mut parser = Parser::new(tokens); let result = get_order_by(&mut parser); + println!("{:?}", result); assert!(result.is_ok()); let order_by_clause = result.unwrap(); let expected = Some(OrderByClause { diff --git a/src/interpreter/ast/helpers/select_statement.rs b/src/interpreter/ast/helpers/select_statement.rs index 3739800..aaaa28c 100644 --- a/src/interpreter/ast/helpers/select_statement.rs +++ b/src/interpreter/ast/helpers/select_statement.rs @@ -239,19 +239,51 @@ mod tests { let result = get_statement(&mut parser); assert!(result.is_ok()); let statement = result.unwrap(); + let expected = SelectStatement { + table_name: "guests".to_string(), + column_names: vec![SelectStatementColumn::new("id".to_string())], + mode: SelectMode::Distinct, + columns: SelectableStack { + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("id".to_string()))], + }, + where_clause: None, + order_by_clause: None, + limit_clause: None, + }; assert_eq!( - statement, - SelectStatement { - table_name: "guests".to_string(), - column_names: vec![SelectStatementColumn::new("id".to_string())], - mode: SelectMode::Distinct, - columns: SelectableStack { - selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("id".to_string()))], - }, - where_clause: None, - order_by_clause: None, - limit_clause: None, - } + expected, + statement ); } + + #[test] + fn select_statement_with_column_alias_is_generated_correctly() { + // SELECT id AS user_id FROM guests; + let tokens = vec![ + token(TokenTypes::Select, "SELECT"), + token(TokenTypes::Identifier, "id"), + token(TokenTypes::As, "AS"), + token(TokenTypes::Identifier, "user_id"), + token(TokenTypes::From, "FROM"), + token(TokenTypes::Identifier, "guests"), + token(TokenTypes::SemiColon, ";"), + ]; + let mut parser = Parser::new(tokens); + let result = get_statement(&mut parser); + println!("{:?}", result); + assert!(result.is_ok()); + let statement = result.unwrap(); + let expected = SelectStatement { + table_name: "guests".to_string(), + mode: SelectMode::All, + columns: SelectableStack { + selectables: vec![SelectableStackElement::Column(SelectStatementColumn{column_name: "id".to_string(), alias: Some("user_id".to_string()), table_name: None})], + }, + column_names: vec![SelectStatementColumn{column_name: "id".to_string(), alias: Some("user_id".to_string()), table_name: None}], + where_clause: None, + order_by_clause: None, + limit_clause: None, + }; + assert_eq!(expected, statement); + } } diff --git a/tatus b/tatus new file mode 100644 index 0000000..77c4601 --- /dev/null +++ b/tatus @@ -0,0 +1,50 @@ +diff --git a/src/interpreter/ast/helpers/common.rs b/src/interpreter/ast/helpers/common.rs +index fd4ea86..997724c 100644 +--- a/src/interpreter/ast/helpers/common.rs ++++ b/src/interpreter/ast/helpers/common.rs +@@ -85,7 +85,15 @@ pub fn get_selectables( + if !allow_multiple { + return Err("Unexpected token: COMMA".to_string()); + } else if let Some(selectable_names_vector) = selectable_names { +- selectable_names_vector.push(SelectStatementColumn::new(current_name)); ++ let mut column = SelectStatementColumn::new(current_name); ++ let peek_token = parser.peek_token()?; ++ if peek_token.token_type == TokenTypes::As { ++ parser.advance()?; ++ parser.advance()?; ++ expect_token_type(parser, TokenTypes::Identifier)?; ++ column.alias = Some(parser.current_token()?.value.to_string()); ++ } ++ selectable_names_vector.push(column); + } + // Default ordering is ASC + if !expect_new_value && let Some(order_by_directions_vector) = order_by_directions { +@@ -245,7 +253,15 @@ pub fn get_selectables( + } +  + if let Some(selectable_names_vector) = selectable_names { +- selectable_names_vector.push(SelectStatementColumn::new(current_name)); ++ let mut column = SelectStatementColumn::new(current_name); ++ let peek_token = parser.peek_token()?; ++ if peek_token.token_type == TokenTypes::As { ++ parser.advance()?; ++ parser.advance()?; ++ expect_token_type(parser, TokenTypes::Identifier)?; ++ column.alias = Some(parser.current_token()?.value.to_string()); ++ } ++ selectable_names_vector.push(column); + } +  + Ok(SelectableStack { +diff --git a/src/interpreter/ast/helpers/order_by_clause.rs b/src/interpreter/ast/helpers/order_by_clause.rs +index fd20703..05c988c 100644 +--- a/src/interpreter/ast/helpers/order_by_clause.rs ++++ b/src/interpreter/ast/helpers/order_by_clause.rs +@@ -48,6 +48,7 @@ mod tests { + ]; + let mut parser = Parser::new(tokens); + let result = get_order_by(&mut parser); ++ println!("{:?}", result); + assert!(result.is_ok()); + let order_by_clause = result.unwrap(); + let expected = Some(OrderByClause { From 58d9fed2891e5b19e06f95888848ae1eafa8d2dd Mon Sep 17 00:00:00 2001 From: Fletcher555 Date: Sun, 21 Sep 2025 12:17:40 -0400 Subject: [PATCH 3/6] Add additional test and fix formatting error --- src/db/table/operations/delete/mod.rs | 8 +- .../operations/helpers/order_by_clause.rs | 8 +- src/db/table/operations/select/mod.rs | 14 +- .../operations/select/select_statement.rs | 18 ++- src/db/table/operations/update/mod.rs | 4 +- src/interpreter/ast/delete_statement.rs | 6 +- src/interpreter/ast/helpers/common.rs | 5 +- .../ast/helpers/order_by_clause.rs | 9 +- .../ast/helpers/select_statement.rs | 134 +++++++++++++++--- src/interpreter/ast/mod.rs | 6 +- src/interpreter/ast/parser.rs | 5 +- src/interpreter/ast/select_statement_stack.rs | 18 ++- src/interpreter/ast/statement_builder.rs | 4 +- src/interpreter/ast/update_statement.rs | 6 +- 14 files changed, 195 insertions(+), 50 deletions(-) diff --git a/src/db/table/operations/delete/mod.rs b/src/db/table/operations/delete/mod.rs index 84e8a82..2e7e3fe 100644 --- a/src/db/table/operations/delete/mod.rs +++ b/src/db/table/operations/delete/mod.rs @@ -166,7 +166,9 @@ mod tests { })]), order_by_clause: Some(OrderByClause { columns: SelectableStack { - selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("id".to_string()))], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new( + "id".to_string(), + ))], }, column_names: vec![SelectStatementColumn::new("id".to_string())], directions: vec![OrderByDirection::Desc], @@ -318,7 +320,9 @@ mod tests { })]), order_by_clause: Some(OrderByClause { columns: SelectableStack { - selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("id".to_string()))], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new( + "id".to_string(), + ))], }, column_names: vec![SelectStatementColumn::new("id".to_string())], directions: vec![OrderByDirection::Desc], diff --git a/src/db/table/operations/helpers/order_by_clause.rs b/src/db/table/operations/helpers/order_by_clause.rs index 902dffe..7f66cb8 100644 --- a/src/db/table/operations/helpers/order_by_clause.rs +++ b/src/db/table/operations/helpers/order_by_clause.rs @@ -40,10 +40,10 @@ mod tests { use crate::db::table::core::{row::Row, value::Value}; use crate::interpreter::ast::OrderByClause; use crate::interpreter::ast::OrderByDirection; + use crate::interpreter::ast::SelectStatementColumn; use crate::interpreter::ast::SelectableStack; use crate::interpreter::ast::SelectableStackElement; - use crate::interpreter::ast::SelectStatementColumn; - + #[test] fn apply_order_by_from_precomputed_single_column_asc() { let mut to_order = vec!["second", "fourth", "third", "first"]; @@ -57,7 +57,9 @@ mod tests { let order_by_clause = OrderByClause { columns: SelectableStack { - selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("age".to_string()))], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new( + "age".to_string(), + ))], }, column_names: vec![SelectStatementColumn::new("age".to_string())], directions: vec![OrderByDirection::Asc], diff --git a/src/db/table/operations/select/mod.rs b/src/db/table/operations/select/mod.rs index 1046285..93c7e6f 100644 --- a/src/db/table/operations/select/mod.rs +++ b/src/db/table/operations/select/mod.rs @@ -21,8 +21,14 @@ pub fn select_statement_stack( match element { SelectStatementStackElement::SelectStatement(select_statement) => { let table = database.get_table(&select_statement.table_name)?; - let expanded_column_names = - expand_all_column_names(table, select_statement.column_names.iter().map(|column| &column.column_name).collect::>())?; + let expanded_column_names = expand_all_column_names( + table, + select_statement + .column_names + .iter() + .map(|column| &column.column_name) + .collect::>(), + )?; match &column_names { Some(column_names) => { if expanded_column_names.len() != column_names.len() { @@ -144,8 +150,8 @@ mod tests { use crate::db::table::core::value::Value; use crate::db::table::test_utils::default_database; use crate::interpreter::ast::{ - LogicalOperator, Operand, Operator, SelectMode, SelectStatement, SelectableStack, SelectStatementColumn, - SelectableStackElement, WhereCondition, WhereStackElement, + LogicalOperator, Operand, Operator, SelectMode, SelectStatement, SelectStatementColumn, + SelectableStack, SelectableStackElement, WhereCondition, WhereStackElement, }; #[test] diff --git a/src/db/table/operations/select/select_statement.rs b/src/db/table/operations/select/select_statement.rs index c47cfd1..0e24cdf 100644 --- a/src/db/table/operations/select/select_statement.rs +++ b/src/db/table/operations/select/select_statement.rs @@ -143,7 +143,10 @@ mod tests { SelectableStackElement::Column(SelectStatementColumn::new("age".to_string())), ], }, - column_names: vec![SelectStatementColumn::new("name".to_string()), SelectStatementColumn::new("age".to_string())], + column_names: vec![ + SelectStatementColumn::new("name".to_string()), + SelectStatementColumn::new("age".to_string()), + ], where_clause: None, order_by_clause: None, limit_clause: None, @@ -200,7 +203,10 @@ mod tests { SelectableStackElement::Column(SelectStatementColumn::new("age".to_string())), ], }, - column_names: vec![SelectStatementColumn::new("name".to_string()), SelectStatementColumn::new("age".to_string())], + column_names: vec![ + SelectStatementColumn::new("name".to_string()), + SelectStatementColumn::new("age".to_string()), + ], where_clause: Some(vec![WhereStackElement::Condition(WhereCondition { l_side: Operand::Identifier("money".to_string()), operator: Operator::Equals, @@ -285,7 +291,9 @@ mod tests { where_clause: None, order_by_clause: Some(OrderByClause { columns: SelectableStack { - selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("money".to_string()))], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new( + "money".to_string(), + ))], }, column_names: vec![SelectStatementColumn::new("money".to_string())], directions: vec![OrderByDirection::Desc], @@ -351,7 +359,9 @@ mod tests { column_names: vec![SelectStatementColumn::new("name".to_string())], mode: SelectMode::Distinct, columns: SelectableStack { - selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("name".to_string()))], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new( + "name".to_string(), + ))], }, where_clause: None, order_by_clause: None, diff --git a/src/db/table/operations/update/mod.rs b/src/db/table/operations/update/mod.rs index c18ac6d..de3021b 100644 --- a/src/db/table/operations/update/mod.rs +++ b/src/db/table/operations/update/mod.rs @@ -167,7 +167,9 @@ mod tests { })]), order_by_clause: Some(OrderByClause { columns: SelectableStack { - selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("id".to_string()))], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new( + "id".to_string(), + ))], }, column_names: vec![SelectStatementColumn::new("id".to_string())], directions: vec![OrderByDirection::Desc], diff --git a/src/interpreter/ast/delete_statement.rs b/src/interpreter/ast/delete_statement.rs index 541cd66..096bfc6 100644 --- a/src/interpreter/ast/delete_statement.rs +++ b/src/interpreter/ast/delete_statement.rs @@ -36,9 +36,9 @@ mod tests { use crate::interpreter::ast::Operator; use crate::interpreter::ast::OrderByClause; use crate::interpreter::ast::OrderByDirection; + use crate::interpreter::ast::SelectStatementColumn; use crate::interpreter::ast::SelectableStack; use crate::interpreter::ast::SelectableStackElement; - use crate::interpreter::ast::SelectStatementColumn; use crate::interpreter::ast::WhereCondition; use crate::interpreter::ast::WhereStackElement; use crate::interpreter::ast::test_utils::token; @@ -99,7 +99,9 @@ mod tests { })]), order_by_clause: Some(OrderByClause { columns: SelectableStack { - selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("id".to_string()))], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new( + "id".to_string(), + ))], }, column_names: vec![SelectStatementColumn::new("id".to_string())], directions: vec![OrderByDirection::Asc], diff --git a/src/interpreter/ast/helpers/common.rs b/src/interpreter/ast/helpers/common.rs index b63ed46..eb64114 100644 --- a/src/interpreter/ast/helpers/common.rs +++ b/src/interpreter/ast/helpers/common.rs @@ -1,7 +1,8 @@ use crate::interpreter::{ ast::{ - ExistenceCheck, LogicalOperator, MathOperator, Operator, OrderByDirection, SelectableStack, SelectStatementColumn, - SelectableStackElement, helpers::token::token_to_value, parser::Parser, + ExistenceCheck, LogicalOperator, MathOperator, Operator, OrderByDirection, + SelectStatementColumn, SelectableStack, SelectableStackElement, + helpers::token::token_to_value, parser::Parser, }, tokenizer::token::TokenTypes, }; diff --git a/src/interpreter/ast/helpers/order_by_clause.rs b/src/interpreter/ast/helpers/order_by_clause.rs index 05c988c..3fc7707 100644 --- a/src/interpreter/ast/helpers/order_by_clause.rs +++ b/src/interpreter/ast/helpers/order_by_clause.rs @@ -53,7 +53,9 @@ mod tests { let order_by_clause = result.unwrap(); let expected = Some(OrderByClause { columns: SelectableStack { - selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("id".to_string()))], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new( + "id".to_string(), + ))], }, column_names: vec![SelectStatementColumn::new("id".to_string())], directions: vec![OrderByDirection::Asc], @@ -102,7 +104,10 @@ mod tests { SelectableStackElement::Column(SelectStatementColumn::new("name".to_string())), ], }, - column_names: vec![SelectStatementColumn::new("id".to_string()), SelectStatementColumn::new("name".to_string())], + column_names: vec![ + SelectStatementColumn::new("id".to_string()), + SelectStatementColumn::new("name".to_string()), + ], directions: vec![OrderByDirection::Asc, OrderByDirection::Desc], }); assert_eq!(expected, order_by_clause); diff --git a/src/interpreter/ast/helpers/select_statement.rs b/src/interpreter/ast/helpers/select_statement.rs index aaaa28c..a9f325a 100644 --- a/src/interpreter/ast/helpers/select_statement.rs +++ b/src/interpreter/ast/helpers/select_statement.rs @@ -1,6 +1,6 @@ use crate::interpreter::{ ast::{ - SelectMode, SelectStatement, SelectableStack, SelectStatementColumn, WhereStackElement, + SelectMode, SelectStatement, SelectStatementColumn, SelectableStack, WhereStackElement, helpers::{ common::{get_selectables, get_table_name}, limit_clause::get_limit, @@ -41,7 +41,9 @@ pub fn get_statement(parser: &mut Parser) -> Result { }); } -fn get_columns_and_names(parser: &mut Parser) -> Result<(SelectableStack, Vec), String> { +fn get_columns_and_names( + parser: &mut Parser, +) -> Result<(SelectableStack, Vec), String> { let mut column_names: Vec = vec![]; Ok(( get_selectables(parser, true, &mut None, &mut Some(&mut column_names))?, @@ -113,7 +115,9 @@ mod tests { table_name: "guests".to_string(), mode: SelectMode::All, columns: SelectableStack { - selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("id".to_string()))], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new( + "id".to_string() + ))], }, column_names: vec![SelectStatementColumn::new("id".to_string())], where_clause: None, @@ -146,11 +150,18 @@ mod tests { mode: SelectMode::All, columns: SelectableStack { selectables: vec![ - SelectableStackElement::Column(SelectStatementColumn::new("id".to_string())), - SelectableStackElement::Column(SelectStatementColumn::new("name".to_string())), + SelectableStackElement::Column(SelectStatementColumn::new( + "id".to_string() + )), + SelectableStackElement::Column(SelectStatementColumn::new( + "name".to_string() + )), ], }, - column_names: vec![SelectStatementColumn::new("id".to_string()), SelectStatementColumn::new("name".to_string())], + column_names: vec![ + SelectStatementColumn::new("id".to_string()), + SelectStatementColumn::new("name".to_string()) + ], where_clause: None, order_by_clause: None, limit_clause: None, @@ -193,7 +204,9 @@ mod tests { table_name: "guests".to_string(), mode: SelectMode::All, columns: SelectableStack { - selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("id".to_string()))], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new( + "id".to_string(), + ))], }, column_names: vec![SelectStatementColumn::new("id".to_string())], where_clause: Some(vec![WhereStackElement::Condition(WhereCondition { @@ -204,12 +217,22 @@ mod tests { order_by_clause: Some(OrderByClause { columns: SelectableStack { selectables: vec![ - SelectableStackElement::Column(SelectStatementColumn::new("id".to_string())), - SelectableStackElement::Column(SelectStatementColumn::new("name".to_string())), - SelectableStackElement::Column(SelectStatementColumn::new("age".to_string())), + SelectableStackElement::Column(SelectStatementColumn::new( + "id".to_string(), + )), + SelectableStackElement::Column(SelectStatementColumn::new( + "name".to_string(), + )), + SelectableStackElement::Column(SelectStatementColumn::new( + "age".to_string(), + )), ], }, - column_names: vec![SelectStatementColumn::new("id".to_string()), SelectStatementColumn::new("name".to_string()), SelectStatementColumn::new("age".to_string())], + column_names: vec![ + SelectStatementColumn::new("id".to_string()), + SelectStatementColumn::new("name".to_string()), + SelectStatementColumn::new("age".to_string()), + ], directions: vec![ OrderByDirection::Asc, OrderByDirection::Desc, @@ -244,16 +267,15 @@ mod tests { column_names: vec![SelectStatementColumn::new("id".to_string())], mode: SelectMode::Distinct, columns: SelectableStack { - selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("id".to_string()))], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new( + "id".to_string(), + ))], }, where_clause: None, order_by_clause: None, limit_clause: None, }; - assert_eq!( - expected, - statement - ); + assert_eq!(expected, statement); } #[test] @@ -277,9 +299,85 @@ mod tests { table_name: "guests".to_string(), mode: SelectMode::All, columns: SelectableStack { - selectables: vec![SelectableStackElement::Column(SelectStatementColumn{column_name: "id".to_string(), alias: Some("user_id".to_string()), table_name: None})], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn { + column_name: "id".to_string(), + alias: Some("user_id".to_string()), + table_name: None, + })], }, - column_names: vec![SelectStatementColumn{column_name: "id".to_string(), alias: Some("user_id".to_string()), table_name: None}], + column_names: vec![SelectStatementColumn { + column_name: "id".to_string(), + alias: Some("user_id".to_string()), + table_name: None, + }], + where_clause: None, + order_by_clause: None, + limit_clause: None, + }; + assert_eq!(expected, statement); + } + + #[test] + fn select_statement_with_column_alias_and_table_name_is_generated_correctly() { + // SELECT name as some, id2, id AS user_id FROM guests; + let tokens = vec![ + token(TokenTypes::Select, "SELECT"), + token(TokenTypes::Identifier, "name"), + token(TokenTypes::As, "AS"), + token(TokenTypes::Identifier, "some"), + token(TokenTypes::Comma, ","), + token(TokenTypes::Identifier, "id2"), + token(TokenTypes::Comma, ","), + token(TokenTypes::Identifier, "id"), + token(TokenTypes::As, "AS"), + token(TokenTypes::Identifier, "user_id"), + token(TokenTypes::From, "FROM"), + token(TokenTypes::Identifier, "guests"), + token(TokenTypes::SemiColon, ";"), + ]; + let mut parser = Parser::new(tokens); + let result = get_statement(&mut parser); + assert!(result.is_ok()); + let statement = result.unwrap(); + let expected = SelectStatement { + table_name: "guests".to_string(), + mode: SelectMode::All, + columns: SelectableStack { + selectables: vec![ + SelectableStackElement::Column(SelectStatementColumn { + column_name: "name".to_string(), + alias: Some("some".to_string()), + table_name: None, + }), + SelectableStackElement::Column(SelectStatementColumn { + column_name: "id2".to_string(), + alias: None, + table_name: None, + }), + SelectableStackElement::Column(SelectStatementColumn { + column_name: "id".to_string(), + alias: Some("user_id".to_string()), + table_name: None, + }), + ], + }, + column_names: vec![ + SelectStatementColumn { + column_name: "name".to_string(), + alias: Some("some".to_string()), + table_name: None, + }, + SelectStatementColumn { + column_name: "id2".to_string(), + alias: None, + table_name: None, + }, + SelectStatementColumn { + column_name: "id".to_string(), + alias: Some("user_id".to_string()), + table_name: None, + }, + ], where_clause: None, order_by_clause: None, limit_clause: None, diff --git a/src/interpreter/ast/mod.rs b/src/interpreter/ast/mod.rs index 724fee6..0b35df0 100644 --- a/src/interpreter/ast/mod.rs +++ b/src/interpreter/ast/mod.rs @@ -122,7 +122,11 @@ pub struct SelectStatementColumn { impl SelectStatementColumn { pub fn new(column_name: String) -> Self { - Self { column_name, alias: None, table_name: None } + Self { + column_name, + alias: None, + table_name: None, + } } } diff --git a/src/interpreter/ast/parser.rs b/src/interpreter/ast/parser.rs index 94e455d..c0777dd 100644 --- a/src/interpreter/ast/parser.rs +++ b/src/interpreter/ast/parser.rs @@ -130,8 +130,9 @@ mod tests { use crate::interpreter::ast::statement_builder::MockStatementBuilder; use crate::interpreter::ast::test_utils::{token, token_with_location}; use crate::interpreter::ast::{ - CreateTableStatement, InsertIntoStatement, SelectMode, SelectStatement, SelectStatementColumn, - SelectStatementStack, SelectStatementStackElement, SelectableStack, SelectableStackElement, + CreateTableStatement, InsertIntoStatement, SelectMode, SelectStatement, + SelectStatementColumn, SelectStatementStack, SelectStatementStackElement, SelectableStack, + SelectableStackElement, }; #[test] diff --git a/src/interpreter/ast/select_statement_stack.rs b/src/interpreter/ast/select_statement_stack.rs index cab348e..904a467 100644 --- a/src/interpreter/ast/select_statement_stack.rs +++ b/src/interpreter/ast/select_statement_stack.rs @@ -168,10 +168,10 @@ mod tests { use crate::interpreter::ast::OrderByDirection; use crate::interpreter::ast::SelectMode; use crate::interpreter::ast::SelectStatement; + use crate::interpreter::ast::SelectStatementColumn; use crate::interpreter::ast::SelectableStack; use crate::interpreter::ast::SelectableStackElement; use crate::interpreter::ast::SetOperator; - use crate::interpreter::ast::SelectStatementColumn; use crate::interpreter::ast::WhereCondition; use crate::interpreter::ast::WhereStackElement; use crate::interpreter::ast::test_utils::token; @@ -360,7 +360,9 @@ mod tests { table_name: "employees".to_string(), mode: SelectMode::All, columns: SelectableStack { - selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("name".to_string()))], + selectables: vec![SelectableStackElement::Column( + SelectStatementColumn::new("name".to_string()), + )], }, column_names: vec![SelectStatementColumn::new("name".to_string())], where_clause: Some(vec![WhereStackElement::Condition(WhereCondition { @@ -375,7 +377,9 @@ mod tests { table_name: "employees".to_string(), mode: SelectMode::All, columns: SelectableStack { - selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("name".to_string()))], + selectables: vec![SelectableStackElement::Column( + SelectStatementColumn::new("name".to_string()), + )], }, column_names: vec![SelectStatementColumn::new("name".to_string())], where_clause: Some(vec![WhereStackElement::Condition(WhereCondition { @@ -390,7 +394,9 @@ mod tests { ], order_by_clause: Some(OrderByClause { columns: SelectableStack { - selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("name".to_string()))], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new( + "name".to_string(), + ))], }, column_names: vec![SelectStatementColumn::new("name".to_string())], directions: vec![OrderByDirection::Asc], @@ -432,7 +438,9 @@ mod tests { ], order_by_clause: Some(OrderByClause { columns: SelectableStack { - selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("name".to_string()))], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new( + "name".to_string(), + ))], }, column_names: vec![SelectStatementColumn::new("name".to_string())], directions: vec![OrderByDirection::Asc], diff --git a/src/interpreter/ast/statement_builder.rs b/src/interpreter/ast/statement_builder.rs index 27adec1..cd23c40 100644 --- a/src/interpreter/ast/statement_builder.rs +++ b/src/interpreter/ast/statement_builder.rs @@ -76,8 +76,8 @@ impl StatementBuilder for DefaultStatementBuilder { pub struct MockStatementBuilder; #[cfg(test)] use crate::interpreter::ast::{ - CreateTableStatement, InsertIntoStatement, SelectMode, SelectStatement, SelectStatementColumn, SelectStatementStack, - SelectStatementStackElement, SelectableStack, SelectableStackElement, + CreateTableStatement, InsertIntoStatement, SelectMode, SelectStatement, SelectStatementColumn, + SelectStatementStack, SelectStatementStackElement, SelectableStack, SelectableStackElement, }; #[cfg(test)] diff --git a/src/interpreter/ast/update_statement.rs b/src/interpreter/ast/update_statement.rs index 5b6d0ff..08acb59 100644 --- a/src/interpreter/ast/update_statement.rs +++ b/src/interpreter/ast/update_statement.rs @@ -70,9 +70,9 @@ mod tests { use crate::interpreter::ast::Operator; use crate::interpreter::ast::OrderByClause; use crate::interpreter::ast::OrderByDirection; + use crate::interpreter::ast::SelectStatementColumn; use crate::interpreter::ast::SelectableStack; use crate::interpreter::ast::SelectableStackElement; - use crate::interpreter::ast::SelectStatementColumn; use crate::interpreter::ast::WhereCondition; use crate::interpreter::ast::WhereStackElement; use crate::interpreter::ast::test_utils::token; @@ -231,7 +231,9 @@ mod tests { })]), order_by_clause: Some(OrderByClause { columns: SelectableStack { - selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new("id".to_string()))], + selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new( + "id".to_string(), + ))], }, column_names: vec![SelectStatementColumn::new("id".to_string())], directions: vec![OrderByDirection::Asc], From f0ee8329b70721f88c2db77de0eda5e5fcc3c3b1 Mon Sep 17 00:00:00 2001 From: Fletcher555 Date: Sun, 21 Sep 2025 13:18:15 -0400 Subject: [PATCH 4/6] Add table name aliasing --- src/db/database.rs | 4 +- src/db/table/operations/delete/mod.rs | 15 ++--- src/db/table/operations/select/mod.rs | 15 ++--- .../operations/select/select_statement.rs | 17 +++--- src/db/table/operations/update/mod.rs | 18 +++--- src/db/transactions/mod.rs | 4 +- src/interpreter/ast/alter_table_statement.rs | 2 +- src/interpreter/ast/create_statement.rs | 2 +- src/interpreter/ast/delete_statement.rs | 7 ++- src/interpreter/ast/drop_statement.rs | 2 +- src/interpreter/ast/helpers/common.rs | 19 +++++-- .../ast/helpers/select_statement.rs | 55 ++++++++++++++++--- src/interpreter/ast/insert_statement.rs | 2 +- src/interpreter/ast/mod.rs | 25 +++++++-- src/interpreter/ast/parser.rs | 4 +- src/interpreter/ast/select_statement_stack.rs | 7 ++- src/interpreter/ast/statement_builder.rs | 4 +- src/interpreter/ast/update_statement.rs | 11 ++-- src/interpreter/tokenizer/mod.rs | 21 +++++++ 19 files changed, 162 insertions(+), 72 deletions(-) diff --git a/src/db/database.rs b/src/db/database.rs index f94d684..e46991f 100644 --- a/src/db/database.rs +++ b/src/db/database.rs @@ -41,7 +41,7 @@ impl Database { } SqlStatement::UpdateStatement(statement) => { let is_transaction = self.transaction.in_transaction(); - let table = self.get_table_mut(&statement.table_name)?; + let table = self.get_table_mut(&statement.table_name.table_name)?; let rows_updated = update::update(table, statement, is_transaction)?; self.transaction .append_entry(sql_statement_clone, rows_updated)?; @@ -49,7 +49,7 @@ impl Database { } SqlStatement::DeleteStatement(statement) => { let is_transaction = self.transaction.in_transaction(); - let table = self.get_table_mut(&statement.table_name)?; + let table = self.get_table_mut(&statement.table_name.table_name)?; let rows_deleted = delete::delete(table, statement, is_transaction)?; self.transaction .append_entry(sql_statement_clone, rows_deleted)?; diff --git a/src/db/table/operations/delete/mod.rs b/src/db/table/operations/delete/mod.rs index 2e7e3fe..3130a28 100644 --- a/src/db/table/operations/delete/mod.rs +++ b/src/db/table/operations/delete/mod.rs @@ -66,6 +66,7 @@ mod tests { use crate::db::table::core::{row::Row, value::Value}; use crate::db::table::test_utils::{assert_table_rows_eq_unordered, default_table}; use crate::interpreter::ast::LimitClause; + use crate::interpreter::ast::SelectStatementTable; use crate::interpreter::ast::SelectStatementColumn; use crate::interpreter::ast::{ Operand, Operator, OrderByClause, OrderByDirection, SelectableStack, @@ -76,7 +77,7 @@ mod tests { fn delete_from_table_works_correctly() { let mut table = default_table(); let statement = DeleteStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), where_clause: Some(vec![WhereStackElement::Condition(WhereCondition { l_side: Operand::Identifier("id".to_string()), operator: Operator::Equals, @@ -158,7 +159,7 @@ mod tests { ]), ]); let statement = DeleteStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), where_clause: Some(vec![WhereStackElement::Condition(WhereCondition { l_side: Operand::Identifier("name".to_string()), operator: Operator::Equals, @@ -225,7 +226,7 @@ mod tests { fn delete_multiple_rows_works_correctly() { let mut table = default_table(); let statement = DeleteStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), where_clause: Some(vec![WhereStackElement::Condition(WhereCondition { l_side: Operand::Identifier("id".to_string()), operator: Operator::GreaterThan, @@ -251,7 +252,7 @@ mod tests { fn delete_all_rows_works_correctly() { let mut table = default_table(); let statement = DeleteStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), where_clause: None, order_by_clause: None, limit_clause: None, @@ -267,7 +268,7 @@ mod tests { let mut table = default_table(); table.set_rows(vec![]); let statement = DeleteStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), where_clause: None, order_by_clause: None, limit_clause: None, @@ -312,7 +313,7 @@ mod tests { ]), ]); let statement = DeleteStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), where_clause: Some(vec![WhereStackElement::Condition(WhereCondition { l_side: Operand::Identifier("age".to_string()), operator: Operator::GreaterEquals, @@ -369,7 +370,7 @@ mod tests { Value::Real(123.45), ])]); let statement = DeleteStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), where_clause: Some(vec![WhereStackElement::Condition(WhereCondition { l_side: Operand::Identifier("id".to_string()), operator: Operator::Equals, diff --git a/src/db/table/operations/select/mod.rs b/src/db/table/operations/select/mod.rs index 93c7e6f..3f269b7 100644 --- a/src/db/table/operations/select/mod.rs +++ b/src/db/table/operations/select/mod.rs @@ -20,7 +20,7 @@ pub fn select_statement_stack( for element in statement.elements { match element { SelectStatementStackElement::SelectStatement(select_statement) => { - let table = database.get_table(&select_statement.table_name)?; + let table = database.get_table(&select_statement.table_name.table_name)?; let expanded_column_names = expand_all_column_names( table, select_statement @@ -152,6 +152,7 @@ mod tests { use crate::interpreter::ast::{ LogicalOperator, Operand, Operator, SelectMode, SelectStatement, SelectStatementColumn, SelectableStack, SelectableStackElement, WhereCondition, WhereStackElement, + SelectStatementTable, }; #[test] @@ -160,7 +161,7 @@ mod tests { let statement = SelectStatementStack { elements: vec![SelectStatementStackElement::SelectStatement( SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::All], @@ -211,7 +212,7 @@ mod tests { let statement = SelectStatementStack { elements: vec![ SelectStatementStackElement::SelectStatement(SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::All], @@ -226,7 +227,7 @@ mod tests { limit_clause: None, }), SelectStatementStackElement::SelectStatement(SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::All], @@ -258,7 +259,7 @@ mod tests { let statement = SelectStatementStack { elements: vec![ SelectStatementStackElement::SelectStatement(SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::All], @@ -269,7 +270,7 @@ mod tests { limit_clause: None, }), SelectStatementStackElement::SelectStatement(SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::All], @@ -293,7 +294,7 @@ mod tests { }), SelectStatementStackElement::SetOperator(SetOperator::Intersect), SelectStatementStackElement::SelectStatement(SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::All], diff --git a/src/db/table/operations/select/select_statement.rs b/src/db/table/operations/select/select_statement.rs index 0e24cdf..39d2c53 100644 --- a/src/db/table/operations/select/select_statement.rs +++ b/src/db/table/operations/select/select_statement.rs @@ -79,6 +79,7 @@ mod tests { use crate::interpreter::ast::Operand; use crate::interpreter::ast::SelectMode; use crate::interpreter::ast::SelectStatementColumn; + use crate::interpreter::ast::SelectStatementTable; use crate::interpreter::ast::WhereCondition; use crate::interpreter::ast::WhereStackElement; use crate::interpreter::ast::{ @@ -90,7 +91,7 @@ mod tests { fn select_with_all_tokens_is_generated_correctly() { let table = default_table(); let statement = SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::All], @@ -135,7 +136,7 @@ mod tests { fn select_specific_columns_is_generated_correctly() { let table = default_table(); let statement = SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![ @@ -166,7 +167,7 @@ mod tests { fn select_with_where_clause_is_generated_correctly() { let table = default_table(); let statement = SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::All], @@ -195,7 +196,7 @@ mod tests { fn select_with_where_clause_using_column_not_included_in_selected_columns() { let table = default_table(); let statement = SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![ @@ -228,7 +229,7 @@ mod tests { fn select_with_limit_clause_is_generated_correctly() { let table = default_table(); let statement = SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::All], @@ -256,7 +257,7 @@ mod tests { fn select_with_where_clause_using_column_not_included_in_table_returns_error() { let table = default_table(); let statement = SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::All], @@ -282,7 +283,7 @@ mod tests { fn select_with_order_by_clause_is_generated_correctly() { let table = default_table(); let statement = SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::All], @@ -355,7 +356,7 @@ mod tests { Row(vec![Value::Integer(4), Value::Null]), ]); let statement = SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), column_names: vec![SelectStatementColumn::new("name".to_string())], mode: SelectMode::Distinct, columns: SelectableStack { diff --git a/src/db/table/operations/update/mod.rs b/src/db/table/operations/update/mod.rs index de3021b..3a9c97e 100644 --- a/src/db/table/operations/update/mod.rs +++ b/src/db/table/operations/update/mod.rs @@ -60,14 +60,14 @@ mod tests { use crate::interpreter::ast::SelectStatementColumn; use crate::interpreter::ast::{ LimitClause, Operand, Operator, OrderByClause, OrderByDirection, SelectableStack, - SelectableStackElement, WhereCondition, WhereStackElement, + SelectableStackElement, WhereCondition, WhereStackElement, SelectStatementTable, }; #[test] fn update_works_correctly() { let mut table = default_table(); let statement = UpdateStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), update_values: vec![ColumnValue { column: "name".to_string(), value: Value::Text("John".to_string()), @@ -155,7 +155,7 @@ mod tests { ]), ]); let statement = UpdateStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), update_values: vec![ColumnValue { column: "name".to_string(), value: Value::Text("Fletcher".to_string()), @@ -232,7 +232,7 @@ mod tests { fn update_multiple_columns_and_rows_works_correctly() { let mut table = default_table(); let statement = UpdateStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), update_values: vec![ ColumnValue { column: "name".to_string(), @@ -296,7 +296,7 @@ mod tests { ); table.set_rows(vec![]); let statement = UpdateStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), update_values: vec![ColumnValue { column: "name".to_string(), value: Value::Text("Fletcher".to_string()), @@ -316,7 +316,7 @@ mod tests { fn update_with_invalid_column_works_correctly() { let mut table = default_table(); let statement = UpdateStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), update_values: vec![ColumnValue { column: "invalid".to_string(), value: Value::Text("Fletcher".to_string()), @@ -337,7 +337,7 @@ mod tests { fn update_with_invalid_value_works_correctly() { let mut table = default_table(); let statement = UpdateStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), update_values: vec![ColumnValue { column: "name".to_string(), value: Value::Integer(1), @@ -358,7 +358,7 @@ mod tests { fn update_with_null_value_works_correctly() { let mut table = default_table(); let statement = UpdateStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), update_values: vec![ColumnValue { column: "money".to_string(), value: Value::Null, @@ -403,7 +403,7 @@ mod tests { fn update_with_transaction_works_correctly() { let mut table = default_table(); let statement = UpdateStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), update_values: vec![ColumnValue { column: "name".to_string(), value: Value::Text("Fletcher".to_string()), diff --git a/src/db/transactions/mod.rs b/src/db/transactions/mod.rs index 26d1ccd..7ea19dc 100644 --- a/src/db/transactions/mod.rs +++ b/src/db/transactions/mod.rs @@ -41,8 +41,8 @@ impl TransactionLog { let table_name = match &sql_statement { SqlStatement::CreateTable(statement) => statement.table_name.clone(), SqlStatement::InsertInto(statement) => statement.table_name.clone(), - SqlStatement::UpdateStatement(statement) => statement.table_name.clone(), - SqlStatement::DeleteStatement(statement) => statement.table_name.clone(), + SqlStatement::UpdateStatement(statement) => statement.table_name.table_name.clone(), + SqlStatement::DeleteStatement(statement) => statement.table_name.table_name.clone(), SqlStatement::DropTable(statement) => statement.table_name.clone(), SqlStatement::AlterTable(statement) => statement.table_name.clone(), SqlStatement::Savepoint(statement) => { diff --git a/src/interpreter/ast/alter_table_statement.rs b/src/interpreter/ast/alter_table_statement.rs index d2ca647..b310b4d 100644 --- a/src/interpreter/ast/alter_table_statement.rs +++ b/src/interpreter/ast/alter_table_statement.rs @@ -10,7 +10,7 @@ pub fn build(parser: &mut Parser) -> Result { parser.advance()?; expect_token_type(parser, TokenTypes::Table)?; parser.advance()?; - let table_name = get_table_name(parser)?; + let table_name = get_table_name(parser, false)?.table_name; let action = get_action(parser)?; parser.advance()?; expect_token_type(parser, TokenTypes::SemiColon)?; diff --git a/src/interpreter/ast/create_statement.rs b/src/interpreter/ast/create_statement.rs index 92bee21..4b5f3d5 100644 --- a/src/interpreter/ast/create_statement.rs +++ b/src/interpreter/ast/create_statement.rs @@ -31,7 +31,7 @@ fn table_statement(parser: &mut Parser) -> Result { parser.advance()?; let existence_check = exists_clause(parser, ExistenceCheck::IfNotExists)?; - let table_name = get_table_name(parser)?; + let table_name = get_table_name(parser, false)?.table_name; let column_definitions = column_definitions(parser)?; return Ok(CreateTable(CreateTableStatement { diff --git a/src/interpreter/ast/delete_statement.rs b/src/interpreter/ast/delete_statement.rs index 096bfc6..2ecf407 100644 --- a/src/interpreter/ast/delete_statement.rs +++ b/src/interpreter/ast/delete_statement.rs @@ -14,7 +14,7 @@ pub fn build(parser: &mut Parser) -> Result { parser.advance()?; expect_token_type(parser, TokenTypes::From)?; parser.advance()?; - let table_name = get_table_name(parser)?; + let table_name = get_table_name(parser, true)?; let where_clause = get_where_clause(parser)?; let order_by_clause = get_order_by(parser)?; let limit_clause = get_limit(parser)?; @@ -39,6 +39,7 @@ mod tests { use crate::interpreter::ast::SelectStatementColumn; use crate::interpreter::ast::SelectableStack; use crate::interpreter::ast::SelectableStackElement; + use crate::interpreter::ast::SelectStatementTable; use crate::interpreter::ast::WhereCondition; use crate::interpreter::ast::WhereStackElement; use crate::interpreter::ast::test_utils::token; @@ -57,7 +58,7 @@ mod tests { assert!(result.is_ok()); let statement = result.unwrap(); let expected = SqlStatement::DeleteStatement(DeleteStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), where_clause: None, order_by_clause: None, limit_clause: None, @@ -91,7 +92,7 @@ mod tests { assert!(result.is_ok()); let statement = result.unwrap(); let expected = SqlStatement::DeleteStatement(DeleteStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), where_clause: Some(vec![WhereStackElement::Condition(WhereCondition { l_side: Operand::Identifier("id".to_string()), operator: Operator::Equals, diff --git a/src/interpreter/ast/drop_statement.rs b/src/interpreter/ast/drop_statement.rs index dfa63ac..cb077a1 100644 --- a/src/interpreter/ast/drop_statement.rs +++ b/src/interpreter/ast/drop_statement.rs @@ -14,7 +14,7 @@ pub fn build(parser: &mut Parser) -> Result { parser.advance()?; let existence_check = exists_clause(parser, ExistenceCheck::IfExists)?; - let table_name = get_table_name(parser)?; + let table_name = get_table_name(parser, false)?.table_name; return Ok(SqlStatement::DropTable(DropTableStatement { table_name: table_name, existence_check: existence_check, diff --git a/src/interpreter/ast/helpers/common.rs b/src/interpreter/ast/helpers/common.rs index eb64114..c32f8ae 100644 --- a/src/interpreter/ast/helpers/common.rs +++ b/src/interpreter/ast/helpers/common.rs @@ -1,7 +1,7 @@ use crate::interpreter::{ ast::{ ExistenceCheck, LogicalOperator, MathOperator, Operator, OrderByDirection, - SelectStatementColumn, SelectableStack, SelectableStackElement, + SelectStatementColumn, SelectableStack, SelectableStackElement, SelectStatementTable, helpers::token::token_to_value, parser::Parser, }, tokenizer::token::TokenTypes, @@ -17,12 +17,21 @@ pub fn expect_token_type(parser: &Parser, token_type: TokenTypes) -> Result<(), Ok(()) } -pub fn get_table_name(parser: &mut Parser) -> Result { - let token = parser.current_token()?; +pub fn get_table_name(parser: &mut Parser, allow_alias: bool) -> Result { expect_token_type(parser, TokenTypes::Identifier)?; - let result = token.value.to_string(); + let table_name = parser.current_token()?.value.to_string(); parser.advance()?; - Ok(result) + if allow_alias && parser.current_token()?.token_type == TokenTypes::As { + parser.advance()?; + expect_token_type(parser, TokenTypes::Identifier)?; + let alias = parser.current_token()?.value.to_string(); + parser.advance()?; + return Ok(SelectStatementTable { + table_name: table_name, + alias: Some(alias), + }); + } + Ok(SelectStatementTable::new(table_name)) } pub const TOKENS_NEEDING_SPECIAL_HANDLING: [TokenTypes; 5] = [ diff --git a/src/interpreter/ast/helpers/select_statement.rs b/src/interpreter/ast/helpers/select_statement.rs index a9f325a..d37c7b8 100644 --- a/src/interpreter/ast/helpers/select_statement.rs +++ b/src/interpreter/ast/helpers/select_statement.rs @@ -25,7 +25,7 @@ pub fn get_statement(parser: &mut Parser) -> Result { let (columns, column_names) = get_columns_and_names(parser)?; expect_token_type(parser, TokenTypes::From)?; // TODO: this is not true, you can do SELECT 1; parser.advance()?; - let table_name = get_table_name(parser)?; + let table_name = get_table_name(parser, true)?; let where_clause: Option> = get_where_clause(parser)?; let order_by_clause = get_order_by(parser)?; let limit_clause = get_limit(parser)?; @@ -61,6 +61,7 @@ mod tests { use crate::interpreter::ast::OrderByClause; use crate::interpreter::ast::OrderByDirection; use crate::interpreter::ast::SelectableStackElement; + use crate::interpreter::ast::SelectStatementTable; use crate::interpreter::ast::WhereCondition; use crate::interpreter::ast::WhereStackElement; use crate::interpreter::ast::test_utils::token; @@ -82,7 +83,7 @@ mod tests { assert_eq!( statement, SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::All], @@ -112,7 +113,7 @@ mod tests { assert_eq!( statement, SelectStatement { - table_name: "guests".to_string(), + table_name: SelectStatementTable::new("guests".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new( @@ -146,7 +147,7 @@ mod tests { assert_eq!( statement, SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![ @@ -201,7 +202,7 @@ mod tests { assert!(result.is_ok()); let statement = result.unwrap(); let expected = SelectStatement { - table_name: "guests".to_string(), + table_name: SelectStatementTable::new("guests".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::Column(SelectStatementColumn::new( @@ -263,7 +264,7 @@ mod tests { assert!(result.is_ok()); let statement = result.unwrap(); let expected = SelectStatement { - table_name: "guests".to_string(), + table_name: SelectStatementTable::new("guests".to_string()), column_names: vec![SelectStatementColumn::new("id".to_string())], mode: SelectMode::Distinct, columns: SelectableStack { @@ -296,7 +297,7 @@ mod tests { assert!(result.is_ok()); let statement = result.unwrap(); let expected = SelectStatement { - table_name: "guests".to_string(), + table_name: SelectStatementTable::new("guests".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::Column(SelectStatementColumn { @@ -340,7 +341,7 @@ mod tests { assert!(result.is_ok()); let statement = result.unwrap(); let expected = SelectStatement { - table_name: "guests".to_string(), + table_name: SelectStatementTable::new("guests".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![ @@ -384,4 +385,42 @@ mod tests { }; assert_eq!(expected, statement); } + + #[test] + fn select_statement_with_table_name_alias_is_generated_correctly() { + // SELECT id FROM guests AS g; + let tokens = vec![ + token(TokenTypes::Select, "SELECT"), + token(TokenTypes::Identifier, "id"), + token(TokenTypes::From, "FROM"), + token(TokenTypes::Identifier, "guests"), + token(TokenTypes::As, "AS"), + token(TokenTypes::Identifier, "g"), + token(TokenTypes::SemiColon, ";"), + ]; + let mut parser = Parser::new(tokens); + let result = get_statement(&mut parser); + assert!(result.is_ok()); + let statement = result.unwrap(); + let expected = SelectStatement { + table_name: SelectStatementTable{table_name: "guests".to_string(), alias: Some("g".to_string())}, + mode: SelectMode::All, + columns: SelectableStack { + selectables: vec![SelectableStackElement::Column(SelectStatementColumn { + column_name: "id".to_string(), + alias: None, + table_name: None, + })], + }, + column_names: vec![SelectStatementColumn { + column_name: "id".to_string(), + alias: None, + table_name: None, + }], + where_clause: None, + order_by_clause: None, + limit_clause: None, + }; + assert_eq!(expected, statement); + } } diff --git a/src/interpreter/ast/insert_statement.rs b/src/interpreter/ast/insert_statement.rs index 215ad10..692793c 100644 --- a/src/interpreter/ast/insert_statement.rs +++ b/src/interpreter/ast/insert_statement.rs @@ -32,7 +32,7 @@ pub fn build(parser: &mut Parser) -> Result { fn into_statement(parser: &mut Parser) -> Result { parser.advance()?; - let table_name = get_table_name(parser)?; + let table_name = get_table_name(parser, false)?.table_name; 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 0b35df0..12f4bc3 100644 --- a/src/interpreter/ast/mod.rs +++ b/src/interpreter/ast/mod.rs @@ -104,7 +104,7 @@ impl SetOperator { #[derive(Debug, PartialEq, Clone)] pub struct SelectStatement { - pub table_name: String, + pub table_name: SelectStatementTable, pub column_names: Vec, pub mode: SelectMode, pub columns: SelectableStack, @@ -113,6 +113,21 @@ pub struct SelectStatement { pub limit_clause: Option, } +#[derive(Debug, PartialEq, Clone)] +pub struct SelectStatementTable { + pub table_name: String, + pub alias: Option, +} + +impl SelectStatementTable { + pub fn new(table_name: String) -> Self { + Self { + table_name, + alias: None, + } + } +} + #[derive(Debug, PartialEq, Clone)] pub struct SelectStatementColumn { pub column_name: String, @@ -132,7 +147,7 @@ impl SelectStatementColumn { #[derive(Debug, PartialEq, Clone)] pub struct DeleteStatement { - pub table_name: String, + pub table_name: SelectStatementTable, pub where_clause: Option>, pub order_by_clause: Option, pub limit_clause: Option, @@ -140,7 +155,7 @@ pub struct DeleteStatement { #[derive(Debug, PartialEq, Clone)] pub struct UpdateStatement { - pub table_name: String, + pub table_name: SelectStatementTable, pub update_values: Vec, pub where_clause: Option>, pub order_by_clause: Option, @@ -495,7 +510,7 @@ mod tests { sql_statement: SqlStatement::Select(SelectStatementStack { elements: vec![SelectStatementStackElement::SelectStatement( SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::All], @@ -589,7 +604,7 @@ mod tests { sql_statement: SqlStatement::Select(SelectStatementStack { elements: vec![SelectStatementStackElement::SelectStatement( SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::All], diff --git a/src/interpreter/ast/parser.rs b/src/interpreter/ast/parser.rs index c0777dd..261f704 100644 --- a/src/interpreter/ast/parser.rs +++ b/src/interpreter/ast/parser.rs @@ -132,7 +132,7 @@ mod tests { use crate::interpreter::ast::{ CreateTableStatement, InsertIntoStatement, SelectMode, SelectStatement, SelectStatementColumn, SelectStatementStack, SelectStatementStackElement, SelectableStack, - SelectableStackElement, + SelectableStackElement, SelectStatementTable, }; #[test] @@ -194,7 +194,7 @@ mod tests { let expected = Some(Ok(SqlStatement::Select(SelectStatementStack { elements: vec![SelectStatementStackElement::SelectStatement( SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::All], diff --git a/src/interpreter/ast/select_statement_stack.rs b/src/interpreter/ast/select_statement_stack.rs index 904a467..ba30488 100644 --- a/src/interpreter/ast/select_statement_stack.rs +++ b/src/interpreter/ast/select_statement_stack.rs @@ -172,6 +172,7 @@ mod tests { use crate::interpreter::ast::SelectableStack; use crate::interpreter::ast::SelectableStackElement; use crate::interpreter::ast::SetOperator; + use crate::interpreter::ast::SelectStatementTable; use crate::interpreter::ast::WhereCondition; use crate::interpreter::ast::WhereStackElement; use crate::interpreter::ast::test_utils::token; @@ -193,7 +194,7 @@ mod tests { fn expected_simple_select_statement(id: i64) -> SelectStatementStackElement { SelectStatementStackElement::SelectStatement(SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::All], @@ -357,7 +358,7 @@ mod tests { let expected = SqlStatement::Select(SelectStatementStack { elements: vec![ SelectStatementStackElement::SelectStatement(SelectStatement { - table_name: "employees".to_string(), + table_name: SelectStatementTable::new("employees".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::Column( @@ -374,7 +375,7 @@ mod tests { limit_clause: None, }), SelectStatementStackElement::SelectStatement(SelectStatement { - table_name: "employees".to_string(), + table_name: SelectStatementTable::new("employees".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::Column( diff --git a/src/interpreter/ast/statement_builder.rs b/src/interpreter/ast/statement_builder.rs index cd23c40..35e3e61 100644 --- a/src/interpreter/ast/statement_builder.rs +++ b/src/interpreter/ast/statement_builder.rs @@ -76,7 +76,7 @@ impl StatementBuilder for DefaultStatementBuilder { pub struct MockStatementBuilder; #[cfg(test)] use crate::interpreter::ast::{ - CreateTableStatement, InsertIntoStatement, SelectMode, SelectStatement, SelectStatementColumn, + CreateTableStatement, InsertIntoStatement, SelectMode, SelectStatement, SelectStatementColumn, SelectStatementTable, SelectStatementStack, SelectStatementStackElement, SelectableStack, SelectableStackElement, }; @@ -108,7 +108,7 @@ impl StatementBuilder for MockStatementBuilder { return Ok(SqlStatement::Select(SelectStatementStack { elements: vec![SelectStatementStackElement::SelectStatement( SelectStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), mode: SelectMode::All, columns: SelectableStack { selectables: vec![SelectableStackElement::All], diff --git a/src/interpreter/ast/update_statement.rs b/src/interpreter/ast/update_statement.rs index 08acb59..d47bb48 100644 --- a/src/interpreter/ast/update_statement.rs +++ b/src/interpreter/ast/update_statement.rs @@ -10,7 +10,7 @@ use crate::interpreter::tokenizer::token::TokenTypes; pub fn build(parser: &mut Parser) -> Result { parser.advance()?; - let table_name = get_table_name(parser)?; + let table_name = get_table_name(parser, true)?; // Ensure Set expect_token_type(parser, TokenTypes::Set)?; let update_values = get_update_values(parser)?; @@ -73,6 +73,7 @@ mod tests { use crate::interpreter::ast::SelectStatementColumn; use crate::interpreter::ast::SelectableStack; use crate::interpreter::ast::SelectableStackElement; + use crate::interpreter::ast::SelectStatementTable; use crate::interpreter::ast::WhereCondition; use crate::interpreter::ast::WhereStackElement; use crate::interpreter::ast::test_utils::token; @@ -94,7 +95,7 @@ mod tests { assert!(result.is_ok()); let statement = result.unwrap(); let expected = SqlStatement::UpdateStatement(UpdateStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), update_values: vec![ColumnValue { column: "column".to_string(), value: Value::Text("value".to_string()), @@ -127,7 +128,7 @@ mod tests { assert!(result.is_ok()); let statement = result.unwrap(); let expected = SqlStatement::UpdateStatement(UpdateStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), update_values: vec![ColumnValue { column: "column".to_string(), value: Value::Integer(1), @@ -168,7 +169,7 @@ mod tests { assert!(result.is_ok()); let statement = result.unwrap(); let expected = SqlStatement::UpdateStatement(UpdateStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), update_values: vec![ ColumnValue { column: "column".to_string(), @@ -219,7 +220,7 @@ mod tests { assert!(result.is_ok()); let statement = result.unwrap(); let expected = SqlStatement::UpdateStatement(UpdateStatement { - table_name: "users".to_string(), + table_name: SelectStatementTable::new("users".to_string()), update_values: vec![ColumnValue { column: "column".to_string(), value: Value::Integer(1), diff --git a/src/interpreter/tokenizer/mod.rs b/src/interpreter/tokenizer/mod.rs index b2c9254..1ac7681 100644 --- a/src/interpreter/tokenizer/mod.rs +++ b/src/interpreter/tokenizer/mod.rs @@ -363,4 +363,25 @@ mod tests { ]; assert_eq!(expected, result); } + + #[test] + fn tokenizer_parses_table_name_with_alias() { + let result = tokenize( + "SELECT u.id AS user_id FROM users AS u" + ); + let expected = vec![ + token(TokenTypes::Select, "SELECT", 0, 1), + token(TokenTypes::Identifier, "u", 7, 1), + token(TokenTypes::Dot, ".", 8, 1), + token(TokenTypes::Identifier, "id", 9, 1), + token(TokenTypes::As, "AS", 12, 1), + token(TokenTypes::Identifier, "user_id", 15, 1), + token(TokenTypes::From, "FROM", 23, 1), + token(TokenTypes::Identifier, "users", 28, 1), + token(TokenTypes::As, "AS", 34, 1), + token(TokenTypes::Identifier, "u", 37, 1), + token(TokenTypes::EOF, "", 0, 0), + ]; + assert_eq!(expected, result); + } } From f7c7dc4c70ecd6c8c9d77d424f6f5814cf1c2eb0 Mon Sep 17 00:00:00 2001 From: Fletcher555 Date: Sun, 21 Sep 2025 13:54:24 -0400 Subject: [PATCH 5/6] add simple single column alias with table name . --- src/interpreter/ast/helpers/common.rs | 57 ++++++++++++----- .../ast/helpers/select_statement.rs | 62 +++++++++++++++++++ 2 files changed, 103 insertions(+), 16 deletions(-) diff --git a/src/interpreter/ast/helpers/common.rs b/src/interpreter/ast/helpers/common.rs index c32f8ae..7bdd88e 100644 --- a/src/interpreter/ast/helpers/common.rs +++ b/src/interpreter/ast/helpers/common.rs @@ -1,10 +1,8 @@ use crate::interpreter::{ ast::{ - ExistenceCheck, LogicalOperator, MathOperator, Operator, OrderByDirection, - SelectStatementColumn, SelectableStack, SelectableStackElement, SelectStatementTable, - helpers::token::token_to_value, parser::Parser, + helpers::token::token_to_value, parser::Parser, ExistenceCheck, LogicalOperator, MathOperator, Operator, OrderByDirection, SelectStatementColumn, SelectStatementTable, SelectableStack, SelectableStackElement }, - tokenizer::token::TokenTypes, + tokenizer::{token::TokenTypes}, }; use std::cmp::Ordering; @@ -58,6 +56,7 @@ pub fn get_selectables( let mut depth = 0; let mut current_name = "".to_string(); let mut current_alias: Option = None; + let mut current_table_name: Option = None; let mut first = true; let mut expect_new_value = false; // Will be set after a valid ASC or DESC to ensure proper syntax @@ -241,18 +240,11 @@ pub fn get_selectables( TokenTypes::Null => SelectableStackElement::Value(token_to_value(parser)?), // TODO: handle ValueList (arrays) TokenTypes::Identifier => { - let mut column = SelectStatementColumn::new(token.value.to_string()); - if let Ok(peek_token) = parser.peek_token() { - if peek_token.token_type == TokenTypes::As { - parser.advance()?; - parser.advance()?; - expect_token_type(parser, TokenTypes::Identifier)?; - let alias = parser.current_token()?.value.to_string(); - column.alias = Some(alias.clone()); - current_alias = Some(alias); - } - } - SelectableStackElement::Column(column) // TODO: verify it's a column, AND handle multi-tokens columns with AS (table_name.column_name) + let column = get_select_statement_column(parser)?; + current_alias = column.alias.clone(); + current_table_name = column.table_name.clone(); + current_name = column.column_name.clone(); + SelectableStackElement::Column(column) } _ => return Err(parser.format_error()), // TODO: better error handling }; @@ -274,6 +266,7 @@ pub fn get_selectables( if let Some(selectable_names_vector) = selectable_names { let mut column = SelectStatementColumn::new(current_name); column.alias = current_alias; + column.table_name = current_table_name; selectable_names_vector.push(column); } @@ -319,6 +312,38 @@ pub fn compare_precedence( }; } +pub fn get_select_statement_column<'a>(parser: &mut Parser<'a>) -> Result { + let mut current_token = parser.current_token()?; + let table_name = if parser.peek_token()?.token_type == TokenTypes::Dot { + let table_name = current_token.value.to_string(); + parser.advance()?; + parser.advance()?; + current_token = parser.current_token()?; + Some(table_name) + } + else { + None + }; + + let column_name = current_token.value.to_string(); + let alias = if let Ok(peek_token) = parser.peek_token() && peek_token.token_type == TokenTypes::As { + parser.advance()?; + parser.advance()?; + expect_token_type(parser, TokenTypes::Identifier)?; + let alias = Some(parser.current_token()?.value.to_string()); + alias + } + else { + None + }; + let column = SelectStatementColumn{ + column_name: column_name, + alias: alias, + table_name: table_name, + }; + Ok(column) +} + fn get_precedence(operator: &SelectableStackElement) -> Result { let result = match operator { SelectableStackElement::MathOperator(MathOperator::Multiply) => 40, diff --git a/src/interpreter/ast/helpers/select_statement.rs b/src/interpreter/ast/helpers/select_statement.rs index d37c7b8..edfe2f9 100644 --- a/src/interpreter/ast/helpers/select_statement.rs +++ b/src/interpreter/ast/helpers/select_statement.rs @@ -423,4 +423,66 @@ mod tests { }; assert_eq!(expected, statement); } + + #[test] + fn select_statement_with_column_table_name_and_alias_is_generated_correctly() { + // SELECT u.id as user_id, t.id as ticket_id FROM guests AS u; + let tokens = vec![ + token(TokenTypes::Select, "SELECT"), + token(TokenTypes::Identifier, "u"), + token(TokenTypes::Dot, "."), + token(TokenTypes::Identifier, "id"), + token(TokenTypes::As, "AS"), + token(TokenTypes::Identifier, "user_id"), + token(TokenTypes::Comma, ","), + token(TokenTypes::Identifier, "t"), + token(TokenTypes::Dot, "."), + token(TokenTypes::Identifier, "id"), + token(TokenTypes::As, "AS"), + token(TokenTypes::Identifier, "ticket_id"), + token(TokenTypes::From, "FROM"), + token(TokenTypes::Identifier, "guests"), + token(TokenTypes::As, "AS"), + token(TokenTypes::Identifier, "u"), + token(TokenTypes::SemiColon, ";"), + ]; + let mut parser = Parser::new(tokens); + let result = get_statement(&mut parser); + assert!(result.is_ok()); + let statement = result.unwrap(); + let expected = SelectStatement { + table_name: SelectStatementTable{table_name: "guests".to_string(), alias: Some("u".to_string())}, + mode: SelectMode::All, + columns: SelectableStack { + selectables: vec![ + SelectableStackElement::Column(SelectStatementColumn { + column_name: "id".to_string(), + alias: Some("user_id".to_string()), + table_name: Some("u".to_string()), + }), + SelectableStackElement::Column(SelectStatementColumn { + column_name: "id".to_string(), + alias: Some("ticket_id".to_string()), + table_name: Some("t".to_string()), + }) + ], + }, + column_names: vec![ + SelectStatementColumn { + column_name: "id".to_string(), + alias: Some("user_id".to_string()), + table_name: Some("u".to_string()), + }, + SelectStatementColumn { + column_name: "id".to_string(), + alias: Some("ticket_id".to_string()), + table_name: Some("t".to_string()), + } + ], + where_clause: None, + order_by_clause: None, + limit_clause: None, + }; + assert_eq!(expected, statement); + } } From 7e85a6787bbcfa91162be7f69184464cbc23df91 Mon Sep 17 00:00:00 2001 From: Fletcher555 Date: Sun, 21 Sep 2025 14:00:12 -0400 Subject: [PATCH 6/6] Fix bug with table name and add tests for different configurations with multiple columns + formatting --- src/db/table/operations/delete/mod.rs | 2 +- src/db/table/operations/select/mod.rs | 4 +- src/db/table/operations/update/mod.rs | 4 +- src/interpreter/ast/delete_statement.rs | 2 +- src/interpreter/ast/helpers/common.rs | 28 ++++-- .../ast/helpers/select_statement.rs | 95 +++++++++++++++++-- src/interpreter/ast/parser.rs | 4 +- src/interpreter/ast/select_statement_stack.rs | 2 +- src/interpreter/ast/statement_builder.rs | 5 +- src/interpreter/ast/update_statement.rs | 2 +- src/interpreter/tokenizer/mod.rs | 4 +- 11 files changed, 120 insertions(+), 32 deletions(-) diff --git a/src/db/table/operations/delete/mod.rs b/src/db/table/operations/delete/mod.rs index 3130a28..42fcc72 100644 --- a/src/db/table/operations/delete/mod.rs +++ b/src/db/table/operations/delete/mod.rs @@ -66,8 +66,8 @@ mod tests { use crate::db::table::core::{row::Row, value::Value}; use crate::db::table::test_utils::{assert_table_rows_eq_unordered, default_table}; use crate::interpreter::ast::LimitClause; - use crate::interpreter::ast::SelectStatementTable; use crate::interpreter::ast::SelectStatementColumn; + use crate::interpreter::ast::SelectStatementTable; use crate::interpreter::ast::{ Operand, Operator, OrderByClause, OrderByDirection, SelectableStack, SelectableStackElement, WhereCondition, WhereStackElement, diff --git a/src/db/table/operations/select/mod.rs b/src/db/table/operations/select/mod.rs index 3f269b7..d5c60a4 100644 --- a/src/db/table/operations/select/mod.rs +++ b/src/db/table/operations/select/mod.rs @@ -151,8 +151,8 @@ mod tests { use crate::db::table::test_utils::default_database; use crate::interpreter::ast::{ LogicalOperator, Operand, Operator, SelectMode, SelectStatement, SelectStatementColumn, - SelectableStack, SelectableStackElement, WhereCondition, WhereStackElement, - SelectStatementTable, + SelectStatementTable, SelectableStack, SelectableStackElement, WhereCondition, + WhereStackElement, }; #[test] diff --git a/src/db/table/operations/update/mod.rs b/src/db/table/operations/update/mod.rs index 3a9c97e..3614555 100644 --- a/src/db/table/operations/update/mod.rs +++ b/src/db/table/operations/update/mod.rs @@ -59,8 +59,8 @@ mod tests { use crate::interpreter::ast::ColumnValue; use crate::interpreter::ast::SelectStatementColumn; use crate::interpreter::ast::{ - LimitClause, Operand, Operator, OrderByClause, OrderByDirection, SelectableStack, - SelectableStackElement, WhereCondition, WhereStackElement, SelectStatementTable, + LimitClause, Operand, Operator, OrderByClause, OrderByDirection, SelectStatementTable, + SelectableStack, SelectableStackElement, WhereCondition, WhereStackElement, }; #[test] diff --git a/src/interpreter/ast/delete_statement.rs b/src/interpreter/ast/delete_statement.rs index 2ecf407..188408d 100644 --- a/src/interpreter/ast/delete_statement.rs +++ b/src/interpreter/ast/delete_statement.rs @@ -37,9 +37,9 @@ mod tests { use crate::interpreter::ast::OrderByClause; use crate::interpreter::ast::OrderByDirection; use crate::interpreter::ast::SelectStatementColumn; + use crate::interpreter::ast::SelectStatementTable; use crate::interpreter::ast::SelectableStack; use crate::interpreter::ast::SelectableStackElement; - use crate::interpreter::ast::SelectStatementTable; use crate::interpreter::ast::WhereCondition; use crate::interpreter::ast::WhereStackElement; use crate::interpreter::ast::test_utils::token; diff --git a/src/interpreter/ast/helpers/common.rs b/src/interpreter/ast/helpers/common.rs index 7bdd88e..59c7a7a 100644 --- a/src/interpreter/ast/helpers/common.rs +++ b/src/interpreter/ast/helpers/common.rs @@ -1,8 +1,10 @@ use crate::interpreter::{ ast::{ - helpers::token::token_to_value, parser::Parser, ExistenceCheck, LogicalOperator, MathOperator, Operator, OrderByDirection, SelectStatementColumn, SelectStatementTable, SelectableStack, SelectableStackElement + ExistenceCheck, LogicalOperator, MathOperator, Operator, OrderByDirection, + SelectStatementColumn, SelectStatementTable, SelectableStack, SelectableStackElement, + helpers::token::token_to_value, parser::Parser, }, - tokenizer::{token::TokenTypes}, + tokenizer::token::TokenTypes, }; use std::cmp::Ordering; @@ -15,7 +17,10 @@ pub fn expect_token_type(parser: &Parser, token_type: TokenTypes) -> Result<(), Ok(()) } -pub fn get_table_name(parser: &mut Parser, allow_alias: bool) -> Result { +pub fn get_table_name( + parser: &mut Parser, + allow_alias: bool, +) -> Result { expect_token_type(parser, TokenTypes::Identifier)?; let table_name = parser.current_token()?.value.to_string(); parser.advance()?; @@ -97,6 +102,7 @@ pub fn get_selectables( } else if let Some(selectable_names_vector) = selectable_names { let mut column = SelectStatementColumn::new(current_name.clone()); column.alias = current_alias.clone(); + column.table_name = current_table_name.clone(); selectable_names_vector.push(column); } // Default ordering is ASC @@ -312,7 +318,9 @@ pub fn compare_precedence( }; } -pub fn get_select_statement_column<'a>(parser: &mut Parser<'a>) -> Result { +pub fn get_select_statement_column<'a>( + parser: &mut Parser<'a>, +) -> Result { let mut current_token = parser.current_token()?; let table_name = if parser.peek_token()?.token_type == TokenTypes::Dot { let table_name = current_token.value.to_string(); @@ -320,23 +328,23 @@ pub fn get_select_statement_column<'a>(parser: &mut Parser<'a>) -> Result