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