1use std::collections::HashMap;
48
49use sqlparser::{
50 ast::{
51 BinaryOperator, CaseWhen, CeilFloorKind, CreateTable, DateTimeField, DuplicateTreatment,
52 Expr, Function, FunctionArg, FunctionArgExpr, FunctionArguments, GroupByExpr, Insert,
53 JoinConstraint, JoinOperator, LimitClause, ObjectName, ObjectNamePart, ObjectType,
54 OrderByExpr, OrderByKind, Query, Select, SelectItem, SetExpr, Statement, TableFactor,
55 TableObject, TrimWhereField, UnaryOperator, Value as SqlValue, Values,
56 },
57 dialect::GenericDialect,
58 parser::Parser,
59};
60
61use mq_lang::{DefaultEngine, parse_markdown_input};
62
63use crate::{
64 DocumentStore, MqdbError,
65 block::{Block, BlockType, Properties, PropertyValue},
66 document::{Document, ZoneMaps},
67 indexes::{DocumentIndex, IndexHint},
68 store::CustomTableState,
69};
70
71#[derive(Debug, Clone, PartialEq)]
72pub enum Value {
73 Str(String),
74 Int(i64),
75 Float(f64),
76 Bool(bool),
77 Null,
78}
79
80impl Value {
81 fn as_str(&self) -> Option<&str> {
82 if let Value::Str(s) = self {
83 Some(s)
84 } else {
85 None
86 }
87 }
88 fn as_i64(&self) -> Option<i64> {
89 match self {
90 Value::Int(n) => Some(*n),
91 Value::Float(f) => Some(*f as i64),
92 _ => None,
93 }
94 }
95 fn as_f64(&self) -> Option<f64> {
96 match self {
97 Value::Float(f) => Some(*f),
98 Value::Int(n) => Some(*n as f64),
99 _ => None,
100 }
101 }
102 fn is_truthy(&self) -> bool {
103 match self {
104 Value::Bool(b) => *b,
105 Value::Int(n) => *n != 0,
106 Value::Float(f) => *f != 0.0,
107 Value::Str(s) => !s.is_empty(),
108 Value::Null => false,
109 }
110 }
111 fn display(&self) -> String {
112 match self {
113 Value::Str(s) => s.clone(),
114 Value::Int(n) => n.to_string(),
115 Value::Float(f) => f.to_string(),
116 Value::Bool(b) => b.to_string(),
117 Value::Null => "NULL".to_string(),
118 }
119 }
120 fn cmp_val(&self, other: &Value) -> Option<std::cmp::Ordering> {
121 match (self, other) {
122 (Value::Int(a), Value::Int(b)) => Some(a.cmp(b)),
123 (Value::Float(a), Value::Float(b)) => a.partial_cmp(b),
124 (Value::Int(a), Value::Float(b)) => (*a as f64).partial_cmp(b),
125 (Value::Float(a), Value::Int(b)) => a.partial_cmp(&(*b as f64)),
126 (Value::Str(a), Value::Str(b)) => Some(a.cmp(b)),
127 (Value::Null, Value::Null) => Some(std::cmp::Ordering::Equal),
128 _ => None,
129 }
130 }
131}
132
133#[derive(PartialEq, Eq, Hash)]
136enum JoinKey {
137 Str(String),
138 Int(i64),
139 Bool(bool),
140 FloatBits(u64),
141 Null,
142}
143
144fn value_join_key(v: &Value) -> Option<JoinKey> {
145 match v {
146 Value::Str(s) => Some(JoinKey::Str(s.clone())),
147 Value::Int(i) => Some(JoinKey::Int(*i)),
148 Value::Bool(b) => Some(JoinKey::Bool(*b)),
149 Value::Null => Some(JoinKey::Null),
150 Value::Float(f) if f.is_nan() => None, Value::Float(f) => {
152 let normalized = if *f == 0.0 { 0.0 } else { *f };
153 Some(JoinKey::FloatBits(normalized.to_bits()))
154 }
155 }
156}
157
158#[derive(Debug, Clone)]
159struct Row {
160 columns: Vec<String>,
161 values: Vec<Value>,
162}
163
164impl Row {
165 fn get(&self, col: &str) -> Option<&Value> {
166 let col_lower = col.to_lowercase();
167 if let Some(i) = self
168 .columns
169 .iter()
170 .position(|c| c.to_lowercase() == col_lower)
171 {
172 return self.values.get(i);
173 }
174 let short = col_lower.split('.').next_back().unwrap_or(&col_lower);
176 self.columns
178 .iter()
179 .position(|c| {
180 let cl = c.to_lowercase();
181 cl == col_lower || cl.split('.').next_back().unwrap_or(&cl) == short
182 })
183 .and_then(|i| self.values.get(i))
184 }
185}
186
187fn json_value_str(s: &str) -> String {
188 if let Ok(n) = s.parse::<i64>() {
189 return n.to_string();
190 }
191 if let Ok(f) = s.parse::<f64>() {
192 return f.to_string();
193 }
194 if s == "true" || s == "false" || s == "null" || s == "NULL" {
195 return s.to_lowercase();
196 }
197 format!("\"{}\"", s.replace('\\', "\\\\").replace('"', "\\\""))
199}
200
201fn csv_cell(s: &str) -> String {
202 if s.contains(',') || s.contains('"') || s.contains('\n') || s.contains('\r') {
203 format!("\"{}\"", s.replace('"', "\"\""))
204 } else {
205 s.to_string()
206 }
207}
208
209fn csv_row(fields: &[String]) -> String {
210 let mut row = fields
211 .iter()
212 .map(|f| csv_cell(f))
213 .collect::<Vec<_>>()
214 .join(",");
215 row.push('\n');
216 row
217}
218
219pub fn html_escape(s: &str) -> String {
220 s.replace('&', "&")
221 .replace('<', "<")
222 .replace('>', ">")
223 .replace('"', """)
224}
225
226#[derive(Debug)]
228pub struct QueryOutput {
229 pub columns: Vec<String>,
230 pub rows: Vec<Vec<String>>,
231}
232
233impl QueryOutput {
234 pub fn to_json(&self) -> String {
236 if self.rows.is_empty() {
237 return "[]\n".to_string();
238 }
239 let objects: Vec<String> = self
240 .rows
241 .iter()
242 .map(|row| {
243 let pairs: Vec<String> = self
244 .columns
245 .iter()
246 .zip(row.iter())
247 .map(|(col, val)| {
248 format!(
249 "\"{}\":{}",
250 col.replace('\\', "\\\\").replace('"', "\\\""),
251 json_value_str(val)
252 )
253 })
254 .collect();
255 format!("{{{}}}", pairs.join(","))
256 })
257 .collect();
258 format!("[{}]\n", objects.join(","))
259 }
260
261 pub fn to_csv(&self) -> String {
263 let mut out = String::new();
264 if !self.columns.is_empty() {
265 out.push_str(&csv_row(&self.columns));
266 }
267 for row in &self.rows {
268 out.push_str(&csv_row(row));
269 }
270 out
271 }
272
273 pub fn to_tsv(&self) -> String {
275 let mut out = String::new();
276 if !self.columns.is_empty() {
277 out.push_str(&self.columns.join("\t"));
278 out.push('\n');
279 }
280 for row in &self.rows {
281 out.push_str(&row.join("\t"));
282 out.push('\n');
283 }
284 out
285 }
286
287 pub fn to_markdown_table(&self) -> String {
289 if self.columns.is_empty() {
290 return String::new();
291 }
292 let mut widths: Vec<usize> = self.columns.iter().map(|h| h.len().max(3)).collect();
293 for row in &self.rows {
294 for (i, cell) in row.iter().enumerate() {
295 if i < widths.len() {
296 widths[i] = widths[i].max(cell.len());
297 }
298 }
299 }
300
301 let mut out = String::new();
302 out.push('|');
303 for (i, h) in self.columns.iter().enumerate() {
304 out.push_str(&format!(" {:<w$} |", h, w = widths[i]));
305 }
306 out.push('\n');
307
308 out.push('|');
309 for &w in &widths {
310 out.push_str(&format!(" {} |", "-".repeat(w)));
311 }
312 out.push('\n');
313
314 for row in &self.rows {
315 out.push('|');
316 for (i, &w) in widths.iter().enumerate() {
317 let cell = row.get(i).map(String::as_str).unwrap_or("");
318 let escaped = cell
319 .replace('|', "\\|")
320 .replace('\n', " ")
321 .replace('\r', "");
322 out.push_str(&format!(" {:<w$} |", escaped, w = w));
323 }
324 out.push('\n');
325 }
326 out
327 }
328
329 pub fn to_html_table(&self) -> String {
331 let mut out = String::from("<table>\n");
332 if !self.columns.is_empty() {
333 out.push_str("<thead><tr>");
334 for h in &self.columns {
335 out.push_str(&format!("<th>{}</th>", html_escape(h)));
336 }
337 out.push_str("</tr></thead>\n");
338 }
339 out.push_str("<tbody>\n");
340 for row in &self.rows {
341 out.push_str("<tr>");
342 for (i, _) in self.columns.iter().enumerate() {
343 let cell = row.get(i).map(String::as_str).unwrap_or("");
344 out.push_str(&format!("<td>{}</td>", html_escape(cell)));
345 }
346 out.push_str("</tr>\n");
347 }
348 out.push_str("</tbody>\n</table>\n");
349 out
350 }
351
352 pub fn to_table(&self) -> String {
354 const MAX_CELL: usize = 60;
355
356 if self.columns.is_empty() {
357 return "(no columns)\n".to_string();
358 }
359 if self.rows.is_empty() {
360 return "(0 rows)\n".to_string();
361 }
362
363 let mut widths: Vec<usize> = self.columns.iter().map(|h| h.len()).collect();
364 for row in &self.rows {
365 for (i, cell) in row.iter().enumerate() {
366 if i < widths.len() {
367 let display_len = cell.replace('\r', "").replace('\n', " ").chars().count();
368 widths[i] = widths[i].max(display_len.min(MAX_CELL));
369 }
370 }
371 }
372
373 let col_count = self.columns.len();
374 let mut out = String::new();
375
376 out.push('┌');
377 for (i, &w) in widths.iter().enumerate() {
378 out.push_str(&"─".repeat(w + 2));
379 out.push(if i + 1 < col_count { '┬' } else { '┐' });
380 }
381 out.push('\n');
382
383 out.push('│');
384 for (i, h) in self.columns.iter().enumerate() {
385 out.push_str(&format!(" {:<width$} │", h, width = widths[i]));
386 }
387 out.push('\n');
388
389 out.push('├');
390 for (i, &w) in widths.iter().enumerate() {
391 out.push_str(&"─".repeat(w + 2));
392 out.push(if i + 1 < col_count { '┼' } else { '┤' });
393 }
394 out.push('\n');
395
396 for row in &self.rows {
397 out.push('│');
398 for (i, &w) in widths.iter().enumerate() {
399 let cell = row.get(i).map(String::as_str).unwrap_or("");
400 let cell = cell.replace('\r', "").replace('\n', " ");
401 let truncated: String = if cell.chars().count() > MAX_CELL {
402 let mut s: String = cell.chars().take(MAX_CELL - 1).collect();
403 s.push('…');
404 s
405 } else {
406 cell
407 };
408 out.push_str(&format!(" {:<width$} │", truncated, width = w));
409 }
410 out.push('\n');
411 }
412
413 out.push('└');
414 for (i, &w) in widths.iter().enumerate() {
415 out.push_str(&"─".repeat(w + 2));
416 out.push(if i + 1 < col_count { '┴' } else { '┘' });
417 }
418 out.push('\n');
419 out.push_str(&format!(
420 "({} row{})\n",
421 self.rows.len(),
422 if self.rows.len() == 1 { "" } else { "s" }
423 ));
424 out
425 }
426}
427
428fn pv_to_json(pv: &PropertyValue) -> String {
429 match pv {
430 PropertyValue::String(s) => {
431 format!("\"{}\"", s.replace('\\', "\\\\").replace('"', "\\\""))
432 }
433 PropertyValue::Int(n) => n.to_string(),
434 PropertyValue::Float(f) => f.to_string(),
435 PropertyValue::Bool(b) => b.to_string(),
436 PropertyValue::Array(arr) => {
437 format!(
438 "[{}]",
439 arr.iter().map(pv_to_json).collect::<Vec<_>>().join(",")
440 )
441 }
442 PropertyValue::Null => "null".to_string(),
443 }
444}
445
446fn properties_to_json(props: &Properties) -> String {
447 let pairs: Vec<String> = props
448 .iter()
449 .map(|(k, v)| {
450 format!(
451 "\"{}\":{}",
452 k.replace('\\', "\\\\").replace('"', "\\\""),
453 pv_to_json(v)
454 )
455 })
456 .collect();
457 format!("{{{}}}", pairs.join(","))
458}
459
460fn block_to_row(doc_id: u32, block: &Block, block_idx: u32) -> Row {
461 Row {
462 columns: vec![
463 "id".into(),
464 "document_id".into(),
465 "block_type".into(),
466 "content".into(),
467 "pre".into(),
468 "post".into(),
469 "depth".into(),
470 "lang".into(),
471 "properties".into(),
472 ],
473 values: vec![
474 Value::Int(block_idx as i64),
475 Value::Int(doc_id as i64),
476 Value::Str(block.block_type.as_str().to_string()),
477 Value::Str(block.content.clone()),
478 Value::Int(block.pre as i64),
479 Value::Int(block.post as i64),
480 Value::Int(block.heading_depth().unwrap_or(0) as i64),
481 Value::Str(block.code_lang().unwrap_or("").to_string()),
482 Value::Str(properties_to_json(&block.properties)),
483 ],
484 }
485}
486
487fn doc_to_row(doc: &Document) -> Row {
488 let tags_json = {
489 let items: Vec<String> = doc
490 .zone_maps
491 .tags
492 .iter()
493 .map(|t| format!("\"{}\"", t.replace('"', "\\\"")))
494 .collect();
495 format!("[{}]", items.join(","))
496 };
497 Row {
498 columns: vec!["id".into(), "path".into(), "title".into(), "tags".into()],
499 values: vec![
500 Value::Int(doc.id as i64),
501 Value::Str(
502 doc.path
503 .as_ref()
504 .and_then(|p| p.to_str())
505 .unwrap_or("")
506 .to_string(),
507 ),
508 Value::Str(doc.zone_maps.title.clone().unwrap_or_default()),
509 Value::Str(tags_json),
510 ],
511 }
512}
513
514fn qualify_row(row: Row, prefix: &str) -> Row {
515 Row {
516 columns: row
517 .columns
518 .iter()
519 .map(|c| format!("{}.{}", prefix, c))
520 .collect(),
521 values: row.values,
522 }
523}
524
525fn cross_join(left: Vec<Row>, right: Vec<Row>) -> Vec<Row> {
526 let mut out = Vec::with_capacity(left.len() * right.len());
527 for l in &left {
528 for r in &right {
529 let mut cols = l.columns.clone();
530 cols.extend(r.columns.iter().cloned());
531 let mut vals = l.values.clone();
532 vals.extend(r.values.iter().cloned());
533 out.push(Row {
534 columns: cols,
535 values: vals,
536 });
537 }
538 }
539 out
540}
541
542fn hash_equi_join(
547 left: Vec<Row>,
548 right: Vec<Row>,
549 left_key_expr: &Expr,
550 right_key_expr: &Expr,
551 full_predicate: &Expr,
552) -> Vec<Row> {
553 let mut buckets: HashMap<JoinKey, Vec<usize>> = HashMap::new();
554 for (i, r) in right.iter().enumerate() {
555 if let Some(key) = value_join_key(&eval_expr(right_key_expr, r)) {
556 buckets.entry(key).or_default().push(i);
557 }
558 }
559
560 let mut out = Vec::new();
561 for l in &left {
562 let Some(key) = value_join_key(&eval_expr(left_key_expr, l)) else {
563 continue;
564 };
565 let Some(candidates) = buckets.get(&key) else {
566 continue;
567 };
568 for &i in candidates {
569 let r = &right[i];
570 let mut cols = l.columns.clone();
571 cols.extend(r.columns.iter().cloned());
572 let mut vals = l.values.clone();
573 vals.extend(r.values.iter().cloned());
574 let combined = Row {
575 columns: cols,
576 values: vals,
577 };
578 if eval_expr(full_predicate, &combined).is_truthy() {
579 out.push(combined);
580 }
581 }
582 }
583 out
584}
585
586fn eval_sql_value(v: &SqlValue) -> Value {
587 match v {
588 SqlValue::Number(n, _) => {
589 if let Ok(i) = n.parse::<i64>() {
590 Value::Int(i)
591 } else if let Ok(f) = n.parse::<f64>() {
592 Value::Float(f)
593 } else {
594 Value::Null
595 }
596 }
597 SqlValue::SingleQuotedString(s) | SqlValue::DoubleQuotedString(s) => Value::Str(s.clone()),
598 SqlValue::Boolean(b) => Value::Bool(*b),
599 SqlValue::Null => Value::Null,
600 _ => Value::Null,
601 }
602}
603
604fn ident_value(part: &ObjectNamePart) -> &str {
605 match part {
606 ObjectNamePart::Identifier(i) => &i.value,
607 ObjectNamePart::Function(_) => "",
608 }
609}
610
611fn eval_expr(expr: &Expr, row: &Row) -> Value {
612 match expr {
613 Expr::Value(v) => eval_sql_value(&v.value),
614 Expr::Identifier(i) => row.get(&i.value).cloned().unwrap_or(Value::Null),
615 Expr::CompoundIdentifier(parts) => {
616 let full = parts
618 .iter()
619 .map(|i| i.value.as_str())
620 .collect::<Vec<_>>()
621 .join(".");
622 let short = parts.last().map(|i| i.value.as_str()).unwrap_or("");
623 row.get(&full)
624 .or_else(|| row.get(short))
625 .cloned()
626 .unwrap_or(Value::Null)
627 }
628 Expr::BinaryOp { left, op, right } => eval_binary(left, op, right, row),
629 Expr::UnaryOp { op, expr } => match op {
630 UnaryOperator::Not => Value::Bool(!eval_expr(expr, row).is_truthy()),
631 UnaryOperator::Minus => match eval_expr(expr, row) {
632 Value::Int(n) => Value::Int(-n),
633 Value::Float(f) => Value::Float(-f),
634 _ => Value::Null,
635 },
636 _ => Value::Null,
637 },
638 Expr::IsNull(inner) => Value::Bool(matches!(eval_expr(inner, row), Value::Null)),
639 Expr::IsNotNull(inner) => Value::Bool(!matches!(eval_expr(inner, row), Value::Null)),
640 Expr::InList {
641 expr,
642 list,
643 negated,
644 } => {
645 let val = eval_expr(expr, row);
646 let found = list.iter().any(|e| eval_expr(e, row) == val);
647 Value::Bool(if *negated { !found } else { found })
648 }
649 Expr::Between {
650 expr,
651 negated,
652 low,
653 high,
654 } => {
655 let val = eval_expr(expr, row);
656 let lo = eval_expr(low, row);
657 let hi = eval_expr(high, row);
658 let in_range = lo.cmp_val(&val).map(|o| o.is_le()).unwrap_or(false)
659 && val.cmp_val(&hi).map(|o| o.is_le()).unwrap_or(false);
660 Value::Bool(if *negated { !in_range } else { in_range })
661 }
662 Expr::Like {
663 expr,
664 negated,
665 pattern,
666 ..
667 } => {
668 let val = eval_expr(expr, row);
669 let pat = eval_expr(pattern, row);
670 if let (Value::Str(s), Value::Str(p)) = (val, pat) {
671 let matched = like_match_str(&s, &p);
672 Value::Bool(if *negated { !matched } else { matched })
673 } else {
674 Value::Bool(false)
675 }
676 }
677 Expr::Function(f) => eval_function_call(f, row),
678 Expr::Nested(inner) => eval_expr(inner, row),
679 Expr::Cast { expr, .. } => eval_expr(expr, row),
680 Expr::Case {
681 operand,
682 conditions,
683 else_result,
684 ..
685 } => eval_case(operand.as_deref(), conditions, else_result.as_deref(), row),
686 Expr::Trim {
687 expr,
688 trim_where,
689 trim_what,
690 trim_characters,
691 } => eval_trim(expr, trim_where, trim_what, trim_characters, row),
692 Expr::Substring {
693 expr,
694 substring_from,
695 substring_for,
696 ..
697 } => eval_substring(expr, substring_from, substring_for, row),
698 Expr::Position { expr, r#in } => eval_position(expr, r#in, row),
699 Expr::Ceil { expr, field } => eval_ceil_floor(expr, field, row, true),
700 Expr::Floor { expr, field } => eval_ceil_floor(expr, field, row, false),
701 _ => Value::Null,
703 }
704}
705
706fn eval_case(
707 operand: Option<&Expr>,
708 conditions: &[CaseWhen],
709 else_result: Option<&Expr>,
710 row: &Row,
711) -> Value {
712 let operand_val = operand.map(|o| eval_expr(o, row));
713 for when in conditions {
714 let matched = match &operand_val {
715 Some(ov) => *ov == eval_expr(&when.condition, row),
716 None => eval_expr(&when.condition, row).is_truthy(),
717 };
718 if matched {
719 return eval_expr(&when.result, row);
720 }
721 }
722 else_result
723 .map(|e| eval_expr(e, row))
724 .unwrap_or(Value::Null)
725}
726
727fn eval_trim(
728 expr: &Expr,
729 trim_where: &Option<TrimWhereField>,
730 trim_what: &Option<Box<Expr>>,
731 trim_characters: &Option<Vec<Expr>>,
732 row: &Row,
733) -> Value {
734 let s = match eval_expr(expr, row).as_str() {
735 Some(s) => s.to_string(),
736 None => return Value::Null,
737 };
738 let chars: Vec<char> = if let Some(w) = trim_what {
739 eval_expr(w, row)
740 .as_str()
741 .map(|s| s.chars().collect())
742 .unwrap_or_default()
743 } else if let Some(cs) = trim_characters {
744 cs.iter()
745 .filter_map(|e| eval_expr(e, row).as_str().map(|s| s.to_string()))
746 .collect::<String>()
747 .chars()
748 .collect()
749 } else {
750 vec![' ', '\t', '\n', '\r']
751 };
752 let is_trim_char = |c: char| chars.contains(&c);
753 let trimmed = match trim_where {
754 Some(TrimWhereField::Leading) => s.trim_start_matches(is_trim_char).to_string(),
755 Some(TrimWhereField::Trailing) => s.trim_end_matches(is_trim_char).to_string(),
756 _ => s.trim_matches(is_trim_char).to_string(),
757 };
758 Value::Str(trimmed)
759}
760
761fn eval_substring(
762 expr: &Expr,
763 substring_from: &Option<Box<Expr>>,
764 substring_for: &Option<Box<Expr>>,
765 row: &Row,
766) -> Value {
767 let s = match eval_expr(expr, row).as_str() {
768 Some(s) => s.to_string(),
769 None => return Value::Null,
770 };
771 let chars: Vec<char> = s.chars().collect();
772 let len = chars.len() as i64;
773 let start_1based = substring_from
774 .as_ref()
775 .map(|e| eval_expr(e, row).as_i64().unwrap_or(1))
776 .unwrap_or(1);
777 let take = substring_for
778 .as_ref()
779 .map(|e| eval_expr(e, row).as_i64().unwrap_or(len));
780 let start_0based = (start_1based - 1).max(0) as usize;
783 let end_0based = match take {
784 Some(n) => {
785 let end = start_1based - 1 + n.max(0);
786 end.clamp(0, len) as usize
787 }
788 None => len as usize,
789 };
790 if start_0based >= chars.len() || end_0based <= start_0based {
791 return Value::Str(String::new());
792 }
793 Value::Str(chars[start_0based..end_0based].iter().collect())
794}
795
796fn eval_position(expr: &Expr, r#in: &Expr, row: &Row) -> Value {
797 let needle = eval_expr(expr, row);
798 let haystack = eval_expr(r#in, row);
799 match (needle.as_str(), haystack.as_str()) {
800 (Some(needle), Some(haystack)) => {
801 let hay_chars: Vec<char> = haystack.chars().collect();
802 let needle_chars: Vec<char> = needle.chars().collect();
803 if needle_chars.is_empty() {
804 return Value::Int(0);
805 }
806 for i in 0..=hay_chars.len().saturating_sub(needle_chars.len()) {
807 if hay_chars[i..i + needle_chars.len()] == needle_chars[..] {
808 return Value::Int(i as i64 + 1);
809 }
810 }
811 Value::Int(0)
812 }
813 _ => Value::Null,
814 }
815}
816
817fn eval_ceil_floor(expr: &Expr, field: &CeilFloorKind, row: &Row, is_ceil: bool) -> Value {
818 let n = match eval_expr(expr, row).as_f64() {
819 Some(n) => n,
820 None => return Value::Null,
821 };
822 let scale = match field {
823 CeilFloorKind::Scale(v) => match &v.value {
824 SqlValue::Number(s, _) => s.parse::<i32>().unwrap_or(0),
825 _ => 0,
826 },
827 CeilFloorKind::DateTimeField(DateTimeField::NoDateTime) => 0,
828 _ => return Value::Null,
830 };
831 let factor = 10f64.powi(scale);
832 let scaled = n * factor;
833 let rounded = if is_ceil {
834 scaled.ceil()
835 } else {
836 scaled.floor()
837 };
838 let result = rounded / factor;
839 if scale <= 0 && result.fract() == 0.0 {
840 Value::Int(result as i64)
841 } else {
842 Value::Float(result)
843 }
844}
845
846fn eval_binary(left: &Expr, op: &BinaryOperator, right: &Expr, row: &Row) -> Value {
847 match op {
848 BinaryOperator::And => {
849 if !eval_expr(left, row).is_truthy() {
850 return Value::Bool(false);
851 }
852 Value::Bool(eval_expr(right, row).is_truthy())
853 }
854 BinaryOperator::Or => {
855 if eval_expr(left, row).is_truthy() {
856 return Value::Bool(true);
857 }
858 Value::Bool(eval_expr(right, row).is_truthy())
859 }
860 BinaryOperator::Eq => Value::Bool(eval_expr(left, row) == eval_expr(right, row)),
861 BinaryOperator::NotEq => Value::Bool(eval_expr(left, row) != eval_expr(right, row)),
862 BinaryOperator::Lt => cmp_op(left, right, row, |o| o.is_lt()),
863 BinaryOperator::LtEq => cmp_op(left, right, row, |o| o.is_le()),
864 BinaryOperator::Gt => cmp_op(left, right, row, |o| o.is_gt()),
865 BinaryOperator::GtEq => cmp_op(left, right, row, |o| o.is_ge()),
866 BinaryOperator::Plus => arith_op(left, right, row, |a, b| a + b, |a, b| a + b),
867 BinaryOperator::Minus => arith_op(left, right, row, |a, b| a - b, |a, b| a - b),
868 BinaryOperator::Multiply => arith_op(left, right, row, |a, b| a * b, |a, b| a * b),
869 BinaryOperator::Divide => {
870 let (l, r) = (eval_expr(left, row), eval_expr(right, row));
871 match (&l, &r) {
872 (Value::Int(a), Value::Int(b)) if *b != 0 => Value::Int(a / b),
873 _ => match (l.as_f64(), r.as_f64()) {
874 (Some(a), Some(b)) if b != 0.0 => Value::Float(a / b),
875 _ => Value::Null,
876 },
877 }
878 }
879 BinaryOperator::StringConcat => {
880 let l = eval_expr(left, row);
881 let r = eval_expr(right, row);
882 Value::Str(format!("{}{}", l.display(), r.display()))
883 }
884 _ => Value::Null,
885 }
886}
887
888fn cmp_op(l: &Expr, r: &Expr, row: &Row, f: impl Fn(std::cmp::Ordering) -> bool) -> Value {
889 Value::Bool(
890 eval_expr(l, row)
891 .cmp_val(&eval_expr(r, row))
892 .map(f)
893 .unwrap_or(false),
894 )
895}
896
897fn arith_op(
898 l: &Expr,
899 r: &Expr,
900 row: &Row,
901 int_f: impl Fn(i64, i64) -> i64,
902 flt_f: impl Fn(f64, f64) -> f64,
903) -> Value {
904 let (lv, rv) = (eval_expr(l, row), eval_expr(r, row));
905 match (&lv, &rv) {
906 (Value::Int(a), Value::Int(b)) => Value::Int(int_f(*a, *b)),
907 _ => match (lv.as_f64(), rv.as_f64()) {
908 (Some(a), Some(b)) => Value::Float(flt_f(a, b)),
909 _ => Value::Null,
910 },
911 }
912}
913
914fn eval_function_call(f: &Function, row: &Row) -> Value {
915 let name = f.name.0.last().map(ident_value).unwrap_or("");
916 if is_aggregate_name(&name.to_lowercase()) {
918 return Value::Int(1);
919 }
920 let args: Vec<Value> = match &f.args {
921 FunctionArguments::List(al) => al
922 .args
923 .iter()
924 .filter_map(|a| match a {
925 FunctionArg::Unnamed(FunctionArgExpr::Expr(e)) => Some(eval_expr(e, row)),
926 _ => None,
927 })
928 .collect(),
929 _ => vec![],
930 };
931 eval_scalar_function(name, &args)
932}
933
934fn eval_scalar_function(name: &str, args: &[Value]) -> Value {
935 match name.to_lowercase().as_str() {
936 "under" => {
937 if args.len() < 4 {
938 return Value::Bool(false);
939 }
940 let (pre, post) = (args[0].as_i64().unwrap_or(0), args[1].as_i64().unwrap_or(0));
941 let (ap, aq) = (args[2].as_i64().unwrap_or(0), args[3].as_i64().unwrap_or(0));
942 Value::Bool(pre > ap && post < aq)
943 }
944 "json_extract" => {
945 if args.len() < 2 {
946 return Value::Null;
947 }
948 let json = args[0].as_str().unwrap_or("");
949 let path = args[1].as_str().unwrap_or("");
950 let key = path.trim_start_matches("$.").trim_matches('"');
951 extract_json_key(json, key)
952 }
953 "mq" => {
954 if args.len() < 2 {
955 return Value::Null;
956 }
957 let program = match args[0].as_str() {
958 Some(s) => s.to_string(),
959 None => return Value::Null,
960 };
961 let content = match args[1].as_str() {
962 Some(s) => s.to_string(),
963 None => return Value::Null,
964 };
965 eval_mq_scalar(&program, &content)
966 }
967
968 "lower" => str_fn(args, |s| s.to_lowercase()),
970 "upper" => str_fn(args, |s| s.to_uppercase()),
971 "length" | "len" | "char_length" | "character_length" => args
972 .first()
973 .and_then(|v| v.as_str())
974 .map(|s| Value::Int(s.chars().count() as i64))
975 .unwrap_or(Value::Null),
976 "trim" => str_fn(args, |s| s.trim().to_string()),
977 "ltrim" => {
978 let chars = trim_char_set(args, 1);
979 str_fn(args, |s| {
980 s.trim_start_matches(|c| chars.contains(&c)).to_string()
981 })
982 }
983 "rtrim" => {
984 let chars = trim_char_set(args, 1);
985 str_fn(args, |s| {
986 s.trim_end_matches(|c| chars.contains(&c)).to_string()
987 })
988 }
989 "concat" => Value::Str(
990 args.iter()
991 .map(|v| v.display())
992 .collect::<Vec<_>>()
993 .join(""),
994 ),
995 "concat_ws" => {
996 let sep = match args.first().and_then(|v| v.as_str()) {
997 Some(s) => s,
998 None => return Value::Null,
999 };
1000 Value::Str(
1001 args[1..]
1002 .iter()
1003 .filter(|v| !matches!(v, Value::Null))
1004 .map(|v| v.display())
1005 .collect::<Vec<_>>()
1006 .join(sep),
1007 )
1008 }
1009 "replace" => {
1010 if args.len() < 3 {
1011 return Value::Null;
1012 }
1013 match (args[0].as_str(), args[1].as_str(), args[2].as_str()) {
1014 (Some(s), Some(from), Some(to)) => Value::Str(s.replace(from, to)),
1015 _ => Value::Null,
1016 }
1017 }
1018 "left" => str_int_fn(args, |chars, n| {
1019 chars[..(n.max(0) as usize).min(chars.len())]
1020 .iter()
1021 .collect()
1022 }),
1023 "right" => str_int_fn(args, |chars, n| {
1024 let n = (n.max(0) as usize).min(chars.len());
1025 chars[chars.len() - n..].iter().collect()
1026 }),
1027 "lpad" => pad_fn(args, true),
1028 "rpad" => pad_fn(args, false),
1029 "reverse" => str_fn(args, |s| s.chars().rev().collect()),
1030 "repeat" => {
1031 if args.len() < 2 {
1032 return Value::Null;
1033 }
1034 match (args[0].as_str(), args[1].as_i64()) {
1035 (Some(s), Some(n)) => Value::Str(s.repeat(n.max(0) as usize)),
1036 _ => Value::Null,
1037 }
1038 }
1039 "initcap" => str_fn(args, |s| {
1040 s.split(' ')
1041 .map(|word| {
1042 let mut c = word.chars();
1043 match c.next() {
1044 Some(first) => {
1045 first.to_uppercase().collect::<String>() + &c.as_str().to_lowercase()
1046 }
1047 None => String::new(),
1048 }
1049 })
1050 .collect::<Vec<_>>()
1051 .join(" ")
1052 }),
1053 "ascii" => args
1054 .first()
1055 .and_then(|v| v.as_str())
1056 .and_then(|s| s.chars().next())
1057 .map(|c| Value::Int(c as i64))
1058 .unwrap_or(Value::Null),
1059 "chr" => args
1060 .first()
1061 .and_then(|v| v.as_i64())
1062 .and_then(|n| u32::try_from(n).ok())
1063 .and_then(char::from_u32)
1064 .map(|c| Value::Str(c.to_string()))
1065 .unwrap_or(Value::Null),
1066 "instr" => {
1067 if args.len() < 2 {
1068 return Value::Null;
1069 }
1070 match (args[0].as_str(), args[1].as_str()) {
1071 (Some(haystack), Some(needle)) => {
1072 let hay_chars: Vec<char> = haystack.chars().collect();
1073 let needle_chars: Vec<char> = needle.chars().collect();
1074 if needle_chars.is_empty() {
1075 return Value::Int(0);
1076 }
1077 for i in 0..=hay_chars.len().saturating_sub(needle_chars.len()) {
1078 if hay_chars[i..i + needle_chars.len()] == needle_chars[..] {
1079 return Value::Int(i as i64 + 1);
1080 }
1081 }
1082 Value::Int(0)
1083 }
1084 _ => Value::Null,
1085 }
1086 }
1087 "split_part" => {
1088 if args.len() < 3 {
1089 return Value::Null;
1090 }
1091 match (args[0].as_str(), args[1].as_str(), args[2].as_i64()) {
1092 (Some(s), Some(delim), Some(n)) if n > 0 => s
1093 .split(delim)
1094 .nth((n - 1) as usize)
1095 .map(|p| Value::Str(p.to_string()))
1096 .unwrap_or(Value::Null),
1097 _ => Value::Null,
1098 }
1099 }
1100
1101 "abs" => num_fn(args, |n| n.abs(), |n| n.abs()),
1103 "round" => {
1104 let n = match args.first().and_then(|v| v.as_f64()) {
1105 Some(n) => n,
1106 None => return Value::Null,
1107 };
1108 let scale = args.get(1).and_then(|v| v.as_i64()).unwrap_or(0);
1109 let factor = 10f64.powi(scale as i32);
1110 let result = (n * factor).round() / factor;
1111 if scale <= 0 {
1112 Value::Int(result as i64)
1113 } else {
1114 Value::Float(result)
1115 }
1116 }
1117 "ceil" | "ceiling" => float_fn(args, |n| n.ceil()),
1118 "floor" => float_fn(args, |n| n.floor()),
1119 "trunc" | "truncate" => {
1120 let n = match args.first().and_then(|v| v.as_f64()) {
1121 Some(n) => n,
1122 None => return Value::Null,
1123 };
1124 let scale = args.get(1).and_then(|v| v.as_i64()).unwrap_or(0);
1125 let factor = 10f64.powi(scale as i32);
1126 let result = (n * factor).trunc() / factor;
1127 if scale <= 0 {
1128 Value::Int(result as i64)
1129 } else {
1130 Value::Float(result)
1131 }
1132 }
1133 "mod" => {
1134 if args.len() < 2 {
1135 return Value::Null;
1136 }
1137 match (&args[0], &args[1]) {
1138 (Value::Int(a), Value::Int(b)) if *b != 0 => Value::Int(a % b),
1139 _ => match (args[0].as_f64(), args[1].as_f64()) {
1140 (Some(a), Some(b)) if b != 0.0 => Value::Float(a % b),
1141 _ => Value::Null,
1142 },
1143 }
1144 }
1145 "power" | "pow" => {
1146 if args.len() < 2 {
1147 return Value::Null;
1148 }
1149 match (args[0].as_f64(), args[1].as_f64()) {
1150 (Some(a), Some(b)) => Value::Float(a.powf(b)),
1151 _ => Value::Null,
1152 }
1153 }
1154 "sqrt" => float_fn(args, |n| n.sqrt()),
1155 "sign" => float_fn(args, |n| {
1156 if n > 0.0 {
1157 1.0
1158 } else if n < 0.0 {
1159 -1.0
1160 } else {
1161 0.0
1162 }
1163 }),
1164 "exp" => float_fn(args, |n| n.exp()),
1165 "ln" => float_fn(args, |n| n.ln()),
1166 "log10" => float_fn(args, |n| n.log10()),
1167 "log2" => float_fn(args, |n| n.log2()),
1168 "log" => {
1169 let n = match args.first().and_then(|v| v.as_f64()) {
1170 Some(n) => n,
1171 None => return Value::Null,
1172 };
1173 match args.get(1).and_then(|v| v.as_f64()) {
1174 Some(base) => Value::Float(n.log(base)),
1175 None => Value::Float(n.log10()),
1176 }
1177 }
1178 "pi" => Value::Float(std::f64::consts::PI),
1179 "greatest" => args
1180 .iter()
1181 .filter(|v| !matches!(v, Value::Null))
1182 .cloned()
1183 .max_by(|a, b| a.cmp_val(b).unwrap_or(std::cmp::Ordering::Equal))
1184 .unwrap_or(Value::Null),
1185 "least" => args
1186 .iter()
1187 .filter(|v| !matches!(v, Value::Null))
1188 .cloned()
1189 .min_by(|a, b| a.cmp_val(b).unwrap_or(std::cmp::Ordering::Equal))
1190 .unwrap_or(Value::Null),
1191
1192 "coalesce" | "ifnull" => args
1194 .iter()
1195 .find(|v| !matches!(v, Value::Null))
1196 .cloned()
1197 .unwrap_or(Value::Null),
1198 "nullif" => {
1199 if args.len() < 2 {
1200 return Value::Null;
1201 }
1202 if args[0] == args[1] {
1203 Value::Null
1204 } else {
1205 args[0].clone()
1206 }
1207 }
1208
1209 "typeof" => Value::Str(
1211 match args.first() {
1212 Some(Value::Str(_)) => "text",
1213 Some(Value::Int(_)) => "integer",
1214 Some(Value::Float(_)) => "float",
1215 Some(Value::Bool(_)) => "boolean",
1216 Some(Value::Null) | None => "null",
1217 }
1218 .to_string(),
1219 ),
1220 "now" | "current_timestamp" => Value::Str(current_datetime_utc(true, true)),
1221 "current_date" => Value::Str(current_datetime_utc(true, false)),
1222 "current_time" => Value::Str(current_datetime_utc(false, true)),
1223 _ => Value::Null,
1224 }
1225}
1226
1227fn str_fn(args: &[Value], f: impl Fn(&str) -> String) -> Value {
1228 args.first()
1229 .and_then(|v| v.as_str())
1230 .map(|s| Value::Str(f(s)))
1231 .unwrap_or(Value::Null)
1232}
1233
1234fn str_int_fn(args: &[Value], f: impl Fn(&[char], i64) -> String) -> Value {
1235 if args.len() < 2 {
1236 return Value::Null;
1237 }
1238 match (args[0].as_str(), args[1].as_i64()) {
1239 (Some(s), Some(n)) => {
1240 let chars: Vec<char> = s.chars().collect();
1241 Value::Str(f(&chars, n))
1242 }
1243 _ => Value::Null,
1244 }
1245}
1246
1247fn num_fn(args: &[Value], int_f: impl Fn(i64) -> i64, flt_f: impl Fn(f64) -> f64) -> Value {
1248 match args.first() {
1249 Some(Value::Int(n)) => Value::Int(int_f(*n)),
1250 Some(v) => v
1251 .as_f64()
1252 .map(|n| Value::Float(flt_f(n)))
1253 .unwrap_or(Value::Null),
1254 None => Value::Null,
1255 }
1256}
1257
1258fn float_fn(args: &[Value], f: impl Fn(f64) -> f64) -> Value {
1259 args.first()
1260 .and_then(|v| v.as_f64())
1261 .map(|n| Value::Float(f(n)))
1262 .unwrap_or(Value::Null)
1263}
1264
1265fn trim_char_set(args: &[Value], chars_idx: usize) -> Vec<char> {
1268 args.get(chars_idx)
1269 .and_then(|v| v.as_str())
1270 .map(|s| s.chars().collect())
1271 .unwrap_or_else(|| vec![' ', '\t', '\n', '\r'])
1272}
1273
1274fn pad_fn(args: &[Value], left: bool) -> Value {
1275 if args.len() < 2 {
1276 return Value::Null;
1277 }
1278 let s = match args[0].as_str() {
1279 Some(s) => s,
1280 None => return Value::Null,
1281 };
1282 let target_len = match args[1].as_i64() {
1283 Some(n) => n.max(0) as usize,
1284 None => return Value::Null,
1285 };
1286 let pad_str = args.get(2).and_then(|v| v.as_str()).unwrap_or(" ");
1287 let mut chars: Vec<char> = s.chars().collect();
1288 if chars.len() >= target_len {
1289 chars.truncate(target_len);
1290 return Value::Str(chars.into_iter().collect());
1291 }
1292 if pad_str.is_empty() {
1293 return Value::Str(s.to_string());
1294 }
1295 let pad_chars: Vec<char> = pad_str.chars().collect();
1296 let needed = target_len - chars.len();
1297 let padding: Vec<char> = pad_chars.iter().cycle().take(needed).copied().collect();
1298 if left {
1299 Value::Str(padding.into_iter().chain(chars).collect())
1300 } else {
1301 chars.extend(padding);
1302 Value::Str(chars.into_iter().collect())
1303 }
1304}
1305
1306fn current_datetime_utc(with_date: bool, with_time: bool) -> String {
1310 let secs = std::time::SystemTime::now()
1311 .duration_since(std::time::UNIX_EPOCH)
1312 .map(|d| d.as_secs())
1313 .unwrap_or(0);
1314 let days = (secs / 86400) as i64;
1315 let time_of_day = secs % 86400;
1316 let (y, m, d) = civil_from_days(days);
1317 let (h, mi, s) = (
1318 time_of_day / 3600,
1319 (time_of_day / 60) % 60,
1320 time_of_day % 60,
1321 );
1322 match (with_date, with_time) {
1323 (true, true) => format!("{y:04}-{m:02}-{d:02} {h:02}:{mi:02}:{s:02}"),
1324 (true, false) => format!("{y:04}-{m:02}-{d:02}"),
1325 _ => format!("{h:02}:{mi:02}:{s:02}"),
1326 }
1327}
1328
1329fn civil_from_days(z: i64) -> (i64, u32, u32) {
1332 let z = z + 719468;
1333 let era = if z >= 0 { z } else { z - 146096 } / 146097;
1334 let doe = (z - era * 146097) as u64;
1335 let yoe = (doe - doe / 1460 + doe / 36524 - doe / 146096) / 365;
1336 let y = yoe as i64 + era * 400;
1337 let doy = doe - (365 * yoe + yoe / 4 - yoe / 100);
1338 let mp = (5 * doy + 2) / 153;
1339 let d = (doy - (153 * mp + 2) / 5 + 1) as u32;
1340 let m = if mp < 10 { mp + 3 } else { mp - 9 } as u32;
1341 let y = if m <= 2 { y + 1 } else { y };
1342 (y, m, d)
1343}
1344
1345fn eval_mq_scalar(program: &str, content: &str) -> Value {
1346 let mut engine = DefaultEngine::default();
1347 engine.load_builtin_module();
1348 let input = match parse_markdown_input(content) {
1349 Ok(i) => i,
1350 Err(_) => return Value::Null,
1351 };
1352 match engine.eval(program, input.into_iter()) {
1353 Ok(output) => {
1354 let parts: Vec<String> = output
1355 .compact()
1356 .into_iter()
1357 .map(|v| v.to_string())
1358 .collect();
1359 if parts.is_empty() {
1360 Value::Null
1361 } else {
1362 Value::Str(parts.join("\n"))
1363 }
1364 }
1365 Err(_) => Value::Null,
1366 }
1367}
1368
1369fn extract_json_key(json: &str, key: &str) -> Value {
1370 let s = json.trim();
1371 if !s.starts_with('{') {
1372 return Value::Null;
1373 }
1374 let target = format!("\"{}\":", key);
1375 if let Some(pos) = s.find(&target) {
1376 let after = s[pos + target.len()..].trim_start();
1377 if let Some(inner) = after.strip_prefix('"') {
1378 if let Some(end) = inner.find('"') {
1379 return Value::Str(inner[..end].to_string());
1380 }
1381 } else if let Some(end) = after.find([',', '}']) {
1382 let raw = after[..end].trim();
1383 if let Ok(n) = raw.parse::<i64>() {
1384 return Value::Int(n);
1385 }
1386 if let Ok(f) = raw.parse::<f64>() {
1387 return Value::Float(f);
1388 }
1389 if raw == "true" {
1390 return Value::Bool(true);
1391 }
1392 if raw == "false" {
1393 return Value::Bool(false);
1394 }
1395 if raw == "null" {
1396 return Value::Null;
1397 }
1398 }
1399 }
1400 Value::Null
1401}
1402
1403fn like_match_str(s: &str, pattern: &str) -> bool {
1405 let s: Vec<char> = s.to_lowercase().chars().collect();
1406 let p: Vec<char> = pattern.to_lowercase().chars().collect();
1407 like_dp(&s, &p, 0, 0)
1408}
1409
1410fn like_dp(s: &[char], p: &[char], si: usize, pi: usize) -> bool {
1411 if pi == p.len() {
1412 return si == s.len();
1413 }
1414 if p[pi] == '%' {
1415 let mut npi = pi + 1;
1417 while npi < p.len() && p[npi] == '%' {
1418 npi += 1;
1419 }
1420 for k in si..=s.len() {
1421 if like_dp(s, p, k, npi) {
1422 return true;
1423 }
1424 }
1425 return false;
1426 }
1427 if si >= s.len() {
1428 return false;
1429 }
1430 let matches = p[pi] == '_' || p[pi] == s[si];
1431 matches && like_dp(s, p, si + 1, pi + 1)
1432}
1433
1434pub struct SqlEngine<'a> {
1440 store: &'a DocumentStore,
1441 indexes: Vec<DocumentIndex>,
1443}
1444
1445impl<'a> SqlEngine<'a> {
1446 pub fn new(store: &'a DocumentStore) -> Result<Self, MqdbError> {
1451 let indexes = store
1452 .documents()
1453 .iter()
1454 .enumerate()
1455 .map(|(i, doc)| {
1456 if let Some(idx) = store.get_doc_index(i) {
1457 idx.clone()
1458 } else {
1459 DocumentIndex::build(&doc.blocks)
1460 }
1461 })
1462 .collect();
1463 Ok(Self { store, indexes })
1464 }
1465
1466 fn documents_with_indexes(&self) -> impl Iterator<Item = (&Document, &DocumentIndex)> {
1467 self.store.documents().iter().zip(self.indexes.iter())
1468 }
1469
1470 pub fn execute(&self, sql: &str) -> Result<QueryOutput, MqdbError> {
1475 let trimmed = sql.trim().trim_end_matches(';');
1477 let upper = trimmed.to_ascii_uppercase();
1478 if upper.starts_with("DESC ") || upper.starts_with("DESCRIBE ") {
1479 let name = trimmed
1480 .split_whitespace()
1481 .nth(1)
1482 .unwrap_or("")
1483 .to_lowercase();
1484 return self.exec_desc(&name);
1485 }
1486 if upper == "SHOW TABLES" {
1487 return self.exec_show_tables();
1488 }
1489
1490 let stmts = Parser::parse_sql(&GenericDialect {}, sql)
1491 .map_err(|e| MqdbError::SqlParse(e.to_string()))?;
1492 let stmt = stmts
1493 .into_iter()
1494 .next()
1495 .ok_or_else(|| MqdbError::SqlParse("empty query".into()))?;
1496 match stmt {
1497 Statement::Query(q) => self.exec_query(&q),
1498 Statement::CreateTable(ct) => self.exec_create_table(&ct),
1499 Statement::Insert(ins) => self.exec_insert(&ins),
1500 Statement::Drop {
1501 object_type: ObjectType::Table,
1502 names,
1503 if_exists,
1504 ..
1505 } => self.exec_drop_tables(&names, if_exists),
1506 _ => Err(MqdbError::SqlExec(
1507 "unsupported statement; supported: SELECT, CREATE TABLE, INSERT INTO, DROP TABLE, DESC, SHOW TABLES".into(),
1508 )),
1509 }
1510 }
1511
1512 fn exec_desc(&self, table_name: &str) -> Result<QueryOutput, MqdbError> {
1513 let schema: Option<Vec<(&str, &str)>> = match table_name {
1514 "blocks" => Some(vec![
1515 ("id", "integer"),
1516 ("document_id", "integer"),
1517 ("block_type", "text"),
1518 ("content", "text"),
1519 ("pre", "integer"),
1520 ("post", "integer"),
1521 ("depth", "integer"),
1522 ("lang", "text"),
1523 ("properties", "text"),
1524 ]),
1525 "documents" => Some(vec![
1526 ("id", "integer"),
1527 ("path", "text"),
1528 ("title", "text"),
1529 ("tags", "text"),
1530 ]),
1531 _ => None,
1532 };
1533 if let Some(rows) = schema {
1534 return Ok(QueryOutput {
1535 columns: vec!["column".to_string(), "type".to_string()],
1536 rows: rows
1537 .iter()
1538 .map(|(c, t)| vec![c.to_string(), t.to_string()])
1539 .collect(),
1540 });
1541 }
1542 let guard = self.store.custom_tables.read().unwrap();
1543 if let Some(state) = guard.get(table_name) {
1544 let rows = state
1545 .columns
1546 .iter()
1547 .map(|c| vec![c.clone(), "text".to_string()])
1548 .collect();
1549 return Ok(QueryOutput {
1550 columns: vec!["column".to_string(), "type".to_string()],
1551 rows,
1552 });
1553 }
1554 Err(MqdbError::SqlExec(format!("unknown table: {table_name}")))
1555 }
1556
1557 fn exec_show_tables(&self) -> Result<QueryOutput, MqdbError> {
1558 let mut rows = vec![
1559 vec!["blocks".to_string(), "built-in".to_string()],
1560 vec!["documents".to_string(), "built-in".to_string()],
1561 ];
1562 let guard = self.store.custom_tables.read().unwrap();
1563 let mut custom: Vec<String> = guard.keys().cloned().collect();
1564 drop(guard);
1565 custom.sort();
1566 rows.extend(custom.into_iter().map(|n| vec![n, "custom".to_string()]));
1567 Ok(QueryOutput {
1568 columns: vec!["table".to_string(), "kind".to_string()],
1569 rows,
1570 })
1571 }
1572
1573 fn exec_create_table(&self, ct: &CreateTable) -> Result<QueryOutput, MqdbError> {
1574 let table_name = ct
1575 .name
1576 .0
1577 .last()
1578 .map(ident_value)
1579 .unwrap_or("")
1580 .to_lowercase();
1581 if matches!(table_name.as_str(), "blocks" | "documents") {
1582 return Err(MqdbError::SqlExec(format!(
1583 "cannot override built-in table '{table_name}'"
1584 )));
1585 }
1586
1587 if let Some(query) = &ct.query {
1588 let result = self.exec_query(query)?;
1590 let n = result.rows.len();
1591 self.store.custom_tables.write().unwrap().insert(
1592 table_name,
1593 CustomTableState {
1594 columns: result.columns,
1595 rows: result.rows,
1596 first_row_page: 0,
1597 last_row_page: 0,
1598 },
1599 );
1600 self.store.try_flush_catalog_to_storage();
1601 return Ok(QueryOutput {
1602 columns: vec!["rows".to_string()],
1603 rows: vec![vec![n.to_string()]],
1604 });
1605 }
1606
1607 let columns: Vec<String> = ct.columns.iter().map(|c| c.name.value.clone()).collect();
1609 if columns.is_empty() {
1610 return Err(MqdbError::SqlExec(
1611 "CREATE TABLE requires at least one column or AS SELECT".into(),
1612 ));
1613 }
1614 let already_exists = self
1615 .store
1616 .custom_tables
1617 .read()
1618 .unwrap()
1619 .contains_key(&table_name);
1620 if already_exists {
1621 if ct.if_not_exists {
1622 return Ok(QueryOutput {
1623 columns: vec!["result".to_string()],
1624 rows: vec![vec!["already exists".to_string()]],
1625 });
1626 }
1627 return Err(MqdbError::SqlExec(format!(
1628 "table '{table_name}' already exists"
1629 )));
1630 }
1631 self.store.custom_tables.write().unwrap().insert(
1632 table_name,
1633 CustomTableState {
1634 columns,
1635 rows: vec![],
1636 first_row_page: 0,
1637 last_row_page: 0,
1638 },
1639 );
1640 self.store.try_flush_catalog_to_storage();
1641 Ok(QueryOutput {
1642 columns: vec!["result".to_string()],
1643 rows: vec![vec!["ok".to_string()]],
1644 })
1645 }
1646
1647 fn exec_insert(&self, ins: &Insert) -> Result<QueryOutput, MqdbError> {
1648 let table_name = match &ins.table {
1649 TableObject::TableName(name) => {
1650 name.0.last().map(ident_value).unwrap_or("").to_lowercase()
1651 }
1652 _ => return Err(MqdbError::SqlExec("unsupported INSERT target".into())),
1653 };
1654
1655 let source = ins
1656 .source
1657 .as_ref()
1658 .ok_or_else(|| MqdbError::SqlExec("INSERT requires VALUES or SELECT".into()))?;
1659 let values_out = self.exec_query(source)?;
1660
1661 let col_indices: Option<Vec<usize>> = if ins.columns.is_empty() {
1663 None } else {
1665 let guard = self.store.custom_tables.read().unwrap();
1666 let table_cols = guard
1667 .get(&table_name)
1668 .map(|state| state.columns.clone())
1669 .ok_or_else(|| MqdbError::SqlExec(format!("unknown table: {table_name}")))?;
1670 drop(guard);
1671 let indices: Result<Vec<usize>, _> = ins
1672 .columns
1673 .iter()
1674 .map(|col_name| {
1675 let name = col_name.0.last().map(ident_value).unwrap_or("");
1676 table_cols
1677 .iter()
1678 .position(|c| c.eq_ignore_ascii_case(name))
1679 .ok_or_else(|| MqdbError::SqlExec(format!("unknown column '{name}'")))
1680 })
1681 .collect();
1682 Some(indices?)
1683 };
1684
1685 let new_rows = {
1686 let mut guard = self.store.custom_tables.write().unwrap();
1687 let state = guard
1688 .get_mut(&table_name)
1689 .ok_or_else(|| MqdbError::SqlExec(format!("unknown table: {table_name}")))?;
1690 let ncols = state.columns.len();
1691
1692 let mut new_rows = Vec::with_capacity(values_out.rows.len());
1693 for src_row in &values_out.rows {
1694 let mut row = vec![String::new(); ncols];
1695 match &col_indices {
1696 None => {
1697 if src_row.len() != ncols {
1698 return Err(MqdbError::SqlExec(format!(
1699 "expected {ncols} columns, got {}",
1700 src_row.len()
1701 )));
1702 }
1703 row = src_row.clone();
1704 }
1705 Some(idx_map) => {
1706 for (dst_idx, &src_idx) in idx_map.iter().enumerate() {
1707 if let Some(v) = src_row.get(dst_idx) {
1708 row[src_idx] = v.clone();
1709 }
1710 }
1711 }
1712 }
1713 state.rows.push(row.clone());
1714 new_rows.push(row);
1715 }
1716 new_rows
1717 }; let inserted = new_rows.len();
1719 self.store
1723 .try_append_table_rows_to_storage(&table_name, &new_rows);
1724 Ok(QueryOutput {
1725 columns: vec!["rows_affected".to_string()],
1726 rows: vec![vec![inserted.to_string()]],
1727 })
1728 }
1729
1730 fn exec_drop_tables(
1731 &self,
1732 names: &[ObjectName],
1733 if_exists: bool,
1734 ) -> Result<QueryOutput, MqdbError> {
1735 let dropped = {
1736 let mut guard = self.store.custom_tables.write().unwrap();
1737 let mut dropped = 0usize;
1738 for name in names {
1739 let table_name = name.0.last().map(ident_value).unwrap_or("").to_lowercase();
1740 if matches!(table_name.as_str(), "blocks" | "documents") {
1741 return Err(MqdbError::SqlExec(format!(
1742 "cannot drop built-in table '{table_name}'"
1743 )));
1744 }
1745 if guard.remove(&table_name).is_some() {
1746 dropped += 1;
1747 } else if !if_exists {
1748 return Err(MqdbError::SqlExec(format!(
1749 "table '{table_name}' does not exist"
1750 )));
1751 }
1752 }
1753 dropped
1754 }; self.store.try_flush_catalog_to_storage();
1756 Ok(QueryOutput {
1757 columns: vec!["result".to_string()],
1758 rows: vec![vec![format!("{dropped} table(s) dropped")]],
1759 })
1760 }
1761
1762 fn exec_query(&self, query: &Query) -> Result<QueryOutput, MqdbError> {
1763 let select = match query.body.as_ref() {
1764 SetExpr::Select(s) => s,
1765 SetExpr::Values(Values { rows, .. }) => {
1766 let empty = Row {
1767 columns: vec![],
1768 values: vec![],
1769 };
1770 let out: Vec<Vec<String>> = rows
1771 .iter()
1772 .map(|row| row.iter().map(|e| eval_expr(e, &empty).display()).collect())
1773 .collect();
1774 return Ok(QueryOutput {
1775 columns: vec![],
1776 rows: out,
1777 });
1778 }
1779 _ => return Err(MqdbError::SqlExec("unsupported query type".into())),
1780 };
1781
1782 let where_expr = select.selection.as_ref();
1784 let hint = where_expr
1785 .map(analyze_where_for_index)
1786 .unwrap_or(IndexHint::FullScan);
1787 let zone_filter =
1790 where_expr.filter(|_| select.from.len() == 1 && select.from[0].joins.is_empty());
1791 let mut rows = self.materialise_from_with_hint(&select.from, &hint, zone_filter)?;
1792
1793 if let Some(where_expr) = &select.selection {
1795 let resolved = self.resolve_subqueries(where_expr)?;
1796 rows.retain(|row| eval_expr(&resolved, row).is_truthy());
1797 }
1798
1799 let limit_expr = query.limit_clause.as_ref().and_then(|lc| match lc {
1801 LimitClause::LimitOffset { limit, .. } => limit.clone(),
1802 LimitClause::OffsetCommaLimit { limit, .. } => Some(limit.clone()),
1803 });
1804
1805 self.project_and_aggregate(select, rows, &query.order_by, limit_expr.as_ref())
1806 }
1807
1808 fn resolve_subqueries(&self, expr: &Expr) -> Result<Expr, MqdbError> {
1809 match expr {
1810 Expr::BinaryOp { left, op, right } => Ok(Expr::BinaryOp {
1811 left: Box::new(self.resolve_subqueries(left)?),
1812 op: op.clone(),
1813 right: Box::new(self.resolve_subqueries(right)?),
1814 }),
1815 Expr::Subquery(q) => {
1816 let out = self.exec_query(q)?;
1817 let val = out
1818 .rows
1819 .first()
1820 .and_then(|r| r.first())
1821 .map(|s| {
1822 if let Ok(n) = s.parse::<i64>() {
1823 Expr::Value(SqlValue::Number(n.to_string(), false).with_empty_span())
1824 } else {
1825 Expr::Value(SqlValue::SingleQuotedString(s.clone()).with_empty_span())
1826 }
1827 })
1828 .unwrap_or(Expr::Value(SqlValue::Null.with_empty_span()));
1829 Ok(val)
1830 }
1831 Expr::Nested(inner) => Ok(Expr::Nested(Box::new(self.resolve_subqueries(inner)?))),
1832 Expr::Function(f) => {
1833 let new_args = match &f.args {
1834 FunctionArguments::List(al) => {
1835 let resolved: Result<Vec<_>, _> = al
1836 .args
1837 .iter()
1838 .map(|a| match a {
1839 FunctionArg::Unnamed(FunctionArgExpr::Expr(e)) => {
1840 Ok::<FunctionArg, MqdbError>(FunctionArg::Unnamed(
1841 FunctionArgExpr::Expr(self.resolve_subqueries(e)?),
1842 ))
1843 }
1844 _ => Ok(a.clone()),
1845 })
1846 .collect();
1847 FunctionArguments::List(sqlparser::ast::FunctionArgumentList {
1848 args: resolved?,
1849 ..al.clone()
1850 })
1851 }
1852 other => other.clone(),
1853 };
1854 Ok(Expr::Function(Function {
1855 args: new_args,
1856 ..f.clone()
1857 }))
1858 }
1859 other => Ok(other.clone()),
1860 }
1861 }
1862
1863 fn materialise_from_with_hint(
1864 &self,
1865 from: &[sqlparser::ast::TableWithJoins],
1866 hint: &IndexHint,
1867 zone_filter: Option<&Expr>,
1868 ) -> Result<Vec<Row>, MqdbError> {
1869 if from.is_empty() {
1870 return Ok(vec![Row {
1871 columns: vec![],
1872 values: vec![],
1873 }]);
1874 }
1875 let mut rows = self.table_rows_with_hint(&from[0].relation, hint, zone_filter)?;
1876 for join in &from[0].joins {
1877 let right = self.table_rows_with_hint(&join.relation, &IndexHint::FullScan, None)?;
1879 match &join.join_operator {
1880 JoinOperator::Inner(JoinConstraint::On(on))
1881 | JoinOperator::Join(JoinConstraint::On(on))
1882 | JoinOperator::Left(JoinConstraint::On(on))
1883 | JoinOperator::LeftOuter(JoinConstraint::On(on)) => {
1884 let resolved = self.resolve_subqueries(on)?;
1885 let left_cols = rows.first().map(|r| r.columns.clone()).unwrap_or_default();
1886 let right_cols = right.first().map(|r| r.columns.clone()).unwrap_or_default();
1887 rows = match find_equi_join_exprs(&resolved, &left_cols, &right_cols) {
1888 Some((left_key, right_key)) => {
1889 hash_equi_join(rows, right, left_key, right_key, &resolved)
1890 }
1891 None => {
1892 let mut combined = cross_join(rows, right);
1893 combined.retain(|row| eval_expr(&resolved, row).is_truthy());
1894 combined
1895 }
1896 };
1897 }
1898 _ => {
1899 rows = cross_join(rows, right);
1900 }
1901 }
1902 }
1903 for twj in from.iter().skip(1) {
1904 let right = self.table_rows_with_hint(&twj.relation, &IndexHint::FullScan, None)?;
1905 rows = cross_join(rows, right);
1906 for join in &twj.joins {
1907 let right2 =
1908 self.table_rows_with_hint(&join.relation, &IndexHint::FullScan, None)?;
1909 rows = cross_join(rows, right2);
1910 }
1911 }
1912 Ok(rows)
1913 }
1914
1915 fn table_rows_with_hint(
1916 &self,
1917 factor: &TableFactor,
1918 hint: &IndexHint,
1919 zone_filter: Option<&Expr>,
1920 ) -> Result<Vec<Row>, MqdbError> {
1921 let (table_name, alias) = match factor {
1922 TableFactor::Table { name, alias, .. } => {
1923 let n = name.0.last().map(ident_value).unwrap_or("").to_lowercase();
1924 let a = alias.as_ref().map(|a| a.name.value.clone());
1925 (n, a)
1926 }
1927 _ => return Err(MqdbError::SqlExec("unsupported FROM clause".into())),
1928 };
1929
1930 match table_name.as_str() {
1931 "blocks" => {
1932 let prefix = alias.as_deref().unwrap_or("blocks");
1933 let mut rows = Vec::new();
1934 let mut global_idx: u32 = 0;
1935
1936 for (doc, doc_idx) in self.documents_with_indexes() {
1937 if let Some(we) = zone_filter
1940 && zone_map_skip(&doc.zone_maps, we)
1941 {
1942 global_idx += doc.blocks.len() as u32;
1943 continue;
1944 }
1945 if let Some(local_indices) = hint.resolve(doc_idx) {
1947 for local_i in local_indices {
1949 if let Some(block) = doc.blocks.get(local_i as usize) {
1950 let block_global_idx = global_idx + local_i;
1951 rows.push(qualify_row(
1952 block_to_row(doc.id, block, block_global_idx),
1953 prefix,
1954 ));
1955 }
1956 }
1957 } else {
1958 for (i, block) in doc.blocks.iter().enumerate() {
1960 rows.push(qualify_row(
1961 block_to_row(doc.id, block, global_idx + i as u32),
1962 prefix,
1963 ));
1964 }
1965 }
1966 global_idx += doc.blocks.len() as u32;
1967 }
1968 Ok(rows)
1969 }
1970 "documents" => {
1971 let prefix = alias.as_deref().unwrap_or("documents");
1972 Ok(self
1973 .store
1974 .documents()
1975 .iter()
1976 .map(|doc| qualify_row(doc_to_row(doc), prefix))
1977 .collect())
1978 }
1979 other => {
1980 let guard = self.store.custom_tables.read().unwrap();
1981 if let Some(state) = guard.get(other) {
1982 let prefix = alias.as_deref().unwrap_or(other);
1983 let rows = state
1984 .rows
1985 .iter()
1986 .map(|row_vals| {
1987 qualify_row(
1988 Row {
1989 columns: state.columns.clone(),
1990 values: row_vals
1991 .iter()
1992 .map(|v| Value::Str(v.clone()))
1993 .collect(),
1994 },
1995 prefix,
1996 )
1997 })
1998 .collect();
1999 return Ok(rows);
2000 }
2001 drop(guard);
2002 Err(MqdbError::SqlExec(format!("unknown table: {other}")))
2003 }
2004 }
2005 }
2006
2007 fn project_and_aggregate(
2008 &self,
2009 select: &Select,
2010 rows: Vec<Row>,
2011 order_by: &Option<sqlparser::ast::OrderBy>,
2012 limit: Option<&Expr>,
2013 ) -> Result<QueryOutput, MqdbError> {
2014 let group_by_exprs: Vec<Expr> = match &select.group_by {
2015 GroupByExpr::Expressions(exprs, _) => exprs.clone(),
2016 _ => vec![],
2017 };
2018 let is_agg = has_aggregate(&select.projection);
2019
2020 if is_agg || !group_by_exprs.is_empty() {
2021 return self.aggregate(select, rows, limit, &group_by_exprs);
2022 }
2023
2024 let columns = projection_columns(&select.projection, rows.first());
2026 let mut result: Vec<(Row, Vec<String>)> = rows
2027 .into_iter()
2028 .map(|row| {
2029 let cells = project_row(&select.projection, &row);
2030 (row, cells)
2031 })
2032 .collect();
2033
2034 if let Some(ob) = order_by {
2036 apply_order_by(&mut result, &ob.kind);
2037 }
2038
2039 let result: Vec<Vec<String>> = if select.distinct.is_some() {
2041 let mut seen = std::collections::HashSet::new();
2042 result
2043 .into_iter()
2044 .filter_map(|(_, cells)| {
2045 if seen.insert(cells.clone()) {
2046 Some(cells)
2047 } else {
2048 None
2049 }
2050 })
2051 .collect()
2052 } else {
2053 result.into_iter().map(|(_, cells)| cells).collect()
2054 };
2055
2056 Ok(QueryOutput {
2057 columns,
2058 rows: apply_limit(result, limit),
2059 })
2060 }
2061
2062 fn aggregate(
2063 &self,
2064 select: &Select,
2065 rows: Vec<Row>,
2066 limit: Option<&Expr>,
2067 group_by_exprs: &[Expr],
2068 ) -> Result<QueryOutput, MqdbError> {
2069 let columns: Vec<String> = select
2070 .projection
2071 .iter()
2072 .enumerate()
2073 .map(|(i, item)| projection_col_name(item, i))
2074 .collect();
2075
2076 let mut groups: Vec<(Vec<Value>, Vec<&Row>)> = Vec::new();
2078 let mut key_index: HashMap<Vec<String>, usize> = HashMap::new();
2079
2080 let owned: Vec<Row> = rows;
2082
2083 if group_by_exprs.is_empty() {
2084 let all: Vec<&Row> = owned.iter().collect();
2086 let out_row = eval_agg_row(&select.projection, group_by_exprs, &[], &all);
2087 return Ok(QueryOutput {
2088 columns,
2089 rows: apply_limit(vec![out_row], limit),
2090 });
2091 }
2092
2093 for row in &owned {
2094 let key: Vec<Value> = group_by_exprs.iter().map(|e| eval_expr(e, row)).collect();
2095 let key_str: Vec<String> = key.iter().map(|v| v.display()).collect();
2096 let idx = key_index.entry(key_str.clone()).or_insert_with(|| {
2097 groups.push((key, Vec::new()));
2098 groups.len() - 1
2099 });
2100 groups[*idx].1.push(row);
2101 }
2102
2103 let out_rows: Vec<Vec<String>> = groups
2104 .iter()
2105 .map(|(key_vals, group_rows)| {
2106 eval_agg_row(&select.projection, group_by_exprs, key_vals, group_rows)
2107 })
2108 .collect();
2109
2110 Ok(QueryOutput {
2111 columns,
2112 rows: apply_limit(out_rows, limit),
2113 })
2114 }
2115}
2116
2117fn projection_columns(projection: &[SelectItem], first_row: Option<&Row>) -> Vec<String> {
2118 if projection.len() == 1 && matches!(projection[0], SelectItem::Wildcard(_)) {
2119 return first_row
2120 .map(|r| {
2121 r.columns
2122 .iter()
2123 .map(|c| c.split('.').next_back().unwrap_or(c).to_string())
2124 .collect()
2125 })
2126 .unwrap_or_default();
2127 }
2128 projection
2129 .iter()
2130 .enumerate()
2131 .map(|(i, item)| projection_col_name(item, i))
2132 .collect()
2133}
2134
2135fn projection_col_name(item: &SelectItem, idx: usize) -> String {
2136 match item {
2137 SelectItem::UnnamedExpr(Expr::Identifier(i)) => i.value.clone(),
2138 SelectItem::UnnamedExpr(Expr::CompoundIdentifier(parts)) => parts
2139 .last()
2140 .map(|i| i.value.as_str())
2141 .unwrap_or("")
2142 .to_string(),
2143 SelectItem::UnnamedExpr(Expr::Function(f)) => {
2144 f.name.0.last().map(ident_value).unwrap_or("").to_string()
2145 }
2146 SelectItem::ExprWithAlias { alias, .. } => alias.value.clone(),
2147 SelectItem::Wildcard(_) => "*".to_string(),
2148 _ => format!("col{}", idx),
2149 }
2150}
2151
2152fn project_row(projection: &[SelectItem], row: &Row) -> Vec<String> {
2153 if projection.len() == 1 && matches!(projection[0], SelectItem::Wildcard(_)) {
2154 return row.values.iter().map(|v| v.display()).collect();
2155 }
2156 projection
2157 .iter()
2158 .map(|item| match item {
2159 SelectItem::UnnamedExpr(e) | SelectItem::ExprWithAlias { expr: e, .. } => {
2160 eval_expr(e, row).display()
2161 }
2162 SelectItem::ExprWithAliases { expr: e, .. } => eval_expr(e, row).display(),
2163 SelectItem::Wildcard(_) => row
2164 .values
2165 .iter()
2166 .map(|v| v.display())
2167 .collect::<Vec<_>>()
2168 .join(","),
2169 SelectItem::QualifiedWildcard(kind, _) => {
2170 let prefix = match kind {
2171 sqlparser::ast::SelectItemQualifiedWildcardKind::ObjectName(name) => {
2172 name.0.last().map(ident_value).unwrap_or("").to_string()
2173 }
2174 _ => String::new(),
2175 };
2176 row.columns
2177 .iter()
2178 .zip(row.values.iter())
2179 .filter(|(c, _)| c.starts_with(&format!("{}.", prefix)))
2180 .map(|(_, v)| v.display())
2181 .collect::<Vec<_>>()
2182 .join(",")
2183 }
2184 })
2185 .collect()
2186}
2187
2188fn has_aggregate(projection: &[SelectItem]) -> bool {
2189 projection.iter().any(|item| match item {
2190 SelectItem::UnnamedExpr(e) | SelectItem::ExprWithAlias { expr: e, .. } => is_agg_expr(e),
2191 _ => false,
2192 })
2193}
2194
2195fn is_agg_expr(expr: &Expr) -> bool {
2196 matches!(expr, Expr::Function(f) if {
2197 let name = f.name.0.last().map(ident_value).unwrap_or("").to_lowercase();
2198 is_aggregate_name(&name)
2199 })
2200}
2201
2202fn is_aggregate_name(name: &str) -> bool {
2203 matches!(
2204 name,
2205 "count" | "sum" | "min" | "max" | "avg" | "group_concat" | "string_agg"
2206 )
2207}
2208
2209fn eval_agg_row(
2210 projection: &[SelectItem],
2211 group_by_exprs: &[Expr],
2212 key_vals: &[Value],
2213 group_rows: &[&Row],
2214) -> Vec<String> {
2215 projection
2216 .iter()
2217 .map(|item| {
2218 let expr = match item {
2219 SelectItem::UnnamedExpr(e) | SelectItem::ExprWithAlias { expr: e, .. } => e,
2220 _ => return String::new(),
2221 };
2222 match expr {
2223 Expr::Function(f) => {
2224 let name = f
2225 .name
2226 .0
2227 .last()
2228 .map(ident_value)
2229 .unwrap_or("")
2230 .to_lowercase();
2231 match name.as_str() {
2232 "count" if is_distinct(f) => {
2233 let mut seen: Vec<Value> = Vec::new();
2234 for r in group_rows {
2235 let v = agg_arg(f, r);
2236 if !matches!(v, Value::Null) && !seen.contains(&v) {
2237 seen.push(v);
2238 }
2239 }
2240 seen.len().to_string()
2241 }
2242 "count" => group_rows.len().to_string(),
2243 "group_concat" | "string_agg" => {
2244 let sep = agg_separator(f);
2245 group_rows
2246 .iter()
2247 .map(|r| agg_arg(f, r))
2248 .filter(|v| !matches!(v, Value::Null))
2249 .map(|v| v.display())
2250 .collect::<Vec<_>>()
2251 .join(&sep)
2252 }
2253 "sum" => {
2254 let sum: f64 = group_rows
2255 .iter()
2256 .filter_map(|r| agg_arg(f, r).as_f64())
2257 .sum();
2258 sum.to_string()
2259 }
2260 "min" => group_rows
2261 .iter()
2262 .map(|r| agg_arg(f, r))
2263 .min_by(|a, b| a.cmp_val(b).unwrap_or(std::cmp::Ordering::Equal))
2264 .map(|v| v.display())
2265 .unwrap_or_else(|| "NULL".into()),
2266 "max" => group_rows
2267 .iter()
2268 .map(|r| agg_arg(f, r))
2269 .max_by(|a, b| a.cmp_val(b).unwrap_or(std::cmp::Ordering::Equal))
2270 .map(|v| v.display())
2271 .unwrap_or_else(|| "NULL".into()),
2272 "avg" => {
2273 let vals: Vec<f64> = group_rows
2274 .iter()
2275 .filter_map(|r| agg_arg(f, r).as_f64())
2276 .collect();
2277 if vals.is_empty() {
2278 "NULL".into()
2279 } else {
2280 (vals.iter().sum::<f64>() / vals.len() as f64).to_string()
2281 }
2282 }
2283 _ => String::new(),
2284 }
2285 }
2286 other => {
2287 if let Some(ki) = group_by_exprs
2288 .iter()
2289 .position(|e| expr_structurally_eq(e, other))
2290 {
2291 key_vals.get(ki).map(|v| v.display()).unwrap_or_default()
2292 } else {
2293 group_rows
2294 .first()
2295 .map(|r| eval_expr(other, r).display())
2296 .unwrap_or_default()
2297 }
2298 }
2299 }
2300 })
2301 .collect()
2302}
2303
2304fn agg_arg(f: &Function, row: &Row) -> Value {
2305 match &f.args {
2306 FunctionArguments::List(al) => al.args.iter().find_map(|a| match a {
2307 FunctionArg::Unnamed(FunctionArgExpr::Expr(e)) => Some(eval_expr(e, row)),
2308 FunctionArg::Unnamed(FunctionArgExpr::Wildcard) => Some(Value::Int(1)),
2309 _ => None,
2310 }),
2311 _ => None,
2312 }
2313 .unwrap_or(Value::Null)
2314}
2315
2316fn is_distinct(f: &Function) -> bool {
2317 matches!(
2318 &f.args,
2319 FunctionArguments::List(al) if al.duplicate_treatment == Some(DuplicateTreatment::Distinct)
2320 )
2321}
2322
2323fn agg_separator(f: &Function) -> String {
2327 if let FunctionArguments::List(al) = &f.args
2328 && let Some(FunctionArg::Unnamed(FunctionArgExpr::Expr(Expr::Value(v)))) = al.args.get(1)
2329 && let Value::Str(s) = eval_sql_value(&v.value)
2330 {
2331 return s;
2332 }
2333 ",".to_string()
2334}
2335
2336fn expr_structurally_eq(a: &Expr, b: &Expr) -> bool {
2337 format!("{:?}", a) == format!("{:?}", b)
2338}
2339
2340fn apply_order_by(rows: &mut [(Row, Vec<String>)], kind: &OrderByKind) {
2341 let exprs: &[OrderByExpr] = match kind {
2342 OrderByKind::Expressions(exprs) => exprs,
2343 _ => return,
2344 };
2345 rows.sort_by(|(ra, _), (rb, _)| {
2346 for ob in exprs {
2347 let va = eval_expr(&ob.expr, ra);
2348 let vb = eval_expr(&ob.expr, rb);
2349 let ord = va.cmp_val(&vb).unwrap_or(std::cmp::Ordering::Equal);
2350 let ord = if ob.options.asc == Some(false) {
2352 ord.reverse()
2353 } else {
2354 ord
2355 };
2356 if ord != std::cmp::Ordering::Equal {
2357 return ord;
2358 }
2359 }
2360 std::cmp::Ordering::Equal
2361 });
2362}
2363
2364fn apply_limit(mut rows: Vec<Vec<String>>, limit: Option<&Expr>) -> Vec<Vec<String>> {
2365 if let Some(lim) = limit {
2366 let dummy = Row {
2367 columns: vec![],
2368 values: vec![],
2369 };
2370 if let Value::Int(n) = eval_expr(lim, &dummy) {
2371 rows.truncate(n as usize);
2372 }
2373 }
2374 rows
2375}
2376
2377fn flatten_and_conjuncts(expr: &Expr) -> Vec<&Expr> {
2380 match expr {
2381 Expr::BinaryOp {
2382 left,
2383 op: BinaryOperator::And,
2384 right,
2385 } => {
2386 let mut out = flatten_and_conjuncts(left);
2387 out.extend(flatten_and_conjuncts(right));
2388 out
2389 }
2390 Expr::Nested(inner) => flatten_and_conjuncts(inner),
2391 other => vec![other],
2392 }
2393}
2394
2395fn schema_has_short_col(schema: &[String], short: &str) -> bool {
2398 schema.iter().any(|c| {
2399 let cl = c.to_lowercase();
2400 cl == short || cl.split('.').next_back().unwrap_or(&cl) == short
2401 })
2402}
2403
2404fn find_equi_join_exprs<'a>(
2409 on: &'a Expr,
2410 left_cols: &[String],
2411 right_cols: &[String],
2412) -> Option<(&'a Expr, &'a Expr)> {
2413 for conjunct in flatten_and_conjuncts(on) {
2414 let Expr::BinaryOp {
2415 left,
2416 op: BinaryOperator::Eq,
2417 right,
2418 } = conjunct
2419 else {
2420 continue;
2421 };
2422 let (Some(lname), Some(rname)) = (expr_col_name(left), expr_col_name(right)) else {
2423 continue;
2424 };
2425 if schema_has_short_col(left_cols, &lname) && schema_has_short_col(right_cols, &rname) {
2426 return Some((left, right));
2427 }
2428 if schema_has_short_col(right_cols, &lname) && schema_has_short_col(left_cols, &rname) {
2429 return Some((right, left));
2430 }
2431 }
2432 None
2433}
2434
2435fn zone_map_skip(zone_maps: &ZoneMaps, where_expr: &Expr) -> bool {
2440 let mut eq_block_type: Option<BlockType> = None;
2441 let mut eq_content: Option<String> = None;
2442 let mut eq_lang: Option<String> = None;
2443 let mut eq_depth: Option<u8> = None;
2444
2445 for conjunct in flatten_and_conjuncts(where_expr) {
2446 let Expr::BinaryOp {
2447 left,
2448 op: BinaryOperator::Eq,
2449 right,
2450 } = conjunct
2451 else {
2452 continue;
2453 };
2454 let col = expr_col_name(left).or_else(|| expr_col_name(right));
2455 let val = expr_str_val(right).or_else(|| expr_str_val(left));
2456 let int_val = expr_int_val(right).or_else(|| expr_int_val(left));
2457
2458 match col.as_deref() {
2459 Some("block_type") => {
2460 if let Some(s) = val.as_deref()
2461 && let Some(bt) = BlockType::from_str(s)
2462 {
2463 eq_block_type = Some(bt);
2464 }
2465 }
2466 Some("content") => eq_content = val,
2467 Some("lang") => {
2470 if let Some(s) = val
2471 && !s.is_empty()
2472 {
2473 eq_lang = Some(s);
2474 }
2475 }
2476 Some("depth") => {
2479 if let Some(n) = int_val
2480 && let Ok(n) = u8::try_from(n)
2481 && n > 0
2482 {
2483 eq_depth = Some(n);
2484 }
2485 }
2486 _ => {}
2487 }
2488 }
2489
2490 if let Some(lang) = &eq_lang
2491 && !zone_maps.code_languages.contains(lang)
2492 {
2493 return true;
2494 }
2495 if let Some(depth) = eq_depth
2496 && depth > zone_maps.max_heading_depth
2497 {
2498 return true;
2499 }
2500 if let Some(content) = &eq_content
2503 && eq_block_type == Some(BlockType::Heading)
2504 && !zone_maps
2505 .heading_contents
2506 .iter()
2507 .any(|h| h.eq_ignore_ascii_case(content))
2508 {
2509 return true;
2510 }
2511
2512 false
2513}
2514
2515fn analyze_where_for_index(expr: &Expr) -> IndexHint {
2531 match expr {
2532 Expr::BinaryOp {
2534 left,
2535 op: BinaryOperator::And,
2536 right,
2537 } => {
2538 let lh = analyze_where_for_index(left);
2539 let rh = analyze_where_for_index(right);
2540 pick_better_hint(lh, rh)
2541 }
2542 Expr::BinaryOp {
2544 left,
2545 op: BinaryOperator::Eq,
2546 right,
2547 } => {
2548 let col = expr_col_name(left).or_else(|| expr_col_name(right));
2549 let val = expr_str_val(right).or_else(|| expr_str_val(left));
2550 let int_val = expr_int_val(right).or_else(|| expr_int_val(left));
2551
2552 match col.as_deref() {
2553 Some("block_type") => {
2554 if let Some(s) = val
2555 && let Some(bt) = BlockType::from_str(&s)
2556 {
2557 return IndexHint::BlockType(vec![bt]);
2558 }
2559 IndexHint::FullScan
2560 }
2561 Some("pre") => {
2562 if let Some(n) = int_val {
2563 return IndexHint::PreExact(n as u32);
2564 }
2565 IndexHint::FullScan
2566 }
2567 Some("content") => {
2568 if let Some(s) = val {
2569 return IndexHint::ContentExact(s);
2570 }
2571 IndexHint::FullScan
2572 }
2573 Some("lang") => {
2574 if let Some(s) = val
2575 && !s.is_empty()
2576 {
2577 return IndexHint::LangExact(s);
2578 }
2579 IndexHint::FullScan
2580 }
2581 Some("depth") => {
2582 if let Some(n) = int_val {
2583 if n > 0 {
2585 return IndexHint::DepthExact(n as u8);
2586 }
2587 }
2588 IndexHint::FullScan
2589 }
2590 _ => IndexHint::FullScan,
2591 }
2592 }
2593 Expr::InList {
2595 expr,
2596 list,
2597 negated: false,
2598 } => {
2599 if expr_col_name(expr).as_deref() == Some("block_type") {
2600 let types: Vec<BlockType> = list
2601 .iter()
2602 .filter_map(expr_str_val)
2603 .filter_map(|s| BlockType::from_str(&s))
2604 .collect();
2605 if !types.is_empty() {
2606 return IndexHint::BlockType(types);
2607 }
2608 }
2609 IndexHint::FullScan
2610 }
2611 Expr::Between {
2613 expr,
2614 negated: false,
2615 low,
2616 high,
2617 } => {
2618 if expr_col_name(expr).as_deref() == Some("pre")
2619 && let (Some(lo), Some(hi)) = (expr_int_val(low), expr_int_val(high))
2620 {
2621 return IndexHint::PreRange(lo as u32, hi as u32);
2622 }
2623 IndexHint::FullScan
2624 }
2625 Expr::Nested(inner) => analyze_where_for_index(inner),
2626 _ => IndexHint::FullScan,
2627 }
2628}
2629
2630fn expr_col_name(expr: &Expr) -> Option<String> {
2632 match expr {
2633 Expr::Identifier(i) => Some(i.value.to_lowercase()),
2634 Expr::CompoundIdentifier(parts) => parts.last().map(|i| i.value.to_lowercase()),
2635 _ => None,
2636 }
2637}
2638
2639fn expr_str_val(expr: &Expr) -> Option<String> {
2640 match expr {
2641 Expr::Value(v) => match &v.value {
2642 SqlValue::SingleQuotedString(s) | SqlValue::DoubleQuotedString(s) => Some(s.clone()),
2643 _ => None,
2644 },
2645 _ => None,
2646 }
2647}
2648
2649fn expr_int_val(expr: &Expr) -> Option<i64> {
2650 match expr {
2651 Expr::Value(v) => match &v.value {
2652 SqlValue::Number(n, _) => n.parse::<i64>().ok(),
2653 _ => None,
2654 },
2655 _ => None,
2656 }
2657}
2658
2659fn pick_better_hint(a: IndexHint, b: IndexHint) -> IndexHint {
2661 match (&a, &b) {
2662 (IndexHint::FullScan, _) => b,
2663 (_, IndexHint::FullScan) => a,
2664 (IndexHint::BlockType(ta), IndexHint::BlockType(tb)) => {
2667 if ta.len() <= tb.len() {
2668 a
2669 } else {
2670 b
2671 }
2672 }
2673 (IndexHint::PreExact(_), _) => a,
2675 (_, IndexHint::PreExact(_)) => b,
2676 _ => a,
2677 }
2678}
2679
2680impl BlockType {
2681 fn from_str(s: &str) -> Option<Self> {
2682 match s {
2683 "heading" => Some(BlockType::Heading),
2684 "paragraph" => Some(BlockType::Paragraph),
2685 "code" => Some(BlockType::Code),
2686 "list" => Some(BlockType::List),
2687 "table_cell" => Some(BlockType::TableCell),
2688 "table_row" => Some(BlockType::TableRow),
2689 "table_align" => Some(BlockType::TableAlign),
2690 "blockquote" => Some(BlockType::Blockquote),
2691 "horizontal_rule" => Some(BlockType::HorizontalRule),
2692 "html" => Some(BlockType::Html),
2693 "yaml" => Some(BlockType::Yaml),
2694 "toml" => Some(BlockType::Toml),
2695 "math" => Some(BlockType::Math),
2696 "definition" => Some(BlockType::Definition),
2697 "footnote" => Some(BlockType::Footnote),
2698 _ => None,
2699 }
2700 }
2701}
2702
2703#[cfg(test)]
2704mod tests {
2705 use super::*;
2706 use crate::DocumentStore;
2707 use rstest::rstest;
2708
2709 fn make_store() -> DocumentStore {
2710 let mut s = DocumentStore::new();
2711 s.add_str(
2712 "# Doc\n\n## Architecture\n\nDetails\n\n```rust\nfn main(){}\n```\n\n## Other\n\nOther\n",
2713 )
2714 .unwrap();
2715 s
2716 }
2717
2718 fn make_multi_doc_store() -> DocumentStore {
2720 let mut s = DocumentStore::new();
2721 s.add_str("# A\n\n```rust\nfn a(){}\n```\n").unwrap();
2722 s.add_str("# B\n\nParagraph\n").unwrap();
2723 s.add_str("# C\n\n## C2\n\n### C3\n\n```rust\nfn c(){}\n```\n")
2724 .unwrap();
2725 s
2726 }
2727
2728 #[test]
2729 fn test_sql_select_all_blocks() {
2730 let store = make_store();
2731 let engine = SqlEngine::new(&store).unwrap();
2732 let out = engine
2733 .execute("SELECT block_type, content FROM blocks ORDER BY pre")
2734 .unwrap();
2735 assert!(!out.rows.is_empty());
2736 }
2737
2738 #[test]
2739 fn test_sql_heading_filter() {
2740 let store = make_store();
2741 let engine = SqlEngine::new(&store).unwrap();
2742 let out = engine
2743 .execute("SELECT content FROM blocks WHERE block_type = 'heading' ORDER BY pre")
2744 .unwrap();
2745 assert_eq!(out.rows.len(), 3);
2746 }
2747
2748 #[test]
2749 fn test_sql_under_function() {
2750 let store = make_store();
2751 let engine = SqlEngine::new(&store).unwrap();
2752 let out = engine
2753 .execute(
2754 "SELECT b.content FROM blocks b
2755 WHERE under(b.pre, b.post,
2756 (SELECT pre FROM blocks WHERE block_type='heading' AND content='Architecture'),
2757 (SELECT post FROM blocks WHERE block_type='heading' AND content='Architecture')
2758 )",
2759 )
2760 .unwrap();
2761 assert_eq!(out.rows.len(), 2);
2762 }
2763
2764 #[test]
2765 fn test_query_output_table() {
2766 let out = QueryOutput {
2767 columns: vec!["id".to_string(), "type".to_string()],
2768 rows: vec![
2769 vec!["1".to_string(), "heading".to_string()],
2770 vec!["2".to_string(), "paragraph".to_string()],
2771 ],
2772 };
2773 let table = out.to_table();
2774 assert!(table.contains("heading"));
2775 assert!(table.contains("paragraph"));
2776 assert!(table.contains("2 rows"));
2777 }
2778
2779 #[test]
2780 fn test_sql_count_aggregate() {
2781 let store = make_store();
2782 let engine = SqlEngine::new(&store).unwrap();
2783 let out = engine
2784 .execute("SELECT count(*) FROM blocks WHERE block_type = 'heading'")
2785 .unwrap();
2786 assert_eq!(out.rows.len(), 1);
2787 assert_eq!(out.rows[0][0], "3");
2788 }
2789
2790 #[test]
2791 fn test_sql_limit() {
2792 let store = make_store();
2793 let engine = SqlEngine::new(&store).unwrap();
2794 let out = engine
2795 .execute("SELECT content FROM blocks LIMIT 2")
2796 .unwrap();
2797 assert_eq!(out.rows.len(), 2);
2798 }
2799
2800 #[test]
2801 fn test_sql_like() {
2802 let store = make_store();
2803 let engine = SqlEngine::new(&store).unwrap();
2804 let out = engine
2805 .execute("SELECT content FROM blocks WHERE content LIKE '%chitect%'")
2806 .unwrap();
2807 assert!(!out.rows.is_empty());
2808 }
2809
2810 #[test]
2811 fn test_sql_order_by_desc() {
2812 let store = make_store();
2813 let engine = SqlEngine::new(&store).unwrap();
2814 let out = engine
2815 .execute("SELECT content FROM blocks ORDER BY pre DESC LIMIT 1")
2816 .unwrap();
2817 assert_eq!(out.rows.len(), 1);
2818 }
2819
2820 #[test]
2821 fn test_sql_engine_zero_copy() {
2822 let mut store = DocumentStore::new();
2823 for _ in 0..100 {
2824 store.add_str("# Heading\n\nParagraph text\n").unwrap();
2825 }
2826 let start = std::time::Instant::now();
2827 let _engine = SqlEngine::new(&store).unwrap();
2828 let elapsed = start.elapsed();
2829 assert!(
2830 elapsed.as_millis() < 1,
2831 "SqlEngine::new took {}ms — should be O(1)",
2832 elapsed.as_millis()
2833 );
2834 }
2835
2836 #[rstest]
2841 #[case("SELECT content FROM blocks WHERE block_type = 'heading'", 3)]
2842 #[case("SELECT content FROM blocks WHERE block_type = 'paragraph'", 2)]
2843 #[case("SELECT content FROM blocks WHERE block_type = 'code'", 1)]
2844 #[case("SELECT content FROM blocks WHERE block_type = 'list'", 0)]
2845 fn test_sql_where_block_type_param(#[case] sql: &str, #[case] expected: usize) {
2846 let store = make_store();
2847 let engine = SqlEngine::new(&store).unwrap();
2848 assert_eq!(engine.execute(sql).unwrap().rows.len(), expected);
2849 }
2850
2851 #[rstest]
2852 #[case("SELECT content FROM blocks WHERE content LIKE '%Doc%'", 1)]
2853 #[case("SELECT content FROM blocks WHERE content LIKE '%chitect%'", 1)]
2854 #[case("SELECT content FROM blocks WHERE content LIKE '%Other%'", 2)]
2855 #[case("SELECT content FROM blocks WHERE content LIKE '%Details%'", 1)]
2856 #[case("SELECT content FROM blocks WHERE content LIKE '%nonexistent%'", 0)]
2857 fn test_sql_like_pattern_param(#[case] sql: &str, #[case] expected: usize) {
2858 let store = make_store();
2859 let engine = SqlEngine::new(&store).unwrap();
2860 assert_eq!(engine.execute(sql).unwrap().rows.len(), expected);
2861 }
2862
2863 #[rstest]
2864 #[case("SELECT content FROM blocks LIMIT 1", 1)]
2865 #[case("SELECT content FROM blocks LIMIT 3", 3)]
2866 #[case("SELECT content FROM blocks LIMIT 5", 5)]
2867 #[case("SELECT content FROM blocks LIMIT 1000", 6)]
2868 fn test_sql_limit_row_count_param(#[case] sql: &str, #[case] expected: usize) {
2869 let store = make_store();
2870 let engine = SqlEngine::new(&store).unwrap();
2871 assert_eq!(engine.execute(sql).unwrap().rows.len(), expected);
2872 }
2873
2874 #[rstest]
2875 #[case("SELECT count(*) FROM blocks", "6")]
2876 #[case("SELECT count(*) FROM blocks WHERE block_type = 'heading'", "3")]
2877 #[case("SELECT count(*) FROM blocks WHERE block_type = 'code'", "1")]
2878 fn test_sql_count_aggregate_param(#[case] sql: &str, #[case] expected: &str) {
2879 let store = make_store();
2880 let engine = SqlEngine::new(&store).unwrap();
2881 let out = engine.execute(sql).unwrap();
2882 assert_eq!(out.rows.len(), 1);
2883 assert_eq!(out.rows[0][0], expected);
2884 }
2885
2886 #[test]
2888 fn test_sql_depth_zero_returns_non_headings() {
2889 let store = make_store();
2890 let engine = SqlEngine::new(&store).unwrap();
2891 let out = engine
2892 .execute("SELECT content FROM blocks WHERE depth = 0")
2893 .unwrap();
2894 assert_eq!(out.rows.len(), 3, "depth=0 must return non-heading blocks");
2896 }
2897
2898 #[test]
2900 fn test_sql_empty_lang_returns_non_code_blocks() {
2901 let store = make_store();
2902 let engine = SqlEngine::new(&store).unwrap();
2903 let out = engine
2904 .execute("SELECT block_type FROM blocks WHERE lang = ''")
2905 .unwrap();
2906 assert_eq!(out.rows.len(), 5, "lang='' must return non-code blocks");
2908 }
2909
2910 #[test]
2912 fn test_to_table_newline_in_cell() {
2913 let out = QueryOutput {
2914 columns: vec!["content".to_string()],
2915 rows: vec![
2916 vec!["line one\nline two".to_string()],
2917 vec!["plain".to_string()],
2918 ],
2919 };
2920 let table = out.to_table();
2921 let bar_lines: Vec<&str> = table.lines().filter(|l| l.starts_with('│')).collect();
2923 assert_eq!(
2924 bar_lines.len(),
2925 3,
2926 "newline in cell must not produce extra table rows"
2927 );
2928 assert!(bar_lines[1].contains("line one line two"));
2930 }
2931
2932 #[test]
2934 fn test_custom_table_query() {
2935 let mut store = DocumentStore::new();
2936 store.register_table(
2937 "kv",
2938 vec!["key".to_string(), "value".to_string()],
2939 vec![
2940 vec!["foo".to_string(), "bar".to_string()],
2941 vec!["hello".to_string(), "world".to_string()],
2942 ],
2943 );
2944 let engine = SqlEngine::new(&store).unwrap();
2945 let out = engine
2946 .execute("SELECT key, value FROM kv WHERE key = 'hello'")
2947 .unwrap();
2948 assert_eq!(out.rows.len(), 1);
2949 assert_eq!(out.rows[0][1], "world");
2950 }
2951
2952 #[test]
2954 fn test_ddl_create_insert_select() {
2955 let store = DocumentStore::new();
2956 let engine = SqlEngine::new(&store).unwrap();
2957
2958 engine
2960 .execute("CREATE TABLE notes (id TEXT, body TEXT)")
2961 .unwrap();
2962 engine
2964 .execute("INSERT INTO notes VALUES ('1', 'hello')")
2965 .unwrap();
2966 engine
2967 .execute("INSERT INTO notes VALUES ('2', 'world')")
2968 .unwrap();
2969 let out = engine
2971 .execute("SELECT body FROM notes WHERE id = '1'")
2972 .unwrap();
2973 assert_eq!(out.rows.len(), 1);
2974 assert_eq!(out.rows[0][0], "hello");
2975 let all = engine.execute("SELECT * FROM notes").unwrap();
2977 assert_eq!(all.rows.len(), 2);
2978 }
2979
2980 #[test]
2982 fn test_ddl_create_as_select() {
2983 let store = {
2984 let mut s = DocumentStore::new();
2985 s.add_str("# H1\n\n## H2\n\nParagraph\n").unwrap();
2986 s
2987 };
2988 let engine = SqlEngine::new(&store).unwrap();
2989 engine
2990 .execute(
2991 "CREATE TABLE headings AS \
2992 SELECT block_type, content FROM blocks WHERE block_type = 'heading'",
2993 )
2994 .unwrap();
2995 let out = engine.execute("SELECT content FROM headings").unwrap();
2996 assert_eq!(out.rows.len(), 2);
2997 }
2998
2999 #[test]
3001 fn test_ddl_drop_table() {
3002 let store = DocumentStore::new();
3003 let engine = SqlEngine::new(&store).unwrap();
3004 engine.execute("CREATE TABLE tmp (x TEXT)").unwrap();
3005 engine.execute("DROP TABLE tmp").unwrap();
3006 let err = engine.execute("SELECT * FROM tmp").unwrap_err();
3007 assert!(err.to_string().contains("unknown table"));
3008 }
3009
3010 #[test]
3012 fn test_ddl_drop_if_exists() {
3013 let store = DocumentStore::new();
3014 let engine = SqlEngine::new(&store).unwrap();
3015 engine
3016 .execute("DROP TABLE IF EXISTS no_such_table")
3017 .unwrap();
3018 }
3019
3020 #[test]
3022 fn test_desc_builtin() {
3023 let store = DocumentStore::new();
3024 let engine = SqlEngine::new(&store).unwrap();
3025 let out = engine.execute("DESC blocks").unwrap();
3026 assert_eq!(out.columns, vec!["column", "type"]);
3027 assert!(out.rows.iter().any(|r| r[0] == "block_type"));
3028 assert!(out.rows.iter().any(|r| r[0] == "content"));
3029 }
3030
3031 #[test]
3033 fn test_desc_custom() {
3034 let store = DocumentStore::new();
3035 let engine = SqlEngine::new(&store).unwrap();
3036 engine
3037 .execute("CREATE TABLE meta (k TEXT, v TEXT)")
3038 .unwrap();
3039 let out = engine.execute("DESC meta").unwrap();
3040 assert_eq!(out.rows.len(), 2);
3041 assert_eq!(out.rows[0][0], "k");
3042 assert_eq!(out.rows[1][0], "v");
3043 }
3044
3045 #[test]
3047 fn test_show_tables() {
3048 let store = DocumentStore::new();
3049 let engine = SqlEngine::new(&store).unwrap();
3050 engine.execute("CREATE TABLE extra (a TEXT)").unwrap();
3051 let out = engine.execute("SHOW TABLES").unwrap();
3052 let names: Vec<&str> = out.rows.iter().map(|r| r[0].as_str()).collect();
3053 assert!(names.contains(&"blocks"));
3054 assert!(names.contains(&"documents"));
3055 assert!(names.contains(&"extra"));
3056 }
3057
3058 #[test]
3060 fn test_mq_scalar_function() {
3061 let store = make_store();
3062 let engine = SqlEngine::new(&store).unwrap();
3063 let out = engine
3064 .execute(
3065 "SELECT mq('.h1 | to_text', '# Hello\n\nWorld\n') AS title FROM blocks LIMIT 1",
3066 )
3067 .unwrap();
3068 assert_eq!(out.rows.len(), 1);
3069 assert_eq!(out.rows[0][0], "Hello");
3070 }
3071
3072 #[test]
3074 fn test_mq_scalar_null_on_no_match() {
3075 let store = make_store();
3076 let engine = SqlEngine::new(&store).unwrap();
3077 let out = engine
3078 .execute("SELECT mq('.h1', '## No h1 here\n') FROM blocks LIMIT 1")
3079 .unwrap();
3080 assert_eq!(out.rows.len(), 1);
3081 assert_eq!(out.rows[0][0], "NULL");
3082 }
3083
3084 fn eval_one(sql: &str) -> String {
3085 let store = DocumentStore::new();
3086 let engine = SqlEngine::new(&store).unwrap();
3087 engine.execute(sql).unwrap().rows[0][0].clone()
3088 }
3089
3090 #[rstest]
3091 #[case("SELECT lower('Hello')", "hello")]
3093 #[case("SELECT upper('Hello')", "HELLO")]
3094 #[case("SELECT length('héllo')", "5")]
3095 #[case("SELECT trim(' hi ')", "hi")]
3096 #[case("SELECT ltrim(' hi ')", "hi ")]
3097 #[case("SELECT rtrim(' hi ')", " hi")]
3098 #[case("SELECT trim(LEADING 'x' FROM 'xxhixx')", "hixx")]
3099 #[case("SELECT trim(TRAILING 'x' FROM 'xxhixx')", "xxhi")]
3100 #[case("SELECT trim('x' FROM 'xxhixx')", "hi")]
3101 #[case("SELECT concat('a', 'b', 'c')", "abc")]
3102 #[case("SELECT concat_ws('-', 'a', 'b', NULL, 'c')", "a-b-c")]
3103 #[case("SELECT replace('foobar', 'o', '0')", "f00bar")]
3104 #[case("SELECT left('hello', 3)", "hel")]
3105 #[case("SELECT right('hello', 3)", "llo")]
3106 #[case("SELECT lpad('7', 3, '0')", "007")]
3107 #[case("SELECT rpad('7', 3, '0')", "700")]
3108 #[case("SELECT reverse('hello')", "olleh")]
3109 #[case("SELECT repeat('ab', 3)", "ababab")]
3110 #[case("SELECT initcap('hello world')", "Hello World")]
3111 #[case("SELECT ascii('A')", "65")]
3112 #[case("SELECT chr(65)", "A")]
3113 #[case("SELECT instr('hello world', 'world')", "7")]
3114 #[case("SELECT position('world' in 'hello world')", "7")]
3115 #[case("SELECT split_part('a,b,c', ',', 2)", "b")]
3116 #[case("SELECT substring('hello world', 1, 5)", "hello")]
3117 #[case("SELECT substring('hello world' from 7)", "world")]
3118 #[case("SELECT substr('hello world', 7, 5)", "world")]
3119 #[case("SELECT abs(-5)", "5")]
3121 #[case("SELECT abs(-5.5)", "5.5")]
3122 #[case("SELECT round(3.456, 2)", "3.46")]
3123 #[case("SELECT round(3.5)", "4")]
3124 #[case("SELECT ceil(3.1)", "4")]
3125 #[case("SELECT floor(3.9)", "3")]
3126 #[case("SELECT trunc(3.789, 1)", "3.7")]
3127 #[case("SELECT mod(10, 3)", "1")]
3128 #[case("SELECT power(2, 10)", "1024")]
3129 #[case("SELECT sqrt(16)", "4")]
3130 #[case("SELECT sign(-3)", "-1")]
3131 #[case("SELECT greatest(3, 7, 2)", "7")]
3132 #[case("SELECT least(3, 7, 2)", "2")]
3133 #[case("SELECT coalesce(NULL, NULL, 'x')", "x")]
3135 #[case("SELECT ifnull(NULL, 'y')", "y")]
3136 #[case("SELECT nullif('a', 'a')", "NULL")]
3137 #[case("SELECT nullif('a', 'b')", "a")]
3138 #[case("SELECT typeof('x')", "text")]
3140 #[case("SELECT typeof(1)", "integer")]
3141 #[case(
3143 "SELECT CASE WHEN 1 = 2 THEN 'a' WHEN 1 = 1 THEN 'b' ELSE 'c' END",
3144 "b"
3145 )]
3146 #[case("SELECT CASE 2 WHEN 1 THEN 'a' WHEN 2 THEN 'b' ELSE 'c' END", "b")]
3147 #[case("SELECT CASE WHEN 1 = 2 THEN 'a' ELSE 'c' END", "c")]
3148 fn test_sql_scalar_functions(#[case] sql: &str, #[case] expected: &str) {
3149 assert_eq!(eval_one(sql), expected);
3150 }
3151
3152 #[test]
3153 fn test_sql_group_concat() {
3154 let store = make_store();
3155 let engine = SqlEngine::new(&store).unwrap();
3156 let out = engine
3157 .execute("SELECT group_concat(content) FROM blocks WHERE block_type = 'heading'")
3158 .unwrap();
3159 assert_eq!(out.rows[0][0], "Doc,Architecture,Other");
3160 }
3161
3162 #[test]
3163 fn test_sql_string_agg_custom_separator() {
3164 let store = make_store();
3165 let engine = SqlEngine::new(&store).unwrap();
3166 let out = engine
3167 .execute("SELECT string_agg(content, ' | ') FROM blocks WHERE block_type = 'heading'")
3168 .unwrap();
3169 assert_eq!(out.rows[0][0], "Doc | Architecture | Other");
3170 }
3171
3172 #[test]
3173 fn test_sql_count_distinct() {
3174 let store = make_store();
3175 let engine = SqlEngine::new(&store).unwrap();
3176 let out = engine
3177 .execute("SELECT count(DISTINCT block_type) FROM blocks")
3178 .unwrap();
3179 assert_eq!(out.rows[0][0], "3");
3180 }
3181
3182 #[test]
3184 fn test_sql_zone_map_skip_by_lang() {
3185 let store = make_multi_doc_store();
3186 let engine = SqlEngine::new(&store).unwrap();
3187 let out = engine
3188 .execute("SELECT content FROM blocks WHERE lang = 'rust' ORDER BY content")
3189 .unwrap();
3190 let contents: Vec<&str> = out.rows.iter().map(|r| r[0].as_str()).collect();
3191 assert_eq!(contents, vec!["fn a(){}", "fn c(){}"]);
3192 }
3193
3194 #[test]
3196 fn test_sql_zone_map_skip_by_depth() {
3197 let store = make_multi_doc_store();
3198 let engine = SqlEngine::new(&store).unwrap();
3199 let out = engine
3200 .execute("SELECT content FROM blocks WHERE depth = 3")
3201 .unwrap();
3202 assert_eq!(out.rows.len(), 1);
3203 assert_eq!(out.rows[0][0], "C3");
3204 }
3205
3206 #[test]
3208 fn test_sql_zone_map_skip_by_heading_content() {
3209 let store = make_multi_doc_store();
3210 let engine = SqlEngine::new(&store).unwrap();
3211 let out = engine
3212 .execute("SELECT content FROM blocks WHERE block_type = 'heading' AND content = 'B'")
3213 .unwrap();
3214 assert_eq!(out.rows.len(), 1);
3215 assert_eq!(out.rows[0][0], "B");
3216 }
3217
3218 #[test]
3220 fn test_sql_zone_map_no_skip_on_empty_lang() {
3221 let store = make_multi_doc_store();
3222 let engine = SqlEngine::new(&store).unwrap();
3223 let out = engine
3224 .execute("SELECT content FROM blocks WHERE lang = ''")
3225 .unwrap();
3226 let contents: Vec<&str> = out.rows.iter().map(|r| r[0].as_str()).collect();
3227 assert!(contents.contains(&"B"), "doc B must not be skipped");
3228 assert!(contents.contains(&"Paragraph"));
3229 }
3230
3231 #[test]
3233 fn test_sql_zone_map_skip_preserves_block_ids() {
3234 let store = make_multi_doc_store();
3235 let engine = SqlEngine::new(&store).unwrap();
3236 let full = engine.execute("SELECT id, content FROM blocks").unwrap();
3237 let filtered = engine
3238 .execute("SELECT id, content FROM blocks WHERE lang = 'rust'")
3239 .unwrap();
3240 assert_eq!(filtered.rows.len(), 2);
3241 for row in &filtered.rows {
3242 let same_id = full.rows.iter().find(|r| r[0] == row[0]).unwrap();
3243 assert_eq!(
3244 same_id[1], row[1],
3245 "id {} must reference the same block content in both queries",
3246 row[0]
3247 );
3248 }
3249 }
3250
3251 #[test]
3254 fn test_sql_zone_map_skip_disabled_for_joins() {
3255 let store = make_multi_doc_store();
3256 let engine = SqlEngine::new(&store).unwrap();
3257 let out = engine
3258 .execute(
3259 "SELECT h.content, c.content FROM blocks h
3260 JOIN blocks c ON c.document_id = h.document_id AND c.block_type = 'code'
3261 WHERE h.block_type = 'heading'",
3262 )
3263 .unwrap();
3264 let headings: Vec<&str> = out.rows.iter().map(|r| r[0].as_str()).collect();
3265 assert_eq!(headings, vec!["A", "C", "C2", "C3"]);
3266 }
3267}