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 match groups.iter_mut().find(|(k, _)| k == &key) {
165 Some((_k, group)) => group.push(row),
166 None => groups.push((key, vec![row])),
167 }
168 }
169
170 let mut result_rows = Vec::with_capacity(groups.len());
171 for (key_values, group_rows) in &groups {
172 let mut result_row = Record::new();
173
174 for (i, gb) in query.group_by.iter().enumerate() {
176 result_row.insert(gb.clone(), key_values[i].clone());
177 }
178
179 for agg in &query.aggregates {
181 let value = compute_aggregate(agg, group_rows);
182 result_row.insert(agg.alias.clone(), value);
183 }
184
185 result_rows.push(result_row);
186 }
187
188 let count = result_rows.len();
189 QueryResult {
190 rows: result_rows,
191 metadata: ExecutionMetadata {
192 debug_query: None,
193 backend: "memory".to_owned(),
194 operation: DataServiceOperation::Query,
195 started_at,
196 ended_at: SystemTime::now(),
197 affected_rows: None,
198 result_count: Some(count),
199 trace_chain: Vec::new(),
200 comment: None,
201 backend_request_id: None,
202 },
203 }
204 }
205}
206
207pub struct ExprEvaluator;
209
210impl ExprEvaluator {
211 pub fn eval(expr: &Expr, row: &Record) -> bool {
213 match expr {
214 Expr::Binary { left, op, right } => {
215 let lv = Self::resolve(left, row);
216 let rv = Self::resolve(right, row);
217 Self::compare_op(&lv, op, &rv)
218 }
219 Expr::And(parts) => parts.iter().all(|p| Self::eval(p, row)),
220 Expr::Or(parts) => parts.iter().any(|p| Self::eval(p, row)),
221 Expr::Not(inner) => !Self::eval(inner, row),
222 Expr::IsNull(inner) => Self::resolve(inner, row) == Value::Null,
223 Expr::IsNotNull(inner) => Self::resolve(inner, row) != Value::Null,
224 Expr::Between {
225 expr: inner,
226 lower,
227 upper,
228 } => {
229 let v = Self::resolve(inner, row);
230 let lo = Self::resolve(lower, row);
231 let hi = Self::resolve(upper, row);
232 compare_values(&v, &lo) != Ordering::Less
233 && compare_values(&v, &hi) != Ordering::Greater
234 }
235 Expr::SubQuery { .. } => false,
237 Expr::Function { .. } => false,
239 Expr::Column(_) | Expr::Value(_) => {
241 matches!(Self::resolve(expr, row), Value::Bool(true))
242 }
243 }
244 }
245
246 pub fn resolve(expr: &Expr, row: &Record) -> Value {
248 match expr {
249 Expr::Column(name) => row.get(name).cloned().unwrap_or(Value::Null),
250 Expr::Value(v) => v.clone(),
251 Expr::Binary { left, op, right } => {
252 let lv = Self::resolve(left, row);
253 let rv = Self::resolve(right, row);
254 Value::Bool(Self::compare_op(&lv, op, &rv))
255 }
256 Expr::And(parts) => Value::Bool(parts.iter().all(|p| Self::eval(p, row))),
257 Expr::Or(parts) => Value::Bool(parts.iter().any(|p| Self::eval(p, row))),
258 Expr::Not(inner) => Value::Bool(!Self::eval(inner, row)),
259 Expr::IsNull(inner) => Value::Bool(Self::resolve(inner, row) == Value::Null),
260 Expr::IsNotNull(inner) => Value::Bool(Self::resolve(inner, row) != Value::Null),
261 Expr::Between {
262 expr: inner,
263 lower,
264 upper,
265 } => {
266 let v = Self::resolve(inner, row);
267 let lo = Self::resolve(lower, row);
268 let hi = Self::resolve(upper, row);
269 Value::Bool(
270 compare_values(&v, &lo) != Ordering::Less
271 && compare_values(&v, &hi) != Ordering::Greater,
272 )
273 }
274 Expr::SubQuery { .. } => Value::Null,
275 Expr::Function { .. } => Value::Null,
276 }
277 }
278
279 fn compare_op(left: &Value, op: &BinaryOp, right: &Value) -> bool {
281 match op {
282 BinaryOp::Eq => left == right,
283 BinaryOp::Ne => left != right,
284 BinaryOp::Gt => compare_values(left, right) == Ordering::Greater,
285 BinaryOp::Gte => matches!(
286 compare_values(left, right),
287 Ordering::Greater | Ordering::Equal
288 ),
289 BinaryOp::Lt => compare_values(left, right) == Ordering::Less,
290 BinaryOp::Lte => matches!(
291 compare_values(left, right),
292 Ordering::Less | Ordering::Equal
293 ),
294 BinaryOp::Like => match (left, right) {
295 (Value::Text(text), Value::Text(pattern)) => Self::like_match(text, pattern),
296 _ => false,
297 },
298 BinaryOp::NotLike => match (left, right) {
299 (Value::Text(text), Value::Text(pattern)) => !Self::like_match(text, pattern),
300 _ => true,
301 },
302 BinaryOp::In | BinaryOp::InLarge => match right {
303 Value::List(items) => items.contains(left),
304 _ => left == right,
305 },
306 BinaryOp::NotIn | BinaryOp::NotInLarge => match right {
307 Value::List(items) => !items.contains(left),
308 _ => left != right,
309 },
310 }
311 }
312
313 fn like_match(text: &str, pattern: &str) -> bool {
318 let text_chars: Vec<char> = text.chars().collect();
319 let pattern_chars: Vec<char> = pattern.chars().collect();
320 like_match_recursive(&text_chars, 0, &pattern_chars, 0)
321 }
322}
323
324fn like_match_recursive(text: &[char], ti: usize, pattern: &[char], pi: usize) -> bool {
327 let mut ti = ti;
328 let mut pi = pi;
329
330 loop {
331 if pi == pattern.len() {
332 return ti == text.len();
333 }
334
335 match pattern[pi] {
336 '%' => {
337 while pi < pattern.len() && pattern[pi] == '%' {
339 pi += 1;
340 }
341 if pi == pattern.len() {
343 return true;
344 }
345 for start in ti..=text.len() {
347 if like_match_recursive(text, start, pattern, pi) {
348 return true;
349 }
350 }
351 return false;
352 }
353 '_' => {
354 if ti >= text.len() {
355 return false;
356 }
357 ti += 1;
358 pi += 1;
359 }
360 ch => {
361 if ti >= text.len() || text[ti] != ch {
362 return false;
363 }
364 ti += 1;
365 pi += 1;
366 }
367 }
368 }
369}
370
371fn compare_i64_u64(a: i64, b: u64) -> Ordering {
373 match a < 0 {
374 true => Ordering::Less,
375 false => (a as u64).cmp(&b),
376 }
377}
378
379fn compare_values(a: &Value, b: &Value) -> Ordering {
381 match (a, b) {
382 (Value::Null, Value::Null) => Ordering::Equal,
383 (Value::Null, _) => Ordering::Less,
384 (_, Value::Null) => Ordering::Greater,
385 (Value::Bool(a), Value::Bool(b)) => a.cmp(b),
386 (Value::I64(a), Value::I64(b)) => a.cmp(b),
387 (Value::U64(a), Value::U64(b)) => a.cmp(b),
388 (Value::I64(a), Value::U64(b)) => compare_i64_u64(*a, *b),
389 (Value::U64(a), Value::I64(b)) => compare_i64_u64(*b, *a).reverse(),
390 (Value::F64(a), Value::F64(b)) => a.partial_cmp(b).unwrap_or(Ordering::Equal),
391 (Value::Decimal(a), Value::Decimal(b)) => a.cmp(b),
392 (Value::Text(a), Value::Text(b)) => a.cmp(b),
393 (Value::Date(a), Value::Date(b)) => a.cmp(b),
394 (Value::Timestamp(a), Value::Timestamp(b)) => a.cmp(b),
395 _ => value_to_f64(a)
397 .zip(value_to_f64(b))
398 .and_then(|(fa, fb)| fa.partial_cmp(&fb))
399 .unwrap_or(Ordering::Equal),
400 }
401}
402
403fn value_to_f64(v: &Value) -> Option<f64> {
405 v.try_f64()
406}
407
408fn count_rows(rows: &[Record], field: &str) -> Value {
410 let count = match field {
411 "*" => rows.len(),
412 _ => rows
413 .iter()
414 .filter(|r| {
415 r.get(field)
416 .map(|v| v != &Value::Null)
417 .unwrap_or(false)
418 })
419 .count(),
420 };
421 Value::I64(count as i64)
422}
423
424fn compute_aggregate(agg: &Aggregate, rows: &[Record]) -> Value {
426 match agg.function {
427 AggregateFunction::Count => count_rows(rows, &agg.field),
428 AggregateFunction::Sum => {
429 let mut sum: f64 = 0.0;
430 let mut found = false;
431 for row in rows {
432 if let Some(v) = row.get(&agg.field) {
433 if let Some(f) = v.try_f64() {
434 sum += f;
435 found = true;
436 }
437 }
438 }
439 found.then(|| Value::F64(sum)).unwrap_or(Value::Null)
440 }
441 AggregateFunction::Avg => {
442 let mut sum: f64 = 0.0;
443 let mut count: u64 = 0;
444 for row in rows {
445 if let Some(v) = row.get(&agg.field) {
446 if let Some(f) = v.try_f64() {
447 sum += f;
448 count += 1;
449 }
450 }
451 }
452 (count > 0)
453 .then(|| Value::F64(sum / count as f64))
454 .unwrap_or(Value::Null)
455 }
456 AggregateFunction::Max => {
457 let mut max: Option<&Value> = None;
458 for row in rows {
459 if let Some(v) = row.get(&agg.field) {
460 if v == &Value::Null {
461 continue;
462 }
463 max = Some(match max {
464 Some(current) if compare_values(v, current) == Ordering::Greater => v,
465 Some(current) => current,
466 None => v,
467 });
468 }
469 }
470 max.cloned().unwrap_or(Value::Null)
471 }
472 AggregateFunction::Min => {
473 let mut min: Option<&Value> = None;
474 for row in rows {
475 if let Some(v) = row.get(&agg.field) {
476 if v == &Value::Null {
477 continue;
478 }
479 min = Some(match min {
480 Some(current) if compare_values(v, current) == Ordering::Less => v,
481 Some(current) => current,
482 None => v,
483 });
484 }
485 }
486 min.cloned().unwrap_or(Value::Null)
487 }
488 AggregateFunction::Stddev
490 | AggregateFunction::StddevPop
491 | AggregateFunction::VarSamp
492 | AggregateFunction::VarPop
493 | AggregateFunction::BitAnd
494 | AggregateFunction::BitOr
495 | AggregateFunction::BitXor => Value::Null,
496 }
497}
498
499#[cfg(test)]
500mod tests {
501 use super::*;
502 use teaql_core::{Aggregate, AggregateFunction, Record, SelectQuery, Value};
503
504 fn make_row(pairs: Vec<(&str, Value)>) -> Record {
505 pairs
506 .into_iter()
507 .map(|(k, v)| (k.to_owned(), v))
508 .collect()
509 }
510
511 fn sample_rows() -> Vec<Record> {
512 vec![
513 make_row(vec![
514 ("id", Value::U64(1)),
515 ("name", Value::Text("Alice".to_owned())),
516 ("age", Value::I64(30)),
517 ]),
518 make_row(vec![
519 ("id", Value::U64(2)),
520 ("name", Value::Text("Bob".to_owned())),
521 ("age", Value::I64(25)),
522 ]),
523 make_row(vec![
524 ("id", Value::U64(3)),
525 ("name", Value::Text("Charlie".to_owned())),
526 ("age", Value::I64(35)),
527 ]),
528 ]
529 }
530
531 #[test]
532 fn test_execute_no_filter() {
533 let query = SelectQuery::new("User");
534 let result = InMemoryQueryEngine::execute(&query, sample_rows());
535 assert_eq!(result.rows.len(), 3);
536 assert_eq!(result.metadata.backend, "memory");
537 }
538
539 #[test]
540 fn test_execute_with_eq_filter() {
541 let query = SelectQuery::new("User").filter(Expr::eq("name", "Bob"));
542 let result = InMemoryQueryEngine::execute(&query, sample_rows());
543 assert_eq!(result.rows.len(), 1);
544 assert_eq!(
545 result.rows[0].get("name"),
546 Some(&Value::Text("Bob".to_owned()))
547 );
548 }
549
550 #[test]
551 fn test_execute_with_gt_filter() {
552 let query = SelectQuery::new("User").filter(Expr::gt("age", 28_i64));
553 let result = InMemoryQueryEngine::execute(&query, sample_rows());
554 assert_eq!(result.rows.len(), 2); }
556
557 #[test]
558 fn test_sort_ascending() {
559 let query = SelectQuery::new("User").order_by(teaql_core::OrderBy::asc("age"));
560 let result = InMemoryQueryEngine::execute(&query, sample_rows());
561 let ages: Vec<_> = result
562 .rows
563 .iter()
564 .map(|r| r.get("age").unwrap().clone())
565 .collect();
566 assert_eq!(ages, vec![Value::I64(25), Value::I64(30), Value::I64(35)]);
567 }
568
569 #[test]
570 fn test_sort_descending() {
571 let query = SelectQuery::new("User").order_by(teaql_core::OrderBy::desc("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(35), Value::I64(30), Value::I64(25)]);
579 }
580
581 #[test]
582 fn test_paginate() {
583 let query = SelectQuery::new("User").page(1, 1);
584 let result = InMemoryQueryEngine::execute(&query, sample_rows());
585 assert_eq!(result.rows.len(), 1);
586 assert_eq!(
587 result.rows[0].get("name"),
588 Some(&Value::Text("Bob".to_owned()))
589 );
590 }
591
592 #[test]
593 fn test_projection() {
594 let query = SelectQuery::new("User").projects(["name"]);
595 let result = InMemoryQueryEngine::execute(&query, sample_rows());
596 for row in &result.rows {
597 assert!(row.contains_key("name"));
598 assert!(!row.contains_key("id"));
599 assert!(!row.contains_key("age"));
600 }
601 }
602
603 #[test]
604 fn test_count_aggregate() {
605 let query = SelectQuery::new("User").aggregate(Aggregate::count("total"));
606 let result = InMemoryQueryEngine::execute(&query, sample_rows());
607 assert_eq!(result.rows.len(), 1);
608 assert_eq!(result.rows[0].get("total"), Some(&Value::I64(3)));
609 }
610
611 #[test]
612 fn test_sum_aggregate() {
613 let query =
614 SelectQuery::new("User").aggregate(Aggregate::sum("age", "age_sum"));
615 let result = InMemoryQueryEngine::execute(&query, sample_rows());
616 assert_eq!(result.rows[0].get("age_sum"), Some(&Value::F64(90.0)));
617 }
618
619 #[test]
620 fn test_avg_aggregate() {
621 let query =
622 SelectQuery::new("User").aggregate(Aggregate::avg("age", "age_avg"));
623 let result = InMemoryQueryEngine::execute(&query, sample_rows());
624 assert_eq!(result.rows[0].get("age_avg"), Some(&Value::F64(30.0)));
625 }
626
627 #[test]
628 fn test_max_aggregate() {
629 let query =
630 SelectQuery::new("User").aggregate(Aggregate::max("age", "age_max"));
631 let result = InMemoryQueryEngine::execute(&query, sample_rows());
632 assert_eq!(result.rows[0].get("age_max"), Some(&Value::I64(35)));
633 }
634
635 #[test]
636 fn test_min_aggregate() {
637 let query =
638 SelectQuery::new("User").aggregate(Aggregate::min("age", "age_min"));
639 let result = InMemoryQueryEngine::execute(&query, sample_rows());
640 assert_eq!(result.rows[0].get("age_min"), Some(&Value::I64(25)));
641 }
642
643 #[test]
644 fn test_like_match_percent() {
645 assert!(ExprEvaluator::like_match("hello world", "%world"));
646 assert!(ExprEvaluator::like_match("hello world", "hello%"));
647 assert!(ExprEvaluator::like_match("hello world", "%lo wo%"));
648 assert!(ExprEvaluator::like_match("hello world", "%"));
649 assert!(!ExprEvaluator::like_match("hello world", "%xyz%"));
650 }
651
652 #[test]
653 fn test_like_match_underscore() {
654 assert!(ExprEvaluator::like_match("abc", "a_c"));
655 assert!(!ExprEvaluator::like_match("abbc", "a_c"));
656 assert!(ExprEvaluator::like_match("abc", "___"));
657 assert!(!ExprEvaluator::like_match("ab", "___"));
658 }
659
660 #[test]
661 fn test_like_match_combined() {
662 assert!(ExprEvaluator::like_match("foobar", "f%r"));
663 assert!(ExprEvaluator::like_match("foobar", "f__b%"));
664 assert!(!ExprEvaluator::like_match("foobar", "f__x%"));
665 }
666
667 #[test]
668 fn test_and_or_not() {
669 let row = make_row(vec![
670 ("a", Value::I64(10)),
671 ("b", Value::I64(20)),
672 ]);
673 let expr_and = Expr::and([Expr::eq("a", 10_i64), Expr::eq("b", 20_i64)]);
674 assert!(ExprEvaluator::eval(&expr_and, &row));
675
676 let expr_or = Expr::or([Expr::eq("a", 99_i64), Expr::eq("b", 20_i64)]);
677 assert!(ExprEvaluator::eval(&expr_or, &row));
678
679 let expr_not = Expr::negate(Expr::eq("a", 99_i64));
680 assert!(ExprEvaluator::eval(&expr_not, &row));
681 }
682
683 #[test]
684 fn test_is_null_is_not_null() {
685 let row = make_row(vec![("x", Value::Null), ("y", Value::I64(1))]);
686 assert!(ExprEvaluator::eval(&Expr::is_null("x"), &row));
687 assert!(!ExprEvaluator::eval(&Expr::is_not_null("x"), &row));
688 assert!(ExprEvaluator::eval(&Expr::is_not_null("y"), &row));
689 }
690
691 #[test]
692 fn test_between() {
693 let row = make_row(vec![("age", Value::I64(30))]);
694 assert!(ExprEvaluator::eval(
695 &Expr::between("age", Value::I64(25), Value::I64(35)),
696 &row
697 ));
698 assert!(!ExprEvaluator::eval(
699 &Expr::between("age", Value::I64(31), Value::I64(35)),
700 &row
701 ));
702 }
703
704 #[test]
705 fn test_in_list() {
706 let row = make_row(vec![("status", Value::Text("active".to_owned()))]);
707 let expr = Expr::in_list(
708 "status",
709 vec![
710 Value::Text("active".to_owned()),
711 Value::Text("pending".to_owned()),
712 ],
713 );
714 assert!(ExprEvaluator::eval(&expr, &row));
715
716 let expr_miss = Expr::in_list(
717 "status",
718 vec![Value::Text("closed".to_owned())],
719 );
720 assert!(!ExprEvaluator::eval(&expr_miss, &row));
721 }
722}