Skip to main content

basalt/
engine.rs

1//! SQL executor and the first query-planning layer.
2//!
3//! Queries are resolved into a flat row/schema representation. This keeps
4//! joins, grouping, scalar expressions, and three-valued predicates on one
5//! execution path while retaining the simple `State` API used by the storage
6//! tests.
7
8use std::cmp::Ordering;
9use std::collections::HashSet;
10
11use crate::db::{Column, DbError, DbErrorKind, Row, State, StatementResult, Table, dberr};
12use crate::eval::{self, ColumnBinding};
13use crate::planner::{self, AccessPath};
14use crate::sql::ast::{ColumnDef, Expr, JoinKind, SelectItems, Statement};
15use crate::types::Value;
16
17pub(crate) const MCP_EXECUTION_WORK_LIMIT: usize = 1_000_000;
18
19#[derive(Debug)]
20pub(crate) struct ExecutionBudget {
21    limit: Option<usize>,
22    remaining: Option<usize>,
23}
24
25impl ExecutionBudget {
26    pub(crate) fn unlimited() -> Self {
27        Self {
28            limit: None,
29            remaining: None,
30        }
31    }
32
33    pub(crate) fn bounded(limit: usize) -> Self {
34        Self {
35            limit: Some(limit),
36            remaining: Some(limit),
37        }
38    }
39
40    fn is_unlimited(&self) -> bool {
41        self.remaining.is_none()
42    }
43
44    fn consume(&mut self, units: usize, operation: &str) -> Result<(), DbError> {
45        let Some(remaining) = &mut self.remaining else {
46            return Ok(());
47        };
48        if units > *remaining {
49            *remaining = 0;
50            let limit = self.limit.expect("bounded budgets have a limit");
51            return Err(dberr(
52                DbErrorKind::Limit,
53                format!(
54                    "execution exceeded the {limit}-unit work limit while {operation}; narrow the query or use the CLI for larger jobs"
55                ),
56            ));
57        }
58        *remaining -= units;
59        Ok(())
60    }
61
62    fn row(&mut self, row: &Row, operation: &str) -> Result<(), DbError> {
63        if self.is_unlimited() {
64            return Ok(());
65        }
66        self.consume(row_work_units(row), operation)
67    }
68
69    fn table_clone(&mut self, table: &Table, operation: &str) -> Result<(), DbError> {
70        if self.is_unlimited() {
71            return Ok(());
72        }
73        let units = table
74            .scan()
75            .fold(table.columns.len().max(1), |total, (_, row)| {
76                total.saturating_add(row_work_units(row))
77            });
78        self.consume(units, operation)
79    }
80
81    pub(crate) fn state_clone(&mut self, state: &State, operation: &str) -> Result<(), DbError> {
82        if self.is_unlimited() {
83            return Ok(());
84        }
85        for table in state.tables.values() {
86            self.table_clone(table, operation)?;
87        }
88        Ok(())
89    }
90}
91
92fn value_work_units(value: &Value) -> usize {
93    match value {
94        Value::Text(value) => value.len().saturating_add(1023) / 1024 + 1,
95        _ => 1,
96    }
97}
98
99fn row_work_units(row: &Row) -> usize {
100    row.iter().map(value_work_units).sum::<usize>().max(1)
101}
102
103fn joined_row_work_units(left: &Row, right: &Row) -> usize {
104    row_work_units(left)
105        .saturating_add(row_work_units(right))
106        .saturating_add(1)
107}
108
109pub fn execute(state: &mut State, stmt: &Statement) -> Result<StatementResult, DbError> {
110    let mut budget = ExecutionBudget::unlimited();
111    execute_with_budget(state, stmt, &mut budget)
112}
113
114pub(crate) fn execute_with_budget(
115    state: &mut State,
116    stmt: &Statement,
117    budget: &mut ExecutionBudget,
118) -> Result<StatementResult, DbError> {
119    match stmt {
120        Statement::CreateTable {
121            name,
122            if_not_exists,
123            columns,
124        } => {
125            if state.contains_table(name) {
126                if *if_not_exists {
127                    return Ok(StatementResult::CreateTable { name: name.clone() });
128                }
129                return Err(dberr(
130                    DbErrorKind::Constraint,
131                    format!("table '{name}' already exists"),
132                ));
133            }
134            let columns: Vec<Column> = columns.iter().map(col_from_def).collect();
135            state
136                .tables
137                .insert(name.clone(), Table::new(name, columns)?);
138            Ok(StatementResult::CreateTable { name: name.clone() })
139        }
140        Statement::DropTable { name, if_exists } => {
141            if let Some(table) = state.table(name) {
142                budget.table_clone(table, "dropping a table")?;
143            }
144            if state.remove_table(name).is_none() {
145                if *if_exists {
146                    return Ok(StatementResult::DropTable { name: name.clone() });
147                }
148                return Err(dberr(
149                    DbErrorKind::UnknownTable,
150                    format!("no such table: {name}"),
151                ));
152            }
153            Ok(StatementResult::DropTable { name: name.clone() })
154        }
155        Statement::CreateIndex {
156            name,
157            table,
158            column,
159            unique,
160            if_not_exists,
161        } => {
162            if state.contains_index(name) {
163                if *if_not_exists {
164                    return Ok(StatementResult::CreateIndex {
165                        name: name.clone(),
166                        table: table.clone(),
167                        column: column.clone(),
168                    });
169                }
170                return Err(dberr(
171                    DbErrorKind::Constraint,
172                    format!("index '{name}' already exists"),
173                ));
174            }
175            let table_ref = state.table(table).ok_or_else(|| {
176                dberr(DbErrorKind::UnknownTable, format!("no such table: {table}"))
177            })?;
178            if table_ref.has_index(name) {
179                if *if_not_exists {
180                    return Ok(StatementResult::CreateIndex {
181                        name: name.clone(),
182                        table: table.clone(),
183                        column: column.clone(),
184                    });
185                }
186                return Err(dberr(
187                    DbErrorKind::Constraint,
188                    format!("index '{name}' already exists"),
189                ));
190            }
191            budget.table_clone(table_ref, "building an index")?;
192            let table_ref = state.table_mut(table).unwrap();
193            let column_index = table_ref.column_index(column)?;
194            table_ref.create_index(name, column_index, *unique)?;
195            Ok(StatementResult::CreateIndex {
196                name: name.clone(),
197                table: table.clone(),
198                column: column.clone(),
199            })
200        }
201        Statement::DropIndex { name, if_exists } => {
202            let dropped = state
203                .tables
204                .values_mut()
205                .any(|table| table.drop_index(name));
206            if !dropped && !if_exists {
207                return Err(dberr(
208                    DbErrorKind::Constraint,
209                    format!("no such index: {name}"),
210                ));
211            }
212            Ok(StatementResult::DropIndex { name: name.clone() })
213        }
214        Statement::Insert {
215            table,
216            columns,
217            rows,
218        } => exec_insert(state, table, columns, rows, budget),
219        Statement::InsertSelect {
220            table,
221            columns,
222            query,
223        } => exec_insert_select(state, table, columns, query, budget),
224        Statement::Select { .. } => exec_select(state, stmt, budget),
225        Statement::Explain(inner) => exec_explain(state, inner, budget),
226        Statement::Update {
227            table,
228            assignments,
229            where_clause,
230        } => exec_update(state, table, assignments, where_clause, budget),
231        Statement::Delete {
232            table,
233            where_clause,
234        } => exec_delete(state, table, where_clause, budget),
235        Statement::Begin => Ok(StatementResult::Begin),
236        Statement::Commit => Ok(StatementResult::Commit),
237        Statement::Rollback => Ok(StatementResult::Rollback),
238        Statement::Checkpoint => Ok(StatementResult::Checkpoint),
239    }
240}
241
242fn col_from_def(definition: &ColumnDef) -> Column {
243    Column {
244        name: definition.name.clone(),
245        ty: definition.ty.clone(),
246        not_null: definition.not_null,
247        unique: definition.unique,
248        primary_key: definition.primary_key,
249    }
250}
251
252fn eval_const(expr: &Expr) -> Result<Value, DbError> {
253    // VALUES expressions are evaluated without a row.  This supports
254    // constant arithmetic, NULL checks, and scalar functions while naturally
255    // rejecting column references and aggregate expressions.
256    let empty_row: Row = Vec::new();
257    eval::eval_with_schema(&[], &empty_row, expr)
258}
259
260fn exec_insert(
261    state: &mut State,
262    table_name: &str,
263    columns: &Option<Vec<String>>,
264    rows: &[Vec<Expr>],
265    budget: &mut ExecutionBudget,
266) -> Result<StatementResult, DbError> {
267    let table = state.table(table_name).ok_or_else(|| {
268        dberr(
269            DbErrorKind::UnknownTable,
270            format!("no such table: {table_name}"),
271        )
272    })?;
273    let column_indices = match columns {
274        Some(names) => {
275            let mut result = Vec::with_capacity(names.len());
276            for name in names {
277                result.push(table.column_index(name)?);
278            }
279            ensure_unique_columns(&result, "INSERT column list")?;
280            result
281        }
282        None => (0..table.columns.len()).collect(),
283    };
284    if rows.is_empty() {
285        return Ok(StatementResult::Insert { rows_affected: 0 });
286    }
287    let expected = column_indices.len();
288    budget.table_clone(table, "preparing an insert")?;
289    let table = state.table_mut(table_name).unwrap();
290    let mut candidate = table.clone();
291    for values_expr in rows {
292        budget.consume(expected.max(1), "materializing inserted rows")?;
293        if values_expr.len() != expected {
294            return Err(dberr(
295                DbErrorKind::ColumnCount,
296                format!(
297                    "column count mismatch: expected {expected}, got {}",
298                    values_expr.len()
299                ),
300            ));
301        }
302        let mut values = vec![Value::Null; candidate.columns.len()];
303        for (column, expr) in column_indices.iter().zip(values_expr) {
304            values[*column] = candidate.coerce_val(&eval_const(expr)?, *column)?;
305        }
306        budget.row(&values, "materializing inserted rows")?;
307        candidate.insert_row(values)?;
308    }
309    let count = rows.len();
310    *table = candidate;
311    Ok(StatementResult::Insert {
312        rows_affected: count,
313    })
314}
315
316fn exec_insert_select(
317    state: &mut State,
318    table_name: &str,
319    columns: &Option<Vec<String>>,
320    query: &Statement,
321    budget: &mut ExecutionBudget,
322) -> Result<StatementResult, DbError> {
323    if !matches!(query, Statement::Select { .. }) {
324        return Err(dberr(
325            DbErrorKind::Syntax("INSERT SELECT requires a SELECT query".into()),
326            "INSERT SELECT requires a SELECT query",
327        ));
328    }
329    let selected = match exec_select(state, query, budget)? {
330        StatementResult::Select { rows, .. } => rows,
331        _ => unreachable!(),
332    };
333    let table = state.table(table_name).ok_or_else(|| {
334        dberr(
335            DbErrorKind::UnknownTable,
336            format!("no such table: {table_name}"),
337        )
338    })?;
339    let column_indices = match columns {
340        Some(names) => {
341            let result = names
342                .iter()
343                .map(|name| table.column_index(name))
344                .collect::<Result<Vec<_>, _>>()?;
345            ensure_unique_columns(&result, "INSERT column list")?;
346            result
347        }
348        None => (0..table.columns.len()).collect(),
349    };
350    let expected = column_indices.len();
351    budget.table_clone(table, "preparing an insert")?;
352    let table = state.table_mut(table_name).unwrap();
353    let mut candidate = table.clone();
354    for source in &selected {
355        budget.row(source, "materializing inserted rows")?;
356        if source.len() != expected {
357            return Err(dberr(
358                DbErrorKind::ColumnCount,
359                format!(
360                    "column count mismatch: expected {expected}, got {}",
361                    source.len()
362                ),
363            ));
364        }
365        let mut values = vec![Value::Null; candidate.columns.len()];
366        for (target, value) in column_indices.iter().zip(source) {
367            values[*target] = candidate.coerce_val(value, *target)?;
368        }
369        candidate.insert_row(values)?;
370    }
371    *table = candidate;
372    Ok(StatementResult::Insert {
373        rows_affected: selected.len(),
374    })
375}
376
377fn exec_explain(
378    state: &State,
379    stmt: &Statement,
380    budget: &mut ExecutionBudget,
381) -> Result<StatementResult, DbError> {
382    let Statement::Select {
383        from,
384        from_alias,
385        where_clause,
386        ..
387    } = stmt
388    else {
389        return Ok(StatementResult::Explain(
390            "EXPLAIN supports SELECT statements".into(),
391        ));
392    };
393    if from.is_empty() {
394        return Ok(StatementResult::Explain(
395            "ConstantScan estimated_rows=1".into(),
396        ));
397    }
398    let table = state
399        .table(from)
400        .ok_or_else(|| dberr(DbErrorKind::UnknownTable, format!("no such table: {from}")))?;
401    let plan = planner::choose(
402        table,
403        from_alias.as_deref().or(Some(from)),
404        where_clause.as_ref(),
405    );
406    let candidate_count = match &plan.access {
407        AccessPath::TableScan => 0,
408        AccessPath::IndexScan { row_ids, .. } | AccessPath::IndexRange { row_ids, .. } => {
409            row_ids.len()
410        }
411    };
412    budget.consume(candidate_count, "building an index access plan")?;
413    let text = match plan.access {
414        AccessPath::TableScan => format!(
415            "TableScan table={from} estimated_rows={}",
416            plan.estimated_rows
417        ),
418        AccessPath::IndexScan {
419            index_name,
420            column,
421            key,
422            row_ids,
423        } => format!(
424            "IndexScan index={index_name} table={from} column={column} key={key} candidates={} estimated_rows={}",
425            row_ids.len(),
426            plan.estimated_rows
427        ),
428        AccessPath::IndexRange {
429            index_name,
430            column,
431            low,
432            high,
433            row_ids,
434        } => format!(
435            "IndexRange index={index_name} table={from} column={column} low={low:?} high={high:?} candidates={} estimated_rows={}",
436            row_ids.len(),
437            plan.estimated_rows
438        ),
439    };
440    Ok(StatementResult::Explain(text))
441}
442
443#[derive(Clone)]
444struct QueryRow {
445    values: Row,
446}
447
448fn relation_schema(table: &Table, alias: Option<&str>) -> Vec<ColumnBinding> {
449    let mut relations = vec![table.name.clone()];
450    if let Some(alias) = alias
451        && !relations
452            .iter()
453            .any(|value| value.eq_ignore_ascii_case(alias))
454    {
455        relations.push(alias.to_string());
456    }
457    table
458        .columns
459        .iter()
460        .map(|column| ColumnBinding::with_relations(column.name.clone(), relations.clone()))
461        .collect()
462}
463
464fn exec_select(
465    state: &State,
466    stmt: &Statement,
467    budget: &mut ExecutionBudget,
468) -> Result<StatementResult, DbError> {
469    let Statement::Select {
470        distinct,
471        columns,
472        from,
473        from_alias,
474        joins,
475        where_clause,
476        group_by,
477        having,
478        order_by,
479        order_by_exprs,
480        limit,
481        offset,
482    } = stmt
483    else {
484        unreachable!()
485    };
486
487    let (mut schema, mut input): (Vec<ColumnBinding>, Vec<QueryRow>) = if from.is_empty() {
488        if !joins.is_empty() {
489            return Err(dberr(
490                DbErrorKind::Syntax("JOIN requires a FROM table".into()),
491                "JOIN requires a FROM table",
492            ));
493        }
494        (Vec::new(), vec![QueryRow { values: Vec::new() }])
495    } else {
496        let base = state
497            .table(from)
498            .ok_or_else(|| dberr(DbErrorKind::UnknownTable, format!("no such table: {from}")))?;
499        let schema = relation_schema(base, from_alias.as_deref());
500        let base_plan = planner::choose(
501            base,
502            from_alias.as_deref().or(Some(from)),
503            where_clause.as_ref(),
504        );
505        let mut base_rows = Vec::new();
506        match base_plan.access {
507            AccessPath::TableScan => {
508                for (_, row) in base.scan() {
509                    budget.row(row, "scanning the base table")?;
510                    base_rows.push(row.clone());
511                }
512            }
513            AccessPath::IndexScan { row_ids, .. } | AccessPath::IndexRange { row_ids, .. } => {
514                for rid in row_ids {
515                    if let Some(row) = base.get_row(rid) {
516                        budget.row(row, "materializing index matches")?;
517                        base_rows.push(row.clone());
518                    }
519                }
520            }
521        }
522        let input = base_rows
523            .into_iter()
524            .map(|values| QueryRow { values })
525            .collect();
526        (schema, input)
527    };
528
529    for join in joins {
530        let right = state.table(&join.table).ok_or_else(|| {
531            dberr(
532                DbErrorKind::UnknownTable,
533                format!("no such table: {}", join.table),
534            )
535        })?;
536        let right_schema = relation_schema(right, join.alias.as_deref());
537        let mut joined_schema = schema.clone();
538        joined_schema.extend(right_schema.iter().cloned());
539        if let Some(on) = &join.on {
540            reject_aggregate(on, "JOIN ON")?;
541            eval::validate_with_schema(&joined_schema, on)?;
542        }
543        let mut right_rows = Vec::new();
544        for (_, row) in right.scan() {
545            budget.row(row, "scanning a joined table")?;
546            right_rows.push(row.clone());
547        }
548        budget.consume(right_rows.len(), "tracking join matches")?;
549        let mut next = Vec::new();
550        let mut matched_right = vec![false; right_rows.len()];
551        if join.kind == JoinKind::Right {
552            for (right_index, right_row) in right_rows.iter().enumerate() {
553                let mut matched = false;
554                for left in &input {
555                    budget.consume(
556                        joined_row_work_units(&left.values, right_row),
557                        "materializing join candidates",
558                    )?;
559                    let mut values = left.values.clone();
560                    values.extend(right_row.iter().cloned());
561                    if join_passes(&join.kind, &join.on, &joined_schema, &values)? {
562                        matched = true;
563                        next.push(QueryRow { values });
564                    }
565                }
566                if !matched {
567                    budget.consume(
568                        input_schema_width(&schema)
569                            .saturating_add(row_work_units(right_row))
570                            .saturating_add(1),
571                        "materializing an unmatched join row",
572                    )?;
573                    let mut values = vec![Value::Null; input_schema_width(&schema)];
574                    values.extend(right_row.iter().cloned());
575                    next.push(QueryRow { values });
576                }
577                matched_right[right_index] = matched;
578            }
579        } else {
580            for left in &input {
581                let mut matched = false;
582                for (right_index, right_row) in right_rows.iter().enumerate() {
583                    budget.consume(
584                        joined_row_work_units(&left.values, right_row),
585                        "materializing join candidates",
586                    )?;
587                    let mut values = left.values.clone();
588                    values.extend(right_row.iter().cloned());
589                    if join_passes(&join.kind, &join.on, &joined_schema, &values)? {
590                        matched = true;
591                        matched_right[right_index] = true;
592                        next.push(QueryRow { values });
593                    }
594                }
595                if (join.kind == JoinKind::Left || join.kind == JoinKind::Full) && !matched {
596                    budget.consume(
597                        row_work_units(&left.values)
598                            .saturating_add(right.columns.len())
599                            .saturating_add(1),
600                        "materializing an unmatched join row",
601                    )?;
602                    let mut values = left.values.clone();
603                    values.extend(std::iter::repeat_n(Value::Null, right.columns.len()));
604                    next.push(QueryRow { values });
605                }
606            }
607            if join.kind == JoinKind::Full {
608                for (right_index, right_row) in right_rows.iter().enumerate() {
609                    if !matched_right[right_index] {
610                        budget.consume(
611                            input_schema_width(&schema)
612                                .saturating_add(row_work_units(right_row))
613                                .saturating_add(1),
614                            "materializing an unmatched join row",
615                        )?;
616                        let mut values = vec![Value::Null; input_schema_width(&schema)];
617                        values.extend(right_row.iter().cloned());
618                        next.push(QueryRow { values });
619                    }
620                }
621            }
622        }
623        schema = joined_schema;
624        input = next;
625    }
626
627    let expanded_items = match columns {
628        SelectItems::Star => Vec::new(),
629        SelectItems::List(items) => expand_select_items(&schema, items)?,
630    };
631    let effective_order: Vec<(Expr, bool)> = if !order_by_exprs.is_empty() {
632        order_by_exprs
633            .iter()
634            .map(|(expr, ascending)| {
635                (
636                    resolve_order_alias(expr.clone(), &expanded_items),
637                    *ascending,
638                )
639            })
640            .collect()
641    } else {
642        order_by
643            .iter()
644            .map(|(name, ascending)| (order_name_expr(name), *ascending))
645            .collect()
646    };
647    if let Some(predicate) = where_clause {
648        reject_aggregate(predicate, "WHERE")?;
649        eval::validate_with_schema(&schema, predicate)?;
650    }
651    for expr in group_by {
652        reject_aggregate(expr, "GROUP BY")?;
653        eval::validate_with_schema(&schema, expr)?;
654    }
655    if let Some(predicate) = having {
656        eval::validate_with_schema(&schema, predicate)?;
657    }
658    for expr in &expanded_items {
659        eval::validate_with_schema(&schema, expr)?;
660    }
661    for (expr, _) in &effective_order {
662        eval::validate_with_schema(&schema, expr)?;
663    }
664
665    let mut filtered = Vec::with_capacity(input.len());
666    for row in input {
667        budget.row(&row.values, "evaluating a filter")?;
668        if where_clause
669            .as_ref()
670            .map(|predicate| eval::where_matches_with_schema(&schema, &row.values, predicate))
671            .transpose()?
672            .unwrap_or(true)
673        {
674            filtered.push(row.values);
675        }
676    }
677
678    let has_aggregate = match columns {
679        SelectItems::Star => false,
680        SelectItems::List(items) => items.iter().any(eval::contains_aggregate),
681    } || having.as_ref().is_some_and(eval::contains_aggregate);
682    let grouped = has_aggregate || !group_by.is_empty() || having.is_some();
683    let groups = make_groups(&schema, filtered, group_by, grouped, budget)?;
684
685    let output_names = select_output_names(&schema, columns, &expanded_items);
686    let mut projected: Vec<(Row, Vec<Value>)> = Vec::new();
687    for group in groups {
688        let group_work = group.iter().fold(1usize, |total, row| {
689            total.saturating_add(row_work_units(row))
690        });
691        budget.consume(group_work, "evaluating a result group")?;
692        if let Some(predicate) = having
693            && !eval::eval_group(&schema, &group, predicate)?
694                .is_truthy()
695                .unwrap_or(false)
696        {
697            continue;
698        }
699        let row = match columns {
700            SelectItems::Star => group.first().cloned().unwrap_or_default(),
701            SelectItems::List(_) => expanded_items
702                .iter()
703                .map(|expr| eval::eval_group(&schema, &group, expr))
704                .collect::<Result<Vec<_>, _>>()?,
705        };
706        let mut keys = Vec::with_capacity(effective_order.len());
707        for (expression, _) in &effective_order {
708            keys.push(eval::eval_group(&schema, &group, expression)?);
709        }
710        budget.consume(
711            row_work_units(&row)
712                .saturating_add(keys.iter().map(value_work_units).sum())
713                .saturating_add(1),
714            "materializing a result row",
715        )?;
716        projected.push((row, keys));
717    }
718
719    if !effective_order.is_empty() {
720        budget.consume(
721            projected.len().saturating_mul(effective_order.len().max(1)),
722            "sorting result rows",
723        )?;
724        projected.sort_by(|left, right| {
725            for (index, (_, ascending)) in effective_order.iter().enumerate() {
726                let mut ordering = left.1[index].cmp_value(&right.1[index]);
727                if !*ascending {
728                    ordering = ordering.reverse();
729                }
730                if ordering != Ordering::Equal {
731                    return ordering;
732                }
733            }
734            Ordering::Equal
735        });
736    }
737    let mut rows: Vec<Row> = projected.into_iter().map(|(row, _)| row).collect();
738    if *distinct {
739        let mut seen = Vec::new();
740        let mut distinct_rows = Vec::new();
741        for row in rows {
742            budget.consume(
743                seen.len()
744                    .saturating_add(row_work_units(&row))
745                    .saturating_add(1),
746                "deduplicating result rows",
747            )?;
748            if !seen.contains(&row) {
749                seen.push(row.clone());
750                distinct_rows.push(row);
751            }
752        }
753        rows = distinct_rows;
754    }
755    if let Some(offset) = offset {
756        if *offset >= rows.len() as u64 {
757            rows.clear();
758        } else {
759            rows = rows.split_off(*offset as usize);
760        }
761    }
762    if let Some(limit) = limit {
763        rows.truncate(*limit as usize);
764    }
765    Ok(StatementResult::Select {
766        columns: output_names,
767        rows,
768    })
769}
770
771fn input_schema_width(schema: &[ColumnBinding]) -> usize {
772    schema.len()
773}
774
775fn join_passes(
776    kind: &JoinKind,
777    on: &Option<Expr>,
778    schema: &[ColumnBinding],
779    row: &Row,
780) -> Result<bool, DbError> {
781    match (kind, on) {
782        (JoinKind::Cross, _) => Ok(true),
783        (_, Some(on)) => eval::where_matches_with_schema(schema, row, on),
784        (_, None) => Ok(true),
785    }
786}
787
788fn make_groups(
789    schema: &[ColumnBinding],
790    rows: Vec<Row>,
791    group_by: &[Expr],
792    grouped: bool,
793    budget: &mut ExecutionBudget,
794) -> Result<Vec<Vec<Row>>, DbError> {
795    if !grouped {
796        return Ok(rows.into_iter().map(|row| vec![row]).collect());
797    }
798    if group_by.is_empty() {
799        return Ok(vec![rows]);
800    }
801    let mut groups: Vec<(Vec<Value>, Vec<Row>)> = Vec::new();
802    for row in rows {
803        budget.row(&row, "evaluating group keys")?;
804        let key = group_by
805            .iter()
806            .map(|expr| eval::eval_with_schema(schema, &row, expr))
807            .collect::<Result<Vec<_>, _>>()?;
808        if let Some((_, values)) = groups.iter_mut().find(|(existing, _)| {
809            existing.len() == key.len()
810                && existing
811                    .iter()
812                    .zip(&key)
813                    .all(|(left, right)| left.cmp_value(right) == Ordering::Equal)
814        }) {
815            values.push(row);
816        } else {
817            groups.push((key, vec![row]));
818        }
819    }
820    Ok(groups.into_iter().map(|(_, rows)| rows).collect())
821}
822
823fn order_name_expr(name: &str) -> Expr {
824    if let Some((relation, column)) = name.split_once('.') {
825        Expr::ColumnRef {
826            relation: relation.to_string(),
827            column: column.to_string(),
828        }
829    } else {
830        Expr::Column(name.to_string())
831    }
832}
833
834fn resolve_order_alias(expr: Expr, items: &[Expr]) -> Expr {
835    let Expr::Column(name) = &expr else {
836        return expr;
837    };
838    for item in items {
839        if let Expr::Alias { expr: inner, alias } = item
840            && alias.eq_ignore_ascii_case(name)
841        {
842            return *inner.clone();
843        }
844    }
845    expr
846}
847
848fn expand_select_items(schema: &[ColumnBinding], items: &[Expr]) -> Result<Vec<Expr>, DbError> {
849    let mut expanded = Vec::new();
850    for item in items {
851        if let Expr::QualifiedWildcard(relation) = item {
852            let mut found = false;
853            for column in schema {
854                if column
855                    .relations
856                    .iter()
857                    .any(|value| value.eq_ignore_ascii_case(relation))
858                {
859                    found = true;
860                    expanded.push(Expr::ColumnRef {
861                        relation: relation.clone(),
862                        column: column.name.clone(),
863                    });
864                }
865            }
866            if !found {
867                return Err(dberr(
868                    DbErrorKind::UnknownTable,
869                    format!("no such table: {relation}"),
870                ));
871            }
872        } else {
873            expanded.push(item.clone());
874        }
875    }
876    Ok(expanded)
877}
878
879fn select_output_names(
880    schema: &[ColumnBinding],
881    columns: &SelectItems,
882    expanded: &[Expr],
883) -> Vec<String> {
884    match columns {
885        SelectItems::Star => schema.iter().map(|column| column.name.clone()).collect(),
886        SelectItems::List(items) => {
887            let mut names = Vec::new();
888            let mut expanded_index = 0usize;
889            for item in items {
890                if let Expr::QualifiedWildcard(relation) = item {
891                    names.extend(
892                        schema
893                            .iter()
894                            .filter(|column| {
895                                column
896                                    .relations
897                                    .iter()
898                                    .any(|value| value.eq_ignore_ascii_case(relation))
899                            })
900                            .map(|column| column.name.clone()),
901                    );
902                    expanded_index += names.len().saturating_sub(expanded_index);
903                } else if let Some(expr) = expanded.get(expanded_index) {
904                    names.push(expr_label(expr));
905                    expanded_index += 1;
906                }
907            }
908            names
909        }
910    }
911}
912
913fn expr_label(expr: &Expr) -> String {
914    match expr {
915        Expr::Column(name) => name.clone(),
916        Expr::ColumnRef { relation, column } => format!("{relation}.{column}"),
917        Expr::QualifiedWildcard(relation) => format!("{relation}.*"),
918        Expr::Alias { alias, .. } => alias.clone(),
919        Expr::Function { name, .. } => name.clone(),
920        Expr::Literal(value) => value.to_string(),
921        _ => "expr".into(),
922    }
923}
924
925fn exec_update(
926    state: &mut State,
927    table_name: &str,
928    assignments: &[(String, Expr)],
929    where_clause: &Option<Expr>,
930    budget: &mut ExecutionBudget,
931) -> Result<StatementResult, DbError> {
932    let targets = {
933        let table = state.table(table_name).ok_or_else(|| {
934            dberr(
935                DbErrorKind::UnknownTable,
936                format!("no such table: {table_name}"),
937            )
938        })?;
939        let targets = assignments
940            .iter()
941            .map(|(name, _)| table.column_index(name))
942            .collect::<Result<Vec<_>, _>>()?;
943        ensure_unique_columns(&targets, "UPDATE assignment list")?;
944        let schema = eval::schema_for_table(table);
945        for (_, expr) in assignments {
946            reject_aggregate(expr, "UPDATE")?;
947            eval::validate_with_schema(&schema, expr)?;
948        }
949        if let Some(predicate) = where_clause {
950            reject_aggregate(predicate, "WHERE")?;
951            eval::validate_with_schema(&schema, predicate)?;
952        }
953        targets
954    };
955    let table = state.table(table_name).unwrap();
956    budget.table_clone(table, "preparing an update")?;
957    let mut plan = Vec::new();
958    {
959        let table = state.table(table_name).unwrap();
960        for (rid, row) in table.scan() {
961            budget.row(row, "scanning rows for update")?;
962            if where_clause
963                .as_ref()
964                .map(|predicate| eval::where_matches(table, row, predicate))
965                .transpose()?
966                .unwrap_or(true)
967            {
968                budget.row(row, "materializing updated rows")?;
969                let mut new_row = row.clone();
970                for (index, (_, expr)) in assignments.iter().enumerate() {
971                    let value = eval::eval(table, row, expr)?;
972                    new_row[targets[index]] = table.coerce_val(&value, targets[index])?;
973                }
974                plan.push((rid, new_row));
975            }
976        }
977    }
978    let table = state.table_mut(table_name).unwrap();
979    let mut candidate = table.clone();
980    let affected = plan.len();
981    for (rid, row) in plan {
982        candidate.replace_row(rid, row)?;
983    }
984    *table = candidate;
985    Ok(StatementResult::Update {
986        rows_affected: affected,
987    })
988}
989
990fn ensure_unique_columns(indices: &[usize], label: &str) -> Result<(), DbError> {
991    let mut seen = HashSet::with_capacity(indices.len());
992    if indices.iter().any(|index| !seen.insert(*index)) {
993        return Err(dberr(
994            DbErrorKind::Constraint,
995            format!("{label} contains a duplicate column"),
996        ));
997    }
998    Ok(())
999}
1000
1001fn reject_aggregate(expr: &Expr, context: &str) -> Result<(), DbError> {
1002    if eval::contains_aggregate(expr) {
1003        return Err(dberr(
1004            DbErrorKind::Syntax(format!("aggregate functions are not allowed in {context}")),
1005            format!("aggregate functions are not allowed in {context}"),
1006        ));
1007    }
1008    Ok(())
1009}
1010
1011fn exec_delete(
1012    state: &mut State,
1013    table_name: &str,
1014    where_clause: &Option<Expr>,
1015    budget: &mut ExecutionBudget,
1016) -> Result<StatementResult, DbError> {
1017    let mut ids = Vec::new();
1018    {
1019        let table = state.table(table_name).ok_or_else(|| {
1020            dberr(
1021                DbErrorKind::UnknownTable,
1022                format!("no such table: {table_name}"),
1023            )
1024        })?;
1025        let schema = eval::schema_for_table(table);
1026        if let Some(predicate) = where_clause {
1027            reject_aggregate(predicate, "WHERE")?;
1028            eval::validate_with_schema(&schema, predicate)?;
1029        }
1030        for (rid, row) in table.scan() {
1031            budget.row(row, "scanning rows for delete")?;
1032            if where_clause
1033                .as_ref()
1034                .map(|predicate| eval::where_matches(table, row, predicate))
1035                .transpose()?
1036                .unwrap_or(true)
1037            {
1038                ids.push(rid);
1039            }
1040        }
1041    }
1042    let table = state.table(table_name).unwrap();
1043    budget.table_clone(table, "preparing a delete")?;
1044    let table = state.table_mut(table_name).unwrap();
1045    let mut candidate = table.clone();
1046    for rid in &ids {
1047        candidate.delete_row(*rid)?;
1048    }
1049    *table = candidate;
1050    Ok(StatementResult::Delete {
1051        rows_affected: ids.len(),
1052    })
1053}
1054
1055#[cfg(test)]
1056mod tests {
1057    use super::*;
1058    use crate::sql::parser::parse;
1059
1060    fn state_with_rows(count: i64) -> State {
1061        let mut state = State::empty();
1062        let mut table = Table::new(
1063            "items",
1064            vec![Column {
1065                name: "id".into(),
1066                ty: crate::types::ColumnType::Integer,
1067                not_null: false,
1068                unique: false,
1069                primary_key: false,
1070            }],
1071        )
1072        .unwrap();
1073        for id in 0..count {
1074            table.insert_row(vec![Value::Integer(id)]).unwrap();
1075        }
1076        state.tables.insert("items".into(), table);
1077        state
1078    }
1079
1080    #[test]
1081    fn bounded_execution_rejects_materialization_before_result_conversion() {
1082        let mut state = state_with_rows(10);
1083        let statement = parse("SELECT * FROM items").unwrap().remove(0);
1084        let mut budget = ExecutionBudget::bounded(5);
1085
1086        let error = execute_with_budget(&mut state, &statement, &mut budget).unwrap_err();
1087
1088        assert_eq!(error.kind, DbErrorKind::Limit);
1089        assert!(error.message.contains("work limit"));
1090    }
1091
1092    #[test]
1093    fn bounded_mutation_does_not_publish_partial_state() {
1094        let mut state = state_with_rows(4);
1095        let statement = parse("UPDATE items SET id = id + 1").unwrap().remove(0);
1096        let mut budget = ExecutionBudget::bounded(10);
1097
1098        let error = execute_with_budget(&mut state, &statement, &mut budget).unwrap_err();
1099
1100        assert_eq!(error.kind, DbErrorKind::Limit);
1101        let query = parse("SELECT id FROM items ORDER BY id").unwrap().remove(0);
1102        let StatementResult::Select { rows, .. } = execute(&mut state, &query).unwrap() else {
1103            panic!("expected select result");
1104        };
1105        assert_eq!(
1106            rows,
1107            vec![
1108                vec![Value::Integer(0)],
1109                vec![Value::Integer(1)],
1110                vec![Value::Integer(2)],
1111                vec![Value::Integer(3)],
1112            ]
1113        );
1114    }
1115}