Skip to main content

sql_cli/sql/parser/expressions/
primary.rs

1// Primary expression parsing
2// Handles literals, identifiers, function calls, and parenthesized expressions
3
4use crate::sql::parser::ast::{ColumnRef, SqlExpression, WindowSpec};
5use crate::sql::parser::lexer::Token;
6use tracing::{debug, trace};
7
8use super::{log_parse_decision, trace_parse_entry, trace_parse_exit, ExpressionParser};
9
10/// Parser context for primary expressions
11pub struct PrimaryExpressionContext<'a> {
12    pub columns: &'a [String],
13    pub in_method_args: bool,
14}
15
16impl<'a> Default for PrimaryExpressionContext<'a> {
17    fn default() -> Self {
18        Self {
19            columns: &[],
20            in_method_args: false,
21        }
22    }
23}
24
25/// Parse a primary expression (literals, identifiers, functions, parentheses)
26/// This is the bottom of the expression hierarchy
27pub fn parse_primary<P>(
28    parser: &mut P,
29    ctx: &PrimaryExpressionContext,
30) -> Result<SqlExpression, String>
31where
32    P: ParsePrimary + ExpressionParser + ?Sized,
33{
34    trace_parse_entry("parse_primary", ExpressionParser::current_token(parser));
35
36    // Special case: check if a number literal could actually be a column name
37    // This handles cases where columns are named with pure numbers like "202204"
38    if let Token::NumberLiteral(num_str) = ExpressionParser::current_token(parser) {
39        if ctx.columns.iter().any(|col| col == num_str) {
40            log_parse_decision(
41                "parse_primary",
42                ExpressionParser::current_token(parser),
43                "Number literal matches column name, treating as column",
44            );
45            let expr = SqlExpression::Column(ColumnRef::unquoted(num_str.clone()));
46            ExpressionParser::advance(parser);
47            let result = Ok(expr);
48            trace_parse_exit("parse_primary", &result);
49            return result;
50        }
51    }
52
53    let result = match ExpressionParser::current_token(parser) {
54        Token::Case => {
55            debug!("Parsing CASE expression");
56            parser.parse_case_expression()
57        }
58
59        Token::DateTime => {
60            debug!("Parsing DateTime constructor");
61            parse_datetime_constructor(parser)
62        }
63
64        Token::Unnest => {
65            debug!("Parsing UNNEST expression");
66            parse_unnest(parser)
67        }
68
69        Token::Identifier(id) => {
70            let id_upper = id.to_uppercase();
71            let id_clone = id.clone();
72
73            // Check for boolean literals first
74            if id_upper == "TRUE" {
75                log_parse_decision(
76                    "parse_primary",
77                    ExpressionParser::current_token(parser),
78                    "Boolean literal TRUE",
79                );
80                ExpressionParser::advance(parser);
81                Ok(SqlExpression::BooleanLiteral(true))
82            } else if id_upper == "FALSE" {
83                log_parse_decision(
84                    "parse_primary",
85                    ExpressionParser::current_token(parser),
86                    "Boolean literal FALSE",
87                );
88                ExpressionParser::advance(parser);
89                Ok(SqlExpression::BooleanLiteral(false))
90            } else {
91                ExpressionParser::advance(parser);
92
93                // Check for table.column notation or method calls
94                if matches!(ExpressionParser::current_token(parser), Token::Dot) {
95                    ExpressionParser::advance(parser); // consume dot
96
97                    if let Token::Identifier(next_id) = ExpressionParser::current_token(parser) {
98                        let next_id = next_id.clone();
99                        ExpressionParser::advance(parser);
100
101                        // Check if this is a method call (followed by parentheses)
102                        if matches!(ExpressionParser::current_token(parser), Token::LeftParen) {
103                            debug!(object = %id_clone, method = %next_id, "Parsing method call");
104                            ExpressionParser::advance(parser); // consume (
105
106                            // Handle empty argument list
107                            let args = if matches!(
108                                ExpressionParser::current_token(parser),
109                                Token::RightParen
110                            ) {
111                                Vec::new()
112                            } else {
113                                parser.parse_expression_list()?
114                            };
115                            ExpressionParser::consume(parser, Token::RightParen)?;
116
117                            log_parse_decision(
118                                "parse_primary",
119                                &Token::Identifier(next_id.clone()),
120                                "Method call",
121                            );
122                            Ok(SqlExpression::MethodCall {
123                                object: id_clone,
124                                method: next_id,
125                                args,
126                            })
127                        } else {
128                            // It's a qualified column reference
129                            let col_ref = ColumnRef::qualified(id_clone, next_id.clone());
130                            log_parse_decision(
131                                "parse_primary",
132                                &Token::Identifier(next_id),
133                                "Qualified column reference",
134                            );
135                            Ok(SqlExpression::Column(col_ref))
136                        }
137                    } else {
138                        Err("Expected identifier after '.'".to_string())
139                    }
140                // CAST(expr AS type) / TRY_CAST(expr AS type) — the `AS type`
141                // form is not a normal argument list, so intercept it here and
142                // lower it into a two-arg function call CAST(expr, 'TYPE').
143                } else if (id_upper == "CAST" || id_upper == "TRY_CAST")
144                    && matches!(ExpressionParser::current_token(parser), Token::LeftParen)
145                {
146                    parse_cast_expression(parser, &id_upper)
147                // Check if this is a function call
148                } else if matches!(ExpressionParser::current_token(parser), Token::LeftParen) {
149                    debug!(function = %id_upper, "Parsing function call");
150                    ExpressionParser::advance(parser); // consume (
151                    let (args, has_distinct) = parser.parse_function_args()?;
152                    ExpressionParser::consume(parser, Token::RightParen)?;
153
154                    // Check for OVER clause for window functions
155                    if matches!(ExpressionParser::current_token(parser), Token::Over) {
156                        debug!(function = %id_upper, "Window function detected");
157                        ExpressionParser::advance(parser); // consume OVER
158                        ExpressionParser::consume(parser, Token::LeftParen)?;
159                        let window_spec = parser.parse_window_spec()?;
160                        ExpressionParser::consume(parser, Token::RightParen)?;
161                        Ok(SqlExpression::WindowFunction {
162                            name: id_upper,
163                            args,
164                            window_spec,
165                        })
166                    } else {
167                        Ok(SqlExpression::FunctionCall {
168                            name: id_upper,
169                            args,
170                            distinct: has_distinct,
171                        })
172                    }
173                } else {
174                    // Otherwise treat as simple column
175                    log_parse_decision(
176                        "parse_primary",
177                        &Token::Identifier(id_clone.clone()),
178                        "Column reference",
179                    );
180                    Ok(SqlExpression::Column(ColumnRef::unquoted(id_clone)))
181                }
182            }
183        }
184
185        Token::QuotedIdentifier(id) => {
186            let expr = if ctx.in_method_args {
187                // In method arguments, treat quoted identifiers as string literals
188                log_parse_decision(
189                    "parse_primary",
190                    ExpressionParser::current_token(parser),
191                    "Quoted identifier in method args - treating as string",
192                );
193                SqlExpression::StringLiteral(id.clone())
194            } else {
195                // Otherwise it's a column name like "Customer Id"
196                log_parse_decision(
197                    "parse_primary",
198                    ExpressionParser::current_token(parser),
199                    "Quoted identifier as column name",
200                );
201                SqlExpression::Column(ColumnRef::quoted(id.clone()))
202            };
203            ExpressionParser::advance(parser);
204            Ok(expr)
205        }
206
207        Token::StringLiteral(s) => {
208            trace!("String literal: {}", s);
209            let expr = SqlExpression::StringLiteral(s.clone());
210            ExpressionParser::advance(parser);
211            Ok(expr)
212        }
213
214        Token::NumberLiteral(n) => {
215            trace!("Number literal: {}", n);
216            let expr = SqlExpression::NumberLiteral(n.clone());
217            ExpressionParser::advance(parser);
218            Ok(expr)
219        }
220
221        Token::Null => {
222            trace!("NULL literal");
223            ExpressionParser::advance(parser);
224            Ok(SqlExpression::Null)
225        }
226
227        // Handle LEFT and RIGHT as function names when followed by parentheses
228        Token::Left | Token::Right => {
229            let func_name = match ExpressionParser::current_token(parser) {
230                Token::Left => "LEFT".to_string(),
231                Token::Right => "RIGHT".to_string(),
232                _ => unreachable!(),
233            };
234
235            ExpressionParser::advance(parser);
236
237            // Check if this is a function call
238            if matches!(ExpressionParser::current_token(parser), Token::LeftParen) {
239                debug!(function = %func_name, "Parsing LEFT/RIGHT function call");
240                ExpressionParser::advance(parser); // consume (
241                let (args, _has_distinct) = parser.parse_function_args()?;
242                ExpressionParser::consume(parser, Token::RightParen)?;
243
244                Ok(SqlExpression::FunctionCall {
245                    name: func_name,
246                    args,
247                    distinct: false,
248                })
249            } else {
250                // If not followed by parenthesis, this is likely a JOIN keyword - error
251                Err(format!(
252                    "{} keyword unexpected in expression context",
253                    func_name
254                ))
255            }
256        }
257
258        Token::LeftParen => {
259            debug!("Parsing parenthesized expression or subquery");
260            ExpressionParser::advance(parser); // consume (
261
262            // Check if this is a subquery. It starts with SELECT, or with WITH
263            // for a CTE in expression position (P12) — parse_subquery() dispatches
264            // a leading WITH to the CTE parser, so both forms flow through here.
265            if matches!(
266                ExpressionParser::current_token(parser),
267                Token::Select | Token::With
268            ) {
269                debug!("Detected subquery - parsing SELECT/WITH statement");
270                let subquery = parser.parse_subquery()?;
271                ExpressionParser::consume(parser, Token::RightParen)?;
272                Ok(SqlExpression::ScalarSubquery {
273                    query: Box::new(subquery),
274                })
275            } else {
276                // Parenthesized expression, possibly a tuple for tuple IN:
277                // (a, b) IN (SELECT x, y FROM ...)
278                let first = parser.parse_logical_or()?;
279
280                if matches!(ExpressionParser::current_token(parser), Token::Comma) {
281                    // Collect the remaining tuple elements
282                    let mut exprs = vec![first];
283                    while matches!(ExpressionParser::current_token(parser), Token::Comma) {
284                        ExpressionParser::advance(parser); // consume ,
285                        exprs.push(parser.parse_logical_or()?);
286                    }
287                    ExpressionParser::consume(parser, Token::RightParen)?;
288
289                    // Expect IN or NOT IN immediately after
290                    match ExpressionParser::current_token(parser) {
291                        Token::In => {
292                            ExpressionParser::advance(parser); // consume IN
293                            ExpressionParser::consume(parser, Token::LeftParen)?;
294                            if !matches!(
295                                ExpressionParser::current_token(parser),
296                                Token::Select | Token::With
297                            ) {
298                                return Err("Tuple IN requires a subquery on the right".to_string());
299                            }
300                            let subquery = parser.parse_subquery()?;
301                            ExpressionParser::consume(parser, Token::RightParen)?;
302                            Ok(SqlExpression::InSubqueryTuple {
303                                exprs,
304                                subquery: Box::new(subquery),
305                            })
306                        }
307                        Token::Not => {
308                            ExpressionParser::advance(parser); // consume NOT
309                            if !matches!(ExpressionParser::current_token(parser), Token::In) {
310                                return Err("Expected IN after NOT for tuple".to_string());
311                            }
312                            ExpressionParser::advance(parser); // consume IN
313                            ExpressionParser::consume(parser, Token::LeftParen)?;
314                            if !matches!(
315                                ExpressionParser::current_token(parser),
316                                Token::Select | Token::With
317                            ) {
318                                return Err(
319                                    "Tuple NOT IN requires a subquery on the right".to_string()
320                                );
321                            }
322                            let subquery = parser.parse_subquery()?;
323                            ExpressionParser::consume(parser, Token::RightParen)?;
324                            Ok(SqlExpression::NotInSubqueryTuple {
325                                exprs,
326                                subquery: Box::new(subquery),
327                            })
328                        }
329                        _ => Err(
330                            "A tuple (expr, expr, ...) may only appear as the left side of IN / NOT IN"
331                                .to_string(),
332                        ),
333                    }
334                } else {
335                    // Regular parenthesized expression
336                    debug!("Regular parenthesized expression");
337                    ExpressionParser::consume(parser, Token::RightParen)?;
338                    Ok(first)
339                }
340            }
341        }
342
343        Token::Not => {
344            debug!("Parsing NOT expression");
345            parse_not_expression(parser)
346        }
347
348        Token::Star => {
349            // Handle * as a literal (like in COUNT(*))
350            trace!("Star token as literal");
351            ExpressionParser::advance(parser);
352            Ok(SqlExpression::StringLiteral("*".to_string()))
353        }
354
355        // Handle window-related keywords that can also be column names
356        Token::Row => {
357            trace!("ROW token treated as identifier in expression context");
358            ExpressionParser::advance(parser);
359            Ok(SqlExpression::Column(ColumnRef::unquoted(
360                "row".to_string(),
361            )))
362        }
363
364        Token::Rows => {
365            trace!("ROWS token treated as identifier in expression context");
366            ExpressionParser::advance(parser);
367            Ok(SqlExpression::Column(ColumnRef::unquoted(
368                "rows".to_string(),
369            )))
370        }
371
372        Token::Range => {
373            trace!("RANGE token treated as identifier in expression context");
374            ExpressionParser::advance(parser);
375            Ok(SqlExpression::Column(ColumnRef::unquoted(
376                "range".to_string(),
377            )))
378        }
379
380        Token::Minus => {
381            // Unary minus: -expr is parsed as 0 - expr
382            debug!("Parsing unary minus expression");
383            ExpressionParser::advance(parser);
384            let operand = parse_primary(parser, ctx)?;
385            Ok(SqlExpression::BinaryOp {
386                left: Box::new(SqlExpression::NumberLiteral("0".to_string())),
387                op: "-".to_string(),
388                right: Box::new(operand),
389            })
390        }
391
392        _ => {
393            let err = format!(
394                "Unexpected token in primary expression: {:?}",
395                ExpressionParser::current_token(parser)
396            );
397            debug!(error = %err);
398            Err(err)
399        }
400    };
401
402    trace_parse_exit("parse_primary", &result);
403    result
404}
405
406/// Parse DateTime constructor
407fn parse_datetime_constructor<P>(parser: &mut P) -> Result<SqlExpression, String>
408where
409    P: ParsePrimary + ExpressionParser + ?Sized,
410{
411    ExpressionParser::advance(parser); // consume DateTime
412    ExpressionParser::consume(parser, Token::LeftParen)?;
413
414    // Check if empty parentheses for DateTime() - today's date
415    if matches!(ExpressionParser::current_token(parser), Token::RightParen) {
416        ExpressionParser::advance(parser); // consume )
417        debug!("DateTime() - today's date");
418        return Ok(SqlExpression::DateTimeToday {
419            hour: None,
420            minute: None,
421            second: None,
422        });
423    }
424
425    // Parse year
426    let year = if let Token::NumberLiteral(n) = ExpressionParser::current_token(parser) {
427        n.parse::<i32>().map_err(|_| "Invalid year")?
428    } else {
429        return Err("Expected year in DateTime constructor".to_string());
430    };
431    ExpressionParser::advance(parser);
432    ExpressionParser::consume(parser, Token::Comma)?;
433
434    // Parse month
435    let month = if let Token::NumberLiteral(n) = ExpressionParser::current_token(parser) {
436        n.parse::<u32>().map_err(|_| "Invalid month")?
437    } else {
438        return Err("Expected month in DateTime constructor".to_string());
439    };
440    ExpressionParser::advance(parser);
441    ExpressionParser::consume(parser, Token::Comma)?;
442
443    // Parse day
444    let day = if let Token::NumberLiteral(n) = ExpressionParser::current_token(parser) {
445        n.parse::<u32>().map_err(|_| "Invalid day")?
446    } else {
447        return Err("Expected day in DateTime constructor".to_string());
448    };
449    ExpressionParser::advance(parser);
450
451    // Check for optional time components
452    let mut hour = None;
453    let mut minute = None;
454    let mut second = None;
455
456    if matches!(ExpressionParser::current_token(parser), Token::Comma) {
457        ExpressionParser::advance(parser); // consume comma
458
459        // Parse hour
460        if let Token::NumberLiteral(n) = ExpressionParser::current_token(parser) {
461            hour = Some(n.parse::<u32>().map_err(|_| "Invalid hour")?);
462            ExpressionParser::advance(parser);
463
464            // Check for minute
465            if matches!(ExpressionParser::current_token(parser), Token::Comma) {
466                ExpressionParser::advance(parser); // consume comma
467
468                if let Token::NumberLiteral(n) = ExpressionParser::current_token(parser) {
469                    minute = Some(n.parse::<u32>().map_err(|_| "Invalid minute")?);
470                    ExpressionParser::advance(parser);
471
472                    // Check for second
473                    if matches!(ExpressionParser::current_token(parser), Token::Comma) {
474                        ExpressionParser::advance(parser); // consume comma
475
476                        if let Token::NumberLiteral(n) = ExpressionParser::current_token(parser) {
477                            second = Some(n.parse::<u32>().map_err(|_| "Invalid second")?);
478                            ExpressionParser::advance(parser);
479                        }
480                    }
481                }
482            }
483        }
484    }
485
486    ExpressionParser::consume(parser, Token::RightParen)?;
487
488    debug!(year = year, month = month, day = day, hour = ?hour, minute = ?minute, second = ?second, "DateTime constructor parsed");
489
490    Ok(SqlExpression::DateTimeConstructor {
491        year,
492        month,
493        day,
494        hour,
495        minute,
496        second,
497    })
498}
499
500/// Parse NOT expression
501fn parse_not_expression<P>(parser: &mut P) -> Result<SqlExpression, String>
502where
503    P: ParsePrimary + ExpressionParser + ?Sized,
504{
505    ExpressionParser::advance(parser); // consume NOT
506
507    // Check if this is a NOT IN expression
508    if let Ok(inner_expr) = parser.parse_comparison() {
509        // After parsing the inner expression, check if we're followed by IN
510        if matches!(ExpressionParser::current_token(parser), Token::In) {
511            debug!("NOT IN expression detected");
512            ExpressionParser::advance(parser); // consume IN
513            ExpressionParser::consume(parser, Token::LeftParen)?;
514            let values = parser.parse_expression_list()?;
515            ExpressionParser::consume(parser, Token::RightParen)?;
516
517            Ok(SqlExpression::NotInList {
518                expr: Box::new(inner_expr),
519                values,
520            })
521        } else {
522            // Regular NOT expression
523            debug!("Regular NOT expression");
524            Ok(SqlExpression::Not {
525                expr: Box::new(inner_expr),
526            })
527        }
528    } else {
529        Err("Expected expression after NOT".to_string())
530    }
531}
532
533/// Parse UNNEST expression
534/// Syntax: UNNEST(column_expr, 'delimiter')
535fn parse_unnest<P>(parser: &mut P) -> Result<SqlExpression, String>
536where
537    P: ParsePrimary + ExpressionParser + ?Sized,
538{
539    debug!("parse_unnest: starting");
540    ExpressionParser::advance(parser); // consume UNNEST
541    ExpressionParser::consume(parser, Token::LeftParen)?;
542
543    // Parse the column expression (first argument)
544    let column = parser.parse_logical_or()?;
545    debug!("parse_unnest: parsed column expression");
546
547    // Expect comma
548    ExpressionParser::consume(parser, Token::Comma)?;
549
550    // Parse the delimiter (second argument - must be a string literal)
551    let delimiter = match ExpressionParser::current_token(parser) {
552        Token::StringLiteral(s) => {
553            let delim = s.clone();
554            ExpressionParser::advance(parser);
555            delim
556        }
557        _ => {
558            return Err("UNNEST delimiter must be a string literal".to_string());
559        }
560    };
561
562    debug!(delimiter = %delimiter, "parse_unnest: parsed delimiter");
563
564    ExpressionParser::consume(parser, Token::RightParen)?;
565
566    debug!("parse_unnest: complete");
567    Ok(SqlExpression::Unnest {
568        column: Box::new(column),
569        delimiter,
570    })
571}
572
573/// Parse a CAST / TRY_CAST expression.
574/// Syntax: `CAST(expr AS type)`.
575/// The current token on entry is the opening `(`. The result is lowered into a
576/// `FunctionCall` so it flows through the existing evaluator and AST machinery:
577/// `CAST(expr, 'TYPE')` where the type name is carried as a string literal.
578fn parse_cast_expression<P>(parser: &mut P, func_name: &str) -> Result<SqlExpression, String>
579where
580    P: ParsePrimary + ExpressionParser + ?Sized,
581{
582    ExpressionParser::advance(parser); // consume (
583
584    let inner = parser.parse_logical_or()?;
585
586    ExpressionParser::consume(parser, Token::As)?;
587
588    let type_name = parse_cast_type_name(parser)?;
589
590    ExpressionParser::consume(parser, Token::RightParen)?;
591
592    let name = if func_name.eq_ignore_ascii_case("TRY_CAST") {
593        "TRY_CAST"
594    } else {
595        "CAST"
596    };
597
598    debug!(target = %type_name, "Parsed CAST expression");
599    Ok(SqlExpression::FunctionCall {
600        name: name.to_string(),
601        args: vec![inner, SqlExpression::StringLiteral(type_name)],
602        distinct: false,
603    })
604}
605
606/// Read a SQL type name for CAST, e.g. `INTEGER`, `VARCHAR`, `DOUBLE`,
607/// `TIMESTAMP`. An optional precision/scale specifier such as `DECIMAL(10, 2)`
608/// or `VARCHAR(50)` is consumed and discarded — we coerce within our own type
609/// confines and do not honour width or scale.
610fn parse_cast_type_name<P>(parser: &mut P) -> Result<String, String>
611where
612    P: ParsePrimary + ExpressionParser + ?Sized,
613{
614    let type_name = match ExpressionParser::current_token(parser) {
615        Token::Identifier(id) => id.clone(),
616        // DATETIME is the one type spelling the lexer reserves as a keyword.
617        Token::DateTime => "DATETIME".to_string(),
618        other => {
619            return Err(format!(
620                "Expected a type name after AS in CAST, got {other:?}"
621            ))
622        }
623    };
624    ExpressionParser::advance(parser);
625
626    // Skip an optional (precision) or (precision, scale) specifier.
627    if matches!(ExpressionParser::current_token(parser), Token::LeftParen) {
628        ExpressionParser::advance(parser); // consume (
629        while !matches!(ExpressionParser::current_token(parser), Token::RightParen) {
630            if matches!(ExpressionParser::current_token(parser), Token::Eof) {
631                return Err("Unterminated type specifier in CAST".to_string());
632            }
633            ExpressionParser::advance(parser);
634        }
635        ExpressionParser::consume(parser, Token::RightParen)?;
636    }
637
638    Ok(type_name)
639}
640
641/// Trait that parsers must implement to use primary expression parsing
642pub trait ParsePrimary {
643    fn current_token(&self) -> &Token;
644    fn advance(&mut self);
645    fn consume(&mut self, expected: Token) -> Result<(), String>;
646
647    // These methods are called from parse_primary
648    fn parse_case_expression(&mut self) -> Result<SqlExpression, String>;
649    fn parse_function_args(&mut self) -> Result<(Vec<SqlExpression>, bool), String>;
650    fn parse_window_spec(&mut self) -> Result<WindowSpec, String>;
651    fn parse_logical_or(&mut self) -> Result<SqlExpression, String>;
652    fn parse_comparison(&mut self) -> Result<SqlExpression, String>;
653    fn parse_expression_list(&mut self) -> Result<Vec<SqlExpression>, String>;
654
655    // For subquery parsing (without parenthesis balance validation)
656    fn parse_subquery(&mut self) -> Result<crate::sql::parser::ast::SelectStatement, String>;
657}