1use regex::Regex;
4use std::sync::LazyLock;
5
6use crate::errors::MdqlError;
7pub use crate::query_ast::*;
8
9static KEYWORDS: &[&str] = &[
12 "SELECT", "FROM", "WHERE", "AND", "OR", "ORDER", "BY",
13 "ASC", "DESC", "LIMIT", "LIKE", "IN", "IS", "NOT", "NULL",
14 "JOIN", "LEFT", "ON", "AS", "GROUP", "HAVING",
15 "INSERT", "INTO", "VALUES", "UPDATE", "SET", "DELETE",
16 "ALTER", "TABLE", "RENAME", "FIELD", "TO", "DROP", "MERGE", "FIELDS",
17 "CASE", "WHEN", "THEN", "ELSE", "END",
18 "INTERVAL", "DAY", "DAYS", "CURRENT_DATE", "CURRENT_TIMESTAMP", "DATEDIFF",
19 "CREATE", "VIEW", "CASCADE", "RESTRICT",
20 "WITH",
21 "OVER", "PARTITION", "ROW_NUMBER", "RANK", "DENSE_RANK", "LAG", "LEAD",
22];
23
24static AGG_FUNCS: &[&str] = &["COUNT", "SUM", "AVG", "MIN", "MAX"];
25static WINDOW_FUNCS: &[&str] = &["ROW_NUMBER", "RANK", "DENSE_RANK", "LAG", "LEAD"];
26
27static TOKEN_RE: LazyLock<Regex> = LazyLock::new(|| {
28 Regex::new(
29 r#"(?x)
30 \s*(?:
31 (?P<backtick>`[^`]+`)
32 | (?P<string>'(?:[^'\\]|\\.)*')
33 | (?P<date>\d{4}-\d{1,2}-\d{1,2}
34 (?:[T\x20]\d{1,2}:\d{2}(?::\d{2})?(?:\.\d+)?(?:Z|[+-]\d{2}:?\d{2})?)?)
35 | (?P<number>\d+(?:\.\d+)?)
36 | (?P<op><=|>=|!=|[=<>,*()+\-/%])
37 | (?P<word>[A-Za-z_][A-Za-z0-9_./-]*)
38 )"#,
39 )
40 .unwrap()
41});
42
43static ISO_DATE_RE: LazyLock<Regex> = LazyLock::new(|| {
44 Regex::new(
45 r"^\d{4}-\d{2}-\d{2}(?:[T ]\d{2}:\d{2}(?::\d{2})?(?:\.\d+)?(?:Z|[+-]\d{2}:?\d{2})?)?$",
46 )
47 .unwrap()
48});
49
50fn date_literal(raw: &str) -> Result<SqlValue, MdqlError> {
54 if ISO_DATE_RE.is_match(raw) {
55 return Ok(SqlValue::String(raw.to_string()));
56 }
57 let padded: Vec<String> = raw
58 .splitn(3, '-')
59 .map(|part| {
60 let (num, rest) = part.split_at(
61 part.find(|c: char| !c.is_ascii_digit()).unwrap_or(part.len()),
62 );
63 format!("{:0>2}{}", num, rest)
64 })
65 .collect();
66 Err(MdqlError::QueryParse(format!(
67 "Ambiguous date literal '{}': dates compare as strings, so use the quoted zero-padded ISO form '{}'",
68 raw,
69 padded.join("-"),
70 )))
71}
72
73#[derive(Debug, Clone)]
74struct Token {
75 token_type: String,
76 value: String,
77 raw: String,
78}
79
80fn tokenize(sql: &str) -> Vec<Token> {
81 let mut tokens = Vec::new();
82 for caps in TOKEN_RE.captures_iter(sql) {
83 if let Some(m) = caps.name("backtick") {
84 let raw = m.as_str();
85 tokens.push(Token {
86 token_type: "ident".into(),
87 value: raw[1..raw.len() - 1].into(),
88 raw: raw.into(),
89 });
90 } else if let Some(m) = caps.name("string") {
91 let raw = m.as_str();
92 tokens.push(Token {
93 token_type: "string".into(),
94 value: raw[1..raw.len() - 1].into(),
95 raw: raw.into(),
96 });
97 } else if let Some(m) = caps.name("date") {
98 let raw = m.as_str();
102 tokens.push(Token {
103 token_type: "date".into(),
104 value: raw.into(),
105 raw: raw.into(),
106 });
107 } else if let Some(m) = caps.name("number") {
108 let raw = m.as_str();
109 tokens.push(Token {
110 token_type: "number".into(),
111 value: raw.into(),
112 raw: raw.into(),
113 });
114 } else if let Some(m) = caps.name("op") {
115 let raw = m.as_str();
116 tokens.push(Token {
117 token_type: "op".into(),
118 value: raw.into(),
119 raw: raw.into(),
120 });
121 } else if let Some(m) = caps.name("word") {
122 let raw = m.as_str();
123 if KEYWORDS.contains(&raw.to_uppercase().as_str()) {
124 tokens.push(Token {
125 token_type: "keyword".into(),
126 value: raw.to_uppercase(),
127 raw: raw.into(),
128 });
129 } else {
130 tokens.push(Token {
131 token_type: "ident".into(),
132 value: raw.into(),
133 raw: raw.into(),
134 });
135 }
136 }
137 }
138 tokens
139}
140
141struct Parser {
144 tokens: Vec<Token>,
145 pos: usize,
146}
147
148impl Parser {
149 fn new(tokens: Vec<Token>) -> Self {
150 Parser { tokens, pos: 0 }
151 }
152
153 fn peek(&self) -> Option<&Token> {
154 self.tokens.get(self.pos)
155 }
156
157 fn advance(&mut self) -> Token {
158 let t = self.tokens[self.pos].clone();
159 self.pos += 1;
160 t
161 }
162
163 fn expect(&mut self, type_: &str, value: Option<&str>) -> Result<Token, MdqlError> {
164 let t = self.peek().ok_or_else(|| {
165 MdqlError::QueryParse(format!(
166 "Unexpected end of query, expected {}",
167 value.unwrap_or(type_)
168 ))
169 })?;
170 let matches_type = t.token_type == type_;
171 let matches_value = value.map_or(true, |v| t.value == v);
172 if !matches_type || !matches_value {
173 return Err(MdqlError::QueryParse(format!(
174 "Expected {}, got '{}' at position {}",
175 value.unwrap_or(type_),
176 t.raw,
177 self.pos
178 )));
179 }
180 Ok(self.advance())
181 }
182
183 fn match_keyword(&mut self, kw: &str) -> bool {
184 if let Some(t) = self.peek() {
185 if t.token_type == "keyword" && t.value == kw {
186 self.advance();
187 return true;
188 }
189 }
190 false
191 }
192
193 fn parse_statement(&mut self) -> Result<Statement, MdqlError> {
194 let t = self.peek().ok_or_else(|| MdqlError::QueryParse("Empty query".into()))?;
195 match (t.token_type.as_str(), t.value.as_str()) {
196 ("keyword", "WITH") => {
197 let ctes = self.parse_ctes()?;
198 let mut q = self.parse_select()?;
199 q.ctes = ctes;
200 self.expect_end()?;
201 Ok(Statement::Select(q))
202 }
203 ("keyword", "SELECT") => {
204 let q = self.parse_select()?;
205 self.expect_end()?;
206 Ok(Statement::Select(q))
207 }
208 ("keyword", "INSERT") => Ok(Statement::Insert(self.parse_insert()?)),
209 ("keyword", "UPDATE") => Ok(Statement::Update(self.parse_update()?)),
210 ("keyword", "DELETE") => Ok(Statement::Delete(self.parse_delete()?)),
211 ("keyword", "ALTER") => self.parse_alter(),
212 ("keyword", "CREATE") => self.parse_create_view(),
213 ("keyword", "DROP") => self.parse_drop_view(),
214 _ => Err(MdqlError::QueryParse(format!(
215 "Expected SELECT, INSERT, UPDATE, DELETE, ALTER, CREATE, or DROP, got '{}'",
216 t.raw
217 ))),
218 }
219 }
220
221 fn parse_ctes(&mut self) -> Result<Vec<CteClause>, MdqlError> {
222 self.expect("keyword", Some("WITH"))?;
223 let mut ctes = Vec::new();
224 loop {
225 let name = self.parse_ident()?;
226 self.expect("keyword", Some("AS"))?;
227 self.expect("op", Some("("))?;
228 let query = self.parse_select()?;
229 self.expect("op", Some(")"))?;
230 ctes.push(CteClause { name, query: Box::new(query) });
231 if !self.peek().map_or(false, |t| t.token_type == "op" && t.value == ",") {
232 break;
233 }
234 self.advance();
235 }
236 Ok(ctes)
237 }
238
239 fn parse_select(&mut self) -> Result<SelectQuery, MdqlError> {
240 self.expect("keyword", Some("SELECT"))?;
241 let distinct = if self
245 .peek()
246 .map_or(false, |t| t.token_type == "ident" && t.value.eq_ignore_ascii_case("distinct"))
247 {
248 self.advance();
249 true
250 } else {
251 false
252 };
253 let columns = self.parse_columns()?;
254 self.expect("keyword", Some("FROM"))?;
255
256 let mut subquery = None;
258 let (table, mut table_alias) = if self.peek().map_or(false, |t| t.token_type == "op" && t.value == "(") {
259 self.advance();
260 let inner = self.parse_select()?;
261 self.expect("op", Some(")"))?;
262 subquery = Some(Box::new(inner));
263 let alias = if let Some(t) = self.peek() {
264 if t.token_type == "ident" && !self.is_clause_keyword(t) {
265 Some(self.advance().value)
266 } else {
267 None
268 }
269 } else {
270 None
271 };
272 ("_subquery".to_string(), alias)
273 } else {
274 let t = self.parse_ident()?;
275 (t, None)
276 };
277
278 if subquery.is_none() {
280 if let Some(t) = self.peek() {
281 if t.token_type == "ident" && !self.is_clause_keyword(t) {
282 table_alias = Some(self.advance().value);
283 }
284 }
285 }
286
287 let mut joins = Vec::new();
289 loop {
290 let jt = if self.match_keyword("LEFT") {
291 self.expect("keyword", Some("JOIN"))?;
292 JoinType::Left
293 } else if self.match_keyword("JOIN") {
294 JoinType::Inner
295 } else {
296 break;
297 };
298 let join_table = self.parse_ident()?;
299 let mut join_alias = None;
300 if let Some(t) = self.peek() {
301 if t.token_type == "ident" && !self.is_clause_keyword(t) {
302 join_alias = Some(self.advance().value);
303 }
304 }
305 self.expect("keyword", Some("ON"))?;
306 let condition = self.parse_or_expr()?;
307 joins.push(JoinClause {
308 join_type: jt,
309 table: join_table,
310 alias: join_alias,
311 condition,
312 });
313 }
314
315 let mut where_clause = None;
316 if self.match_keyword("WHERE") {
317 where_clause = Some(self.parse_or_expr()?);
318 }
319
320 let mut group_by = None;
321 if self.match_keyword("GROUP") {
322 self.expect("keyword", Some("BY"))?;
323 let mut cols = vec![self.parse_ident()?];
324 while self.peek().map_or(false, |t| t.token_type == "op" && t.value == ",") {
325 self.advance();
326 cols.push(self.parse_ident()?);
327 }
328 group_by = Some(cols);
329 }
330
331 let mut having = None;
332 if self.match_keyword("HAVING") {
333 having = Some(self.parse_or_expr()?);
334 }
335
336 let mut order_by = None;
337 if self.match_keyword("ORDER") {
338 self.expect("keyword", Some("BY"))?;
339 order_by = Some(self.parse_order_by()?);
340 }
341
342 let mut limit = None;
343 if self.match_keyword("LIMIT") {
344 let t = self.expect("number", None)?;
345 limit = Some(t.value.parse::<i64>().map_err(|_| {
346 MdqlError::QueryParse(format!("Invalid LIMIT value: {}", t.value))
347 })?);
348 }
349
350 Ok(SelectQuery {
351 distinct,
352 columns,
353 table,
354 table_alias,
355 subquery,
356 joins,
357 where_clause,
358 group_by,
359 having,
360 order_by,
361 limit,
362 ctes: vec![],
363 })
364 }
365
366 fn parse_insert(&mut self) -> Result<InsertQuery, MdqlError> {
367 self.expect("keyword", Some("INSERT"))?;
368 self.expect("keyword", Some("INTO"))?;
369 let table = self.parse_ident()?;
370
371 self.expect("op", Some("("))?;
372 let mut columns = vec![self.parse_ident()?];
373 while self.peek().map_or(false, |t| t.token_type == "op" && t.value == ",") {
374 self.advance();
375 columns.push(self.parse_ident()?);
376 }
377 self.expect("op", Some(")"))?;
378
379 self.expect("keyword", Some("VALUES"))?;
380
381 self.expect("op", Some("("))?;
382 let mut values = vec![self.parse_value()?];
383 while self.peek().map_or(false, |t| t.token_type == "op" && t.value == ",") {
384 self.advance();
385 values.push(self.parse_value()?);
386 }
387 self.expect("op", Some(")"))?;
388
389 if columns.len() != values.len() {
390 return Err(MdqlError::QueryParse(format!(
391 "Column count ({}) does not match value count ({})",
392 columns.len(),
393 values.len()
394 )));
395 }
396
397 self.expect_end()?;
398 Ok(InsertQuery {
399 table,
400 columns,
401 values,
402 })
403 }
404
405 fn parse_update(&mut self) -> Result<UpdateQuery, MdqlError> {
406 self.expect("keyword", Some("UPDATE"))?;
407 let table = self.parse_ident()?;
408 self.expect("keyword", Some("SET"))?;
409
410 let mut assignments = Vec::new();
411 let col = self.parse_ident()?;
412 self.expect("op", Some("="))?;
413 let val = self.parse_value()?;
414 assignments.push((col, val));
415
416 while self.peek().map_or(false, |t| t.token_type == "op" && t.value == ",") {
417 self.advance();
418 let col = self.parse_ident()?;
419 self.expect("op", Some("="))?;
420 let val = self.parse_value()?;
421 assignments.push((col, val));
422 }
423
424 let mut where_clause = None;
425 if self.match_keyword("WHERE") {
426 where_clause = Some(self.parse_or_expr()?);
427 }
428
429 self.expect_end()?;
430 Ok(UpdateQuery {
431 table,
432 assignments,
433 where_clause,
434 })
435 }
436
437 fn parse_delete(&mut self) -> Result<DeleteQuery, MdqlError> {
438 self.expect("keyword", Some("DELETE"))?;
439 self.expect("keyword", Some("FROM"))?;
440 let table = self.parse_ident()?;
441
442 let mut where_clause = None;
443 if self.match_keyword("WHERE") {
444 where_clause = Some(self.parse_or_expr()?);
445 }
446
447 let mode = if self.match_keyword("CASCADE") {
448 DeleteMode::Cascade
449 } else if self.match_keyword("RESTRICT") {
450 DeleteMode::Restrict
451 } else {
452 DeleteMode::Default
453 };
454
455 self.expect_end()?;
456 Ok(DeleteQuery {
457 table,
458 where_clause,
459 mode,
460 })
461 }
462
463 fn parse_alter(&mut self) -> Result<Statement, MdqlError> {
464 self.expect("keyword", Some("ALTER"))?;
465 self.expect("keyword", Some("TABLE"))?;
466 let table = self.parse_ident()?;
467
468 let t = self.peek().ok_or_else(|| {
469 MdqlError::QueryParse("Expected RENAME, DROP, or MERGE after table name".into())
470 })?;
471
472 match (t.token_type.as_str(), t.value.as_str()) {
473 ("keyword", "RENAME") => {
474 self.advance();
475 self.expect("keyword", Some("FIELD"))?;
476 let old_name = self.parse_string_or_ident()?;
477 self.expect("keyword", Some("TO"))?;
478 let new_name = self.parse_string_or_ident()?;
479 self.expect_end()?;
480 Ok(Statement::AlterRename(AlterRenameFieldQuery {
481 table,
482 old_name,
483 new_name,
484 }))
485 }
486 ("keyword", "DROP") => {
487 self.advance();
488 self.expect("keyword", Some("FIELD"))?;
489 let field_name = self.parse_string_or_ident()?;
490 self.expect_end()?;
491 Ok(Statement::AlterDrop(AlterDropFieldQuery {
492 table,
493 field_name,
494 }))
495 }
496 ("keyword", "MERGE") => {
497 self.advance();
498 self.expect("keyword", Some("FIELDS"))?;
499 let mut sources = vec![self.parse_string_or_ident()?];
500 while self.peek().map_or(false, |t| t.token_type == "op" && t.value == ",") {
501 self.advance();
502 sources.push(self.parse_string_or_ident()?);
503 }
504 self.expect("keyword", Some("INTO"))?;
505 let target = self.parse_string_or_ident()?;
506 self.expect_end()?;
507 Ok(Statement::AlterMerge(AlterMergeFieldsQuery {
508 table,
509 sources,
510 into: target,
511 }))
512 }
513 _ => Err(MdqlError::QueryParse(format!(
514 "Expected RENAME, DROP, or MERGE, got '{}'",
515 t.raw
516 ))),
517 }
518 }
519
520 fn parse_create_view(&mut self) -> Result<Statement, MdqlError> {
521 self.expect("keyword", Some("CREATE"))?;
522 self.expect("keyword", Some("VIEW"))?;
523 let view_name = self.parse_ident()?;
524
525 let columns = if self.peek().map_or(false, |t| t.token_type == "op" && t.value == "(") {
526 self.advance();
527 let mut cols = vec![self.parse_ident()?];
528 while self.peek().map_or(false, |t| t.token_type == "op" && t.value == ",") {
529 self.advance();
530 cols.push(self.parse_ident()?);
531 }
532 self.expect("op", Some(")"))?;
533 Some(cols)
534 } else {
535 None
536 };
537
538 self.expect("keyword", Some("AS"))?;
539 let query = Box::new(self.parse_select()?);
540 self.expect_end()?;
541
542 Ok(Statement::CreateView(CreateViewQuery {
543 view_name,
544 columns,
545 query,
546 }))
547 }
548
549 fn parse_drop_view(&mut self) -> Result<Statement, MdqlError> {
550 self.expect("keyword", Some("DROP"))?;
551 self.expect("keyword", Some("VIEW"))?;
552 let view_name = self.parse_ident()?;
553 self.expect_end()?;
554 Ok(Statement::DropView(DropViewQuery { view_name }))
555 }
556
557 fn parse_string_or_ident(&mut self) -> Result<String, MdqlError> {
558 let t = self.peek().ok_or_else(|| {
559 MdqlError::QueryParse("Expected field name, got end of query".into())
560 })?;
561 match t.token_type.as_str() {
562 "string" => {
563 let v = self.advance().value;
564 Ok(v)
565 }
566 "ident" | "keyword" => {
567 let v = self.advance().value;
568 Ok(v)
569 }
570 _ => Err(MdqlError::QueryParse(format!(
571 "Expected field name, got '{}'",
572 t.raw
573 ))),
574 }
575 }
576
577 fn parse_columns(&mut self) -> Result<ColumnList, MdqlError> {
578 if let Some(t) = self.peek() {
579 if t.token_type == "op" && t.value == "*" {
580 self.advance();
581 return Ok(ColumnList::All);
582 }
583 }
584
585 let mut exprs = vec![self.parse_select_expr()?];
586 while self.peek().map_or(false, |t| t.token_type == "op" && t.value == ",") {
587 self.advance();
588 exprs.push(self.parse_select_expr()?);
589 }
590 Ok(ColumnList::Named(exprs))
591 }
592
593 fn peek_is_window_func(&self) -> bool {
594 let t = match self.peek() {
595 Some(t) => t,
596 None => return false,
597 };
598 let name_upper = t.value.to_uppercase();
599 if !WINDOW_FUNCS.contains(&name_upper.as_str()) {
600 return false;
601 }
602 self.tokens
603 .get(self.pos + 1)
604 .map_or(false, |next| next.token_type == "op" && next.value == "(")
605 }
606
607 fn peek_is_agg_func(&self) -> bool {
608 let t = match self.peek() {
609 Some(t) => t,
610 None => return false,
611 };
612 let name_upper = t.value.to_uppercase();
613 if !AGG_FUNCS.contains(&name_upper.as_str()) {
614 return false;
615 }
616 self.tokens
618 .get(self.pos + 1)
619 .map_or(false, |next| next.token_type == "op" && next.value == "(")
620 }
621
622 fn parse_select_expr(&mut self) -> Result<SelectExpr, MdqlError> {
623 let _t = self.peek().ok_or_else(|| {
624 MdqlError::QueryParse("Expected column or aggregate, got end of query".into())
625 })?;
626
627 let expr = self.parse_additive()?;
628
629 let alias = if self.match_keyword("AS") {
630 Some(self.parse_ident()?)
631 } else if self.peek().map_or(false, |t| {
632 t.token_type == "ident" && !self.is_clause_keyword(t)
633 }) {
634 Some(self.advance().value)
635 } else {
636 None
637 };
638
639 if let Expr::Aggregate { func, arg, arg_expr } = expr {
641 return Ok(SelectExpr::Aggregate {
642 func,
643 arg,
644 arg_expr: arg_expr.map(|e| *e),
645 alias,
646 });
647 }
648
649 if alias.is_none() {
650 if let Expr::Column(name) = &expr {
651 return Ok(SelectExpr::Column(name.clone()));
652 }
653 }
654
655 Ok(SelectExpr::Expr { expr, alias })
656 }
657
658 fn peek_is_additive_op(&self) -> bool {
661 self.peek().map_or(false, |t| {
662 t.token_type == "op" && (t.value == "+" || t.value == "-")
663 })
664 }
665
666 fn peek_is_multiplicative_op(&self) -> bool {
667 self.peek().map_or(false, |t| {
668 t.token_type == "op" && (t.value == "*" || t.value == "/" || t.value == "%")
669 })
670 }
671
672 fn parse_additive(&mut self) -> Result<Expr, MdqlError> {
673 let mut left = self.parse_multiplicative()?;
674 while self.peek_is_additive_op() {
675 let op_tok = self.advance();
676 let is_sub = op_tok.value == "-";
677
678 if self.peek().map_or(false, |t| t.token_type == "keyword" && t.value == "INTERVAL") {
680 self.advance(); let days_expr = self.parse_multiplicative()?;
682 if !self.match_keyword("DAY") && !self.match_keyword("DAYS") {
684 return Err(MdqlError::QueryParse("Expected DAY after INTERVAL value".into()));
685 }
686 let days = if is_sub {
687 Expr::UnaryMinus(Box::new(days_expr))
688 } else {
689 days_expr
690 };
691 left = Expr::DateAdd {
692 date: Box::new(left),
693 days: Box::new(days),
694 };
695 continue;
696 }
697
698 let op = match op_tok.value.as_str() {
699 "+" => ArithOp::Add,
700 "-" => ArithOp::Sub,
701 _ => unreachable!(),
702 };
703 let right = self.parse_multiplicative()?;
704 left = Expr::BinaryOp {
705 left: Box::new(left),
706 op,
707 right: Box::new(right),
708 };
709 }
710 Ok(left)
711 }
712
713 fn parse_multiplicative(&mut self) -> Result<Expr, MdqlError> {
714 let mut left = self.parse_unary()?;
715 while self.peek_is_multiplicative_op() {
716 let op_tok = self.advance();
717 let op = match op_tok.value.as_str() {
718 "*" => ArithOp::Mul,
719 "/" => ArithOp::Div,
720 "%" => ArithOp::Mod,
721 _ => unreachable!(),
722 };
723 let right = self.parse_unary()?;
724 left = Expr::BinaryOp {
725 left: Box::new(left),
726 op,
727 right: Box::new(right),
728 };
729 }
730 Ok(left)
731 }
732
733 fn parse_unary(&mut self) -> Result<Expr, MdqlError> {
734 if self.peek().map_or(false, |t| t.token_type == "op" && t.value == "-") {
735 self.advance();
736 let inner = self.parse_atom()?;
737 match inner {
739 Expr::Literal(SqlValue::Int(n)) => Ok(Expr::Literal(SqlValue::Int(-n))),
740 Expr::Literal(SqlValue::Float(f)) => Ok(Expr::Literal(SqlValue::Float(-f))),
741 _ => Ok(Expr::UnaryMinus(Box::new(inner))),
742 }
743 } else {
744 self.parse_atom()
745 }
746 }
747
748 fn parse_atom(&mut self) -> Result<Expr, MdqlError> {
749 if self.peek_is_window_func() {
750 return self.parse_standalone_window();
751 }
752
753 if self.peek_is_agg_func() {
754 let agg = self.parse_agg_expr()?;
755 if self.peek().map_or(false, |t| t.token_type == "keyword" && t.value == "OVER") {
756 self.advance();
757 let over = self.parse_window_spec()?;
758 if let Expr::Aggregate { func, arg, arg_expr } = agg {
759 return Ok(Expr::Window {
760 func: WindowFunc::Agg(func),
761 args: if arg == "*" {
762 vec![]
763 } else {
764 vec![arg_expr.map(|e| *e).unwrap_or(Expr::Column(arg))]
765 },
766 over,
767 });
768 }
769 }
770 return Ok(agg);
771 }
772
773 let t = self.peek().ok_or_else(|| {
774 MdqlError::QueryParse("Expected expression, got end of query".into())
775 })?;
776
777 match t.token_type.as_str() {
778 "number" => {
779 let v = self.advance().value;
780 if v.contains('.') {
781 let f: f64 = v.parse().map_err(|_| {
782 MdqlError::QueryParse(format!("Invalid float: {}", v))
783 })?;
784 Ok(Expr::Literal(SqlValue::Float(f)))
785 } else {
786 let n: i64 = v.parse().map_err(|_| {
787 MdqlError::QueryParse(format!("Invalid int: {}", v))
788 })?;
789 Ok(Expr::Literal(SqlValue::Int(n)))
790 }
791 }
792 "string" => {
793 let v = self.advance().value;
794 Ok(Expr::Literal(SqlValue::String(v)))
795 }
796 "date" => {
797 let v = self.advance().value;
798 Ok(Expr::Literal(date_literal(&v)?))
799 }
800 "keyword" if t.value == "NULL" => {
801 self.advance();
802 Ok(Expr::Literal(SqlValue::Null))
803 }
804 "keyword" if t.value == "CASE" => {
805 self.parse_case_expr()
806 }
807 "keyword" if t.value == "CURRENT_DATE" => {
808 self.advance();
809 Ok(Expr::CurrentDate)
810 }
811 "keyword" if t.value == "CURRENT_TIMESTAMP" => {
812 self.advance();
813 Ok(Expr::CurrentTimestamp)
814 }
815 "keyword" if t.value == "DATEDIFF" => {
816 self.advance();
817 self.expect("op", Some("("))?;
818 let left = self.parse_additive()?;
819 self.expect("op", Some(","))?;
820 let right = self.parse_additive()?;
821 self.expect("op", Some(")"))?;
822 Ok(Expr::DateDiff { left: Box::new(left), right: Box::new(right) })
823 }
824 "op" if t.value == "(" => {
825 let next_is_select = self.tokens.get(self.pos + 1)
826 .map_or(false, |t| t.token_type == "keyword" && t.value == "SELECT");
827 if next_is_select {
828 self.advance();
829 let sq = self.parse_select()?;
830 self.expect("op", Some(")"))?;
831 Ok(Expr::Subquery(Box::new(sq)))
832 } else {
833 self.advance();
834 let expr = self.parse_additive()?;
835 self.expect("op", Some(")"))?;
836 Ok(expr)
837 }
838 }
839 "ident" if t.value.eq_ignore_ascii_case("true") => {
842 self.advance();
843 Ok(Expr::Literal(SqlValue::Bool(true)))
844 }
845 "ident" if t.value.eq_ignore_ascii_case("false") => {
846 self.advance();
847 Ok(Expr::Literal(SqlValue::Bool(false)))
848 }
849 "ident" => {
850 let name = self.advance().value;
851 Ok(Expr::Column(name))
852 }
853 "keyword" if !Self::is_reserved_keyword(&t.value) => {
854 let name = self.advance().value;
855 Ok(Expr::Column(name))
856 }
857 _ => Err(MdqlError::QueryParse(format!(
858 "Expected expression, got '{}'",
859 t.raw
860 ))),
861 }
862 }
863
864 fn parse_case_expr(&mut self) -> Result<Expr, MdqlError> {
865 self.expect("keyword", Some("CASE"))?;
866 let mut whens = Vec::new();
867 while self.match_keyword("WHEN") {
868 let condition = self.parse_or_expr()?;
869 self.expect("keyword", Some("THEN"))?;
870 let result = self.parse_additive()?;
871 whens.push((condition, Box::new(result)));
872 }
873 if whens.is_empty() {
874 return Err(MdqlError::QueryParse("CASE requires at least one WHEN clause".into()));
875 }
876 let else_expr = if self.match_keyword("ELSE") {
877 Some(Box::new(self.parse_additive()?))
878 } else {
879 None
880 };
881 self.expect("keyword", Some("END"))?;
882 Ok(Expr::Case { whens, else_expr })
883 }
884
885 fn parse_agg_expr(&mut self) -> Result<Expr, MdqlError> {
886 let func_name = self.advance().value.to_uppercase();
887 let func = match func_name.as_str() {
888 "COUNT" => AggFunc::Count,
889 "SUM" => AggFunc::Sum,
890 "AVG" => AggFunc::Avg,
891 "MIN" => AggFunc::Min,
892 "MAX" => AggFunc::Max,
893 _ => unreachable!(),
894 };
895 self.expect("op", Some("("))?;
896 let (arg, arg_expr) = if self.peek().map_or(false, |t| t.token_type == "op" && t.value == "*") {
897 self.advance();
898 ("*".to_string(), None)
899 } else {
900 let expr = self.parse_additive()?;
901 if let Expr::Column(name) = &expr {
902 (name.clone(), None)
903 } else {
904 (expr.display_name(), Some(Box::new(expr)))
905 }
906 };
907 self.expect("op", Some(")"))?;
908 Ok(Expr::Aggregate { func, arg, arg_expr })
909 }
910
911 fn parse_standalone_window(&mut self) -> Result<Expr, MdqlError> {
912 let func_name = self.advance().value.to_uppercase();
913 let func = match func_name.as_str() {
914 "ROW_NUMBER" => WindowFunc::RowNumber,
915 "RANK" => WindowFunc::Rank,
916 "DENSE_RANK" => WindowFunc::DenseRank,
917 "LAG" => WindowFunc::Lag,
918 "LEAD" => WindowFunc::Lead,
919 _ => unreachable!(),
920 };
921 self.expect("op", Some("("))?;
922 let mut args = Vec::new();
923 if !self.peek().map_or(false, |t| t.token_type == "op" && t.value == ")") {
924 args.push(self.parse_additive()?);
925 while self.peek().map_or(false, |t| t.token_type == "op" && t.value == ",") {
926 self.advance();
927 args.push(self.parse_additive()?);
928 }
929 }
930 self.expect("op", Some(")"))?;
931 self.expect("keyword", Some("OVER"))?;
932 let over = self.parse_window_spec()?;
933 Ok(Expr::Window { func, args, over })
934 }
935
936 fn parse_window_spec(&mut self) -> Result<WindowSpec, MdqlError> {
937 self.expect("op", Some("("))?;
938 let mut partition_by = Vec::new();
939 if self.match_keyword("PARTITION") {
940 self.expect("keyword", Some("BY"))?;
941 partition_by.push(self.parse_ident()?);
942 while self.peek().map_or(false, |t| t.token_type == "op" && t.value == ",") {
943 self.advance();
944 partition_by.push(self.parse_ident()?);
945 }
946 }
947 let mut order_by = Vec::new();
948 if self.match_keyword("ORDER") {
949 self.expect("keyword", Some("BY"))?;
950 order_by = self.parse_order_by()?;
951 }
952 self.expect("op", Some(")"))?;
953 Ok(WindowSpec { partition_by, order_by })
954 }
955
956 fn parse_ident(&mut self) -> Result<String, MdqlError> {
957 let t = self.peek().ok_or_else(|| {
958 MdqlError::QueryParse("Expected identifier, got end of query".into())
959 })?;
960 match t.token_type.as_str() {
961 "ident" | "keyword" => {
962 let v = self.advance().value;
963 Ok(v)
964 }
965 _ => Err(MdqlError::QueryParse(format!(
966 "Expected identifier, got '{}'",
967 t.raw
968 ))),
969 }
970 }
971
972 fn parse_or_expr(&mut self) -> Result<WhereClause, MdqlError> {
973 let mut left = self.parse_and_expr()?;
974 while self.match_keyword("OR") {
975 let right = self.parse_and_expr()?;
976 left = WhereClause::BoolOp(BoolOp {
977 op: BoolOpKind::Or,
978 left: Box::new(left),
979 right: Box::new(right),
980 });
981 }
982 Ok(left)
983 }
984
985 fn parse_and_expr(&mut self) -> Result<WhereClause, MdqlError> {
986 let mut left = self.parse_comparison()?;
987 while self.match_keyword("AND") {
988 let right = self.parse_comparison()?;
989 left = WhereClause::BoolOp(BoolOp {
990 op: BoolOpKind::And,
991 left: Box::new(left),
992 right: Box::new(right),
993 });
994 }
995 Ok(left)
996 }
997
998 fn parse_comparison(&mut self) -> Result<WhereClause, MdqlError> {
999 if self.peek().map_or(false, |t| t.token_type == "op" && t.value == "(") {
1001 let saved_pos = self.pos;
1003 self.advance();
1004 let result = self.parse_or_expr();
1006 if result.is_ok() && self.peek().map_or(false, |t| t.token_type == "op" && t.value == ")") {
1007 self.advance();
1008 return result;
1009 }
1010 self.pos = saved_pos;
1012 }
1013
1014 let left_expr = self.parse_additive()?;
1016
1017 let col = left_expr.as_column().unwrap_or("").to_string();
1019
1020 if self.match_keyword("IS") {
1022 if self.match_keyword("NOT") {
1023 self.expect("keyword", Some("NULL"))?;
1024 return Ok(WhereClause::Comparison(Comparison {
1025 column: col,
1026 op: CmpOp::IsNotNull,
1027 value: None,
1028 left_expr: Some(left_expr),
1029 right_expr: None,
1030 }));
1031 }
1032 self.expect("keyword", Some("NULL"))?;
1033 return Ok(WhereClause::Comparison(Comparison {
1034 column: col,
1035 op: CmpOp::IsNull,
1036 value: None,
1037 left_expr: Some(left_expr),
1038 right_expr: None,
1039 }));
1040 }
1041
1042 if self.match_keyword("IN") {
1044 self.expect("op", Some("("))?;
1045 let is_subquery = self.peek().map_or(false, |t| t.token_type == "keyword" && t.value == "SELECT");
1046 if is_subquery {
1047 let sq = self.parse_select()?;
1048 self.expect("op", Some(")"))?;
1049 return Ok(WhereClause::Comparison(Comparison {
1050 column: col,
1051 op: CmpOp::In,
1052 value: None,
1053 left_expr: Some(left_expr),
1054 right_expr: Some(Expr::Subquery(Box::new(sq))),
1055 }));
1056 }
1057 let mut values = vec![self.parse_value()?];
1058 while self.peek().map_or(false, |t| t.token_type == "op" && t.value == ",") {
1059 self.advance();
1060 values.push(self.parse_value()?);
1061 }
1062 self.expect("op", Some(")"))?;
1063 return Ok(WhereClause::Comparison(Comparison {
1064 column: col,
1065 op: CmpOp::In,
1066 value: Some(SqlValue::List(values)),
1067 left_expr: Some(left_expr),
1068 right_expr: None,
1069 }));
1070 }
1071
1072 if self.match_keyword("LIKE") {
1074 let val = self.parse_value()?;
1075 return Ok(WhereClause::Comparison(Comparison {
1076 column: col,
1077 op: CmpOp::Like,
1078 value: Some(val),
1079 left_expr: Some(left_expr),
1080 right_expr: None,
1081 }));
1082 }
1083
1084 if self.match_keyword("NOT") {
1086 if self.match_keyword("LIKE") {
1087 let val = self.parse_value()?;
1088 return Ok(WhereClause::Comparison(Comparison {
1089 column: col,
1090 op: CmpOp::NotLike,
1091 value: Some(val),
1092 left_expr: Some(left_expr),
1093 right_expr: None,
1094 }));
1095 }
1096 return Err(MdqlError::QueryParse("Expected LIKE after NOT".into()));
1097 }
1098
1099 if let Some(t) = self.peek() {
1101 if t.token_type == "op" && ["=", "!=", "<", ">", "<=", ">="].contains(&t.value.as_str())
1102 {
1103 let op_str = self.advance().value;
1104 let op = match op_str.as_str() {
1105 "=" => CmpOp::Eq,
1106 "!=" => CmpOp::Ne,
1107 "<" => CmpOp::Lt,
1108 ">" => CmpOp::Gt,
1109 "<=" => CmpOp::Le,
1110 ">=" => CmpOp::Ge,
1111 _ => unreachable!(),
1112 };
1113 let right_expr = self.parse_additive()?;
1115 let value = match &right_expr {
1117 Expr::Literal(v) => Some(v.clone()),
1118 _ => None,
1119 };
1120 return Ok(WhereClause::Comparison(Comparison {
1121 column: col,
1122 op,
1123 value,
1124 left_expr: Some(left_expr),
1125 right_expr: Some(right_expr),
1126 }));
1127 }
1128 }
1129
1130 let got = self.peek().map_or("end".to_string(), |t| t.raw.clone());
1131 Err(MdqlError::QueryParse(format!(
1132 "Expected operator after '{}', got '{}'",
1133 left_expr.display_name(), got
1134 )))
1135 }
1136
1137 fn parse_value(&mut self) -> Result<SqlValue, MdqlError> {
1138 let t = self.peek().ok_or_else(|| {
1139 MdqlError::QueryParse("Expected value, got end of query".into())
1140 })?;
1141 match t.token_type.as_str() {
1142 "string" => {
1143 let v = self.advance().value;
1144 Ok(SqlValue::String(v))
1145 }
1146 "date" => {
1147 let v = self.advance().value;
1148 date_literal(&v)
1149 }
1150 "number" => {
1151 let v = self.advance().value;
1152 if v.contains('.') {
1153 Ok(SqlValue::Float(v.parse().map_err(|_| {
1154 MdqlError::QueryParse(format!("Invalid float: {}", v))
1155 })?))
1156 } else {
1157 Ok(SqlValue::Int(v.parse().map_err(|_| {
1158 MdqlError::QueryParse(format!("Invalid int: {}", v))
1159 })?))
1160 }
1161 }
1162 "keyword" if t.value == "NULL" => {
1163 self.advance();
1164 Ok(SqlValue::Null)
1165 }
1166 "ident" if t.value.eq_ignore_ascii_case("true") => {
1170 self.advance();
1171 Ok(SqlValue::Bool(true))
1172 }
1173 "ident" if t.value.eq_ignore_ascii_case("false") => {
1174 self.advance();
1175 Ok(SqlValue::Bool(false))
1176 }
1177 _ => Err(MdqlError::QueryParse(format!(
1178 "Expected value, got '{}'",
1179 t.raw
1180 ))),
1181 }
1182 }
1183
1184 fn parse_order_by(&mut self) -> Result<Vec<OrderSpec>, MdqlError> {
1185 let mut specs = vec![self.parse_order_spec()?];
1186 while self.peek().map_or(false, |t| t.token_type == "op" && t.value == ",") {
1187 self.advance();
1188 specs.push(self.parse_order_spec()?);
1189 }
1190 Ok(specs)
1191 }
1192
1193 fn parse_order_spec(&mut self) -> Result<OrderSpec, MdqlError> {
1194 let expr = self.parse_additive()?;
1195 let col = expr.as_column().unwrap_or("").to_string();
1196 let descending = if self.match_keyword("DESC") {
1197 true
1198 } else {
1199 self.match_keyword("ASC");
1200 false
1201 };
1202 Ok(OrderSpec {
1203 column: col,
1204 expr: Some(expr),
1205 descending,
1206 })
1207 }
1208
1209 fn is_clause_keyword(&self, t: &Token) -> bool {
1210 t.token_type == "keyword"
1211 && ["WHERE", "ORDER", "LIMIT", "JOIN", "LEFT", "ON", "GROUP"].contains(&t.value.as_str())
1212 }
1213
1214 fn is_reserved_keyword(kw: &str) -> bool {
1216 matches!(kw,
1217 "AS" | "FROM" | "WHERE" | "AND" | "OR" | "ORDER" | "BY"
1218 | "ASC" | "DESC" | "LIMIT" | "JOIN" | "ON" | "GROUP"
1219 | "SELECT" | "INSERT" | "INTO" | "VALUES" | "UPDATE" | "SET"
1220 | "DELETE" | "ALTER" | "TABLE" | "IS" | "NOT" | "IN" | "LIKE"
1221 | "RENAME" | "FIELD" | "TO" | "DROP" | "MERGE" | "FIELDS"
1222 | "CASE" | "WHEN" | "THEN" | "ELSE" | "END"
1223 | "HAVING" | "INTERVAL" | "DAY" | "DAYS"
1224 | "CURRENT_DATE" | "CURRENT_TIMESTAMP" | "DATEDIFF"
1225 | "CREATE" | "VIEW" | "CASCADE" | "RESTRICT"
1226 | "WITH"
1227 | "OVER" | "PARTITION" | "ROW_NUMBER" | "RANK" | "DENSE_RANK" | "LAG" | "LEAD"
1228 )
1229 }
1230
1231 fn expect_end(&self) -> Result<(), MdqlError> {
1232 if let Some(t) = self.peek() {
1233 return Err(MdqlError::QueryParse(format!(
1234 "Unexpected token '{}' at position {}",
1235 t.raw, self.pos
1236 )));
1237 }
1238 Ok(())
1239 }
1240}
1241
1242pub fn parse_query(sql: &str) -> crate::errors::Result<Statement> {
1243 let tokens = tokenize(sql);
1244 if tokens.is_empty() {
1245 return Err(MdqlError::QueryParse("Empty query".into()));
1246 }
1247 let mut parser = Parser::new(tokens);
1248 parser.parse_statement()
1249}
1250
1251#[cfg(test)]
1252mod tests {
1253 use super::*;
1254
1255 #[test]
1256 fn test_simple_select() {
1257 let stmt = parse_query("SELECT title, status FROM strategies").unwrap();
1258 if let Statement::Select(q) = stmt {
1259 assert_eq!(q.columns, ColumnList::Named(vec![SelectExpr::Column("title".into()), SelectExpr::Column("status".into())]));
1260 assert_eq!(q.table, "strategies");
1261 } else {
1262 panic!("Expected Select");
1263 }
1264 }
1265
1266 #[test]
1267 fn test_select_star() {
1268 let stmt = parse_query("SELECT * FROM test").unwrap();
1269 if let Statement::Select(q) = stmt {
1270 assert_eq!(q.columns, ColumnList::All);
1271 } else {
1272 panic!("Expected Select");
1273 }
1274 }
1275
1276 #[test]
1277 fn test_select_distinct_parses_flag_and_column() {
1278 let stmt = parse_query("SELECT DISTINCT strategy FROM backtests").unwrap();
1281 let Statement::Select(q) = stmt else { panic!("Expected Select") };
1282 assert!(q.distinct);
1283 assert_eq!(
1284 q.columns,
1285 ColumnList::Named(vec![SelectExpr::Column("strategy".into())])
1286 );
1287
1288 let Statement::Select(q) = parse_query("select distinct a, b from t").unwrap() else {
1290 panic!("Expected Select")
1291 };
1292 assert!(q.distinct);
1293 let Statement::Select(q) = parse_query("SELECT strategy FROM backtests").unwrap() else {
1294 panic!("Expected Select")
1295 };
1296 assert!(!q.distinct);
1297 }
1298
1299 #[test]
1300 fn test_boolean_literal_parses_as_bool_not_column() {
1301 for (sql, expected) in [
1304 ("SELECT title FROM t WHERE enabled = true", SqlValue::Bool(true)),
1305 ("SELECT title FROM t WHERE enabled = FALSE", SqlValue::Bool(false)),
1306 ] {
1307 let stmt = parse_query(sql).unwrap();
1308 let Statement::Select(q) = stmt else { panic!("Expected Select") };
1309 let WhereClause::Comparison(cmp) = q.where_clause.unwrap() else {
1310 panic!("Expected Comparison")
1311 };
1312 assert_eq!(cmp.value, Some(expected));
1313 }
1314 }
1315
1316 #[test]
1317 fn test_unquoted_date_literal_parses_as_string_not_arithmetic() {
1318 for (bare, quoted) in [
1321 (
1322 "SELECT title FROM t WHERE date >= 2026-01-01",
1323 "SELECT title FROM t WHERE date >= '2026-01-01'",
1324 ),
1325 (
1326 "SELECT title FROM t WHERE ts < 2026-01-01T12:30:00",
1327 "SELECT title FROM t WHERE ts < '2026-01-01T12:30:00'",
1328 ),
1329 (
1330 "SELECT title FROM t WHERE ts = 2026-01-01 12:30:00",
1331 "SELECT title FROM t WHERE ts = '2026-01-01 12:30:00'",
1332 ),
1333 (
1334 "SELECT title FROM t WHERE d IN (2026-01-01, 2026-02-01)",
1335 "SELECT title FROM t WHERE d IN ('2026-01-01', '2026-02-01')",
1336 ),
1337 (
1338 "UPDATE t SET closed = 2026-01-01 WHERE id = 1",
1339 "UPDATE t SET closed = '2026-01-01' WHERE id = 1",
1340 ),
1341 ] {
1342 assert_eq!(
1343 parse_query(bare).unwrap(),
1344 parse_query(quoted).unwrap(),
1345 "bare and quoted forms must parse identically: {bare}"
1346 );
1347 }
1348 }
1349
1350 #[test]
1351 fn test_non_iso_date_literal_errors_instead_of_comparing_wrong() {
1352 let err = parse_query("SELECT title FROM t WHERE date >= 2026-1-1").unwrap_err();
1355 let msg = err.to_string();
1356 assert!(msg.contains("2026-1-1"), "{msg}");
1357 assert!(msg.contains("'2026-01-01'"), "{msg}");
1358 }
1359
1360 #[test]
1361 fn test_spaced_arithmetic_still_subtracts() {
1362 let Statement::Select(q) = parse_query("SELECT n - 1 - 1 AS m FROM t").unwrap() else {
1364 panic!("Expected Select")
1365 };
1366 let ColumnList::Named(cols) = q.columns else { panic!("Expected named columns") };
1367 assert_eq!(cols.len(), 1);
1368 let SelectExpr::Expr { expr, .. } = &cols[0] else { panic!("Expected expression") };
1369 assert!(matches!(expr, Expr::BinaryOp { op: ArithOp::Sub, .. }), "{expr:?}");
1370 }
1371
1372 #[test]
1373 fn test_where_clause() {
1374 let stmt = parse_query("SELECT title FROM test WHERE count > 5").unwrap();
1375 if let Statement::Select(q) = stmt {
1376 assert!(q.where_clause.is_some());
1377 } else {
1378 panic!("Expected Select");
1379 }
1380 }
1381
1382 #[test]
1383 fn test_order_by() {
1384 let stmt =
1385 parse_query("SELECT title FROM test ORDER BY composite DESC, title ASC").unwrap();
1386 if let Statement::Select(q) = stmt {
1387 let ob = q.order_by.unwrap();
1388 assert_eq!(ob.len(), 2);
1389 assert!(ob[0].descending);
1390 assert!(!ob[1].descending);
1391 } else {
1392 panic!("Expected Select");
1393 }
1394 }
1395
1396 #[test]
1397 fn test_limit() {
1398 let stmt = parse_query("SELECT * FROM test LIMIT 10").unwrap();
1399 if let Statement::Select(q) = stmt {
1400 assert_eq!(q.limit, Some(10));
1401 } else {
1402 panic!("Expected Select");
1403 }
1404 }
1405
1406 #[test]
1407 fn test_insert() {
1408 let stmt = parse_query(
1409 "INSERT INTO test (title, count) VALUES ('Hello', 42)",
1410 )
1411 .unwrap();
1412 if let Statement::Insert(q) = stmt {
1413 assert_eq!(q.table, "test");
1414 assert_eq!(q.columns, vec!["title", "count"]);
1415 assert_eq!(q.values[0], SqlValue::String("Hello".into()));
1416 assert_eq!(q.values[1], SqlValue::Int(42));
1417 } else {
1418 panic!("Expected Insert");
1419 }
1420 }
1421
1422 #[test]
1423 fn test_update() {
1424 let stmt = parse_query("UPDATE test SET status = 'KILLED' WHERE path = 'a.md'").unwrap();
1425 if let Statement::Update(q) = stmt {
1426 assert_eq!(q.table, "test");
1427 assert_eq!(q.assignments.len(), 1);
1428 assert!(q.where_clause.is_some());
1429 } else {
1430 panic!("Expected Update");
1431 }
1432 }
1433
1434 #[test]
1435 fn test_delete() {
1436 let stmt = parse_query("DELETE FROM test WHERE status = 'draft'").unwrap();
1437 if let Statement::Delete(q) = stmt {
1438 assert_eq!(q.table, "test");
1439 assert!(q.where_clause.is_some());
1440 } else {
1441 panic!("Expected Delete");
1442 }
1443 }
1444
1445 #[test]
1446 fn test_alter_rename() {
1447 let stmt =
1448 parse_query("ALTER TABLE test RENAME FIELD 'Summary' TO 'Overview'").unwrap();
1449 if let Statement::AlterRename(q) = stmt {
1450 assert_eq!(q.old_name, "Summary");
1451 assert_eq!(q.new_name, "Overview");
1452 } else {
1453 panic!("Expected AlterRename");
1454 }
1455 }
1456
1457 #[test]
1458 fn test_alter_drop() {
1459 let stmt = parse_query("ALTER TABLE test DROP FIELD 'Details'").unwrap();
1460 if let Statement::AlterDrop(q) = stmt {
1461 assert_eq!(q.field_name, "Details");
1462 } else {
1463 panic!("Expected AlterDrop");
1464 }
1465 }
1466
1467 #[test]
1468 fn test_alter_merge() {
1469 let stmt = parse_query(
1470 "ALTER TABLE test MERGE FIELDS 'Entry Rules', 'Exit Rules' INTO 'Trading Rules'",
1471 )
1472 .unwrap();
1473 if let Statement::AlterMerge(q) = stmt {
1474 assert_eq!(q.sources, vec!["Entry Rules", "Exit Rules"]);
1475 assert_eq!(q.into, "Trading Rules");
1476 } else {
1477 panic!("Expected AlterMerge");
1478 }
1479 }
1480
1481 #[test]
1482 fn test_backtick_ident() {
1483 let stmt = parse_query("SELECT `Structural Mechanism` FROM test").unwrap();
1484 if let Statement::Select(q) = stmt {
1485 assert_eq!(
1486 q.columns,
1487 ColumnList::Named(vec![SelectExpr::Column("Structural Mechanism".into())])
1488 );
1489 } else {
1490 panic!("Expected Select");
1491 }
1492 }
1493
1494 #[test]
1495 fn test_like_operator() {
1496 let stmt = parse_query("SELECT title FROM test WHERE categories LIKE '%defi%'").unwrap();
1497 if let Statement::Select(q) = stmt {
1498 if let Some(WhereClause::Comparison(c)) = q.where_clause {
1499 assert_eq!(c.op, CmpOp::Like);
1500 assert_eq!(c.value, Some(SqlValue::String("%defi%".into())));
1501 } else {
1502 panic!("Expected LIKE comparison");
1503 }
1504 } else {
1505 panic!("Expected Select");
1506 }
1507 }
1508
1509 #[test]
1510 fn test_in_operator() {
1511 let stmt =
1512 parse_query("SELECT * FROM test WHERE status IN ('ACTIVE', 'LIVE')").unwrap();
1513 if let Statement::Select(q) = stmt {
1514 if let Some(WhereClause::Comparison(c)) = q.where_clause {
1515 assert_eq!(c.op, CmpOp::In);
1516 } else {
1517 panic!("Expected IN comparison");
1518 }
1519 } else {
1520 panic!("Expected Select");
1521 }
1522 }
1523
1524 #[test]
1525 fn test_is_null() {
1526 let stmt = parse_query("SELECT * FROM test WHERE title IS NULL").unwrap();
1527 if let Statement::Select(q) = stmt {
1528 if let Some(WhereClause::Comparison(c)) = q.where_clause {
1529 assert_eq!(c.op, CmpOp::IsNull);
1530 } else {
1531 panic!("Expected IS NULL comparison");
1532 }
1533 } else {
1534 panic!("Expected Select");
1535 }
1536 }
1537
1538 #[test]
1539 fn test_and_or() {
1540 let stmt = parse_query(
1541 "SELECT * FROM test WHERE status = 'ACTIVE' AND count > 5 OR title LIKE '%test%'",
1542 )
1543 .unwrap();
1544 if let Statement::Select(q) = stmt {
1545 assert!(q.where_clause.is_some());
1546 } else {
1547 panic!("Expected Select");
1548 }
1549 }
1550
1551 #[test]
1552 fn test_join() {
1553 let stmt = parse_query(
1554 "SELECT s.title, b.sharpe FROM strategies s JOIN backtests b ON b.strategy = s.path",
1555 )
1556 .unwrap();
1557 if let Statement::Select(q) = stmt {
1558 assert_eq!(q.table, "strategies");
1559 assert_eq!(q.table_alias, Some("s".into()));
1560 assert_eq!(q.joins.len(), 1);
1561 let join = &q.joins[0];
1562 assert_eq!(join.table, "backtests");
1563 assert_eq!(join.alias, Some("b".into()));
1564 } else {
1565 panic!("Expected Select");
1566 }
1567 }
1568
1569 #[test]
1570 fn test_multi_join() {
1571 let stmt = parse_query(
1572 "SELECT s.title, b.sharpe, c.verdict FROM strategies s JOIN backtests b ON b.strategy = s.path JOIN critiques c ON c.strategy = s.path",
1573 )
1574 .unwrap();
1575 if let Statement::Select(q) = stmt {
1576 assert_eq!(q.table, "strategies");
1577 assert_eq!(q.table_alias, Some("s".into()));
1578 assert_eq!(q.joins.len(), 2);
1579 assert_eq!(q.joins[0].table, "backtests");
1580 assert_eq!(q.joins[0].alias, Some("b".into()));
1581 assert_eq!(where_clause_to_sql(&q.joins[0].condition), "b.strategy = s.path");
1582 assert_eq!(q.joins[1].table, "critiques");
1583 assert_eq!(q.joins[1].alias, Some("c".into()));
1584 assert_eq!(where_clause_to_sql(&q.joins[1].condition), "c.strategy = s.path");
1585 } else {
1586 panic!("Expected Select");
1587 }
1588 }
1589
1590 #[test]
1591 fn test_left_join() {
1592 let stmt = parse_query(
1593 "SELECT s.title, b.sharpe FROM strategies s LEFT JOIN backtests b ON b.strategy = s.path",
1594 )
1595 .unwrap();
1596 if let Statement::Select(q) = stmt {
1597 assert_eq!(q.joins.len(), 1);
1598 assert_eq!(q.joins[0].join_type, JoinType::Left);
1599 assert_eq!(q.joins[0].table, "backtests");
1600 } else {
1601 panic!("Expected Select");
1602 }
1603 }
1604
1605 #[test]
1606 fn test_mixed_join_types() {
1607 let stmt = parse_query(
1608 "SELECT s.title FROM strategies s JOIN backtests b ON b.strategy = s.path LEFT JOIN allocations a ON a.strategy = s.path",
1609 )
1610 .unwrap();
1611 if let Statement::Select(q) = stmt {
1612 assert_eq!(q.joins.len(), 2);
1613 assert_eq!(q.joins[0].join_type, JoinType::Inner);
1614 assert_eq!(q.joins[1].join_type, JoinType::Left);
1615 } else {
1616 panic!("Expected Select");
1617 }
1618 }
1619
1620 #[test]
1621 fn test_join_compound_and() {
1622 let stmt = parse_query(
1623 "SELECT s.title FROM strategies s LEFT JOIN backtests b ON b.strategy = s.path AND b.mode = 'PAPER'",
1624 )
1625 .unwrap();
1626 if let Statement::Select(q) = stmt {
1627 assert_eq!(q.joins.len(), 1);
1628 assert_eq!(q.joins[0].join_type, JoinType::Left);
1629 let sql = where_clause_to_sql(&q.joins[0].condition);
1630 assert!(sql.contains("b.strategy = s.path"));
1631 assert!(sql.contains("AND"));
1632 assert!(sql.contains("b.mode = 'PAPER'"));
1633 } else {
1634 panic!("Expected Select");
1635 }
1636 }
1637
1638 #[test]
1639 fn test_join_compound_or() {
1640 let stmt = parse_query(
1641 "SELECT * FROM a JOIN b ON a.id = b.id OR a.alt = b.id",
1642 )
1643 .unwrap();
1644 if let Statement::Select(q) = stmt {
1645 let sql = where_clause_to_sql(&q.joins[0].condition);
1646 assert!(sql.contains("OR"));
1647 } else {
1648 panic!("Expected Select");
1649 }
1650 }
1651
1652 #[test]
1653 fn test_join_compound_with_where() {
1654 let stmt = parse_query(
1655 "SELECT s.title FROM strategies s JOIN backtests b ON b.strategy = s.path AND b.mode = 'PAPER' WHERE s.title = 'Alpha'",
1656 )
1657 .unwrap();
1658 if let Statement::Select(q) = stmt {
1659 assert_eq!(q.joins.len(), 1);
1660 assert!(q.where_clause.is_some());
1661 let join_sql = where_clause_to_sql(&q.joins[0].condition);
1662 assert!(join_sql.contains("AND"));
1663 } else {
1664 panic!("Expected Select");
1665 }
1666 }
1667
1668 #[test]
1669 fn test_empty_query() {
1670 assert!(parse_query("").is_err());
1671 }
1672
1673 #[test]
1674 fn test_count_star() {
1675 let stmt = parse_query("SELECT status, COUNT(*) AS cnt FROM strategies GROUP BY status").unwrap();
1676 if let Statement::Select(q) = stmt {
1677 if let ColumnList::Named(exprs) = &q.columns {
1678 assert_eq!(exprs.len(), 2);
1679 assert_eq!(exprs[0], SelectExpr::Column("status".into()));
1680 assert!(matches!(&exprs[1], SelectExpr::Aggregate {
1681 func: AggFunc::Count,
1682 arg,
1683 alias: Some(a),
1684 ..
1685 } if arg == "*" && a == "cnt"));
1686 } else {
1687 panic!("Expected Named columns");
1688 }
1689 assert_eq!(q.group_by, Some(vec!["status".into()]));
1690 } else {
1691 panic!("Expected Select");
1692 }
1693 }
1694
1695 #[test]
1696 fn test_count_column_as_ident() {
1697 let stmt = parse_query("INSERT INTO test (title, count) VALUES ('Hello', 42)").unwrap();
1699 if let Statement::Insert(q) = stmt {
1700 assert_eq!(q.columns, vec!["title", "count"]);
1701 } else {
1702 panic!("Expected Insert");
1703 }
1704 }
1705
1706 #[test]
1707 fn test_multiple_aggregates() {
1708 let stmt = parse_query("SELECT MIN(composite), MAX(composite), AVG(composite) FROM strategies").unwrap();
1709 if let Statement::Select(q) = stmt {
1710 if let ColumnList::Named(exprs) = &q.columns {
1711 assert_eq!(exprs.len(), 3);
1712 assert!(matches!(&exprs[0], SelectExpr::Aggregate { func: AggFunc::Min, .. }));
1713 assert!(matches!(&exprs[1], SelectExpr::Aggregate { func: AggFunc::Max, .. }));
1714 assert!(matches!(&exprs[2], SelectExpr::Aggregate { func: AggFunc::Avg, .. }));
1715 } else {
1716 panic!("Expected Named columns");
1717 }
1718 assert_eq!(q.group_by, None);
1719 } else {
1720 panic!("Expected Select");
1721 }
1722 }
1723
1724 #[test]
1727 fn test_select_arithmetic_expr() {
1728 let stmt = parse_query("SELECT a + b FROM test").unwrap();
1729 if let Statement::Select(q) = stmt {
1730 if let ColumnList::Named(exprs) = &q.columns {
1731 assert_eq!(exprs.len(), 1);
1732 assert!(matches!(&exprs[0], SelectExpr::Expr {
1733 expr: Expr::BinaryOp { op: ArithOp::Add, .. },
1734 alias: None,
1735 }));
1736 } else {
1737 panic!("Expected Named columns");
1738 }
1739 } else {
1740 panic!("Expected Select");
1741 }
1742 }
1743
1744 #[test]
1745 fn test_select_arithmetic_with_alias() {
1746 let stmt = parse_query("SELECT a + b AS total FROM test").unwrap();
1747 if let Statement::Select(q) = stmt {
1748 if let ColumnList::Named(exprs) = &q.columns {
1749 assert_eq!(exprs.len(), 1);
1750 assert!(matches!(&exprs[0], SelectExpr::Expr {
1751 alias: Some(a),
1752 ..
1753 } if a == "total"));
1754 assert_eq!(exprs[0].output_name(), "total");
1755 } else {
1756 panic!("Expected Named columns");
1757 }
1758 } else {
1759 panic!("Expected Select");
1760 }
1761 }
1762
1763 #[test]
1764 fn test_select_precedence() {
1765 let stmt = parse_query("SELECT a + b * c FROM test").unwrap();
1767 if let Statement::Select(q) = stmt {
1768 if let ColumnList::Named(exprs) = &q.columns {
1769 if let SelectExpr::Expr { expr, .. } = &exprs[0] {
1770 if let Expr::BinaryOp { left, op, right } = expr {
1771 assert_eq!(*op, ArithOp::Add);
1772 assert!(matches!(left.as_ref(), Expr::Column(n) if n == "a"));
1773 assert!(matches!(right.as_ref(), Expr::BinaryOp { op: ArithOp::Mul, .. }));
1774 } else {
1775 panic!("Expected BinaryOp");
1776 }
1777 } else {
1778 panic!("Expected Expr variant");
1779 }
1780 } else {
1781 panic!("Expected Named columns");
1782 }
1783 } else {
1784 panic!("Expected Select");
1785 }
1786 }
1787
1788 #[test]
1789 fn test_select_parenthesized_expr() {
1790 let stmt = parse_query("SELECT (a + b) * c FROM test").unwrap();
1792 if let Statement::Select(q) = stmt {
1793 if let ColumnList::Named(exprs) = &q.columns {
1794 if let SelectExpr::Expr { expr, .. } = &exprs[0] {
1795 if let Expr::BinaryOp { left, op, .. } = expr {
1796 assert_eq!(*op, ArithOp::Mul);
1797 assert!(matches!(left.as_ref(), Expr::BinaryOp { op: ArithOp::Add, .. }));
1798 } else {
1799 panic!("Expected BinaryOp");
1800 }
1801 } else {
1802 panic!("Expected Expr variant");
1803 }
1804 } else {
1805 panic!("Expected Named columns");
1806 }
1807 } else {
1808 panic!("Expected Select");
1809 }
1810 }
1811
1812 #[test]
1813 fn test_select_unary_minus() {
1814 let stmt = parse_query("SELECT -count FROM test").unwrap();
1815 if let Statement::Select(q) = stmt {
1816 if let ColumnList::Named(exprs) = &q.columns {
1817 assert!(matches!(&exprs[0], SelectExpr::Expr {
1818 expr: Expr::UnaryMinus(_),
1819 ..
1820 }));
1821 } else {
1822 panic!("Expected Named columns");
1823 }
1824 } else {
1825 panic!("Expected Select");
1826 }
1827 }
1828
1829 #[test]
1830 fn test_select_negative_literal() {
1831 let stmt = parse_query("SELECT -42 FROM test").unwrap();
1832 if let Statement::Select(q) = stmt {
1833 if let ColumnList::Named(exprs) = &q.columns {
1834 assert!(matches!(&exprs[0], SelectExpr::Expr {
1836 expr: Expr::Literal(SqlValue::Int(-42)),
1837 ..
1838 }));
1839 } else {
1840 panic!("Expected Named columns");
1841 }
1842 } else {
1843 panic!("Expected Select");
1844 }
1845 }
1846
1847 #[test]
1848 fn test_where_arithmetic_expr() {
1849 let stmt = parse_query("SELECT * FROM test WHERE a + b > 10").unwrap();
1850 if let Statement::Select(q) = stmt {
1851 if let Some(WhereClause::Comparison(c)) = q.where_clause {
1852 assert_eq!(c.op, CmpOp::Gt);
1853 assert!(matches!(&c.left_expr, Some(Expr::BinaryOp { op: ArithOp::Add, .. })));
1854 assert!(matches!(&c.right_expr, Some(Expr::Literal(SqlValue::Int(10)))));
1855 } else {
1856 panic!("Expected comparison");
1857 }
1858 } else {
1859 panic!("Expected Select");
1860 }
1861 }
1862
1863 #[test]
1864 fn test_where_both_sides_expr() {
1865 let stmt = parse_query("SELECT * FROM test WHERE a * 2 > b + 1").unwrap();
1866 if let Statement::Select(q) = stmt {
1867 if let Some(WhereClause::Comparison(c)) = q.where_clause {
1868 assert_eq!(c.op, CmpOp::Gt);
1869 assert!(matches!(&c.left_expr, Some(Expr::BinaryOp { op: ArithOp::Mul, .. })));
1870 assert!(matches!(&c.right_expr, Some(Expr::BinaryOp { op: ArithOp::Add, .. })));
1871 } else {
1872 panic!("Expected comparison");
1873 }
1874 } else {
1875 panic!("Expected Select");
1876 }
1877 }
1878
1879 #[test]
1880 fn test_order_by_expr() {
1881 let stmt = parse_query("SELECT * FROM test ORDER BY a + b DESC").unwrap();
1882 if let Statement::Select(q) = stmt {
1883 let ob = q.order_by.unwrap();
1884 assert_eq!(ob.len(), 1);
1885 assert!(ob[0].descending);
1886 assert!(matches!(&ob[0].expr, Some(Expr::BinaryOp { op: ArithOp::Add, .. })));
1887 } else {
1888 panic!("Expected Select");
1889 }
1890 }
1891
1892 #[test]
1893 fn test_all_arithmetic_ops() {
1894 let stmt = parse_query("SELECT a + b, a - b, a * b, a / b, a % b FROM test").unwrap();
1895 if let Statement::Select(q) = stmt {
1896 if let ColumnList::Named(exprs) = &q.columns {
1897 assert_eq!(exprs.len(), 5);
1898 assert!(matches!(&exprs[0], SelectExpr::Expr { expr: Expr::BinaryOp { op: ArithOp::Add, .. }, .. }));
1899 assert!(matches!(&exprs[1], SelectExpr::Expr { expr: Expr::BinaryOp { op: ArithOp::Sub, .. }, .. }));
1900 assert!(matches!(&exprs[2], SelectExpr::Expr { expr: Expr::BinaryOp { op: ArithOp::Mul, .. }, .. }));
1901 assert!(matches!(&exprs[3], SelectExpr::Expr { expr: Expr::BinaryOp { op: ArithOp::Div, .. }, .. }));
1902 assert!(matches!(&exprs[4], SelectExpr::Expr { expr: Expr::BinaryOp { op: ArithOp::Mod, .. }, .. }));
1903 } else {
1904 panic!("Expected Named columns");
1905 }
1906 } else {
1907 panic!("Expected Select");
1908 }
1909 }
1910
1911 #[test]
1912 fn test_column_with_literal_arithmetic() {
1913 let stmt = parse_query("SELECT count * 2 + 1 FROM test").unwrap();
1914 if let Statement::Select(q) = stmt {
1915 if let ColumnList::Named(exprs) = &q.columns {
1916 if let SelectExpr::Expr { expr, .. } = &exprs[0] {
1918 if let Expr::BinaryOp { left, op, right } = expr {
1919 assert_eq!(*op, ArithOp::Add);
1920 assert!(matches!(right.as_ref(), Expr::Literal(SqlValue::Int(1))));
1921 assert!(matches!(left.as_ref(), Expr::BinaryOp { op: ArithOp::Mul, .. }));
1922 } else {
1923 panic!("Expected BinaryOp");
1924 }
1925 } else {
1926 panic!("Expected Expr");
1927 }
1928 } else {
1929 panic!("Expected Named columns");
1930 }
1931 } else {
1932 panic!("Expected Select");
1933 }
1934 }
1935
1936 #[test]
1937 fn test_mixed_columns_and_exprs() {
1938 let stmt = parse_query("SELECT title, a + b AS sum, count FROM test").unwrap();
1939 if let Statement::Select(q) = stmt {
1940 if let ColumnList::Named(exprs) = &q.columns {
1941 assert_eq!(exprs.len(), 3);
1942 assert_eq!(exprs[0], SelectExpr::Column("title".into()));
1943 assert!(matches!(&exprs[1], SelectExpr::Expr { alias: Some(a), .. } if a == "sum"));
1944 assert_eq!(exprs[2], SelectExpr::Column("count".into()));
1945 } else {
1946 panic!("Expected Named columns");
1947 }
1948 } else {
1949 panic!("Expected Select");
1950 }
1951 }
1952
1953 #[test]
1956 fn test_case_when_basic() {
1957 let stmt = parse_query(
1958 "SELECT CASE WHEN status = 'ACTIVE' THEN 1 ELSE 0 END FROM test"
1959 ).unwrap();
1960 if let Statement::Select(q) = stmt {
1961 if let ColumnList::Named(exprs) = &q.columns {
1962 assert_eq!(exprs.len(), 1);
1963 assert!(matches!(&exprs[0], SelectExpr::Expr {
1964 expr: Expr::Case { .. },
1965 ..
1966 }));
1967 } else {
1968 panic!("Expected Named columns");
1969 }
1970 } else {
1971 panic!("Expected Select");
1972 }
1973 }
1974
1975 #[test]
1976 fn test_case_when_multiple_branches() {
1977 let stmt = parse_query(
1978 "SELECT CASE WHEN x > 10 THEN 'high' WHEN x > 5 THEN 'mid' ELSE 'low' END FROM test"
1979 ).unwrap();
1980 if let Statement::Select(q) = stmt {
1981 if let ColumnList::Named(exprs) = &q.columns {
1982 if let SelectExpr::Expr { expr: Expr::Case { whens, else_expr }, .. } = &exprs[0] {
1983 assert_eq!(whens.len(), 2);
1984 assert!(else_expr.is_some());
1985 } else {
1986 panic!("Expected Case expression");
1987 }
1988 } else {
1989 panic!("Expected Named columns");
1990 }
1991 } else {
1992 panic!("Expected Select");
1993 }
1994 }
1995
1996 #[test]
1997 fn test_case_when_no_else() {
1998 let stmt = parse_query(
1999 "SELECT CASE WHEN x = 1 THEN 'one' END FROM test"
2000 ).unwrap();
2001 if let Statement::Select(q) = stmt {
2002 if let ColumnList::Named(exprs) = &q.columns {
2003 if let SelectExpr::Expr { expr: Expr::Case { whens, else_expr }, .. } = &exprs[0] {
2004 assert_eq!(whens.len(), 1);
2005 assert!(else_expr.is_none());
2006 } else {
2007 panic!("Expected Case expression");
2008 }
2009 } else {
2010 panic!("Expected Named columns");
2011 }
2012 } else {
2013 panic!("Expected Select");
2014 }
2015 }
2016
2017 #[test]
2018 fn test_case_when_in_aggregate() {
2019 let stmt = parse_query(
2020 "SELECT SUM(CASE WHEN side = 'BUY' THEN size ELSE -size END) AS net FROM orders GROUP BY token"
2021 ).unwrap();
2022 if let Statement::Select(q) = stmt {
2023 if let ColumnList::Named(exprs) = &q.columns {
2024 assert_eq!(exprs.len(), 1);
2025 assert!(matches!(&exprs[0], SelectExpr::Aggregate {
2026 func: AggFunc::Sum,
2027 arg_expr: Some(Expr::Case { .. }),
2028 alias: Some(a),
2029 ..
2030 } if a == "net"));
2031 } else {
2032 panic!("Expected Named columns");
2033 }
2034 } else {
2035 panic!("Expected Select");
2036 }
2037 }
2038
2039 #[test]
2040 fn test_case_when_with_alias() {
2041 let stmt = parse_query(
2042 "SELECT CASE WHEN x > 0 THEN 'pos' ELSE 'neg' END AS sign FROM test"
2043 ).unwrap();
2044 if let Statement::Select(q) = stmt {
2045 if let ColumnList::Named(exprs) = &q.columns {
2046 assert!(matches!(&exprs[0], SelectExpr::Expr {
2047 expr: Expr::Case { .. },
2048 alias: Some(a),
2049 } if a == "sign"));
2050 } else {
2051 panic!("Expected Named columns");
2052 }
2053 } else {
2054 panic!("Expected Select");
2055 }
2056 }
2057
2058 #[test]
2059 fn test_create_view() {
2060 let stmt = parse_query("CREATE VIEW live AS SELECT * FROM strategies WHERE status = 'LIVE'").unwrap();
2061 if let Statement::CreateView(cv) = stmt {
2062 assert_eq!(cv.view_name, "live");
2063 assert!(cv.columns.is_none());
2064 assert_eq!(cv.query.table, "strategies");
2065 assert!(cv.query.where_clause.is_some());
2066 } else {
2067 panic!("Expected CreateView, got {:?}", stmt);
2068 }
2069 }
2070
2071 #[test]
2072 fn test_create_view_with_columns() {
2073 let stmt = parse_query("CREATE VIEW v1 (a, b) AS SELECT title, status FROM t").unwrap();
2074 if let Statement::CreateView(cv) = stmt {
2075 assert_eq!(cv.view_name, "v1");
2076 assert_eq!(cv.columns, Some(vec!["a".into(), "b".into()]));
2077 } else {
2078 panic!("Expected CreateView");
2079 }
2080 }
2081
2082 #[test]
2083 fn test_drop_view() {
2084 let stmt = parse_query("DROP VIEW live").unwrap();
2085 if let Statement::DropView(dv) = stmt {
2086 assert_eq!(dv.view_name, "live");
2087 } else {
2088 panic!("Expected DropView, got {:?}", stmt);
2089 }
2090 }
2091
2092 #[test]
2093 fn test_create_view_case_insensitive() {
2094 let stmt = parse_query("create view My_View as select * from t").unwrap();
2095 if let Statement::CreateView(cv) = stmt {
2096 assert_eq!(cv.view_name, "My_View");
2097 } else {
2098 panic!("Expected CreateView");
2099 }
2100 }
2101
2102 #[test]
2105 fn test_aggregate_division() {
2106 let stmt = parse_query(
2107 "SELECT token, SUM(sell) / SUM(buy) as ratio FROM orders GROUP BY token"
2108 ).unwrap();
2109 if let Statement::Select(q) = stmt {
2110 assert_eq!(q.group_by, Some(vec!["token".into()]));
2111 if let ColumnList::Named(exprs) = &q.columns {
2112 assert_eq!(exprs.len(), 2);
2113 assert!(exprs[1].is_aggregate());
2114 } else {
2115 panic!("Expected Named columns");
2116 }
2117 } else {
2118 panic!("Expected Select");
2119 }
2120 }
2121
2122 #[test]
2123 fn test_aggregate_subtraction() {
2124 let stmt = parse_query(
2125 "SELECT token, SUM(sell) - SUM(buy) as net FROM orders GROUP BY token"
2126 ).unwrap();
2127 if let Statement::Select(q) = stmt {
2128 if let ColumnList::Named(exprs) = &q.columns {
2129 assert_eq!(exprs[1].output_name(), "net");
2130 }
2131 } else {
2132 panic!("Expected Select");
2133 }
2134 }
2135
2136 #[test]
2137 fn test_create_view_with_arithmetic() {
2138 let stmt = parse_query(
2139 "CREATE VIEW positions AS SELECT token, SUM(sell) / SUM(buy) as ratio FROM orders GROUP BY token"
2140 ).unwrap();
2141 if let Statement::CreateView(cv) = stmt {
2142 assert_eq!(cv.view_name, "positions");
2143 } else {
2144 panic!("Expected CreateView, got {:?}", stmt);
2145 }
2146 }
2147
2148 #[test]
2151 fn test_subquery_in_from() {
2152 let stmt = parse_query(
2153 "SELECT token, sell_size FROM (SELECT token, SUM(size) as sell_size FROM orders GROUP BY token) LIMIT 5"
2154 ).unwrap();
2155 if let Statement::Select(q) = stmt {
2156 assert!(q.subquery.is_some());
2157 assert_eq!(q.limit, Some(5));
2158 let sub = q.subquery.unwrap();
2159 assert_eq!(sub.table, "orders");
2160 assert!(sub.group_by.is_some());
2161 } else {
2162 panic!("Expected Select");
2163 }
2164 }
2165
2166 #[test]
2169 fn test_create_view_with_having() {
2170 let stmt = parse_query(
2171 "CREATE VIEW positions AS SELECT token, SUM(sell) as sell_size, SUM(buy) as buy_size FROM orders GROUP BY token HAVING sell_size > buy_size"
2172 ).unwrap();
2173 if let Statement::CreateView(cv) = stmt {
2174 assert_eq!(cv.view_name, "positions");
2175 assert!(cv.query.having.is_some());
2176 } else {
2177 panic!("Expected CreateView, got {:?}", stmt);
2178 }
2179 }
2180
2181 #[test]
2184 fn test_aggregate_multiplication() {
2185 let stmt = parse_query(
2186 "SELECT SUM(a) * 2 as doubled FROM test"
2187 ).unwrap();
2188 if let Statement::Select(q) = stmt {
2189 if let ColumnList::Named(exprs) = &q.columns {
2190 assert_eq!(exprs.len(), 1);
2191 assert!(exprs[0].is_aggregate());
2192 assert_eq!(exprs[0].output_name(), "doubled");
2193 } else {
2194 panic!("Expected Named columns");
2195 }
2196 } else {
2197 panic!("Expected Select");
2198 }
2199 }
2200
2201 #[test]
2202 fn test_complex_aggregate_arithmetic() {
2203 let stmt = parse_query(
2204 "SELECT SUM(CASE WHEN side = 'SELL' THEN size ELSE 0 END) / SUM(CASE WHEN side = 'BUY' THEN size ELSE 0 END) as ratio FROM orders GROUP BY token"
2205 ).unwrap();
2206 if let Statement::Select(q) = stmt {
2207 if let ColumnList::Named(exprs) = &q.columns {
2208 assert_eq!(exprs.len(), 1);
2209 assert!(exprs[0].is_aggregate());
2210 assert_eq!(exprs[0].output_name(), "ratio");
2211 } else {
2212 panic!("Expected Named columns");
2213 }
2214 assert_eq!(q.group_by, Some(vec!["token".into()]));
2215 } else {
2216 panic!("Expected Select");
2217 }
2218 }
2219
2220 #[test]
2223 fn test_subquery_with_alias() {
2224 let stmt = parse_query(
2225 "SELECT x FROM (SELECT x FROM t) sub"
2226 ).unwrap();
2227 if let Statement::Select(q) = stmt {
2228 assert!(q.subquery.is_some());
2229 let sub = q.subquery.unwrap();
2230 assert_eq!(sub.table, "t");
2231 if let ColumnList::Named(exprs) = &q.columns {
2232 assert_eq!(exprs.len(), 1);
2233 assert_eq!(exprs[0].output_name(), "x");
2234 } else {
2235 panic!("Expected Named columns");
2236 }
2237 } else {
2238 panic!("Expected Select");
2239 }
2240 }
2241
2242 #[test]
2243 fn test_subquery_with_where() {
2244 let stmt = parse_query(
2245 "SELECT x FROM (SELECT x FROM t WHERE y > 0) LIMIT 5"
2246 ).unwrap();
2247 if let Statement::Select(q) = stmt {
2248 assert!(q.subquery.is_some());
2249 assert_eq!(q.limit, Some(5));
2250 let sub = q.subquery.unwrap();
2251 assert_eq!(sub.table, "t");
2252 assert!(sub.where_clause.is_some());
2253 } else {
2254 panic!("Expected Select");
2255 }
2256 }
2257
2258 #[test]
2261 fn test_create_view_aggregate_subtraction() {
2262 let stmt = parse_query(
2263 "CREATE VIEW v AS SELECT token, SUM(sell) - SUM(buy) as net FROM orders GROUP BY token"
2264 ).unwrap();
2265 if let Statement::CreateView(cv) = stmt {
2266 assert_eq!(cv.view_name, "v");
2267 assert_eq!(cv.query.group_by, Some(vec!["token".into()]));
2268 if let ColumnList::Named(exprs) = &cv.query.columns {
2269 assert_eq!(exprs.len(), 2);
2270 assert_eq!(exprs[1].output_name(), "net");
2271 assert!(exprs[1].is_aggregate());
2272 } else {
2273 panic!("Expected Named columns");
2274 }
2275 } else {
2276 panic!("Expected CreateView, got {:?}", stmt);
2277 }
2278 }
2279
2280 #[test]
2281 fn test_delete_cascade() {
2282 let stmt = parse_query("DELETE FROM strategies WHERE status = 'KILLED' CASCADE").unwrap();
2283 if let Statement::Delete(q) = stmt {
2284 assert_eq!(q.table, "strategies");
2285 assert!(q.where_clause.is_some());
2286 assert_eq!(q.mode, DeleteMode::Cascade);
2287 } else {
2288 panic!("Expected Delete");
2289 }
2290 }
2291
2292 #[test]
2293 fn test_delete_restrict() {
2294 let stmt = parse_query("DELETE FROM strategies WHERE path = 'alpha.md' RESTRICT").unwrap();
2295 if let Statement::Delete(q) = stmt {
2296 assert_eq!(q.table, "strategies");
2297 assert_eq!(q.mode, DeleteMode::Restrict);
2298 } else {
2299 panic!("Expected Delete");
2300 }
2301 }
2302
2303 #[test]
2304 fn test_delete_default_unchanged() {
2305 let stmt = parse_query("DELETE FROM strategies WHERE status = 'KILLED'").unwrap();
2306 if let Statement::Delete(q) = stmt {
2307 assert_eq!(q.mode, DeleteMode::Default);
2308 } else {
2309 panic!("Expected Delete");
2310 }
2311 }
2312
2313 #[test]
2314 fn test_delete_cascade_no_where() {
2315 let stmt = parse_query("DELETE FROM strategies CASCADE").unwrap();
2316 if let Statement::Delete(q) = stmt {
2317 assert_eq!(q.table, "strategies");
2318 assert!(q.where_clause.is_none());
2319 assert_eq!(q.mode, DeleteMode::Cascade);
2320 } else {
2321 panic!("Expected Delete");
2322 }
2323 }
2324
2325 #[test]
2328 fn test_cte_basic() {
2329 let stmt = parse_query(
2330 "WITH live AS (SELECT * FROM strategies WHERE status = 'LIVE') SELECT * FROM live"
2331 ).unwrap();
2332 if let Statement::Select(q) = stmt {
2333 assert_eq!(q.ctes.len(), 1);
2334 assert_eq!(q.ctes[0].name, "live");
2335 assert_eq!(q.ctes[0].query.table, "strategies");
2336 assert!(q.ctes[0].query.where_clause.is_some());
2337 assert_eq!(q.table, "live");
2338 } else {
2339 panic!("Expected Select");
2340 }
2341 }
2342
2343 #[test]
2344 fn test_cte_multi() {
2345 let stmt = parse_query(
2346 "WITH a AS (SELECT * FROM t1), b AS (SELECT * FROM t2) SELECT * FROM a JOIN b ON a.id = b.id"
2347 ).unwrap();
2348 if let Statement::Select(q) = stmt {
2349 assert_eq!(q.ctes.len(), 2);
2350 assert_eq!(q.ctes[0].name, "a");
2351 assert_eq!(q.ctes[0].query.table, "t1");
2352 assert_eq!(q.ctes[1].name, "b");
2353 assert_eq!(q.ctes[1].query.table, "t2");
2354 assert_eq!(q.table, "a");
2355 assert_eq!(q.joins.len(), 1);
2356 } else {
2357 panic!("Expected Select");
2358 }
2359 }
2360
2361 #[test]
2362 fn test_cte_with_aggregation() {
2363 let stmt = parse_query(
2364 "WITH totals AS (SELECT strategy, COUNT(*) AS cnt FROM backtests GROUP BY strategy) SELECT * FROM totals WHERE cnt > 1"
2365 ).unwrap();
2366 if let Statement::Select(q) = stmt {
2367 assert_eq!(q.ctes.len(), 1);
2368 assert_eq!(q.ctes[0].name, "totals");
2369 assert!(q.ctes[0].query.group_by.is_some());
2370 assert_eq!(q.table, "totals");
2371 assert!(q.where_clause.is_some());
2372 } else {
2373 panic!("Expected Select");
2374 }
2375 }
2376
2377 #[test]
2378 fn test_cte_no_ctes_on_plain_select() {
2379 let stmt = parse_query("SELECT * FROM t").unwrap();
2380 if let Statement::Select(q) = stmt {
2381 assert!(q.ctes.is_empty());
2382 } else {
2383 panic!("Expected Select");
2384 }
2385 }
2386
2387 #[test]
2390 fn test_where_in_subquery() {
2391 let stmt = parse_query(
2392 "SELECT * FROM strategies WHERE path IN (SELECT strategy FROM backtests)"
2393 ).unwrap();
2394 if let Statement::Select(q) = stmt {
2395 if let Some(WhereClause::Comparison(c)) = &q.where_clause {
2396 assert_eq!(c.op, CmpOp::In);
2397 assert!(matches!(&c.right_expr, Some(Expr::Subquery(_))));
2398 } else {
2399 panic!("Expected IN comparison");
2400 }
2401 } else {
2402 panic!("Expected Select");
2403 }
2404 }
2405
2406 #[test]
2407 fn test_scalar_subquery_in_where() {
2408 let stmt = parse_query(
2409 "SELECT * FROM backtests WHERE sharpe > (SELECT AVG(sharpe) FROM backtests)"
2410 ).unwrap();
2411 if let Statement::Select(q) = stmt {
2412 if let Some(WhereClause::Comparison(c)) = &q.where_clause {
2413 assert_eq!(c.op, CmpOp::Gt);
2414 assert!(matches!(&c.right_expr, Some(Expr::Subquery(_))));
2415 } else {
2416 panic!("Expected comparison");
2417 }
2418 } else {
2419 panic!("Expected Select");
2420 }
2421 }
2422
2423 #[test]
2424 fn test_scalar_subquery_in_select() {
2425 let stmt = parse_query(
2426 "SELECT title, (SELECT COUNT(*) FROM backtests) AS cnt FROM strategies"
2427 ).unwrap();
2428 if let Statement::Select(q) = stmt {
2429 if let ColumnList::Named(exprs) = &q.columns {
2430 assert_eq!(exprs.len(), 2);
2431 assert!(matches!(&exprs[1], SelectExpr::Expr {
2432 expr: Expr::Subquery(_),
2433 alias: Some(a),
2434 } if a == "cnt"));
2435 } else {
2436 panic!("Expected Named columns");
2437 }
2438 } else {
2439 panic!("Expected Select");
2440 }
2441 }
2442
2443 #[test]
2446 fn test_row_number_over_order_by() {
2447 let stmt = parse_query(
2448 "SELECT title, ROW_NUMBER() OVER (ORDER BY count DESC) AS rn FROM test"
2449 ).unwrap();
2450 if let Statement::Select(q) = stmt {
2451 if let ColumnList::Named(exprs) = &q.columns {
2452 assert_eq!(exprs.len(), 2);
2453 if let SelectExpr::Expr { expr: Expr::Window { func, args, over }, alias } = &exprs[1] {
2454 assert_eq!(*func, WindowFunc::RowNumber);
2455 assert!(args.is_empty());
2456 assert!(over.partition_by.is_empty());
2457 assert_eq!(over.order_by.len(), 1);
2458 assert!(over.order_by[0].descending);
2459 assert_eq!(alias.as_deref(), Some("rn"));
2460 } else {
2461 panic!("Expected Window expression, got {:?}", exprs[1]);
2462 }
2463 } else {
2464 panic!("Expected Named columns");
2465 }
2466 } else {
2467 panic!("Expected Select");
2468 }
2469 }
2470
2471 #[test]
2472 fn test_rank_with_partition_by() {
2473 let stmt = parse_query(
2474 "SELECT RANK() OVER (PARTITION BY category ORDER BY price DESC) AS rnk FROM test"
2475 ).unwrap();
2476 if let Statement::Select(q) = stmt {
2477 if let ColumnList::Named(exprs) = &q.columns {
2478 if let SelectExpr::Expr { expr: Expr::Window { func, over, .. }, .. } = &exprs[0] {
2479 assert_eq!(*func, WindowFunc::Rank);
2480 assert_eq!(over.partition_by, vec!["category"]);
2481 assert_eq!(over.order_by.len(), 1);
2482 } else {
2483 panic!("Expected Window expression");
2484 }
2485 } else {
2486 panic!("Expected Named columns");
2487 }
2488 } else {
2489 panic!("Expected Select");
2490 }
2491 }
2492
2493 #[test]
2494 fn test_agg_over_window() {
2495 let stmt = parse_query(
2496 "SELECT SUM(price) OVER (PARTITION BY category) AS cat_total FROM test"
2497 ).unwrap();
2498 if let Statement::Select(q) = stmt {
2499 if let ColumnList::Named(exprs) = &q.columns {
2500 if let SelectExpr::Expr { expr: Expr::Window { func, args, over }, alias } = &exprs[0] {
2501 assert!(matches!(func, WindowFunc::Agg(AggFunc::Sum)));
2502 assert_eq!(args.len(), 1);
2503 assert_eq!(over.partition_by, vec!["category"]);
2504 assert!(over.order_by.is_empty());
2505 assert_eq!(alias.as_deref(), Some("cat_total"));
2506 } else {
2507 panic!("Expected Window expression");
2508 }
2509 } else {
2510 panic!("Expected Named columns");
2511 }
2512 } else {
2513 panic!("Expected Select");
2514 }
2515 }
2516
2517 #[test]
2518 fn test_lag_with_args() {
2519 let stmt = parse_query(
2520 "SELECT LAG(price, 1) OVER (ORDER BY price) AS prev_price FROM test"
2521 ).unwrap();
2522 if let Statement::Select(q) = stmt {
2523 if let ColumnList::Named(exprs) = &q.columns {
2524 if let SelectExpr::Expr { expr: Expr::Window { func, args, .. }, .. } = &exprs[0] {
2525 assert_eq!(*func, WindowFunc::Lag);
2526 assert_eq!(args.len(), 2);
2527 } else {
2528 panic!("Expected Window expression");
2529 }
2530 } else {
2531 panic!("Expected Named columns");
2532 }
2533 } else {
2534 panic!("Expected Select");
2535 }
2536 }
2537
2538 #[test]
2539 fn test_dense_rank() {
2540 let stmt = parse_query(
2541 "SELECT DENSE_RANK() OVER (ORDER BY count DESC) AS dr FROM test"
2542 ).unwrap();
2543 if let Statement::Select(q) = stmt {
2544 if let ColumnList::Named(exprs) = &q.columns {
2545 if let SelectExpr::Expr { expr: Expr::Window { func, .. }, .. } = &exprs[0] {
2546 assert_eq!(*func, WindowFunc::DenseRank);
2547 } else {
2548 panic!("Expected Window expression");
2549 }
2550 } else {
2551 panic!("Expected Named columns");
2552 }
2553 } else {
2554 panic!("Expected Select");
2555 }
2556 }
2557
2558 #[test]
2559 fn test_sum_without_over_is_aggregate() {
2560 let stmt = parse_query("SELECT SUM(count) FROM test").unwrap();
2561 if let Statement::Select(q) = stmt {
2562 if let ColumnList::Named(exprs) = &q.columns {
2563 assert!(matches!(&exprs[0], SelectExpr::Aggregate { func: AggFunc::Sum, .. }));
2564 } else {
2565 panic!("Expected Named columns");
2566 }
2567 } else {
2568 panic!("Expected Select");
2569 }
2570 }
2571}