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