Skip to main content

laminar_sql/parser/
join_parser.rs

1//! Join query analysis and extraction
2//!
3//! This module analyzes JOIN clauses to extract:
4//! - Join type (INNER, LEFT, RIGHT, FULL)
5//! - Key columns for join condition
6//! - Time bounds for stream-stream joins
7//! - Detection of lookup joins vs stream-stream joins
8
9use std::time::Duration;
10
11use sqlparser::ast::{
12    BinaryOperator, Expr, FunctionArg, FunctionArgExpr, FunctionArguments, JoinConstraint,
13    JoinOperator, ObjectName, ObjectNamePart, Select, TableFactor, TableVersion,
14};
15
16use super::window_rewriter::WindowRewriter;
17use super::ParseError;
18
19/// Join type classification
20#[derive(Debug, Clone, Copy, PartialEq, Eq)]
21pub enum JoinType {
22    /// INNER JOIN
23    Inner,
24    /// LEFT \[OUTER\] JOIN
25    Left,
26    /// RIGHT \[OUTER\] JOIN
27    Right,
28    /// FULL \[OUTER\] JOIN
29    Full,
30    /// LEFT SEMI JOIN — emit left rows with at least one match
31    LeftSemi,
32    /// LEFT ANTI JOIN — emit left rows with no match
33    LeftAnti,
34    /// RIGHT SEMI JOIN — emit right rows with at least one match
35    RightSemi,
36    /// RIGHT ANTI JOIN — emit right rows with no match
37    RightAnti,
38    /// ASOF JOIN
39    AsOf,
40}
41
42/// Direction for ASOF JOIN time matching.
43#[derive(Debug, Clone, Copy, PartialEq, Eq)]
44pub enum AsofSqlDirection {
45    /// `left.ts >= right.ts` — find most recent right row
46    Backward,
47    /// `left.ts <= right.ts` — find next right row
48    Forward,
49    /// Match by minimum absolute time difference
50    Nearest,
51}
52
53impl std::fmt::Display for AsofSqlDirection {
54    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
55        match self {
56            AsofSqlDirection::Backward => write!(f, "BACKWARD"),
57            AsofSqlDirection::Forward => write!(f, "FORWARD"),
58            AsofSqlDirection::Nearest => write!(f, "NEAREST"),
59        }
60    }
61}
62
63/// Unresolved time column refs from a BETWEEN clause.
64#[derive(Debug, Clone)]
65struct RawTimeCols {
66    expr_qualifier: String,
67    expr_col: String,
68    low_qualifier: String,
69    low_col: String,
70}
71
72#[derive(Debug, Clone, Copy, PartialEq, Eq)]
73enum JoinSide {
74    Left,
75    Right,
76}
77
78#[derive(Debug, Clone, Copy)]
79struct JoinSides<'a> {
80    left_table: &'a str,
81    right_table: &'a str,
82    left_alias: Option<&'a str>,
83    right_alias: Option<&'a str>,
84}
85
86impl JoinSides<'_> {
87    fn resolve_qualifier(&self, qualifier: &str, context: &str) -> Result<JoinSide, ParseError> {
88        let is_left =
89            qualifier == self.left_table || self.left_alias.is_some_and(|alias| qualifier == alias);
90        let is_right = qualifier == self.right_table
91            || self.right_alias.is_some_and(|alias| qualifier == alias);
92
93        match (is_left, is_right) {
94            (true, false) => Ok(JoinSide::Left),
95            (false, true) => Ok(JoinSide::Right),
96            (true, true) => Err(ParseError::StreamingError(format!(
97                "{context} must use unambiguous left and right join input names; qualifier '{qualifier}' names both inputs"
98            ))),
99            (false, false) => Err(ParseError::StreamingError(format!(
100                "{context} must use unambiguous left and right join input names; qualifier '{qualifier}' names neither input"
101            ))),
102        }
103    }
104}
105
106/// Resolve BETWEEN time columns to `(left_time_col, right_time_col)` using
107/// table qualifiers.
108fn resolve_time_cols(
109    raw: &RawTimeCols,
110    left_table: &str,
111    right_table: &str,
112    left_alias: Option<&str>,
113    right_alias: Option<&str>,
114) -> Result<(String, String), ParseError> {
115    let sides = JoinSides {
116        left_table,
117        right_table,
118        left_alias,
119        right_alias,
120    };
121    let expr_side = sides.resolve_qualifier(&raw.expr_qualifier, "streaming interval timestamp")?;
122    let low_side = sides.resolve_qualifier(&raw.low_qualifier, "streaming interval timestamp")?;
123
124    match (expr_side, low_side) {
125        (JoinSide::Right, JoinSide::Left) => Ok((raw.low_col.clone(), raw.expr_col.clone())),
126        (JoinSide::Left, JoinSide::Right) => Err(ParseError::StreamingError(
127            "streaming interval joins require the right timestamp BETWEEN the left timestamp and left timestamp + interval"
128                .to_string(),
129        )),
130        _ => Err(ParseError::StreamingError(
131            "streaming interval join timestamps must reference opposite join inputs".to_string(),
132        )),
133    }
134}
135
136/// Analysis result for a JOIN clause
137#[derive(Debug, Clone)]
138pub struct JoinAnalysis {
139    /// Type of join (inner, left, right, full)
140    pub join_type: JoinType,
141    /// Left side table name
142    pub left_table: String,
143    /// Right side table name
144    pub right_table: String,
145    /// Left side key column
146    pub left_key_column: String,
147    /// Right side key column
148    pub right_key_column: String,
149    /// Time bound for stream-stream joins (None for lookup joins)
150    pub time_bound: Option<Duration>,
151    /// Whether this is a lookup join (no time bound)
152    pub is_lookup_join: bool,
153    /// Left side alias (if any)
154    pub left_alias: Option<String>,
155    /// Right side alias (if any)
156    pub right_alias: Option<String>,
157    /// Whether this is an ASOF join
158    pub is_asof_join: bool,
159    /// ASOF join direction (Backward or Forward)
160    pub asof_direction: Option<AsofSqlDirection>,
161    /// Left side time column for ASOF join
162    pub left_time_column: Option<String>,
163    /// Right side time column for ASOF join
164    pub right_time_column: Option<String>,
165    /// ASOF join tolerance (max time difference)
166    pub asof_tolerance: Option<Duration>,
167    /// Whether this is a temporal join (FOR SYSTEM_TIME AS OF)
168    pub is_temporal_join: bool,
169    /// The version column from FOR SYSTEM_TIME AS OF (e.g., `order_time`)
170    pub temporal_version_column: Option<String>,
171    /// Additional key columns for composite join keys (beyond the primary key pair)
172    pub additional_key_columns: Vec<(String, String)>,
173}
174
175impl JoinAnalysis {
176    /// Create a stream-stream join analysis
177    #[must_use]
178    pub fn stream_stream(
179        left_table: String,
180        right_table: String,
181        left_key: String,
182        right_key: String,
183        time_bound: Duration,
184        join_type: JoinType,
185    ) -> Self {
186        Self {
187            join_type,
188            left_table,
189            right_table,
190            left_key_column: left_key,
191            right_key_column: right_key,
192            time_bound: Some(time_bound),
193            is_lookup_join: false,
194            left_alias: None,
195            right_alias: None,
196            is_asof_join: false,
197            asof_direction: None,
198            left_time_column: None,
199            right_time_column: None,
200            asof_tolerance: None,
201            is_temporal_join: false,
202            temporal_version_column: None,
203            additional_key_columns: vec![],
204        }
205    }
206
207    /// Create a lookup join analysis
208    #[must_use]
209    pub fn lookup(
210        left_table: String,
211        right_table: String,
212        left_key: String,
213        right_key: String,
214        join_type: JoinType,
215    ) -> Self {
216        Self {
217            join_type,
218            left_table,
219            right_table,
220            left_key_column: left_key,
221            right_key_column: right_key,
222            time_bound: None,
223            is_lookup_join: true,
224            left_alias: None,
225            right_alias: None,
226            is_asof_join: false,
227            asof_direction: None,
228            left_time_column: None,
229            right_time_column: None,
230            asof_tolerance: None,
231            is_temporal_join: false,
232            temporal_version_column: None,
233            additional_key_columns: vec![],
234        }
235    }
236
237    /// Create an ASOF join analysis
238    #[must_use]
239    #[allow(clippy::too_many_arguments)]
240    pub fn asof(
241        left_table: String,
242        right_table: String,
243        left_key: String,
244        right_key: String,
245        direction: AsofSqlDirection,
246        left_time_col: String,
247        right_time_col: String,
248        tolerance: Option<Duration>,
249    ) -> Self {
250        Self {
251            join_type: JoinType::AsOf,
252            left_table,
253            right_table,
254            left_key_column: left_key,
255            right_key_column: right_key,
256            time_bound: None,
257            is_lookup_join: false,
258            left_alias: None,
259            right_alias: None,
260            is_asof_join: true,
261            asof_direction: Some(direction),
262            left_time_column: Some(left_time_col),
263            right_time_column: Some(right_time_col),
264            asof_tolerance: tolerance,
265            is_temporal_join: false,
266            temporal_version_column: None,
267            additional_key_columns: vec![],
268        }
269    }
270
271    /// Create a temporal join analysis (FOR SYSTEM_TIME AS OF).
272    #[must_use]
273    pub fn temporal(
274        left_table: String,
275        right_table: String,
276        left_key: String,
277        right_key: String,
278        version_column: String,
279        join_type: JoinType,
280    ) -> Self {
281        Self {
282            join_type,
283            left_table,
284            right_table,
285            left_key_column: left_key,
286            right_key_column: right_key,
287            time_bound: None,
288            is_lookup_join: false,
289            left_alias: None,
290            right_alias: None,
291            is_asof_join: false,
292            asof_direction: None,
293            left_time_column: None,
294            right_time_column: None,
295            asof_tolerance: None,
296            is_temporal_join: true,
297            temporal_version_column: Some(version_column),
298            additional_key_columns: vec![],
299        }
300    }
301
302    /// True if this step has any kind of temporal bound — a `BETWEEN`-derived
303    /// time bound, ASOF match condition, or `FOR SYSTEM_TIME AS OF`.
304    #[must_use]
305    pub fn is_bounded(&self) -> bool {
306        self.time_bound.is_some() || self.is_asof_join || self.is_temporal_join
307    }
308}
309
310/// Analyze a SELECT statement for join information.
311///
312/// # Errors
313///
314/// Returns `ParseError::StreamingError` if:
315/// - Join constraint is not supported
316/// - Cannot extract key columns
317pub fn analyze_join(select: &Select) -> Result<Option<JoinAnalysis>, ParseError> {
318    let from = &select.from;
319    if from.is_empty() {
320        return Ok(None);
321    }
322
323    let first_table = &from[0];
324    if first_table.joins.is_empty() {
325        return Ok(None);
326    }
327
328    // Extract left table information
329    let left_table = extract_table_name(&first_table.relation)?;
330    let left_alias = extract_table_alias(&first_table.relation);
331
332    // Analyze the first join
333    let join = &first_table.joins[0];
334    let right_table = extract_table_name(&join.relation)?;
335    let right_alias = extract_table_alias(&join.relation);
336
337    let join_type = map_join_operator(&join.join_operator);
338    let sides = JoinSides {
339        left_table: &left_table,
340        right_table: &right_table,
341        left_alias: left_alias.as_deref(),
342        right_alias: right_alias.as_deref(),
343    };
344
345    // Handle ASOF JOIN specially
346    if let JoinOperator::AsOf {
347        match_condition,
348        constraint,
349    } = &join.join_operator
350    {
351        let (direction, left_time, right_time, tolerance) =
352            analyze_asof_match_condition(match_condition, &sides)?;
353
354        // Extract key columns from the ON constraint
355        let (left_key, right_key) = analyze_asof_constraint(constraint, &sides)?;
356
357        let mut analysis = JoinAnalysis::asof(
358            left_table,
359            right_table,
360            left_key,
361            right_key,
362            direction,
363            left_time,
364            right_time,
365            tolerance,
366        );
367        analysis.left_alias = left_alias;
368        analysis.right_alias = right_alias;
369        return Ok(Some(analysis));
370    }
371
372    // Check for temporal join (FOR SYSTEM_TIME AS OF)
373    if let Some(version_col) = extract_temporal_version(&join.relation) {
374        let (left_key, right_key, additional, _, _) =
375            analyze_join_constraint(&join.join_operator, &sides)?;
376        let mut analysis = JoinAnalysis::temporal(
377            left_table,
378            right_table,
379            left_key,
380            right_key,
381            version_col,
382            join_type,
383        );
384        analysis.left_alias = left_alias;
385        analysis.right_alias = right_alias;
386        analysis.additional_key_columns = additional;
387        return Ok(Some(analysis));
388    }
389
390    // Analyze the join constraint
391    let (left_key, right_key, additional, time_bound, time_cols) =
392        analyze_join_constraint(&join.join_operator, &sides)?;
393
394    let mut analysis = if let Some(tb) = time_bound {
395        JoinAnalysis::stream_stream(left_table, right_table, left_key, right_key, tb, join_type)
396    } else {
397        JoinAnalysis::lookup(left_table, right_table, left_key, right_key, join_type)
398    };
399
400    analysis.left_alias.clone_from(&left_alias);
401    analysis.right_alias.clone_from(&right_alias);
402    analysis.additional_key_columns = additional;
403
404    if let Some(ref raw) = time_cols {
405        let (lt, rt) = resolve_time_cols(
406            raw,
407            &analysis.left_table,
408            &analysis.right_table,
409            left_alias.as_deref(),
410            right_alias.as_deref(),
411        )?;
412        analysis.left_time_column = Some(lt);
413        analysis.right_time_column = Some(rt);
414    }
415
416    Ok(Some(analysis))
417}
418
419/// Extract table name from a TableFactor.
420fn extract_table_name(factor: &TableFactor) -> Result<String, ParseError> {
421    match factor {
422        TableFactor::Table { name, .. } => match name.0.as_slice() {
423            [ObjectNamePart::Identifier(ident)] => Ok(ident.value.clone()),
424            _ => Err(ParseError::StreamingError(
425                "streaming joins require single-part relation names; qualify the relation in the catalog before planning the join"
426                    .to_string(),
427            )),
428        },
429        TableFactor::Derived { alias, .. } => {
430            if let Some(alias) = alias {
431                Ok(alias.name.value.clone())
432            } else {
433                Err(ParseError::StreamingError(
434                    "Derived table without alias not supported".to_string(),
435                ))
436            }
437        }
438        _ => Err(ParseError::StreamingError(
439            "Unsupported table factor type".to_string(),
440        )),
441    }
442}
443
444/// Extract the version column from a temporal join's `FOR SYSTEM_TIME AS OF` clause.
445///
446/// Returns `Some(column_name)` if the table factor has a temporal version qualifier,
447/// `None` otherwise.
448fn extract_temporal_version(factor: &TableFactor) -> Option<String> {
449    if let TableFactor::Table {
450        version: Some(TableVersion::ForSystemTimeAsOf(expr)),
451        ..
452    } = factor
453    {
454        Some(extract_column_name_from_expr(expr))
455    } else {
456        None
457    }
458}
459
460/// Extract a column name from an expression (e.g., `o.order_time` → `order_time`).
461///
462/// Falls back to the full expression string for complex expressions.
463fn extract_column_name_from_expr(expr: &Expr) -> String {
464    match expr {
465        Expr::Identifier(ident) => ident.value.clone(),
466        Expr::CompoundIdentifier(parts) => parts
467            .last()
468            .map_or_else(|| expr.to_string(), |p| p.value.clone()),
469        _ => expr.to_string(),
470    }
471}
472
473/// Extract table alias from a TableFactor.
474fn extract_table_alias(factor: &TableFactor) -> Option<String> {
475    match factor {
476        TableFactor::Table { alias, .. } => alias.as_ref().map(|a| a.name.value.clone()),
477        TableFactor::Derived { alias, .. } => alias.as_ref().map(|a| a.name.value.clone()),
478        _ => None,
479    }
480}
481
482/// Map sqlparser `JoinOperator` to our `JoinType`.
483fn map_join_operator(op: &JoinOperator) -> JoinType {
484    match op {
485        JoinOperator::Inner(_) | JoinOperator::Join(_) | JoinOperator::StraightJoin(_) => {
486            JoinType::Inner
487        }
488        JoinOperator::Left(_) | JoinOperator::LeftOuter(_) => JoinType::Left,
489        JoinOperator::LeftSemi(_) | JoinOperator::Semi(_) => JoinType::LeftSemi,
490        JoinOperator::LeftAnti(_) | JoinOperator::Anti(_) => JoinType::LeftAnti,
491        JoinOperator::AsOf { .. } => JoinType::AsOf,
492        JoinOperator::Right(_) | JoinOperator::RightOuter(_) => JoinType::Right,
493        JoinOperator::RightSemi(_) => JoinType::RightSemi,
494        JoinOperator::RightAnti(_) => JoinType::RightAnti,
495        JoinOperator::FullOuter(_) => JoinType::Full,
496        // CrossJoin, CrossApply, OuterApply are rejected by get_join_constraint()
497        _ => JoinType::Inner,
498    }
499}
500
501/// Analyze join constraint to extract key columns, additional key columns,
502/// time bound, and optional time column pair.
503#[allow(clippy::type_complexity)]
504fn analyze_join_constraint(
505    op: &JoinOperator,
506    sides: &JoinSides<'_>,
507) -> Result<
508    (
509        String,
510        String,
511        Vec<(String, String)>,
512        Option<Duration>,
513        Option<RawTimeCols>,
514    ),
515    ParseError,
516> {
517    let constraint = get_join_constraint(op)?;
518
519    match constraint {
520        JoinConstraint::On(expr) => {
521            let (key_pairs, time_bound, time_cols) = analyze_on_expression(expr, sides)?;
522            if key_pairs.is_empty() {
523                return Ok((String::new(), String::new(), vec![], time_bound, time_cols));
524            }
525            let (first_left, first_right) = key_pairs[0].clone();
526            let additional = key_pairs[1..].to_vec();
527            Ok((first_left, first_right, additional, time_bound, time_cols))
528        }
529        JoinConstraint::Using(cols) => {
530            if cols.is_empty() {
531                return Err(ParseError::StreamingError(
532                    "USING clause requires at least one column".to_string(),
533                ));
534            }
535            // First column is the primary key pair
536            let first_col = extract_using_column(&cols[0])?;
537            // Remaining columns are additional key pairs
538            let additional: Vec<(String, String)> = cols[1..]
539                .iter()
540                .map(|c| {
541                    let col = extract_using_column(c)?;
542                    Ok((col.clone(), col))
543                })
544                .collect::<Result<_, ParseError>>()?;
545            Ok((first_col.clone(), first_col, additional, None, None))
546        }
547        JoinConstraint::Natural => Err(ParseError::StreamingError(
548            "NATURAL JOIN not supported for streaming".to_string(),
549        )),
550        JoinConstraint::None => Err(ParseError::StreamingError(
551            "JOIN without condition not supported for streaming".to_string(),
552        )),
553    }
554}
555
556fn extract_using_column(name: &ObjectName) -> Result<String, ParseError> {
557    match name.0.as_slice() {
558        [ObjectNamePart::Identifier(ident)] => Ok(ident.value.clone()),
559        _ => Err(ParseError::StreamingError(
560            "streaming JOIN USING keys must be single column identifiers".to_string(),
561        )),
562    }
563}
564
565/// Get the JoinConstraint from a JoinOperator.
566fn get_join_constraint(op: &JoinOperator) -> Result<&JoinConstraint, ParseError> {
567    match op {
568        JoinOperator::Inner(constraint)
569        | JoinOperator::Join(constraint)
570        | JoinOperator::Left(constraint)
571        | JoinOperator::LeftOuter(constraint)
572        | JoinOperator::Right(constraint)
573        | JoinOperator::RightOuter(constraint)
574        | JoinOperator::FullOuter(constraint)
575        | JoinOperator::LeftSemi(constraint)
576        | JoinOperator::RightSemi(constraint)
577        | JoinOperator::LeftAnti(constraint)
578        | JoinOperator::RightAnti(constraint)
579        | JoinOperator::Semi(constraint)
580        | JoinOperator::Anti(constraint)
581        | JoinOperator::StraightJoin(constraint)
582        | JoinOperator::AsOf { constraint, .. } => Ok(constraint),
583        JoinOperator::CrossJoin(_) | JoinOperator::CrossApply | JoinOperator::OuterApply => Err(
584            ParseError::StreamingError("CROSS JOIN not supported for streaming".to_string()),
585        ),
586    }
587}
588
589/// Analyze ON expression to extract all key column pairs, time bound,
590/// and optional time column pair for stream-stream joins.
591#[allow(clippy::type_complexity)]
592fn analyze_on_expression(
593    expr: &Expr,
594    sides: &JoinSides<'_>,
595) -> Result<(Vec<(String, String)>, Option<Duration>, Option<RawTimeCols>), ParseError> {
596    // Handle compound expressions (AND)
597    match expr {
598        Expr::BinaryOp {
599            left,
600            op: BinaryOperator::And,
601            right,
602        } => {
603            let (mut key_pairs, left_bound, left_cols) = analyze_on_expression(left, sides)?;
604            let (right_keys, right_bound, right_cols) = analyze_on_expression(right, sides)?;
605
606            if left_bound.is_some() && right_bound.is_some() {
607                return Err(ParseError::StreamingError(
608                    "streaming interval joins require exactly one time-bound predicate".to_string(),
609                ));
610            }
611
612            key_pairs.extend(right_keys);
613            Ok((
614                key_pairs,
615                left_bound.or(right_bound),
616                left_cols.or(right_cols),
617            ))
618        }
619        // Equality condition: a.col = b.col
620        Expr::BinaryOp {
621            left,
622            op: BinaryOperator::Eq,
623            right,
624        } => Ok((
625            vec![orient_equality_columns(
626                left,
627                right,
628                sides,
629                "join equality",
630            )?],
631            None,
632            None,
633        )),
634        // BETWEEN clause for time bound: p.ts BETWEEN o.ts AND o.ts + INTERVAL
635        Expr::Between {
636            expr: between_expr,
637            negated,
638            low,
639            high,
640        } => {
641            if *negated {
642                return Err(ParseError::StreamingError(
643                    "NOT BETWEEN is not supported for streaming interval joins".to_string(),
644                ));
645            }
646
647            let (expr_qualifier, expr_col) = extract_qualified_column_ref(between_expr)
648                .ok_or_else(|| {
649                    ParseError::StreamingError(
650                        "streaming interval join timestamps must be qualified column references"
651                            .to_string(),
652                    )
653                })?;
654            let (low_qualifier, low_col) = extract_qualified_column_ref(low).ok_or_else(|| {
655                ParseError::StreamingError(
656                    "streaming interval join timestamps must be qualified column references"
657                        .to_string(),
658                )
659            })?;
660            let time_bound = extract_strict_interval_bound(high, &low_qualifier, &low_col)?;
661
662            Ok((
663                vec![],
664                Some(time_bound),
665                Some(RawTimeCols {
666                    expr_qualifier,
667                    expr_col,
668                    low_qualifier,
669                    low_col,
670                }),
671            ))
672        }
673        Expr::Nested(inner) => analyze_on_expression(inner, sides),
674        _ => Err(ParseError::StreamingError(format!(
675            "Unsupported join condition expression: {expr:?}"
676        ))),
677    }
678}
679
680fn extract_qualified_column_ref(expr: &Expr) -> Option<(String, String)> {
681    match strip_nested(expr) {
682        Expr::CompoundIdentifier(parts) if parts.len() == 2 => {
683            Some((parts[0].value.clone(), parts[1].value.clone()))
684        }
685        _ => None,
686    }
687}
688
689fn orient_equality_columns(
690    expression_left: &Expr,
691    expression_right: &Expr,
692    sides: &JoinSides<'_>,
693    context: &str,
694) -> Result<(String, String), ParseError> {
695    let (left_qualifier, left_column) =
696        extract_qualified_column_ref(expression_left).ok_or_else(|| {
697            ParseError::StreamingError(format!(
698                "Cannot extract column references from {context}; operands must be qualified column references"
699            ))
700        })?;
701    let (right_qualifier, right_column) = extract_qualified_column_ref(expression_right)
702        .ok_or_else(|| {
703            ParseError::StreamingError(format!(
704                "Cannot extract column references from {context}; operands must be qualified column references"
705            ))
706        })?;
707    let left_side = sides.resolve_qualifier(&left_qualifier, context)?;
708    let right_side = sides.resolve_qualifier(&right_qualifier, context)?;
709
710    match (left_side, right_side) {
711        (JoinSide::Left, JoinSide::Right) => Ok((left_column, right_column)),
712        (JoinSide::Right, JoinSide::Left) => Ok((right_column, left_column)),
713        _ => Err(ParseError::StreamingError(format!(
714            "{context} must compare one left-input column with one right-input column"
715        ))),
716    }
717}
718
719fn strip_nested(mut expr: &Expr) -> &Expr {
720    while let Expr::Nested(inner) = expr {
721        expr = inner;
722    }
723    expr
724}
725
726/// Parse the only admitted interval upper bound: the exact lower timestamp
727/// column plus a positive interval.
728fn extract_strict_interval_bound(
729    high: &Expr,
730    low_qualifier: &str,
731    low_col: &str,
732) -> Result<Duration, ParseError> {
733    let Expr::BinaryOp { left, op, right } = strip_nested(high) else {
734        return Err(ParseError::StreamingError(
735            "streaming interval join upper bound must be the left timestamp plus an interval"
736                .to_string(),
737        ));
738    };
739    if !matches!(op, BinaryOperator::Plus) {
740        return Err(ParseError::StreamingError(
741            "streaming interval join upper bound must use addition".to_string(),
742        ));
743    }
744
745    let Some((high_qualifier, high_col)) = extract_qualified_column_ref(left) else {
746        return Err(ParseError::StreamingError(
747            "streaming interval join upper bound must repeat the qualified lower timestamp"
748                .to_string(),
749        ));
750    };
751    if high_qualifier != low_qualifier || high_col != low_col {
752        return Err(ParseError::StreamingError(
753            "streaming interval join upper bound must use the same timestamp as its lower bound"
754                .to_string(),
755        ));
756    }
757
758    let interval = strip_nested(right);
759    if !matches!(interval, Expr::Interval(_)) {
760        return Err(ParseError::StreamingError(
761            "streaming interval join upper bound must end with an INTERVAL".to_string(),
762        ));
763    }
764    let duration = WindowRewriter::parse_interval_to_duration(interval)?;
765    if duration.is_zero() {
766        return Err(ParseError::StreamingError(
767            "streaming interval joins require a positive finite time bound".to_string(),
768        ));
769    }
770    Ok(duration)
771}
772
773/// Analyze ASOF JOIN MATCH_CONDITION expression.
774///
775/// Extracts direction, time column names, and optional tolerance.
776fn analyze_asof_match_condition(
777    expr: &Expr,
778    sides: &JoinSides<'_>,
779) -> Result<(AsofSqlDirection, String, String, Option<Duration>), ParseError> {
780    if let Expr::BinaryOp {
781        left,
782        op: BinaryOperator::And,
783        right,
784    } = strip_nested(expr)
785    {
786        let left_direction = analyze_asof_direction(left, sides);
787        let right_direction = analyze_asof_direction(right, sides);
788        match (left_direction, right_direction) {
789            (Ok(_), Ok(_)) => Err(ParseError::StreamingError(
790                "ASOF MATCH_CONDITION requires exactly one time-direction predicate".to_string(),
791            )),
792            (Ok((direction, left_time, right_time)), Err(_)) => {
793                if direction == AsofSqlDirection::Nearest {
794                    return Err(ParseError::StreamingError(
795                        "ASOF NEAREST does not support a tolerance predicate".to_string(),
796                    ));
797                }
798                let tolerance = extract_asof_tolerance(
799                    right,
800                    sides,
801                    direction,
802                    &left_time,
803                    &right_time,
804                )?;
805                Ok((direction, left_time, right_time, Some(tolerance)))
806            }
807            (Err(_), Ok((direction, left_time, right_time))) => {
808                if direction == AsofSqlDirection::Nearest {
809                    return Err(ParseError::StreamingError(
810                        "ASOF NEAREST does not support a tolerance predicate".to_string(),
811                    ));
812                }
813                let tolerance = extract_asof_tolerance(
814                    left,
815                    sides,
816                    direction,
817                    &left_time,
818                    &right_time,
819                )?;
820                Ok((direction, left_time, right_time, Some(tolerance)))
821            }
822            (Err(_), Err(_)) => Err(ParseError::StreamingError(
823                "ASOF MATCH_CONDITION must contain exactly one qualified >=, <=, or NEAREST time predicate and at most one valid tolerance"
824                    .to_string(),
825            )),
826        }
827    } else {
828        let (dir, lt, rt) = analyze_asof_direction(expr, sides)?;
829        Ok((dir, lt, rt, None))
830    }
831}
832
833/// Extract ASOF direction and time columns from a comparison expression.
834fn analyze_asof_direction(
835    expr: &Expr,
836    sides: &JoinSides<'_>,
837) -> Result<(AsofSqlDirection, String, String), ParseError> {
838    match strip_nested(expr) {
839        Expr::BinaryOp { left, op, right }
840            if matches!(op, BinaryOperator::GtEq | BinaryOperator::LtEq) =>
841        {
842            let (expression_left_qualifier, expression_left_column) =
843                extract_qualified_column_ref(left).ok_or_else(|| {
844                    ParseError::StreamingError(
845                        "ASOF time operands must be qualified column references".to_string(),
846                    )
847                })?;
848            let (expression_right_qualifier, expression_right_column) =
849                extract_qualified_column_ref(right).ok_or_else(|| {
850                    ParseError::StreamingError(
851                        "ASOF time operands must be qualified column references".to_string(),
852                    )
853                })?;
854            let expression_left_side =
855                sides.resolve_qualifier(&expression_left_qualifier, "ASOF time predicate")?;
856            let expression_right_side =
857                sides.resolve_qualifier(&expression_right_qualifier, "ASOF time predicate")?;
858
859            let (left_time, right_time, left_first) =
860                match (expression_left_side, expression_right_side) {
861                    (JoinSide::Left, JoinSide::Right) => {
862                        (expression_left_column, expression_right_column, true)
863                    }
864                    (JoinSide::Right, JoinSide::Left) => {
865                        (expression_right_column, expression_left_column, false)
866                    }
867                    _ => {
868                        return Err(ParseError::StreamingError(
869                            "ASOF time predicate must compare one left-input timestamp with one right-input timestamp"
870                                .to_string(),
871                        ))
872                    }
873                };
874            let direction = match (op, left_first) {
875                (BinaryOperator::GtEq, true) | (BinaryOperator::LtEq, false) => {
876                    AsofSqlDirection::Backward
877                }
878                (BinaryOperator::LtEq, true) | (BinaryOperator::GtEq, false) => {
879                    AsofSqlDirection::Forward
880                }
881                _ => unreachable!("comparison operator was restricted above"),
882            };
883            Ok((direction, left_time, right_time))
884        }
885        Expr::Function(func) => {
886            let name = func.name.to_string().to_uppercase();
887            if name != "NEAREST" {
888                return Err(ParseError::StreamingError(format!(
889                    "Unknown ASOF MATCH_CONDITION function: {name}"
890                )));
891            }
892            let args = match &func.args {
893                FunctionArguments::List(arg_list) => &arg_list.args,
894                _ => {
895                    return Err(ParseError::StreamingError(
896                        "NEAREST() requires exactly 2 column arguments".to_string(),
897                    ))
898                }
899            };
900            if args.len() != 2 {
901                return Err(ParseError::StreamingError(format!(
902                    "NEAREST() requires exactly 2 arguments, got {}",
903                    args.len()
904                )));
905            }
906            let first = extract_expr_from_function_arg(&args[0]).ok_or_else(|| {
907                ParseError::StreamingError("NEAREST arguments must be columns".to_string())
908            })?;
909            let second = extract_expr_from_function_arg(&args[1]).ok_or_else(|| {
910                ParseError::StreamingError("NEAREST arguments must be columns".to_string())
911            })?;
912            let (left_col, right_col) =
913                orient_equality_columns(first, second, sides, "ASOF NEAREST")?;
914            Ok((AsofSqlDirection::Nearest, left_col, right_col))
915        }
916        _ => Err(ParseError::StreamingError(
917            "ASOF MATCH_CONDITION must be >= or <= comparison, or NEAREST()".to_string(),
918        )),
919    }
920}
921
922fn extract_expr_from_function_arg(arg: &FunctionArg) -> Option<&Expr> {
923    let (FunctionArg::Unnamed(FunctionArgExpr::Expr(expr))
924    | FunctionArg::Named {
925        arg: FunctionArgExpr::Expr(expr),
926        ..
927    }
928    | FunctionArg::ExprNamed {
929        arg: FunctionArgExpr::Expr(expr),
930        ..
931    }) = arg
932    else {
933        return None;
934    };
935    Some(expr)
936}
937
938/// Extract tolerance duration from an ASOF tolerance expression.
939///
940/// Handles: `left - right <= value` or `left - right <= INTERVAL '...'`
941fn extract_asof_tolerance(
942    expr: &Expr,
943    sides: &JoinSides<'_>,
944    direction: AsofSqlDirection,
945    left_time: &str,
946    right_time: &str,
947) -> Result<Duration, ParseError> {
948    let Expr::BinaryOp {
949        left: delta,
950        op: BinaryOperator::LtEq,
951        right: bound,
952    } = strip_nested(expr)
953    else {
954        return Err(ParseError::StreamingError(
955            "ASOF tolerance expression must be <= comparison".to_string(),
956        ));
957    };
958    let Expr::BinaryOp {
959        left: delta_left,
960        op: BinaryOperator::Minus,
961        right: delta_right,
962    } = strip_nested(delta)
963    else {
964        return Err(ParseError::StreamingError(
965            "ASOF tolerance left side must subtract the matched timestamps".to_string(),
966        ));
967    };
968    let (delta_left_qualifier, delta_left_column) = extract_qualified_column_ref(delta_left)
969        .ok_or_else(|| {
970            ParseError::StreamingError(
971                "ASOF tolerance timestamps must be qualified column references".to_string(),
972            )
973        })?;
974    let (delta_right_qualifier, delta_right_column) = extract_qualified_column_ref(delta_right)
975        .ok_or_else(|| {
976            ParseError::StreamingError(
977                "ASOF tolerance timestamps must be qualified column references".to_string(),
978            )
979        })?;
980    let delta_left_side = sides.resolve_qualifier(&delta_left_qualifier, "ASOF tolerance")?;
981    let delta_right_side = sides.resolve_qualifier(&delta_right_qualifier, "ASOF tolerance")?;
982    let matches_direction = match direction {
983        AsofSqlDirection::Backward => {
984            delta_left_side == JoinSide::Left
985                && delta_left_column == left_time
986                && delta_right_side == JoinSide::Right
987                && delta_right_column == right_time
988        }
989        AsofSqlDirection::Forward => {
990            delta_left_side == JoinSide::Right
991                && delta_left_column == right_time
992                && delta_right_side == JoinSide::Left
993                && delta_right_column == left_time
994        }
995        AsofSqlDirection::Nearest => false,
996    };
997    if !matches_direction {
998        return Err(ParseError::StreamingError(
999            "ASOF tolerance must subtract the same timestamps in the match direction".to_string(),
1000        ));
1001    }
1002
1003    let duration = match strip_nested(bound) {
1004        Expr::Value(v) => {
1005            let sqlparser::ast::Value::Number(number, _) = &v.value else {
1006                return Err(ParseError::StreamingError(
1007                    "ASOF tolerance must be a positive number of milliseconds or INTERVAL"
1008                        .to_string(),
1009                ));
1010            };
1011            let milliseconds: u64 = number.parse().map_err(|_| {
1012                ParseError::StreamingError(format!(
1013                    "ASOF tolerance is not a valid millisecond count: {number}"
1014                ))
1015            })?;
1016            Duration::from_millis(milliseconds)
1017        }
1018        interval @ Expr::Interval(_) => WindowRewriter::parse_interval_to_duration(interval)?,
1019        _ => {
1020            return Err(ParseError::StreamingError(
1021                "ASOF tolerance must be a positive number of milliseconds or INTERVAL".to_string(),
1022            ))
1023        }
1024    };
1025    if duration.is_zero() {
1026        return Err(ParseError::StreamingError(
1027            "ASOF tolerance must be positive".to_string(),
1028        ));
1029    }
1030    Ok(duration)
1031}
1032
1033/// Extract key columns from an ASOF JOIN constraint (ON clause).
1034fn analyze_asof_constraint(
1035    constraint: &JoinConstraint,
1036    sides: &JoinSides<'_>,
1037) -> Result<(String, String), ParseError> {
1038    match constraint {
1039        JoinConstraint::On(expr) => {
1040            let Expr::BinaryOp {
1041                left,
1042                op: BinaryOperator::Eq,
1043                right,
1044            } = strip_nested(expr)
1045            else {
1046                return Err(ParseError::StreamingError(
1047                    "ASOF JOIN ON requires exactly one equality condition".to_string(),
1048                ));
1049            };
1050            orient_equality_columns(left, right, sides, "ASOF join equality")
1051        }
1052        JoinConstraint::Using(cols) => {
1053            if cols.len() != 1 {
1054                return Err(ParseError::StreamingError(
1055                    "ASOF JOIN USING requires exactly one key column".to_string(),
1056                ));
1057            }
1058            let col = extract_using_column(&cols[0])?;
1059            Ok((col.clone(), col))
1060        }
1061        _ => Err(ParseError::StreamingError(
1062            "ASOF JOIN requires ON or USING constraint".to_string(),
1063        )),
1064    }
1065}
1066
1067/// Check if a SELECT contains a join.
1068#[must_use]
1069pub fn has_join(select: &Select) -> bool {
1070    !select.from.is_empty() && !select.from[0].joins.is_empty()
1071}
1072
1073/// Count the number of joins in a SELECT.
1074#[must_use]
1075pub fn count_joins(select: &Select) -> usize {
1076    select
1077        .from
1078        .iter()
1079        .map(|table_with_joins| table_with_joins.joins.len())
1080        .sum()
1081}
1082
1083/// Analysis result for multi-way JOINs (e.g., `A JOIN B ... JOIN C ...`).
1084///
1085/// Each step represents one left-deep join: step 0 joins the base table with
1086/// the first right table, step 1 joins the result with the next right table, etc.
1087#[derive(Debug, Clone)]
1088pub struct MultiJoinAnalysis {
1089    /// Ordered join steps (left-to-right)
1090    pub joins: Vec<JoinAnalysis>,
1091    /// All referenced tables in order (base table first, then each right table)
1092    pub tables: Vec<String>,
1093}
1094
1095impl MultiJoinAnalysis {
1096    /// Number of join steps.
1097    #[must_use]
1098    pub fn len(&self) -> usize {
1099        self.joins.len()
1100    }
1101
1102    /// Whether there are no join steps.
1103    #[must_use]
1104    pub fn is_empty(&self) -> bool {
1105        self.joins.is_empty()
1106    }
1107
1108    /// Whether this is a single join (backward-compatible case).
1109    #[must_use]
1110    pub fn is_single(&self) -> bool {
1111        self.joins.len() == 1
1112    }
1113
1114    /// The first join step (convenience for single-join queries).
1115    #[must_use]
1116    pub fn first(&self) -> Option<&JoinAnalysis> {
1117        self.joins.first()
1118    }
1119}
1120
1121/// Analyze a SELECT statement for all join steps (multi-way).
1122///
1123/// Returns `None` if the query has no joins. For a single join this
1124/// returns a `MultiJoinAnalysis` with one step, making it backward
1125/// compatible with `analyze_join()`.
1126///
1127/// # Errors
1128///
1129/// Returns `ParseError::StreamingError` if any join constraint is
1130/// not supported or key columns cannot be extracted.
1131pub fn analyze_joins(select: &Select) -> Result<Option<MultiJoinAnalysis>, ParseError> {
1132    let from = &select.from;
1133    if from.is_empty() {
1134        return Ok(None);
1135    }
1136
1137    let first_table = &from[0];
1138    if first_table.joins.is_empty() {
1139        return Ok(None);
1140    }
1141
1142    // Extract base table
1143    let base_table = extract_table_name(&first_table.relation)?;
1144    let base_alias = extract_table_alias(&first_table.relation);
1145
1146    let mut join_steps = Vec::with_capacity(first_table.joins.len());
1147    let mut tables = vec![base_table.clone()];
1148
1149    // Track the left table name for left-deep chaining
1150    let mut prev_left_table = base_table;
1151    let mut prev_left_alias = base_alias;
1152
1153    for join in &first_table.joins {
1154        let right_table = extract_table_name(&join.relation)?;
1155        let right_alias = extract_table_alias(&join.relation);
1156        tables.push(right_table.clone());
1157
1158        let join_type = map_join_operator(&join.join_operator);
1159        let sides = JoinSides {
1160            left_table: &prev_left_table,
1161            right_table: &right_table,
1162            left_alias: prev_left_alias.as_deref(),
1163            right_alias: right_alias.as_deref(),
1164        };
1165
1166        // Handle ASOF JOIN
1167        if let JoinOperator::AsOf {
1168            match_condition,
1169            constraint,
1170        } = &join.join_operator
1171        {
1172            let (direction, left_time, right_time, tolerance) =
1173                analyze_asof_match_condition(match_condition, &sides)?;
1174            let (left_key, right_key) = analyze_asof_constraint(constraint, &sides)?;
1175
1176            let mut analysis = JoinAnalysis::asof(
1177                prev_left_table.clone(),
1178                right_table.clone(),
1179                left_key,
1180                right_key,
1181                direction,
1182                left_time,
1183                right_time,
1184                tolerance,
1185            );
1186            analysis.left_alias.clone_from(&prev_left_alias);
1187            analysis.right_alias = right_alias;
1188            join_steps.push(analysis);
1189        } else if let Some(version_col) = extract_temporal_version(&join.relation) {
1190            // Temporal join: right side has FOR SYSTEM_TIME AS OF
1191            let (left_key, right_key, additional, _, _) =
1192                analyze_join_constraint(&join.join_operator, &sides)?;
1193
1194            let mut analysis = JoinAnalysis::temporal(
1195                prev_left_table.clone(),
1196                right_table.clone(),
1197                left_key,
1198                right_key,
1199                version_col,
1200                join_type,
1201            );
1202            analysis.left_alias.clone_from(&prev_left_alias);
1203            analysis.right_alias = right_alias;
1204            analysis.additional_key_columns = additional;
1205            join_steps.push(analysis);
1206        } else {
1207            // Regular join (inner, left, right, full)
1208            let (left_key, right_key, additional, time_bound, time_cols) =
1209                analyze_join_constraint(&join.join_operator, &sides)?;
1210
1211            let mut analysis = if let Some(tb) = time_bound {
1212                JoinAnalysis::stream_stream(
1213                    prev_left_table.clone(),
1214                    right_table.clone(),
1215                    left_key,
1216                    right_key,
1217                    tb,
1218                    join_type,
1219                )
1220            } else {
1221                JoinAnalysis::lookup(
1222                    prev_left_table.clone(),
1223                    right_table.clone(),
1224                    left_key,
1225                    right_key,
1226                    join_type,
1227                )
1228            };
1229            analysis.left_alias.clone_from(&prev_left_alias);
1230            analysis.right_alias.clone_from(&right_alias);
1231            analysis.additional_key_columns = additional;
1232
1233            if let Some(ref raw) = time_cols {
1234                let (lt, rt) = resolve_time_cols(
1235                    raw,
1236                    &analysis.left_table,
1237                    &analysis.right_table,
1238                    prev_left_alias.as_deref(),
1239                    right_alias.as_deref(),
1240                )?;
1241                analysis.left_time_column = Some(lt);
1242                analysis.right_time_column = Some(rt);
1243            }
1244            join_steps.push(analysis);
1245        }
1246
1247        // Next step's left table is this step's right table (left-deep)
1248        prev_left_table = right_table;
1249        prev_left_alias = extract_table_alias(&join.relation);
1250    }
1251
1252    Ok(Some(MultiJoinAnalysis {
1253        joins: join_steps,
1254        tables,
1255    }))
1256}
1257
1258#[cfg(test)]
1259mod tests {
1260    use super::*;
1261    use sqlparser::ast::{SetExpr, Statement};
1262    use sqlparser::dialect::GenericDialect;
1263    use sqlparser::parser::Parser;
1264
1265    fn parse_select(sql: &str) -> Select {
1266        let dialect = GenericDialect {};
1267        let statements = Parser::parse_sql(&dialect, sql).unwrap();
1268        if let Statement::Query(query) = &statements[0] {
1269            if let SetExpr::Select(select) = query.body.as_ref() {
1270                return *select.clone();
1271            }
1272        }
1273        panic!("Expected SELECT query");
1274    }
1275
1276    fn join_error(sql: &str) -> String {
1277        analyze_join(&parse_select(sql)).unwrap_err().to_string()
1278    }
1279
1280    #[test]
1281    fn test_analyze_inner_join() {
1282        let sql = "SELECT * FROM orders o INNER JOIN payments p ON o.order_id = p.order_id";
1283        let select = parse_select(sql);
1284
1285        let analysis = analyze_join(&select).unwrap().unwrap();
1286
1287        assert_eq!(analysis.join_type, JoinType::Inner);
1288        assert_eq!(analysis.left_table, "orders");
1289        assert_eq!(analysis.right_table, "payments");
1290        assert_eq!(analysis.left_key_column, "order_id");
1291        assert_eq!(analysis.right_key_column, "order_id");
1292        assert!(analysis.is_lookup_join); // No time bound = lookup join
1293    }
1294
1295    #[test]
1296    fn test_analyze_left_join() {
1297        let sql = "SELECT * FROM orders o LEFT JOIN customers c ON o.customer_id = c.id";
1298        let select = parse_select(sql);
1299
1300        let analysis = analyze_join(&select).unwrap().unwrap();
1301
1302        assert_eq!(analysis.join_type, JoinType::Left);
1303        assert_eq!(analysis.left_key_column, "customer_id");
1304        assert_eq!(analysis.right_key_column, "id");
1305    }
1306
1307    #[test]
1308    fn test_analyze_join_using() {
1309        let sql = "SELECT * FROM orders o JOIN payments p USING (order_id)";
1310        let select = parse_select(sql);
1311
1312        let analysis = analyze_join(&select).unwrap().unwrap();
1313
1314        assert_eq!(analysis.left_key_column, "order_id");
1315        assert_eq!(analysis.right_key_column, "order_id");
1316    }
1317
1318    #[test]
1319    fn test_analyze_stream_stream_join_with_time_bound() {
1320        let sql = "SELECT * FROM orders o
1321                   JOIN payments p ON o.order_id = p.order_id
1322                   AND p.ts BETWEEN o.ts AND o.ts + INTERVAL '1' HOUR";
1323        let select = parse_select(sql);
1324
1325        let analysis = analyze_join(&select).unwrap().unwrap();
1326
1327        assert!(!analysis.is_lookup_join);
1328        assert!(analysis.time_bound.is_some());
1329        assert_eq!(analysis.time_bound.unwrap(), Duration::from_secs(3600));
1330        assert_eq!(analysis.left_time_column.as_deref(), Some("ts"));
1331        assert_eq!(analysis.right_time_column.as_deref(), Some("ts"));
1332    }
1333
1334    #[test]
1335    fn test_interval_join_accepts_table_qualifiers() {
1336        let sql = "SELECT * FROM orders
1337                   JOIN payments ON orders.order_id = payments.order_id
1338                   AND payments.received_at BETWEEN orders.created_at
1339                       AND orders.created_at + INTERVAL '250' MILLISECOND";
1340
1341        let analysis = analyze_join(&parse_select(sql)).unwrap().unwrap();
1342
1343        assert_eq!(analysis.time_bound, Some(Duration::from_millis(250)));
1344        assert_eq!(analysis.left_time_column.as_deref(), Some("created_at"));
1345        assert_eq!(analysis.right_time_column.as_deref(), Some("received_at"));
1346    }
1347
1348    #[test]
1349    fn test_interval_join_preserves_composite_equality_keys() {
1350        let sql = "SELECT * FROM orders o JOIN payments p
1351                   ON o.tenant_id = p.tenant_id
1352                   AND o.order_id = p.order_id
1353                   AND p.ts BETWEEN o.ts AND o.ts + INTERVAL '1' SECOND";
1354
1355        let analysis = analyze_join(&parse_select(sql)).unwrap().unwrap();
1356
1357        assert_eq!(analysis.left_key_column, "tenant_id");
1358        assert_eq!(analysis.right_key_column, "tenant_id");
1359        assert_eq!(
1360            analysis.additional_key_columns,
1361            vec![("order_id".to_string(), "order_id".to_string())]
1362        );
1363    }
1364
1365    #[test]
1366    fn test_join_orients_reversed_different_name_keys() {
1367        let analysis = analyze_join(&parse_select(
1368            "SELECT * FROM orders o JOIN payments p ON p.order_id = o.id",
1369        ))
1370        .unwrap()
1371        .unwrap();
1372
1373        assert_eq!(analysis.left_key_column, "id");
1374        assert_eq!(analysis.right_key_column, "order_id");
1375    }
1376
1377    #[test]
1378    fn test_composite_join_orients_each_key_independently() {
1379        let analysis = analyze_join(&parse_select(
1380            "SELECT * FROM orders o JOIN payments p
1381             ON p.order_id = o.id AND o.tenant = p.account",
1382        ))
1383        .unwrap()
1384        .unwrap();
1385
1386        assert_eq!(analysis.left_key_column, "id");
1387        assert_eq!(analysis.right_key_column, "order_id");
1388        assert_eq!(
1389            analysis.additional_key_columns,
1390            vec![("tenant".to_string(), "account".to_string())]
1391        );
1392    }
1393
1394    #[test]
1395    fn test_join_accepts_quoted_relation_and_key_identity() {
1396        let analysis = analyze_join(&parse_select(
1397            "SELECT * FROM \"Orders\"
1398             JOIN \"Payments\"
1399             ON \"Payments\".\"order id\" = \"Orders\".\"id\"",
1400        ))
1401        .unwrap()
1402        .unwrap();
1403
1404        assert_eq!(analysis.left_table, "Orders");
1405        assert_eq!(analysis.right_table, "Payments");
1406        assert_eq!(analysis.left_key_column, "id");
1407        assert_eq!(analysis.right_key_column, "order id");
1408    }
1409
1410    #[test]
1411    fn test_join_accepts_quoted_alias_identity() {
1412        let analysis = analyze_join(&parse_select(
1413            "SELECT * FROM orders AS \"left input\"
1414             JOIN payments AS \"right input\"
1415             ON \"left input\".id = \"right input\".order_id",
1416        ))
1417        .unwrap()
1418        .unwrap();
1419
1420        assert_eq!(analysis.left_alias.as_deref(), Some("left input"));
1421        assert_eq!(analysis.right_alias.as_deref(), Some("right input"));
1422        assert_eq!(analysis.left_key_column, "id");
1423        assert_eq!(analysis.right_key_column, "order_id");
1424    }
1425
1426    #[test]
1427    fn test_join_rejects_unqualified_key() {
1428        let error = join_error("SELECT * FROM orders o JOIN payments p ON id = p.order_id");
1429        assert!(error.contains("qualified column references"), "{error}");
1430    }
1431
1432    #[test]
1433    fn test_join_rejects_unknown_key_qualifier() {
1434        let error = join_error("SELECT * FROM orders o JOIN payments p ON missing.id = p.order_id");
1435        assert!(error.contains("names neither input"), "{error}");
1436    }
1437
1438    #[test]
1439    fn test_join_rejects_same_side_key_expression() {
1440        let error = join_error("SELECT * FROM orders o JOIN payments p ON o.id = o.parent_id");
1441        assert!(error.contains("one left-input column"), "{error}");
1442    }
1443
1444    #[test]
1445    fn test_join_rejects_ambiguous_qualifier() {
1446        let error = join_error(
1447            "SELECT * FROM orders duplicate JOIN payments duplicate
1448             ON duplicate.id = payments.order_id",
1449        );
1450        assert!(error.contains("names both inputs"), "{error}");
1451    }
1452
1453    #[test]
1454    fn test_join_rejects_compound_relation_identity() {
1455        let error = join_error(
1456            "SELECT * FROM catalog.orders JOIN payments
1457             ON catalog.orders.id = payments.order_id",
1458        );
1459        assert!(error.contains("single-part relation names"), "{error}");
1460    }
1461
1462    #[test]
1463    fn test_interval_join_rejects_unqualified_timestamp() {
1464        let error = join_error(
1465            "SELECT * FROM orders o JOIN payments p ON o.id = p.id
1466             AND p.ts BETWEEN ts AND ts + INTERVAL '1' SECOND",
1467        );
1468        assert!(error.contains("qualified column references"), "{error}");
1469    }
1470
1471    #[test]
1472    fn test_interval_join_rejects_unknown_qualifier() {
1473        let error = join_error(
1474            "SELECT * FROM orders o JOIN payments p ON o.id = p.id
1475             AND unknown.ts BETWEEN o.ts AND o.ts + INTERVAL '1' SECOND",
1476        );
1477        assert!(error.contains("unambiguous left and right"), "{error}");
1478    }
1479
1480    #[test]
1481    fn test_interval_join_rejects_reversed_timestamps() {
1482        let error = join_error(
1483            "SELECT * FROM orders o JOIN payments p ON o.id = p.id
1484             AND o.ts BETWEEN p.ts AND p.ts + INTERVAL '1' SECOND",
1485        );
1486        assert!(
1487            error.contains("right timestamp BETWEEN the left"),
1488            "{error}"
1489        );
1490    }
1491
1492    #[test]
1493    fn test_interval_join_rejects_not_between() {
1494        let error = join_error(
1495            "SELECT * FROM orders o JOIN payments p ON o.id = p.id
1496             AND p.ts NOT BETWEEN o.ts AND o.ts + INTERVAL '1' SECOND",
1497        );
1498        assert!(error.contains("NOT BETWEEN"), "{error}");
1499    }
1500
1501    #[test]
1502    fn test_interval_join_rejects_mismatched_upper_timestamp() {
1503        let error = join_error(
1504            "SELECT * FROM orders o JOIN payments p ON o.id = p.id
1505             AND p.ts BETWEEN o.ts AND o.other_ts + INTERVAL '1' SECOND",
1506        );
1507        assert!(error.contains("same timestamp"), "{error}");
1508    }
1509
1510    #[test]
1511    fn test_interval_join_rejects_subtracted_bound() {
1512        let error = join_error(
1513            "SELECT * FROM orders o JOIN payments p ON o.id = p.id
1514             AND p.ts BETWEEN o.ts AND o.ts - INTERVAL '1' SECOND",
1515        );
1516        assert!(error.contains("must use addition"), "{error}");
1517    }
1518
1519    #[test]
1520    fn test_interval_join_rejects_direct_interval_upper_bound() {
1521        let error = join_error(
1522            "SELECT * FROM orders o JOIN payments p ON o.id = p.id
1523             AND p.ts BETWEEN o.ts AND INTERVAL '1' SECOND",
1524        );
1525        assert!(error.contains("left timestamp plus an interval"), "{error}");
1526    }
1527
1528    #[test]
1529    fn test_interval_join_rejects_zero_bound() {
1530        let error = join_error(
1531            "SELECT * FROM orders o JOIN payments p ON o.id = p.id
1532             AND p.ts BETWEEN o.ts AND o.ts + INTERVAL '0' SECOND",
1533        );
1534        assert!(error.contains("positive finite"), "{error}");
1535    }
1536
1537    #[test]
1538    fn test_interval_join_rejects_negative_bound() {
1539        let error = join_error(
1540            "SELECT * FROM orders o JOIN payments p ON o.id = p.id
1541             AND p.ts BETWEEN o.ts AND o.ts + INTERVAL '-1' SECOND",
1542        );
1543        assert!(error.contains("Invalid interval value"), "{error}");
1544    }
1545
1546    #[test]
1547    fn test_interval_join_rejects_multiple_time_bounds() {
1548        let error = join_error(
1549            "SELECT * FROM orders o JOIN payments p ON o.id = p.id
1550             AND p.ts BETWEEN o.ts AND o.ts + INTERVAL '1' SECOND
1551             AND p.created_at BETWEEN o.created_at
1552                 AND o.created_at + INTERVAL '1' SECOND",
1553        );
1554        assert!(
1555            error.contains("exactly one time-bound predicate"),
1556            "{error}"
1557        );
1558    }
1559
1560    #[test]
1561    fn test_join_rejects_non_equi_residual_conjunct() {
1562        let error = join_error(
1563            "SELECT * FROM orders o JOIN payments p
1564             ON o.id = p.id AND p.amount > o.amount",
1565        );
1566        assert!(error.contains("Unsupported join condition"), "{error}");
1567    }
1568
1569    #[test]
1570    fn test_join_rejects_unsupported_equality_conjunct() {
1571        let error = join_error(
1572            "SELECT * FROM orders o JOIN payments p
1573             ON o.id = p.id AND ABS(o.amount) = ABS(p.amount)",
1574        );
1575        assert!(
1576            error.contains("Cannot extract column references"),
1577            "{error}"
1578        );
1579    }
1580
1581    #[test]
1582    fn test_no_join() {
1583        let sql = "SELECT * FROM orders";
1584        let select = parse_select(sql);
1585
1586        let analysis = analyze_join(&select).unwrap();
1587        assert!(analysis.is_none());
1588    }
1589
1590    #[test]
1591    fn test_has_join() {
1592        let sql_with_join = "SELECT * FROM orders o JOIN payments p ON o.id = p.order_id";
1593        let sql_without_join = "SELECT * FROM orders";
1594
1595        let select_with = parse_select(sql_with_join);
1596        let select_without = parse_select(sql_without_join);
1597
1598        assert!(has_join(&select_with));
1599        assert!(!has_join(&select_without));
1600    }
1601
1602    #[test]
1603    fn test_count_joins() {
1604        let sql_one = "SELECT * FROM a JOIN b ON a.id = b.id";
1605        let sql_two = "SELECT * FROM a JOIN b ON a.id = b.id JOIN c ON b.id = c.id";
1606        let sql_zero = "SELECT * FROM a";
1607
1608        assert_eq!(count_joins(&parse_select(sql_one)), 1);
1609        assert_eq!(count_joins(&parse_select(sql_two)), 2);
1610        assert_eq!(count_joins(&parse_select(sql_zero)), 0);
1611    }
1612
1613    #[test]
1614    fn test_aliases() {
1615        let sql = "SELECT * FROM orders AS o JOIN payments AS p ON o.id = p.order_id";
1616        let select = parse_select(sql);
1617
1618        let analysis = analyze_join(&select).unwrap().unwrap();
1619
1620        assert_eq!(analysis.left_alias, Some("o".to_string()));
1621        assert_eq!(analysis.right_alias, Some("p".to_string()));
1622    }
1623
1624    // -- ASOF JOIN tests --
1625
1626    fn parse_select_snowflake(sql: &str) -> Select {
1627        let dialect = sqlparser::dialect::SnowflakeDialect {};
1628        let statements = Parser::parse_sql(&dialect, sql).unwrap();
1629        if let Statement::Query(query) = &statements[0] {
1630            if let SetExpr::Select(select) = query.body.as_ref() {
1631                return *select.clone();
1632            }
1633        }
1634        panic!("Expected SELECT query");
1635    }
1636
1637    fn parse_select_laminar(sql: &str) -> Select {
1638        let dialect = crate::parser::dialect::LaminarDialect::default();
1639        let statements = Parser::parse_sql(&dialect, sql).unwrap();
1640        if let Statement::Query(query) = &statements[0] {
1641            if let SetExpr::Select(select) = query.body.as_ref() {
1642                return *select.clone();
1643            }
1644        }
1645        panic!("Expected SELECT query");
1646    }
1647
1648    fn asof_join_error(sql: &str) -> String {
1649        analyze_join(&parse_select_snowflake(sql))
1650            .unwrap_err()
1651            .to_string()
1652    }
1653
1654    #[test]
1655    fn test_asof_join_backward() {
1656        let sql = "SELECT * FROM trades t \
1657                    ASOF JOIN quotes q \
1658                    MATCH_CONDITION(t.ts >= q.ts) \
1659                    ON t.symbol = q.symbol";
1660        let select = parse_select_snowflake(sql);
1661        let analysis = analyze_join(&select).unwrap().unwrap();
1662
1663        assert!(analysis.is_asof_join);
1664        assert_eq!(analysis.asof_direction, Some(AsofSqlDirection::Backward));
1665        assert_eq!(analysis.join_type, JoinType::AsOf);
1666        assert!(analysis.asof_tolerance.is_none());
1667    }
1668
1669    #[test]
1670    fn test_asof_join_forward() {
1671        let sql = "SELECT * FROM trades t \
1672                    ASOF JOIN quotes q \
1673                    MATCH_CONDITION(t.ts <= q.ts) \
1674                    ON t.symbol = q.symbol";
1675        let select = parse_select_snowflake(sql);
1676        let analysis = analyze_join(&select).unwrap().unwrap();
1677
1678        assert!(analysis.is_asof_join);
1679        assert_eq!(analysis.asof_direction, Some(AsofSqlDirection::Forward));
1680    }
1681
1682    #[test]
1683    fn test_asof_join_orients_reversed_time_and_key_operands() {
1684        let sql = "SELECT * FROM trades t
1685                   ASOF JOIN quotes q
1686                   MATCH_CONDITION(q.ts <= t.trade_ts)
1687                   ON q.symbol_id = t.symbol";
1688        let analysis = analyze_join(&parse_select_snowflake(sql)).unwrap().unwrap();
1689
1690        assert_eq!(analysis.asof_direction, Some(AsofSqlDirection::Backward));
1691        assert_eq!(analysis.left_time_column.as_deref(), Some("trade_ts"));
1692        assert_eq!(analysis.right_time_column.as_deref(), Some("ts"));
1693        assert_eq!(analysis.left_key_column, "symbol");
1694        assert_eq!(analysis.right_key_column, "symbol_id");
1695    }
1696
1697    #[test]
1698    fn test_asof_join_nearest() {
1699        let sql = "SELECT * FROM trades t \
1700                    ASOF JOIN quotes q \
1701                    MATCH_CONDITION(NEAREST(t.ts, q.ts)) \
1702                    ON t.symbol = q.symbol";
1703        let select = parse_select_snowflake(sql);
1704        let analysis = analyze_join(&select).unwrap().unwrap();
1705
1706        assert!(analysis.is_asof_join);
1707        assert_eq!(analysis.asof_direction, Some(AsofSqlDirection::Nearest));
1708        assert_eq!(analysis.join_type, JoinType::AsOf);
1709        assert!(analysis.asof_tolerance.is_none());
1710    }
1711
1712    #[test]
1713    fn test_asof_join_with_tolerance() {
1714        let sql = "SELECT * FROM trades t \
1715                    ASOF JOIN quotes q \
1716                    MATCH_CONDITION(t.ts >= q.ts AND t.ts - q.ts <= 5000) \
1717                    ON t.symbol = q.symbol";
1718        let select = parse_select_snowflake(sql);
1719        let analysis = analyze_join(&select).unwrap().unwrap();
1720
1721        assert!(analysis.is_asof_join);
1722        assert_eq!(analysis.asof_direction, Some(AsofSqlDirection::Backward));
1723        assert_eq!(analysis.asof_tolerance, Some(Duration::from_secs(5)));
1724    }
1725
1726    #[test]
1727    fn test_asof_join_with_interval_tolerance() {
1728        let sql = "SELECT * FROM trades t \
1729                    ASOF JOIN quotes q \
1730                    MATCH_CONDITION(t.ts >= q.ts AND t.ts - q.ts <= INTERVAL '5' SECOND) \
1731                    ON t.symbol = q.symbol";
1732        let select = parse_select_snowflake(sql);
1733        let analysis = analyze_join(&select).unwrap().unwrap();
1734
1735        assert!(analysis.is_asof_join);
1736        assert_eq!(analysis.asof_direction, Some(AsofSqlDirection::Backward));
1737        assert_eq!(analysis.asof_tolerance, Some(Duration::from_secs(5)));
1738    }
1739
1740    #[test]
1741    fn test_asof_forward_tolerance_uses_forward_difference() {
1742        let sql = "SELECT * FROM trades t
1743                   ASOF JOIN quotes q
1744                   MATCH_CONDITION(t.ts <= q.ts AND q.ts - t.ts <= 5000)
1745                   ON t.symbol = q.symbol";
1746        let analysis = analyze_join(&parse_select_snowflake(sql)).unwrap().unwrap();
1747
1748        assert_eq!(analysis.asof_direction, Some(AsofSqlDirection::Forward));
1749        assert_eq!(analysis.asof_tolerance, Some(Duration::from_secs(5)));
1750    }
1751
1752    #[test]
1753    fn test_asof_rejects_unqualified_time_operand() {
1754        let error = asof_join_error(
1755            "SELECT * FROM trades t ASOF JOIN quotes q
1756             MATCH_CONDITION(ts >= q.ts) ON t.symbol = q.symbol",
1757        );
1758        assert!(error.contains("qualified column references"), "{error}");
1759    }
1760
1761    #[test]
1762    fn test_asof_rejects_unknown_time_qualifier() {
1763        let error = asof_join_error(
1764            "SELECT * FROM trades t ASOF JOIN quotes q
1765             MATCH_CONDITION(missing.ts >= q.ts) ON t.symbol = q.symbol",
1766        );
1767        assert!(error.contains("names neither input"), "{error}");
1768    }
1769
1770    #[test]
1771    fn test_asof_rejects_same_side_time_expression() {
1772        let error = asof_join_error(
1773            "SELECT * FROM trades t ASOF JOIN quotes q
1774             MATCH_CONDITION(t.ts >= t.previous_ts) ON t.symbol = q.symbol",
1775        );
1776        assert!(error.contains("one left-input timestamp"), "{error}");
1777    }
1778
1779    #[test]
1780    fn test_asof_rejects_composite_or_residual_on_clause() {
1781        let error = asof_join_error(
1782            "SELECT * FROM trades t ASOF JOIN quotes q
1783             MATCH_CONDITION(t.ts >= q.ts)
1784             ON t.symbol = q.symbol AND t.venue = q.venue",
1785        );
1786        assert!(error.contains("exactly one equality"), "{error}");
1787    }
1788
1789    #[test]
1790    fn test_asof_rejects_unqualified_on_key() {
1791        let error = asof_join_error(
1792            "SELECT * FROM trades t ASOF JOIN quotes q
1793             MATCH_CONDITION(t.ts >= q.ts) ON symbol = q.symbol",
1794        );
1795        assert!(error.contains("qualified column references"), "{error}");
1796    }
1797
1798    #[test]
1799    fn test_asof_rejects_ignored_match_residual() {
1800        let error = asof_join_error(
1801            "SELECT * FROM trades t ASOF JOIN quotes q
1802             MATCH_CONDITION(t.ts >= q.ts AND q.price > 0)
1803             ON t.symbol = q.symbol",
1804        );
1805        assert!(error.contains("tolerance"), "{error}");
1806    }
1807
1808    #[test]
1809    fn test_asof_rejects_tolerance_with_wrong_time_column() {
1810        let error = asof_join_error(
1811            "SELECT * FROM trades t ASOF JOIN quotes q
1812             MATCH_CONDITION(t.ts >= q.ts AND t.other_ts - q.ts <= 5000)
1813             ON t.symbol = q.symbol",
1814        );
1815        assert!(error.contains("same timestamps"), "{error}");
1816    }
1817
1818    #[test]
1819    fn test_asof_rejects_tolerance_with_wrong_direction() {
1820        let error = asof_join_error(
1821            "SELECT * FROM trades t ASOF JOIN quotes q
1822             MATCH_CONDITION(t.ts <= q.ts AND t.ts - q.ts <= 5000)
1823             ON t.symbol = q.symbol",
1824        );
1825        assert!(error.contains("match direction"), "{error}");
1826    }
1827
1828    #[test]
1829    fn test_asof_rejects_zero_tolerance() {
1830        let error = asof_join_error(
1831            "SELECT * FROM trades t ASOF JOIN quotes q
1832             MATCH_CONDITION(t.ts >= q.ts AND t.ts - q.ts <= 0)
1833             ON t.symbol = q.symbol",
1834        );
1835        assert!(error.contains("must be positive"), "{error}");
1836    }
1837
1838    #[test]
1839    fn test_asof_rejects_tolerance_with_nearest() {
1840        let error = asof_join_error(
1841            "SELECT * FROM trades t ASOF JOIN quotes q
1842             MATCH_CONDITION(NEAREST(t.ts, q.ts) AND t.ts - q.ts <= 5000)
1843             ON t.symbol = q.symbol",
1844        );
1845        assert!(error.contains("NEAREST does not support"), "{error}");
1846    }
1847
1848    #[test]
1849    fn test_asof_join_type_mapping() {
1850        let sql = "SELECT * FROM trades t \
1851                    ASOF JOIN quotes q \
1852                    MATCH_CONDITION(t.ts >= q.ts) \
1853                    ON t.symbol = q.symbol";
1854        let select = parse_select_snowflake(sql);
1855        let analysis = analyze_join(&select).unwrap().unwrap();
1856
1857        assert_eq!(analysis.join_type, JoinType::AsOf);
1858        assert!(!analysis.is_lookup_join);
1859    }
1860
1861    #[test]
1862    fn test_asof_join_extracts_time_columns() {
1863        let sql = "SELECT * FROM trades t \
1864                    ASOF JOIN quotes q \
1865                    MATCH_CONDITION(t.ts >= q.ts) \
1866                    ON t.symbol = q.symbol";
1867        let select = parse_select_snowflake(sql);
1868        let analysis = analyze_join(&select).unwrap().unwrap();
1869
1870        assert_eq!(analysis.left_time_column, Some("ts".to_string()));
1871        assert_eq!(analysis.right_time_column, Some("ts".to_string()));
1872    }
1873
1874    #[test]
1875    fn test_asof_join_extracts_key_columns() {
1876        let sql = "SELECT * FROM trades t \
1877                    ASOF JOIN quotes q \
1878                    MATCH_CONDITION(t.ts >= q.ts) \
1879                    ON t.symbol = q.symbol";
1880        let select = parse_select_snowflake(sql);
1881        let analysis = analyze_join(&select).unwrap().unwrap();
1882
1883        assert_eq!(analysis.left_key_column, "symbol");
1884        assert_eq!(analysis.right_key_column, "symbol");
1885    }
1886
1887    #[test]
1888    fn test_asof_join_aliases() {
1889        let sql = "SELECT * FROM trades AS t \
1890                    ASOF JOIN quotes AS q \
1891                    MATCH_CONDITION(t.ts >= q.ts) \
1892                    ON t.symbol = q.symbol";
1893        let select = parse_select_snowflake(sql);
1894        let analysis = analyze_join(&select).unwrap().unwrap();
1895
1896        assert_eq!(analysis.left_alias, Some("t".to_string()));
1897        assert_eq!(analysis.right_alias, Some("q".to_string()));
1898        assert_eq!(analysis.left_table, "trades");
1899        assert_eq!(analysis.right_table, "quotes");
1900    }
1901
1902    // -- Multi-way JOIN tests --
1903
1904    #[test]
1905    fn test_multi_join_single_backward_compat() {
1906        let sql = "SELECT * FROM orders o JOIN payments p ON o.id = p.order_id";
1907        let select = parse_select(sql);
1908        let multi = analyze_joins(&select).unwrap().unwrap();
1909
1910        assert!(multi.is_single());
1911        assert_eq!(multi.len(), 1);
1912        assert!(!multi.is_empty());
1913        let first = multi.first().unwrap();
1914        assert_eq!(first.left_table, "orders");
1915        assert_eq!(first.right_table, "payments");
1916    }
1917
1918    #[test]
1919    fn test_multi_join_two_way() {
1920        let sql = "SELECT * FROM a JOIN b ON a.id = b.a_id JOIN c ON c.b_id = b.id";
1921        let select = parse_select(sql);
1922        let multi = analyze_joins(&select).unwrap().unwrap();
1923
1924        assert_eq!(multi.len(), 2);
1925        assert!(!multi.is_single());
1926
1927        assert_eq!(multi.joins[0].left_table, "a");
1928        assert_eq!(multi.joins[0].right_table, "b");
1929        assert_eq!(multi.joins[0].left_key_column, "id");
1930        assert_eq!(multi.joins[0].right_key_column, "a_id");
1931
1932        assert_eq!(multi.joins[1].left_table, "b");
1933        assert_eq!(multi.joins[1].right_table, "c");
1934        assert_eq!(multi.joins[1].left_key_column, "id");
1935        assert_eq!(multi.joins[1].right_key_column, "b_id");
1936    }
1937
1938    #[test]
1939    fn test_multi_join_three_way() {
1940        let sql = "SELECT * FROM a \
1941                    JOIN b ON a.id = b.a_id \
1942                    JOIN c ON b.id = c.b_id \
1943                    JOIN d ON c.id = d.c_id";
1944        let select = parse_select(sql);
1945        let multi = analyze_joins(&select).unwrap().unwrap();
1946
1947        assert_eq!(multi.len(), 3);
1948        assert_eq!(multi.tables.len(), 4);
1949        assert_eq!(multi.tables, vec!["a", "b", "c", "d"]);
1950    }
1951
1952    #[test]
1953    fn test_multi_join_mixed_asof_and_lookup() {
1954        // ASOF first, then lookup (use Snowflake dialect for ASOF)
1955        let sql = "SELECT * FROM trades t \
1956                    ASOF JOIN quotes q \
1957                    MATCH_CONDITION(t.ts >= q.ts) \
1958                    ON t.symbol = q.symbol \
1959                    JOIN products p ON q.product_id = p.id";
1960        let select = parse_select_snowflake(sql);
1961        let multi = analyze_joins(&select).unwrap().unwrap();
1962
1963        assert_eq!(multi.len(), 2);
1964        assert!(multi.joins[0].is_asof_join);
1965        assert!(multi.joins[1].is_lookup_join);
1966    }
1967
1968    #[test]
1969    fn test_multi_join_stream_stream_and_lookup() {
1970        let sql = "SELECT * FROM orders o \
1971                    JOIN payments p ON o.id = p.order_id \
1972                        AND p.ts BETWEEN o.ts AND o.ts + INTERVAL '1' HOUR \
1973                    JOIN customers c ON p.customer_id = c.id";
1974        let select = parse_select(sql);
1975        let multi = analyze_joins(&select).unwrap().unwrap();
1976
1977        assert_eq!(multi.len(), 2);
1978        assert!(!multi.joins[0].is_lookup_join); // stream-stream
1979        assert!(multi.joins[0].time_bound.is_some());
1980        assert!(multi.joins[1].is_lookup_join); // lookup
1981    }
1982
1983    #[test]
1984    fn test_multi_join_rejects_key_from_non_current_left_relation() {
1985        let select = parse_select(
1986            "SELECT * FROM a
1987             JOIN b ON a.id = b.a_id
1988             JOIN c ON a.id = c.a_id",
1989        );
1990        let error = analyze_joins(&select).unwrap_err().to_string();
1991
1992        assert!(error.contains("names neither input"), "{error}");
1993    }
1994
1995    #[test]
1996    fn test_multi_join_tables_list() {
1997        let sql = "SELECT * FROM a JOIN b ON a.id = b.a_id JOIN c ON b.id = c.b_id";
1998        let select = parse_select(sql);
1999        let multi = analyze_joins(&select).unwrap().unwrap();
2000
2001        assert_eq!(multi.tables, vec!["a", "b", "c"]);
2002    }
2003
2004    #[test]
2005    fn test_multi_join_aliases() {
2006        let sql = "SELECT * FROM orders AS o \
2007                    JOIN payments AS p ON o.id = p.order_id \
2008                    JOIN refunds AS r ON p.id = r.payment_id";
2009        let select = parse_select(sql);
2010        let multi = analyze_joins(&select).unwrap().unwrap();
2011
2012        assert_eq!(multi.joins[0].left_alias, Some("o".to_string()));
2013        assert_eq!(multi.joins[0].right_alias, Some("p".to_string()));
2014        assert_eq!(multi.joins[1].left_alias, Some("p".to_string()));
2015        assert_eq!(multi.joins[1].right_alias, Some("r".to_string()));
2016    }
2017
2018    #[test]
2019    fn test_multi_join_no_join_returns_none() {
2020        let sql = "SELECT * FROM orders";
2021        let select = parse_select(sql);
2022        let multi = analyze_joins(&select).unwrap();
2023        assert!(multi.is_none());
2024    }
2025
2026    // -- Temporal JOIN tests (FOR SYSTEM_TIME AS OF) --
2027
2028    #[test]
2029    fn test_temporal_join_detected() {
2030        let sql = "SELECT o.*, p.price \
2031                    FROM orders o \
2032                    JOIN products FOR SYSTEM_TIME AS OF o.order_time AS p \
2033                    ON o.product_id = p.id";
2034        let select = parse_select_laminar(sql);
2035        let analysis = analyze_join(&select).unwrap().unwrap();
2036
2037        assert!(analysis.is_temporal_join);
2038        assert_eq!(
2039            analysis.temporal_version_column,
2040            Some("order_time".to_string())
2041        );
2042        assert_eq!(analysis.left_table, "orders");
2043        assert_eq!(analysis.right_table, "products");
2044        assert_eq!(analysis.left_key_column, "product_id");
2045        assert_eq!(analysis.right_key_column, "id");
2046        assert!(!analysis.is_lookup_join);
2047        assert!(!analysis.is_asof_join);
2048    }
2049
2050    #[test]
2051    fn test_temporal_join_via_analyze_joins() {
2052        let sql = "SELECT o.*, p.price \
2053                    FROM orders o \
2054                    JOIN products FOR SYSTEM_TIME AS OF o.order_time AS p \
2055                    ON o.product_id = p.id";
2056        let select = parse_select_laminar(sql);
2057        let multi = analyze_joins(&select).unwrap().unwrap();
2058
2059        assert_eq!(multi.len(), 1);
2060        let first = multi.first().unwrap();
2061        assert!(first.is_temporal_join);
2062        assert_eq!(
2063            first.temporal_version_column,
2064            Some("order_time".to_string())
2065        );
2066    }
2067
2068    #[test]
2069    fn test_non_temporal_join_not_flagged() {
2070        let sql = "SELECT * FROM orders o JOIN payments p ON o.id = p.order_id";
2071        let select = parse_select(sql);
2072        let analysis = analyze_join(&select).unwrap().unwrap();
2073
2074        assert!(!analysis.is_temporal_join);
2075        assert!(analysis.temporal_version_column.is_none());
2076    }
2077
2078    #[test]
2079    fn test_unqualified_anti_maps_to_left_anti() {
2080        let sql = "SELECT * FROM orders o ANTI JOIN returns r ON o.id = r.order_id";
2081        let select = parse_select(sql);
2082        let analysis = analyze_join(&select).unwrap().unwrap();
2083        assert_eq!(analysis.join_type, JoinType::LeftAnti);
2084    }
2085
2086    #[test]
2087    fn test_unqualified_semi_maps_to_left_semi() {
2088        let sql = "SELECT * FROM orders o SEMI JOIN payments p ON o.id = p.order_id";
2089        let select = parse_select(sql);
2090        let analysis = analyze_join(&select).unwrap().unwrap();
2091        assert_eq!(analysis.join_type, JoinType::LeftSemi);
2092    }
2093
2094    #[test]
2095    fn test_composite_join_keys() {
2096        let sql = "SELECT * FROM orders o \
2097                    JOIN shipments s \
2098                    ON o.order_id = s.order_id AND o.region = s.region";
2099        let select = parse_select(sql);
2100        let analysis = analyze_join(&select).unwrap().unwrap();
2101
2102        // First key pair is the primary key
2103        assert_eq!(analysis.left_key_column, "order_id");
2104        assert_eq!(analysis.right_key_column, "order_id");
2105
2106        // Second key pair should be in additional_key_columns
2107        assert_eq!(
2108            analysis.additional_key_columns.len(),
2109            1,
2110            "Should have 1 additional key pair"
2111        );
2112        assert_eq!(analysis.additional_key_columns[0].0, "region");
2113        assert_eq!(analysis.additional_key_columns[0].1, "region");
2114    }
2115
2116    #[test]
2117    fn test_composite_using_clause() {
2118        let sql = "SELECT * FROM orders o JOIN shipments s USING (order_id, region)";
2119        let select = parse_select(sql);
2120        let analysis = analyze_join(&select).unwrap().unwrap();
2121
2122        // First column becomes primary key
2123        assert_eq!(analysis.left_key_column, "order_id");
2124        assert_eq!(analysis.right_key_column, "order_id");
2125
2126        // Additional columns
2127        assert_eq!(
2128            analysis.additional_key_columns.len(),
2129            1,
2130            "USING(order_id, region) should have 1 additional key"
2131        );
2132        assert_eq!(analysis.additional_key_columns[0].0, "region");
2133        assert_eq!(analysis.additional_key_columns[0].1, "region");
2134    }
2135
2136    #[test]
2137    fn test_using_preserves_quoted_key_identity() {
2138        let sql = "SELECT * FROM orders o JOIN shipments s USING (\"order id\")";
2139        let analysis = analyze_join(&parse_select(sql)).unwrap().unwrap();
2140
2141        assert_eq!(analysis.left_key_column, "order id");
2142        assert_eq!(analysis.right_key_column, "order id");
2143    }
2144}