1use std::cmp::Ordering;
2use std::time::SystemTime;
3
4use teaql_core::{
5 Aggregate, AggregateFunction, BinaryOp, Expr, OrderBy, Record, SelectQuery, SortDirection,
6 Value,
7};
8use teaql_data_service::{DataServiceOperation, ExecutionMetadata, QueryResult};
9
10pub struct InMemoryQueryEngine;
13
14impl InMemoryQueryEngine {
15 pub fn execute(query: &SelectQuery, mut rows: Vec<Record>) -> QueryResult {
19 let started_at = SystemTime::now();
20
21 if let Some(filter) = &query.filter {
23 Self::filter(&mut rows, filter);
24 }
25
26 if !query.aggregates.is_empty() {
28 let mut result = Self::aggregate(query, rows);
29 result.metadata.started_at = started_at;
30 result.metadata.ended_at = SystemTime::now();
31 return result;
32 }
33
34 if !query.order_by.is_empty() {
36 Self::sort(&mut rows, &query.order_by);
37 }
38
39 if let Some(slice) = &query.slice {
41 rows = Self::paginate(rows, slice);
42 }
43
44 if !query.projection.is_empty() {
46 rows = Self::project(rows, &query.projection);
47 }
48
49 let count = rows.len();
50 QueryResult {
51 rows,
52 metadata: ExecutionMetadata {
53 debug_query: None,
54 backend: "memory".to_owned(),
55 operation: DataServiceOperation::Query,
56 started_at,
57 ended_at: SystemTime::now(),
58 affected_rows: None,
59 result_count: Some(count),
60 trace_chain: Vec::new(),
61 comment: None,
62 backend_request_id: None,
63 },
64 }
65 }
66
67 fn filter(rows: &mut Vec<Record>, expr: &Expr) {
69 rows.retain(|row| ExprEvaluator::eval(expr, row));
70 }
71
72 fn sort(rows: &mut Vec<Record>, order_by: &[OrderBy]) {
74 rows.sort_by(|a, b| {
75 for ob in order_by {
76 let va = a.get(&ob.field).unwrap_or(&Value::Null);
77 let vb = b.get(&ob.field).unwrap_or(&Value::Null);
78 let ord = compare_values(va, vb);
79 let ord = match ob.direction {
80 SortDirection::Asc => ord,
81 SortDirection::Desc => ord.reverse(),
82 };
83 if ord != Ordering::Equal {
84 return ord;
85 }
86 }
87 Ordering::Equal
88 });
89 }
90
91 fn paginate(rows: Vec<Record>, slice: &teaql_core::Slice) -> Vec<Record> {
93 let offset = slice.offset as usize;
94 let iter = rows.into_iter().skip(offset);
95 match slice.limit {
96 Some(limit) => iter.take(limit as usize).collect(),
97 None => iter.collect(),
98 }
99 }
100
101 fn project(rows: Vec<Record>, projection: &[String]) -> Vec<Record> {
103 rows.into_iter()
104 .map(|row| {
105 row.into_iter()
106 .filter(|(key, _)| projection.contains(key))
107 .collect()
108 })
109 .collect()
110 }
111
112 fn aggregate(query: &SelectQuery, rows: Vec<Record>) -> QueryResult {
115 let started_at = SystemTime::now();
116
117 if !query.group_by.is_empty() {
119 return Self::aggregate_grouped(query, rows, started_at);
120 }
121
122 let mut result_row = Record::new();
124 for agg in &query.aggregates {
125 let value = compute_aggregate(agg, &rows);
126 result_row.insert(agg.alias.clone(), value);
127 }
128
129 let result_rows = vec![result_row];
130 let count = result_rows.len();
131 QueryResult {
132 rows: result_rows,
133 metadata: ExecutionMetadata {
134 debug_query: None,
135 backend: "memory".to_owned(),
136 operation: DataServiceOperation::Query,
137 started_at,
138 ended_at: SystemTime::now(),
139 affected_rows: None,
140 result_count: Some(count),
141 trace_chain: Vec::new(),
142 comment: None,
143 backend_request_id: None,
144 },
145 }
146 }
147
148 fn aggregate_grouped(
150 query: &SelectQuery,
151 rows: Vec<Record>,
152 started_at: SystemTime,
153 ) -> QueryResult {
154 let mut groups: Vec<(Vec<Value>, Vec<Record>)> = Vec::new();
156
157 for row in rows {
158 let key: Vec<Value> = query
159 .group_by
160 .iter()
161 .map(|gb| row.get(gb).cloned().unwrap_or(Value::Null))
162 .collect();
163
164 if let Some((_k, group)) = groups.iter_mut().find(|(k, _)| k == &key) {
165 group.push(row);
166 } else {
167 groups.push((key, vec![row]));
168 }
169 }
170
171 let mut result_rows = Vec::with_capacity(groups.len());
172 for (key_values, group_rows) in &groups {
173 let mut result_row = Record::new();
174
175 for (i, gb) in query.group_by.iter().enumerate() {
177 result_row.insert(gb.clone(), key_values[i].clone());
178 }
179
180 for agg in &query.aggregates {
182 let value = compute_aggregate(agg, group_rows);
183 result_row.insert(agg.alias.clone(), value);
184 }
185
186 result_rows.push(result_row);
187 }
188
189 let count = result_rows.len();
190 QueryResult {
191 rows: result_rows,
192 metadata: ExecutionMetadata {
193 debug_query: None,
194 backend: "memory".to_owned(),
195 operation: DataServiceOperation::Query,
196 started_at,
197 ended_at: SystemTime::now(),
198 affected_rows: None,
199 result_count: Some(count),
200 trace_chain: Vec::new(),
201 comment: None,
202 backend_request_id: None,
203 },
204 }
205 }
206}
207
208pub struct ExprEvaluator;
210
211impl ExprEvaluator {
212 pub fn eval(expr: &Expr, row: &Record) -> bool {
214 match expr {
215 Expr::Binary { left, op, right } => {
216 let lv = Self::resolve(left, row);
217 let rv = Self::resolve(right, row);
218 Self::compare_op(&lv, op, &rv)
219 }
220 Expr::And(parts) => parts.iter().all(|p| Self::eval(p, row)),
221 Expr::Or(parts) => parts.iter().any(|p| Self::eval(p, row)),
222 Expr::Not(inner) => !Self::eval(inner, row),
223 Expr::IsNull(inner) => Self::resolve(inner, row) == Value::Null,
224 Expr::IsNotNull(inner) => Self::resolve(inner, row) != Value::Null,
225 Expr::Between {
226 expr: inner,
227 lower,
228 upper,
229 } => {
230 let v = Self::resolve(inner, row);
231 let lo = Self::resolve(lower, row);
232 let hi = Self::resolve(upper, row);
233 compare_values(&v, &lo) != Ordering::Less
234 && compare_values(&v, &hi) != Ordering::Greater
235 }
236 Expr::SubQuery { .. } => false,
238 Expr::Function { .. } => false,
240 Expr::Column(_) | Expr::Value(_) => {
242 matches!(Self::resolve(expr, row), Value::Bool(true))
243 }
244 }
245 }
246
247 pub fn resolve(expr: &Expr, row: &Record) -> Value {
249 match expr {
250 Expr::Column(name) => row.get(name).cloned().unwrap_or(Value::Null),
251 Expr::Value(v) => v.clone(),
252 Expr::Binary { left, op, right } => {
253 let lv = Self::resolve(left, row);
254 let rv = Self::resolve(right, row);
255 Value::Bool(Self::compare_op(&lv, op, &rv))
256 }
257 Expr::And(parts) => Value::Bool(parts.iter().all(|p| Self::eval(p, row))),
258 Expr::Or(parts) => Value::Bool(parts.iter().any(|p| Self::eval(p, row))),
259 Expr::Not(inner) => Value::Bool(!Self::eval(inner, row)),
260 Expr::IsNull(inner) => Value::Bool(Self::resolve(inner, row) == Value::Null),
261 Expr::IsNotNull(inner) => Value::Bool(Self::resolve(inner, row) != Value::Null),
262 Expr::Between {
263 expr: inner,
264 lower,
265 upper,
266 } => {
267 let v = Self::resolve(inner, row);
268 let lo = Self::resolve(lower, row);
269 let hi = Self::resolve(upper, row);
270 Value::Bool(
271 compare_values(&v, &lo) != Ordering::Less
272 && compare_values(&v, &hi) != Ordering::Greater,
273 )
274 }
275 Expr::SubQuery { .. } => Value::Null,
276 Expr::Function { .. } => Value::Null,
277 }
278 }
279
280 fn compare_op(left: &Value, op: &BinaryOp, right: &Value) -> bool {
282 match op {
283 BinaryOp::Eq => left == right,
284 BinaryOp::Ne => left != right,
285 BinaryOp::Gt => compare_values(left, right) == Ordering::Greater,
286 BinaryOp::Gte => matches!(
287 compare_values(left, right),
288 Ordering::Greater | Ordering::Equal
289 ),
290 BinaryOp::Lt => compare_values(left, right) == Ordering::Less,
291 BinaryOp::Lte => matches!(
292 compare_values(left, right),
293 Ordering::Less | Ordering::Equal
294 ),
295 BinaryOp::Like => match (left, right) {
296 (Value::Text(text), Value::Text(pattern)) => Self::like_match(text, pattern),
297 _ => false,
298 },
299 BinaryOp::NotLike => match (left, right) {
300 (Value::Text(text), Value::Text(pattern)) => !Self::like_match(text, pattern),
301 _ => true,
302 },
303 BinaryOp::In | BinaryOp::InLarge => match right {
304 Value::List(items) => items.contains(left),
305 _ => left == right,
306 },
307 BinaryOp::NotIn | BinaryOp::NotInLarge => match right {
308 Value::List(items) => !items.contains(left),
309 _ => left != right,
310 },
311 }
312 }
313
314 fn like_match(text: &str, pattern: &str) -> bool {
319 let text_chars: Vec<char> = text.chars().collect();
320 let pattern_chars: Vec<char> = pattern.chars().collect();
321 like_match_recursive(&text_chars, 0, &pattern_chars, 0)
322 }
323}
324
325fn like_match_recursive(text: &[char], ti: usize, pattern: &[char], pi: usize) -> bool {
328 let mut ti = ti;
329 let mut pi = pi;
330
331 loop {
332 if pi == pattern.len() {
333 return ti == text.len();
334 }
335
336 match pattern[pi] {
337 '%' => {
338 while pi < pattern.len() && pattern[pi] == '%' {
340 pi += 1;
341 }
342 if pi == pattern.len() {
344 return true;
345 }
346 for start in ti..=text.len() {
348 if like_match_recursive(text, start, pattern, pi) {
349 return true;
350 }
351 }
352 return false;
353 }
354 '_' => {
355 if ti >= text.len() {
356 return false;
357 }
358 ti += 1;
359 pi += 1;
360 }
361 ch => {
362 if ti >= text.len() || text[ti] != ch {
363 return false;
364 }
365 ti += 1;
366 pi += 1;
367 }
368 }
369 }
370}
371
372fn compare_values(a: &Value, b: &Value) -> Ordering {
374 match (a, b) {
375 (Value::Null, Value::Null) => Ordering::Equal,
376 (Value::Null, _) => Ordering::Less,
377 (_, Value::Null) => Ordering::Greater,
378 (Value::Bool(a), Value::Bool(b)) => a.cmp(b),
379 (Value::I64(a), Value::I64(b)) => a.cmp(b),
380 (Value::U64(a), Value::U64(b)) => a.cmp(b),
381 (Value::I64(a), Value::U64(b)) => {
382 if *a < 0 {
383 Ordering::Less
384 } else {
385 (*a as u64).cmp(b)
386 }
387 }
388 (Value::U64(a), Value::I64(b)) => {
389 if *b < 0 {
390 Ordering::Greater
391 } else {
392 a.cmp(&(*b as u64))
393 }
394 }
395 (Value::F64(a), Value::F64(b)) => a.partial_cmp(b).unwrap_or(Ordering::Equal),
396 (Value::Decimal(a), Value::Decimal(b)) => a.cmp(b),
397 (Value::Text(a), Value::Text(b)) => a.cmp(b),
398 (Value::Date(a), Value::Date(b)) => a.cmp(b),
399 (Value::Timestamp(a), Value::Timestamp(b)) => a.cmp(b),
400 _ => {
402 if let (Some(fa), Some(fb)) = (value_to_f64(a), value_to_f64(b)) {
403 fa.partial_cmp(&fb).unwrap_or(Ordering::Equal)
404 } else {
405 Ordering::Equal
406 }
407 }
408 }
409}
410
411fn value_to_f64(v: &Value) -> Option<f64> {
413 v.try_f64()
414}
415
416fn compute_aggregate(agg: &Aggregate, rows: &[Record]) -> Value {
418 match agg.function {
419 AggregateFunction::Count => {
420 if agg.field == "*" {
421 Value::I64(rows.len() as i64)
422 } else {
423 let count = rows
424 .iter()
425 .filter(|r| {
426 r.get(&agg.field)
427 .map(|v| v != &Value::Null)
428 .unwrap_or(false)
429 })
430 .count();
431 Value::I64(count as i64)
432 }
433 }
434 AggregateFunction::Sum => {
435 let mut sum: f64 = 0.0;
436 let mut found = false;
437 for row in rows {
438 if let Some(v) = row.get(&agg.field) {
439 if let Some(f) = v.try_f64() {
440 sum += f;
441 found = true;
442 }
443 }
444 }
445 if found {
446 Value::F64(sum)
447 } else {
448 Value::Null
449 }
450 }
451 AggregateFunction::Avg => {
452 let mut sum: f64 = 0.0;
453 let mut count: u64 = 0;
454 for row in rows {
455 if let Some(v) = row.get(&agg.field) {
456 if let Some(f) = v.try_f64() {
457 sum += f;
458 count += 1;
459 }
460 }
461 }
462 if count > 0 {
463 Value::F64(sum / count as f64)
464 } else {
465 Value::Null
466 }
467 }
468 AggregateFunction::Max => {
469 let mut max: Option<&Value> = None;
470 for row in rows {
471 if let Some(v) = row.get(&agg.field) {
472 if v == &Value::Null {
473 continue;
474 }
475 max = Some(match max {
476 Some(current) if compare_values(v, current) == Ordering::Greater => v,
477 Some(current) => current,
478 None => v,
479 });
480 }
481 }
482 max.cloned().unwrap_or(Value::Null)
483 }
484 AggregateFunction::Min => {
485 let mut min: Option<&Value> = None;
486 for row in rows {
487 if let Some(v) = row.get(&agg.field) {
488 if v == &Value::Null {
489 continue;
490 }
491 min = Some(match min {
492 Some(current) if compare_values(v, current) == Ordering::Less => v,
493 Some(current) => current,
494 None => v,
495 });
496 }
497 }
498 min.cloned().unwrap_or(Value::Null)
499 }
500 AggregateFunction::Stddev
502 | AggregateFunction::StddevPop
503 | AggregateFunction::VarSamp
504 | AggregateFunction::VarPop
505 | AggregateFunction::BitAnd
506 | AggregateFunction::BitOr
507 | AggregateFunction::BitXor => Value::Null,
508 }
509}
510
511#[cfg(test)]
512mod tests {
513 use super::*;
514 use teaql_core::{Aggregate, AggregateFunction, Record, SelectQuery, Value};
515
516 fn make_row(pairs: Vec<(&str, Value)>) -> Record {
517 pairs
518 .into_iter()
519 .map(|(k, v)| (k.to_owned(), v))
520 .collect()
521 }
522
523 fn sample_rows() -> Vec<Record> {
524 vec![
525 make_row(vec![
526 ("id", Value::U64(1)),
527 ("name", Value::Text("Alice".to_owned())),
528 ("age", Value::I64(30)),
529 ]),
530 make_row(vec![
531 ("id", Value::U64(2)),
532 ("name", Value::Text("Bob".to_owned())),
533 ("age", Value::I64(25)),
534 ]),
535 make_row(vec![
536 ("id", Value::U64(3)),
537 ("name", Value::Text("Charlie".to_owned())),
538 ("age", Value::I64(35)),
539 ]),
540 ]
541 }
542
543 #[test]
544 fn test_execute_no_filter() {
545 let query = SelectQuery::new("User");
546 let result = InMemoryQueryEngine::execute(&query, sample_rows());
547 assert_eq!(result.rows.len(), 3);
548 assert_eq!(result.metadata.backend, "memory");
549 }
550
551 #[test]
552 fn test_execute_with_eq_filter() {
553 let query = SelectQuery::new("User").filter(Expr::eq("name", "Bob"));
554 let result = InMemoryQueryEngine::execute(&query, sample_rows());
555 assert_eq!(result.rows.len(), 1);
556 assert_eq!(
557 result.rows[0].get("name"),
558 Some(&Value::Text("Bob".to_owned()))
559 );
560 }
561
562 #[test]
563 fn test_execute_with_gt_filter() {
564 let query = SelectQuery::new("User").filter(Expr::gt("age", 28_i64));
565 let result = InMemoryQueryEngine::execute(&query, sample_rows());
566 assert_eq!(result.rows.len(), 2); }
568
569 #[test]
570 fn test_sort_ascending() {
571 let query = SelectQuery::new("User").order_by(teaql_core::OrderBy::asc("age"));
572 let result = InMemoryQueryEngine::execute(&query, sample_rows());
573 let ages: Vec<_> = result
574 .rows
575 .iter()
576 .map(|r| r.get("age").unwrap().clone())
577 .collect();
578 assert_eq!(ages, vec![Value::I64(25), Value::I64(30), Value::I64(35)]);
579 }
580
581 #[test]
582 fn test_sort_descending() {
583 let query = SelectQuery::new("User").order_by(teaql_core::OrderBy::desc("age"));
584 let result = InMemoryQueryEngine::execute(&query, sample_rows());
585 let ages: Vec<_> = result
586 .rows
587 .iter()
588 .map(|r| r.get("age").unwrap().clone())
589 .collect();
590 assert_eq!(ages, vec![Value::I64(35), Value::I64(30), Value::I64(25)]);
591 }
592
593 #[test]
594 fn test_paginate() {
595 let query = SelectQuery::new("User").page(1, 1);
596 let result = InMemoryQueryEngine::execute(&query, sample_rows());
597 assert_eq!(result.rows.len(), 1);
598 assert_eq!(
599 result.rows[0].get("name"),
600 Some(&Value::Text("Bob".to_owned()))
601 );
602 }
603
604 #[test]
605 fn test_projection() {
606 let query = SelectQuery::new("User").projects(["name"]);
607 let result = InMemoryQueryEngine::execute(&query, sample_rows());
608 for row in &result.rows {
609 assert!(row.contains_key("name"));
610 assert!(!row.contains_key("id"));
611 assert!(!row.contains_key("age"));
612 }
613 }
614
615 #[test]
616 fn test_count_aggregate() {
617 let query = SelectQuery::new("User").aggregate(Aggregate::count("total"));
618 let result = InMemoryQueryEngine::execute(&query, sample_rows());
619 assert_eq!(result.rows.len(), 1);
620 assert_eq!(result.rows[0].get("total"), Some(&Value::I64(3)));
621 }
622
623 #[test]
624 fn test_sum_aggregate() {
625 let query =
626 SelectQuery::new("User").aggregate(Aggregate::sum("age", "age_sum"));
627 let result = InMemoryQueryEngine::execute(&query, sample_rows());
628 assert_eq!(result.rows[0].get("age_sum"), Some(&Value::F64(90.0)));
629 }
630
631 #[test]
632 fn test_avg_aggregate() {
633 let query =
634 SelectQuery::new("User").aggregate(Aggregate::avg("age", "age_avg"));
635 let result = InMemoryQueryEngine::execute(&query, sample_rows());
636 assert_eq!(result.rows[0].get("age_avg"), Some(&Value::F64(30.0)));
637 }
638
639 #[test]
640 fn test_max_aggregate() {
641 let query =
642 SelectQuery::new("User").aggregate(Aggregate::max("age", "age_max"));
643 let result = InMemoryQueryEngine::execute(&query, sample_rows());
644 assert_eq!(result.rows[0].get("age_max"), Some(&Value::I64(35)));
645 }
646
647 #[test]
648 fn test_min_aggregate() {
649 let query =
650 SelectQuery::new("User").aggregate(Aggregate::min("age", "age_min"));
651 let result = InMemoryQueryEngine::execute(&query, sample_rows());
652 assert_eq!(result.rows[0].get("age_min"), Some(&Value::I64(25)));
653 }
654
655 #[test]
656 fn test_like_match_percent() {
657 assert!(ExprEvaluator::like_match("hello world", "%world"));
658 assert!(ExprEvaluator::like_match("hello world", "hello%"));
659 assert!(ExprEvaluator::like_match("hello world", "%lo wo%"));
660 assert!(ExprEvaluator::like_match("hello world", "%"));
661 assert!(!ExprEvaluator::like_match("hello world", "%xyz%"));
662 }
663
664 #[test]
665 fn test_like_match_underscore() {
666 assert!(ExprEvaluator::like_match("abc", "a_c"));
667 assert!(!ExprEvaluator::like_match("abbc", "a_c"));
668 assert!(ExprEvaluator::like_match("abc", "___"));
669 assert!(!ExprEvaluator::like_match("ab", "___"));
670 }
671
672 #[test]
673 fn test_like_match_combined() {
674 assert!(ExprEvaluator::like_match("foobar", "f%r"));
675 assert!(ExprEvaluator::like_match("foobar", "f__b%"));
676 assert!(!ExprEvaluator::like_match("foobar", "f__x%"));
677 }
678
679 #[test]
680 fn test_and_or_not() {
681 let row = make_row(vec![
682 ("a", Value::I64(10)),
683 ("b", Value::I64(20)),
684 ]);
685 let expr_and = Expr::and([Expr::eq("a", 10_i64), Expr::eq("b", 20_i64)]);
686 assert!(ExprEvaluator::eval(&expr_and, &row));
687
688 let expr_or = Expr::or([Expr::eq("a", 99_i64), Expr::eq("b", 20_i64)]);
689 assert!(ExprEvaluator::eval(&expr_or, &row));
690
691 let expr_not = Expr::negate(Expr::eq("a", 99_i64));
692 assert!(ExprEvaluator::eval(&expr_not, &row));
693 }
694
695 #[test]
696 fn test_is_null_is_not_null() {
697 let row = make_row(vec![("x", Value::Null), ("y", Value::I64(1))]);
698 assert!(ExprEvaluator::eval(&Expr::is_null("x"), &row));
699 assert!(!ExprEvaluator::eval(&Expr::is_not_null("x"), &row));
700 assert!(ExprEvaluator::eval(&Expr::is_not_null("y"), &row));
701 }
702
703 #[test]
704 fn test_between() {
705 let row = make_row(vec![("age", Value::I64(30))]);
706 assert!(ExprEvaluator::eval(
707 &Expr::between("age", Value::I64(25), Value::I64(35)),
708 &row
709 ));
710 assert!(!ExprEvaluator::eval(
711 &Expr::between("age", Value::I64(31), Value::I64(35)),
712 &row
713 ));
714 }
715
716 #[test]
717 fn test_in_list() {
718 let row = make_row(vec![("status", Value::Text("active".to_owned()))]);
719 let expr = Expr::in_list(
720 "status",
721 vec![
722 Value::Text("active".to_owned()),
723 Value::Text("pending".to_owned()),
724 ],
725 );
726 assert!(ExprEvaluator::eval(&expr, &row));
727
728 let expr_miss = Expr::in_list(
729 "status",
730 vec![Value::Text("closed".to_owned())],
731 );
732 assert!(!ExprEvaluator::eval(&expr_miss, &row));
733 }
734}