1use crate::db::{DbError, DbErrorKind, Row, Table, dberr};
8use crate::sql::ast::{BinOp, Expr, UnaryOp};
9use crate::types::Value;
10
11#[derive(Debug, Clone, PartialEq, Eq)]
14pub struct ColumnBinding {
15 pub name: String,
16 pub relations: Vec<String>,
17}
18
19impl ColumnBinding {
20 pub fn new(name: impl Into<String>, relation: impl Into<String>) -> ColumnBinding {
21 ColumnBinding {
22 name: name.into(),
23 relations: vec![relation.into()],
24 }
25 }
26
27 pub fn with_relations(name: impl Into<String>, relations: Vec<String>) -> ColumnBinding {
28 ColumnBinding {
29 name: name.into(),
30 relations,
31 }
32 }
33}
34
35pub fn schema_for_table(table: &Table) -> Vec<ColumnBinding> {
36 table
37 .columns
38 .iter()
39 .map(|column| ColumnBinding::new(column.name.clone(), table.name.clone()))
40 .collect()
41}
42
43pub fn eval(table: &Table, row: &Row, expr: &Expr) -> Result<Value, DbError> {
45 let schema = schema_for_table(table);
46 eval_with_schema(&schema, row, expr)
47}
48
49pub fn eval_with_schema(
52 schema: &[ColumnBinding],
53 row: &Row,
54 expr: &Expr,
55) -> Result<Value, DbError> {
56 match expr {
57 Expr::Literal(value) => Ok(value.clone()),
58 Expr::Column(name) => {
59 let index = resolve_column(schema, None, name)?;
60 Ok(row[index].clone())
61 }
62 Expr::QualifiedWildcard(relation) => Err(dberr(
63 DbErrorKind::Syntax(format!("{relation}.* is only valid in a SELECT list")),
64 format!("{relation}.* is only valid in a SELECT list"),
65 )),
66 Expr::ColumnRef { relation, column } => {
67 let index = resolve_column(schema, Some(relation), column)?;
68 Ok(row[index].clone())
69 }
70 Expr::Alias { expr, .. } => eval_with_schema(schema, row, expr),
71 Expr::Function {
72 name,
73 args,
74 distinct: _,
75 } => eval_function(schema, row, name, args),
76 Expr::IsNull { expr, negated } => {
77 let value = eval_with_schema(schema, row, expr)?;
78 Ok(Value::Boolean(matches!(value, Value::Null) != *negated))
79 }
80 Expr::Unary { op, expr } => apply_unary(op, eval_with_schema(schema, row, expr)?),
81 Expr::Binary { left, op, right } => apply_binary(
82 op,
83 eval_with_schema(schema, row, left)?,
84 eval_with_schema(schema, row, right)?,
85 ),
86 }
87}
88
89pub fn validate_with_schema(schema: &[ColumnBinding], expr: &Expr) -> Result<(), DbError> {
93 match expr {
94 Expr::Literal(_) => Ok(()),
95 Expr::Column(name) => resolve_column(schema, None, name).map(|_| ()),
96 Expr::QualifiedWildcard(relation) => Err(dberr(
97 DbErrorKind::Syntax(format!("{relation}.* is only valid in a SELECT list")),
98 format!("{relation}.* is only valid in a SELECT list"),
99 )),
100 Expr::ColumnRef { relation, column } => {
101 resolve_column(schema, Some(relation), column).map(|_| ())
102 }
103 Expr::Alias { expr, .. } => validate_with_schema(schema, expr),
104 Expr::Function {
105 name,
106 args,
107 distinct,
108 } => {
109 let upper = name.to_ascii_uppercase();
110 if is_aggregate_name(name) {
111 match upper.as_str() {
112 "COUNT" if args.len() <= 1 => {}
113 "SUM" | "AVG" | "MIN" | "MAX" if args.len() == 1 => {}
114 "COUNT" => {
115 return Err(dberr(
116 DbErrorKind::Syntax("COUNT expects at most one argument".into()),
117 "COUNT expects at most one argument",
118 ));
119 }
120 _ => {
121 return Err(dberr(
122 DbErrorKind::Syntax(format!("{upper} expects one argument")),
123 format!("{upper} expects one argument"),
124 ));
125 }
126 }
127 for arg in args {
128 if !matches!(arg, Expr::Column(value) if value == "*") {
129 validate_with_schema(schema, arg)?;
130 }
131 }
132 return Ok(());
133 }
134 if *distinct {
135 return Err(dberr(
136 DbErrorKind::Syntax("DISTINCT is only valid for aggregate functions".into()),
137 "DISTINCT is only valid for aggregate functions",
138 ));
139 }
140 match upper.as_str() {
141 "LOWER" | "UPPER" | "LENGTH" | "ABS" if args.len() == 1 => {}
142 "COALESCE" => {}
143 "NULLIF" if args.len() == 2 => {}
144 "LOWER" | "UPPER" | "LENGTH" | "ABS" => {
145 return Err(dberr(
146 DbErrorKind::Syntax(format!("{upper} expects one argument")),
147 format!("{upper} expects one argument"),
148 ));
149 }
150 "NULLIF" => {
151 return Err(dberr(
152 DbErrorKind::Syntax("NULLIF expects two arguments".into()),
153 "NULLIF expects two arguments",
154 ));
155 }
156 _ => {
157 return Err(dberr(
158 DbErrorKind::Syntax(format!("unknown function {name}")),
159 format!("unknown function: {name}"),
160 ));
161 }
162 }
163 for arg in args {
164 validate_with_schema(schema, arg)?;
165 }
166 Ok(())
167 }
168 Expr::Binary { left, right, .. } => {
169 validate_with_schema(schema, left)?;
170 validate_with_schema(schema, right)
171 }
172 Expr::Unary { expr, .. } | Expr::IsNull { expr, .. } => validate_with_schema(schema, expr),
173 }
174}
175
176fn resolve_column(
177 schema: &[ColumnBinding],
178 relation: Option<&str>,
179 name: &str,
180) -> Result<usize, DbError> {
181 let mut matches = Vec::new();
182 for (index, binding) in schema.iter().enumerate() {
183 let name_matches = binding.name.eq_ignore_ascii_case(name);
184 let relation_matches = relation
185 .map(|wanted| {
186 binding
187 .relations
188 .iter()
189 .any(|value| value.eq_ignore_ascii_case(wanted))
190 })
191 .unwrap_or(true);
192 if name_matches && relation_matches {
193 matches.push(index);
194 }
195 }
196 match matches.as_slice() {
197 [index] => Ok(*index),
198 [] => {
199 let label = relation
200 .map(|value| format!("{value}.{name}"))
201 .unwrap_or_else(|| name.to_string());
202 Err(dberr(
203 DbErrorKind::UnknownColumn,
204 format!("no such column: {label}"),
205 ))
206 }
207 _ => Err(dberr(
208 DbErrorKind::UnknownColumn,
209 format!("ambiguous column: {name}"),
210 )),
211 }
212}
213
214fn eval_function(
215 schema: &[ColumnBinding],
216 row: &Row,
217 name: &str,
218 args: &[Expr],
219) -> Result<Value, DbError> {
220 if is_aggregate_name(name) {
221 return Err(dberr(
222 DbErrorKind::TypeMismatch,
223 format!("aggregate function {name} requires a query group"),
224 ));
225 }
226 let values = args
227 .iter()
228 .map(|arg| eval_with_schema(schema, row, arg))
229 .collect::<Result<Vec<_>, _>>()?;
230 apply_scalar_function(&name.to_ascii_uppercase(), name, values)
231}
232
233fn apply_scalar_function(upper: &str, name: &str, values: Vec<Value>) -> Result<Value, DbError> {
234 match upper {
235 "LOWER" => match values.as_slice() {
236 [Value::Text(value)] => Ok(Value::Text(value.to_lowercase())),
237 [Value::Null] => Ok(Value::Null),
238 _ => Err(dberr(DbErrorKind::TypeMismatch, "LOWER expects TEXT")),
239 },
240 "UPPER" => match values.as_slice() {
241 [Value::Text(value)] => Ok(Value::Text(value.to_uppercase())),
242 [Value::Null] => Ok(Value::Null),
243 _ => Err(dberr(DbErrorKind::TypeMismatch, "UPPER expects TEXT")),
244 },
245 "LENGTH" => match values.as_slice() {
246 [Value::Text(value)] => Ok(Value::Integer(value.chars().count() as i64)),
247 [Value::Null] => Ok(Value::Null),
248 _ => Err(dberr(DbErrorKind::TypeMismatch, "LENGTH expects TEXT")),
249 },
250 "ABS" => match values.as_slice() {
251 [Value::Integer(value)] => value
252 .checked_abs()
253 .map(Value::Integer)
254 .ok_or_else(|| dberr(DbErrorKind::TypeMismatch, "integer overflow")),
255 [Value::Real(value)] if value.is_finite() => Ok(Value::Real(value.abs())),
256 [Value::Real(_)] => Err(dberr(DbErrorKind::TypeMismatch, "real overflow")),
257 [Value::Null] => Ok(Value::Null),
258 _ => Err(dberr(DbErrorKind::TypeMismatch, "ABS expects a number")),
259 },
260 "COALESCE" => Ok(values
261 .into_iter()
262 .find(|value| !matches!(value, Value::Null))
263 .unwrap_or(Value::Null)),
264 "NULLIF" => {
265 if values.len() != 2 {
266 return Err(dberr(
267 DbErrorKind::Syntax("NULLIF expects two arguments".into()),
268 "NULLIF expects two arguments",
269 ));
270 }
271 if matches!(values[0], Value::Null) || matches!(values[1], Value::Null) {
272 Ok(values[0].clone())
273 } else if values[0].cmp_value(&values[1]) == std::cmp::Ordering::Equal {
274 Ok(Value::Null)
275 } else {
276 Ok(values[0].clone())
277 }
278 }
279 _ => Err(dberr(
280 DbErrorKind::Syntax(format!("unknown function {name}")),
281 format!("unknown function: {name}"),
282 )),
283 }
284}
285
286fn apply_unary(op: &UnaryOp, value: Value) -> Result<Value, DbError> {
287 match op {
288 UnaryOp::Neg => match value {
289 Value::Integer(value) => value
290 .checked_neg()
291 .map(Value::Integer)
292 .ok_or_else(|| dberr(DbErrorKind::TypeMismatch, "integer overflow")),
293 Value::Real(value) if value.is_finite() => Ok(Value::Real(-value)),
294 Value::Real(_) => Err(dberr(DbErrorKind::TypeMismatch, "real overflow")),
295 Value::Null => Ok(Value::Null),
296 _ => Err(dberr(
297 DbErrorKind::TypeMismatch,
298 "cannot negate non-numeric value",
299 )),
300 },
301 UnaryOp::Not => match value.is_truthy() {
302 Some(value) => Ok(Value::Boolean(!value)),
303 None => Ok(Value::Null),
304 },
305 }
306}
307
308pub(crate) fn apply_binary(op: &BinOp, left: Value, right: Value) -> Result<Value, DbError> {
309 use BinOp::*;
310 match op {
311 And => {
312 let left = left.is_truthy();
313 let right = right.is_truthy();
314 if left == Some(false) || right == Some(false) {
315 Ok(Value::Boolean(false))
316 } else if left.is_none() || right.is_none() {
317 Ok(Value::Null)
318 } else {
319 Ok(Value::Boolean(true))
320 }
321 }
322 Or => {
323 let left = left.is_truthy();
324 let right = right.is_truthy();
325 if left == Some(true) || right == Some(true) {
326 Ok(Value::Boolean(true))
327 } else if left.is_none() || right.is_none() {
328 Ok(Value::Null)
329 } else {
330 Ok(Value::Boolean(false))
331 }
332 }
333 Eq | NotEq | Lt | LtEq | Gt | GtEq => {
334 if matches!(left, Value::Null) || matches!(right, Value::Null) {
335 return Ok(Value::Null);
336 }
337 let ordering = left.cmp_value(&right);
338 let value = match op {
339 Eq => ordering == std::cmp::Ordering::Equal,
340 NotEq => ordering != std::cmp::Ordering::Equal,
341 Lt => ordering == std::cmp::Ordering::Less,
342 LtEq => ordering != std::cmp::Ordering::Greater,
343 Gt => ordering == std::cmp::Ordering::Greater,
344 GtEq => ordering != std::cmp::Ordering::Less,
345 _ => unreachable!(),
346 };
347 Ok(Value::Boolean(value))
348 }
349 Add | Sub | Mul | Div | Mod => apply_arithmetic(op, left, right),
350 }
351}
352
353fn apply_arithmetic(op: &BinOp, left: Value, right: Value) -> Result<Value, DbError> {
354 use BinOp::*;
355 use Value::*;
356 match (left, right) {
357 (Null, _) | (_, Null) => Ok(Null),
358 (Integer(a), Integer(b)) => match op {
359 Add => a
360 .checked_add(b)
361 .map(Integer)
362 .ok_or_else(|| dberr(DbErrorKind::TypeMismatch, "integer overflow")),
363 Sub => a
364 .checked_sub(b)
365 .map(Integer)
366 .ok_or_else(|| dberr(DbErrorKind::TypeMismatch, "integer overflow")),
367 Mul => a
368 .checked_mul(b)
369 .map(Integer)
370 .ok_or_else(|| dberr(DbErrorKind::TypeMismatch, "integer overflow")),
371 Div => {
372 if b == 0 {
373 Ok(Null)
374 } else {
375 a.checked_div(b)
376 .map(Integer)
377 .ok_or_else(|| dberr(DbErrorKind::TypeMismatch, "integer overflow"))
378 }
379 }
380 Mod => {
381 if b == 0 {
382 Ok(Null)
383 } else {
384 a.checked_rem(b)
385 .map(Integer)
386 .ok_or_else(|| dberr(DbErrorKind::TypeMismatch, "integer overflow"))
387 }
388 }
389 _ => unreachable!(),
390 },
391 (Integer(a), Real(b)) => float_arithmetic(op, a as f64, b),
392 (Real(a), Integer(b)) => float_arithmetic(op, a, b as f64),
393 (Real(a), Real(b)) => float_arithmetic(op, a, b),
394 _ => Err(dberr(
395 DbErrorKind::TypeMismatch,
396 "arithmetic on TEXT/BOOLEAN values is not supported",
397 )),
398 }
399}
400
401fn float_arithmetic(op: &BinOp, left: f64, right: f64) -> Result<Value, DbError> {
402 use BinOp::*;
403 if matches!(op, Div | Mod) && right == 0.0 {
404 return Ok(Value::Null);
405 }
406 let value = match op {
407 Add => Value::Real(left + right),
408 Sub => Value::Real(left - right),
409 Mul => Value::Real(left * right),
410 Div => Value::Real(left / right),
411 Mod => Value::Real(left % right),
412 _ => unreachable!(),
413 };
414 if matches!(&value, Value::Real(value) if !value.is_finite()) {
415 return Err(dberr(DbErrorKind::TypeMismatch, "real overflow"));
416 }
417 Ok(value)
418}
419
420pub fn contains_aggregate(expr: &Expr) -> bool {
421 match expr {
422 Expr::Function { name, args, .. } => {
423 is_aggregate_name(name) || args.iter().any(contains_aggregate)
424 }
425 Expr::Binary { left, right, .. } => contains_aggregate(left) || contains_aggregate(right),
426 Expr::Unary { expr, .. } | Expr::IsNull { expr, .. } => contains_aggregate(expr),
427 Expr::Literal(_)
428 | Expr::Column(_)
429 | Expr::QualifiedWildcard(_)
430 | Expr::ColumnRef { .. } => false,
431 Expr::Alias { expr, .. } => contains_aggregate(expr),
432 }
433}
434
435fn is_aggregate_name(name: &str) -> bool {
436 matches!(
437 name.to_ascii_uppercase().as_str(),
438 "COUNT" | "SUM" | "AVG" | "MIN" | "MAX"
439 )
440}
441
442pub fn eval_group(schema: &[ColumnBinding], rows: &[Row], expr: &Expr) -> Result<Value, DbError> {
445 match expr {
446 Expr::Function {
447 name,
448 args,
449 distinct,
450 } if is_aggregate_name(name) => eval_aggregate(schema, rows, name, args, *distinct),
451 Expr::Binary { left, op, right } => apply_binary(
452 op,
453 eval_group(schema, rows, left)?,
454 eval_group(schema, rows, right)?,
455 ),
456 Expr::Unary { op, expr } => apply_unary(op, eval_group(schema, rows, expr)?),
457 Expr::IsNull { expr, negated } => {
458 let value = eval_group(schema, rows, expr)?;
459 Ok(Value::Boolean(matches!(value, Value::Null) != *negated))
460 }
461 Expr::Literal(value) => Ok(value.clone()),
462 Expr::Column(_) | Expr::ColumnRef { .. } => rows
463 .first()
464 .map(|row| eval_with_schema(schema, row, expr))
465 .unwrap_or_else(|| Ok(Value::Null)),
466 Expr::QualifiedWildcard(relation) => Err(dberr(
467 DbErrorKind::Syntax(format!("{relation}.* is only valid in a SELECT list")),
468 format!("{relation}.* is only valid in a SELECT list"),
469 )),
470 Expr::Alias { expr, .. } => eval_group(schema, rows, expr),
471 Expr::Function { .. } => rows
472 .first()
473 .map(|row| eval_with_schema(schema, row, expr))
474 .unwrap_or_else(|| Ok(Value::Null)),
475 }
476}
477
478fn eval_aggregate(
479 schema: &[ColumnBinding],
480 rows: &[Row],
481 name: &str,
482 args: &[Expr],
483 distinct: bool,
484) -> Result<Value, DbError> {
485 let upper = name.to_ascii_uppercase();
486 if upper == "COUNT" {
487 if args.is_empty() {
488 return Ok(Value::Integer(rows.len() as i64));
489 }
490 if args.len() != 1 {
491 return Err(dberr(
492 DbErrorKind::Syntax("COUNT expects at most one argument".into()),
493 "COUNT expects at most one argument",
494 ));
495 }
496 if matches!(args.first(), Some(Expr::Column(value)) if value == "*") {
497 return Ok(Value::Integer(rows.len() as i64));
498 }
499 let values = distinct_values(group_values(schema, rows, args.first().unwrap())?, distinct);
500 return Ok(Value::Integer(
501 values
502 .into_iter()
503 .filter(|value| !matches!(value, Value::Null))
504 .count() as i64,
505 ));
506 }
507 if args.len() != 1 {
508 return Err(dberr(
509 DbErrorKind::Syntax(format!("{upper} expects one argument")),
510 format!("{upper} expects one argument"),
511 ));
512 }
513 let argument = &args[0];
514 let values = distinct_values(group_values(schema, rows, argument)?, distinct)
515 .into_iter()
516 .filter(|value| !matches!(value, Value::Null))
517 .collect::<Vec<_>>();
518 if values.is_empty() {
519 return Ok(Value::Null);
520 }
521 match upper.as_str() {
522 "SUM" => sum_values(values),
523 "AVG" => {
524 let mut total = 0.0;
525 let mut count = 0usize;
526 for value in values {
527 total += match value {
528 Value::Integer(value) => value as f64,
529 Value::Real(value) => value,
530 _ => return Err(dberr(DbErrorKind::TypeMismatch, "AVG expects numbers")),
531 };
532 if !total.is_finite() {
533 return Err(dberr(DbErrorKind::TypeMismatch, "real overflow"));
534 }
535 count += 1;
536 }
537 let average = total / count as f64;
538 if average.is_finite() {
539 Ok(Value::Real(average))
540 } else {
541 Err(dberr(DbErrorKind::TypeMismatch, "real overflow"))
542 }
543 }
544 "MIN" | "MAX" => {
545 let mut result = values[0].clone();
546 for value in values.into_iter().skip(1) {
547 let ordering = value.cmp_value(&result);
548 if (upper == "MIN" && ordering == std::cmp::Ordering::Less)
549 || (upper == "MAX" && ordering == std::cmp::Ordering::Greater)
550 {
551 result = value;
552 }
553 }
554 Ok(result)
555 }
556 _ => unreachable!(),
557 }
558}
559
560fn group_values(
561 schema: &[ColumnBinding],
562 rows: &[Row],
563 expr: &Expr,
564) -> Result<Vec<Value>, DbError> {
565 rows.iter()
566 .map(|row| eval_with_schema(schema, row, expr))
567 .collect()
568}
569
570fn distinct_values(values: Vec<Value>, distinct: bool) -> Vec<Value> {
571 if !distinct {
572 return values;
573 }
574 let mut result: Vec<Value> = Vec::new();
575 for value in values {
576 if !result
577 .iter()
578 .any(|existing| existing.cmp_value(&value) == std::cmp::Ordering::Equal)
579 {
580 result.push(value);
581 }
582 }
583 result
584}
585
586fn sum_values(values: Vec<Value>) -> Result<Value, DbError> {
587 let has_real = values.iter().any(|value| matches!(value, Value::Real(_)));
588 if has_real {
589 let mut total = 0.0;
590 for value in values {
591 total += match value {
592 Value::Integer(value) => value as f64,
593 Value::Real(value) => value,
594 _ => return Err(dberr(DbErrorKind::TypeMismatch, "SUM expects numbers")),
595 };
596 if !total.is_finite() {
597 return Err(dberr(DbErrorKind::TypeMismatch, "real overflow"));
598 }
599 }
600 Ok(Value::Real(total))
601 } else {
602 let mut total = 0i64;
603 for value in values {
604 let Value::Integer(value) = value else {
605 return Err(dberr(DbErrorKind::TypeMismatch, "SUM expects numbers"));
606 };
607 total = total
608 .checked_add(value)
609 .ok_or_else(|| dberr(DbErrorKind::TypeMismatch, "integer overflow"))?;
610 }
611 Ok(Value::Integer(total))
612 }
613}
614
615pub fn where_matches(table: &Table, row: &Row, expr: &Expr) -> Result<bool, DbError> {
617 Ok(eval(table, row, expr)?.is_truthy() == Some(true))
618}
619
620pub fn where_matches_with_schema(
621 schema: &[ColumnBinding],
622 row: &Row,
623 expr: &Expr,
624) -> Result<bool, DbError> {
625 Ok(eval_with_schema(schema, row, expr)?.is_truthy() == Some(true))
626}
627
628#[cfg(test)]
629mod tests {
630 use super::*;
631 use crate::db::Column;
632 use crate::sql::parser::parse;
633 use crate::types::ColumnType as CT;
634
635 fn table() -> Table {
636 Table::new(
637 "t",
638 vec![
639 Column {
640 name: "id".into(),
641 ty: CT::Integer,
642 not_null: true,
643 unique: false,
644 primary_key: true,
645 },
646 Column {
647 name: "score".into(),
648 ty: CT::Real,
649 not_null: false,
650 unique: false,
651 primary_key: false,
652 },
653 Column {
654 name: "name".into(),
655 ty: CT::Text,
656 not_null: true,
657 unique: false,
658 primary_key: false,
659 },
660 ],
661 )
662 .unwrap()
663 }
664
665 fn expr(sql: &str) -> Expr {
666 match parse(sql).unwrap().into_iter().next().unwrap() {
667 crate::sql::ast::Statement::Select {
668 where_clause: Some(expr),
669 ..
670 } => expr,
671 other => panic!("need WHERE: {other:?}"),
672 }
673 }
674
675 fn row() -> Row {
676 vec![Value::Integer(1), Value::Null, Value::Text("a".into())]
677 }
678
679 #[test]
680 fn null_logic_and_comparisons() {
681 let table = table();
682 assert_eq!(
683 eval(&table, &row(), &expr("SELECT * FROM t WHERE score > 5")).unwrap(),
684 Value::Null
685 );
686 assert_eq!(
687 eval(
688 &table,
689 &row(),
690 &expr("SELECT * FROM t WHERE id = 2 AND score > 1")
691 )
692 .unwrap(),
693 Value::Boolean(false)
694 );
695 assert_eq!(
696 eval(
697 &table,
698 &row(),
699 &expr("SELECT * FROM t WHERE id = 1 OR score > 1")
700 )
701 .unwrap(),
702 Value::Boolean(true)
703 );
704 }
705
706 #[test]
707 fn aggregate_group() {
708 let table = table();
709 let schema = schema_for_table(&table);
710 let rows = vec![
711 vec![Value::Integer(1), Value::Real(2.0), Value::Text("a".into())],
712 vec![Value::Integer(2), Value::Real(3.0), Value::Text("b".into())],
713 ];
714 let parsed = parse("SELECT SUM(score) FROM t").unwrap();
715 let crate::sql::ast::Statement::Select {
716 columns: crate::sql::ast::SelectItems::List(items),
717 ..
718 } = &parsed[0]
719 else {
720 panic!()
721 };
722 assert_eq!(
723 eval_group(&schema, &rows, &items[0]).unwrap(),
724 Value::Real(5.0)
725 );
726 }
727}