Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 4 additions & 3 deletions src/db/table/helpers/common.rs
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
use crate::db::table::{Table, Value, DataType};
use crate::interpreter::ast::{SelectStatementColumns, WhereStackElement, OrderByClause, LimitClause};
use crate::db::table::helpers::where_stack::matches_where_stack;
use crate::db::table::helpers::{order_by_clause::get_ordered_row_indicies, limit_clause::get_limited_row_indicies};
use crate::db::table::helpers::{order_by_clause::get_ordered_row_indicies, limit_clause::get_limited_rows};

pub fn validate_and_clone_row(table: &Table, row: &Vec<Value>) -> Result<Vec<Value>, String> {
if row.len() != table.width() {
Expand Down Expand Up @@ -54,7 +54,7 @@ pub fn get_columns_from_row(table: &Table, row: &Vec<Value>, selected_columns: &
} else {
let specific_selected_columns = selected_columns.columns()?;
for (i, column) in table.columns.iter().enumerate() {
if (*specific_selected_columns).contains(&column.name) {
if (*specific_selected_columns).contains(&&column.name) {
row_values.push(row[i].clone());
}
}
Expand All @@ -70,7 +70,8 @@ pub fn get_row_indicies_matching_clauses(table: &Table, where_clause: &Option<Ve
}

if let Some(limit_clause) = limit_clause {
row_indicies = get_limited_row_indicies(row_indicies, &limit_clause)?;
let result = get_limited_rows(row_indicies, &limit_clause)?;
return Ok(result.to_vec());
}

return Ok(row_indicies);
Expand Down
41 changes: 23 additions & 18 deletions src/db/table/helpers/limit_clause.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,31 +4,30 @@ use crate::interpreter::ast::LimitClause;
use crate::db::table::Value;


pub fn get_limited_row_indicies(rows: Vec<usize>, limit_clause: &LimitClause) -> Result<Vec<usize>, String> {
pub fn get_limited_rows<T>(mut rows: Vec<T>, limit_clause: &LimitClause) -> Result<Vec<T>, String> {
let mut index: usize = 0;
if let Some(offset) = &limit_clause.offset && let Value::Integer(offset) = offset {
index = *offset as usize;
}
if index >= rows.len() {
return Ok(vec![]);
if index >= rows.len() {
rows.truncate(0);
return Ok(rows);
}

rows.drain(0..index);

let limit = match limit_clause.limit {
Value::Integer(limit) => {
if limit < 0 {
rows.len()
} else {
min((limit as usize)+index, rows.len())
min(limit as usize, rows.len())
}
},
_ => return Err("Limit must be an integer".to_string()), // The parser should have already validated this
_ => unreachable!() // validated by parser
};
rows.truncate(limit);

let mut limited_rows: Vec<usize> = vec![];
for i in index..limit {
limited_rows.push(rows[i]);
}
return Ok(limited_rows);
return Ok(rows);
}


Expand All @@ -50,23 +49,26 @@ mod tests {
#[test]
fn no_offset_and_limit_is_equal_to_rows_length() {
let limit_clause = generate_limit_clause(10, None);
let result = get_limited_row_indicies(default_rows(), &limit_clause);
let table = default_rows();
let result = get_limited_rows(table, &limit_clause);
assert!(result.is_ok());
assert_eq!(default_rows(), result.unwrap());
}

#[test]
fn no_offset_and_limit_is_greater_than_rows_length() {
let limit_clause = generate_limit_clause(15, None);
let result = get_limited_row_indicies(default_rows(), &limit_clause);
let table = default_rows();
let result = get_limited_rows(table, &limit_clause);
assert!(result.is_ok());
assert_eq!(default_rows(), result.unwrap());
}

#[test]
fn no_offset_and_limit_is_less_than_rows_length() {
let limit_clause = generate_limit_clause(5, None);
let result = get_limited_row_indicies(default_rows(), &limit_clause);
let table = default_rows();
let result = get_limited_rows(table, &limit_clause);
assert!(result.is_ok());
let expected = vec![0, 1, 2, 3, 4];
assert_eq!(expected, result.unwrap());
Expand All @@ -75,24 +77,27 @@ mod tests {
#[test]
fn no_offset_and_negative_limit_returns_all_rows() {
let limit_clause = generate_limit_clause(-1, None);
let result = get_limited_row_indicies(default_rows(), &limit_clause);
let table = default_rows();
let result = get_limited_rows(table, &limit_clause);
assert!(result.is_ok());
assert_eq!(default_rows(), result.unwrap());
}

#[test]
fn offset_and_limit_is_generated_correctly() {
let limit_clause = generate_limit_clause(5, Some(1));
let result = get_limited_row_indicies(default_rows(), &limit_clause);
let table = default_rows();
let result = get_limited_rows(table, &limit_clause);
assert!(result.is_ok());
let expected = vec![1, 2, 3, 4, 5];
let expected: Vec<usize> = vec![1, 2, 3, 4, 5];
assert_eq!(expected, result.unwrap());
}

#[test]
fn offset_is_greater_than_rows_length_returns_empty_rows() {
let limit_clause = generate_limit_clause(5, Some(10));
let result = get_limited_row_indicies(default_rows(), &limit_clause);
let table = default_rows();
let result = get_limited_rows(table, &limit_clause);
assert!(result.is_ok());
let expected: Vec<usize> = vec![];
assert_eq!(expected, result.unwrap());
Expand Down
19 changes: 15 additions & 4 deletions src/db/table/helpers/order_by_clause.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,19 +9,20 @@ use crate::db::table::Value;
// This sorting algorithm will always return a stable sort, this is given by all of the order columns
// then the input order of the rows is maintained with any required tie breaking.
pub fn get_ordered_row_indicies(table: &Table, mut row_indicies: Vec<usize>, order_by_clauses: &Vec<OrderByClause>) -> Result<Vec<usize>, String> {
let columns: Vec<&String> = table.columns.iter().map(|column| &column.name).collect();
row_indicies.sort_by(|a, b| {
perform_comparions(table, &table.rows[*a], &table.rows[*b], order_by_clauses)
perform_comparions(&columns, &table.rows[*a], &table.rows[*b], order_by_clauses)
});
return Ok(row_indicies);
}

fn perform_comparions(table: &Table, row1: &Vec<Value>, row2: &Vec<Value>, order_by_clauses: &Vec<OrderByClause>) -> Ordering {
pub fn perform_comparions(columns: &Vec<&String>, row1: &Vec<Value>, row2: &Vec<Value>, order_by_clauses: &Vec<OrderByClause>) -> Ordering {
let mut result = Ordering::Equal;
for comparison in order_by_clauses {
let index = table.get_index_of_column(&comparison.column);
let index = get_index_of_column(columns, &comparison.column);
let index = match index {
Ok(index) => index,
Err(_) => return Ordering::Equal, // Bad but should never happen because we've validated the columns in the parser
Err(_) => unreachable!(),
};
let ordering = row1[index].compare(&row2[index], &comparison.direction);
if ordering != Ordering::Equal {
Expand All @@ -32,6 +33,16 @@ fn perform_comparions(table: &Table, row1: &Vec<Value>, row2: &Vec<Value>, order
return result;
}

fn get_index_of_column(columns: &Vec<&String>, column_name: &String) -> Result<usize, String> {
let result = columns.iter().position(|column| column_name == *column);
if let Some(index) = result {
return Ok(index);
}
else {
return Err(format!("Column {} does not exist in table", column_name));
}
}

#[cfg(test)]
mod tests {
use super::*;
Expand Down
4 changes: 4 additions & 0 deletions src/db/table/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -175,4 +175,8 @@ impl Table {
}
return Err(format!("Column {} does not exist in table {}", column, self.name));
}

pub fn get_columns(&self) -> Vec<&String> {
self.columns.iter().map(|column| &column.name).collect()
}
}
43 changes: 41 additions & 2 deletions src/db/table/select/mod.rs
Original file line number Diff line number Diff line change
@@ -1,14 +1,37 @@
mod select_statement;
mod set_operator_evaluator;
use crate::db::{database::Database, table::Value};
use crate::interpreter::ast::{SelectStatementStack, SetOperator, SelectStatementStackElement};
use crate::interpreter::ast::{SelectStatementStack, SetOperator, SelectStatementStackElement, SelectStatementColumns};
use crate::db::table::helpers::{order_by_clause::{perform_comparions}, limit_clause::get_limited_rows};


pub fn select_statement_stack(database: &Database, statement: SelectStatementStack) -> Result<Vec<Vec<Value>>, String> {
let mut evaluator = set_operator_evaluator::SetOperatorEvaluator::new();
let statement_columns = statement.columns.columns();
let mut columns: Option<Vec<&String>> = match statement_columns {
Err(_) => None,
Ok(columns_list) => Some(columns_list),
};
for element in statement.elements {
match element {
SelectStatementStackElement::SelectStatement(select_statement) => {
let table = database.get_table(&select_statement.table_name)?;
columns = match columns {
None => Some(table.get_columns()),
Some(columns) => {
if statement.columns == SelectStatementColumns::All {
if table.get_columns() != columns {
return Err(format!("Columns mismatch between SELECT statements in Union"));
}
}
else {
if statement.columns.columns()? != columns {
return Err(format!("Columns mismatch between SELECT statements in Union"));
}
}
Some(columns)
},
};
let rows = select_statement::select_statement(table, &select_statement)?;
evaluator.push(rows);
}
Expand All @@ -30,7 +53,20 @@ pub fn select_statement_stack(database: &Database, statement: SelectStatementSta
}
}
}
let result = evaluator.result()?;
let mut result = evaluator.result()?;
if let Some(order_by_clause) = statement.order_by_clause {
result.sort_by(|a, b| {
if let Some(columns) = &columns {
perform_comparions(&columns, a, b, &order_by_clause)
}
else {
unreachable!()
}
});
}
if let Some(limit_clause) = statement.limit_clause {
result = get_limited_rows(result, &limit_clause)?;
}
Ok(result)
}

Expand All @@ -46,6 +82,7 @@ mod tests {
fn select_statement_stack_with_multiple_set_operators_works_correctly() {
let database = default_database();
let statement = SelectStatementStack {
columns: SelectStatementColumns::All,
elements: vec![SelectStatementStackElement::SelectStatement(SelectStatement {
table_name: "users".to_string(),
columns: SelectStatementColumns::All,
Expand All @@ -71,6 +108,7 @@ mod tests {
fn select_statement_stack_with_set_operator_works_correctly() {
let database = default_database();
let statement = SelectStatementStack {
columns: SelectStatementColumns::All,
elements: vec![
SelectStatementStackElement::SelectStatement(SelectStatement {
table_name: "users".to_string(),
Expand Down Expand Up @@ -107,6 +145,7 @@ mod tests {
fn select_statement_stack_works_correctly_with_multiple_set_operators() {
let database = default_database();
let statement = SelectStatementStack {
columns: SelectStatementColumns::All,
elements: vec![SelectStatementStackElement::SelectStatement(SelectStatement {
table_name: "users".to_string(),
columns: SelectStatementColumns::All,
Expand Down
10 changes: 6 additions & 4 deletions src/interpreter/ast/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,7 @@ pub struct InsertIntoStatement {

#[derive(Debug, PartialEq)]
pub struct SelectStatementStack {
pub columns: SelectStatementColumns,
pub elements: Vec<SelectStatementStackElement>,
pub order_by_clause: Option<Vec<OrderByClause>>,
pub limit_clause: Option<LimitClause>,
Expand Down Expand Up @@ -109,17 +110,17 @@ pub struct ColumnValue {
pub value: Value,
}

#[derive(Debug, PartialEq)]
#[derive(Debug, PartialEq, Clone)]
pub enum SelectStatementColumns {
All,
Specific(Vec<String>),
}

impl SelectStatementColumns {
pub fn columns(&self) -> Result<&Vec<String>, String> {
pub fn columns(&self) -> Result<Vec<&String>, String> {
return match self {
SelectStatementColumns::All => Err("Cannot get columns from all columns".to_string()),
SelectStatementColumns::Specific(columns) => Ok(columns),
SelectStatementColumns::Specific(columns) => Ok(columns.iter().map(|column| column).collect()),
}
}
}
Expand Down Expand Up @@ -355,6 +356,7 @@ mod tests {
let expected = vec![
Ok(DatabaseSqlStatement {
sql_statement: SqlStatement::Select(SelectStatementStack {
columns: SelectStatementColumns::All,
elements: vec![SelectStatementStackElement::SelectStatement(SelectStatement {
table_name: "users".to_string(),
columns: SelectStatementColumns::All,
Expand Down Expand Up @@ -402,7 +404,6 @@ mod tests {
token(TokenTypes::EOF, ""),
];
let result = generate(tokens);
println!("{:?}", result);
assert!(result[0].is_err());
assert!(result[1].is_ok());
let expected = vec![
Expand Down Expand Up @@ -448,6 +449,7 @@ mod tests {
let expected = vec![
Ok(DatabaseSqlStatement {
sql_statement: SqlStatement::Select(SelectStatementStack {
columns: SelectStatementColumns::All,
elements: vec![SelectStatementStackElement::SelectStatement(SelectStatement {
table_name: "users".to_string(),
columns: SelectStatementColumns::All,
Expand Down
2 changes: 2 additions & 0 deletions src/interpreter/ast/parser.rs
Original file line number Diff line number Diff line change
Expand Up @@ -148,6 +148,7 @@ mod tests {
parser.advance()?;
parser.advance_past_semicolon()?;
return Ok(SqlStatement::Select(SelectStatementStack {
columns: SelectStatementColumns::All,
elements: vec![SelectStatementStackElement::SelectStatement(SelectStatement {
table_name: "users".to_string(),
columns: SelectStatementColumns::All,
Expand Down Expand Up @@ -202,6 +203,7 @@ mod tests {
// Select
let result = parser.next_statement(builder);
let expected = Some(Ok(SqlStatement::Select(SelectStatementStack {
columns: SelectStatementColumns::All,
elements: vec![SelectStatementStackElement::SelectStatement(SelectStatement {
table_name: "users".to_string(),
columns: SelectStatementColumns::All,
Expand Down
Loading
Loading