1use crate::RowSchema;
10use crate::ScalarExpr;
11
12pub fn join_conjuncts(expr: &ScalarExpr) -> Vec<&ScalarExpr> {
13 match expr {
14 ScalarExpr::And(items) => {
15 let mut conjuncts = Vec::with_capacity(items.len());
16 for item in items {
17 conjuncts.extend(join_conjuncts(item));
18 }
19 conjuncts
20 }
21 _ => vec![expr],
22 }
23}
24
25pub fn decide_join_sides<'a>(
27 left: &RowSchema,
28 right: &RowSchema,
29 lhs: &'a ScalarExpr,
30 rhs: &'a ScalarExpr,
31) -> Option<(&'a ScalarExpr, &'a ScalarExpr)> {
32 if expression_binds_to(lhs, left) && expression_binds_to(rhs, right) {
33 return Some((lhs, rhs));
34 }
35 if expression_binds_to(rhs, left) && expression_binds_to(lhs, right) {
36 return Some((rhs, lhs));
37 }
38 None
39}
40
41fn expression_binds_to(expression: &ScalarExpr, schema: &RowSchema) -> bool {
42 let (valid, has_column) = expression_binding(expression, schema);
43 valid && has_column
44}
45
46fn expression_binding(expression: &ScalarExpr, schema: &RowSchema) -> (bool, bool) {
47 match expression {
48 ScalarExpr::Column(column) => (schema.unqualified_position(column).is_some(), true),
49 ScalarExpr::QualifiedColumn { qualifier, column } => {
50 (schema.qualified_position(qualifier, column).is_some(), true)
51 }
52 ScalarExpr::Position(position) => (*position < schema.len(), true),
53 ScalarExpr::InternalColumn(column) => (schema.internal_slot(*column).is_some(), true),
54 ScalarExpr::Literal(_) | ScalarExpr::TypedLiteral { .. } | ScalarExpr::Param(_) => {
55 (true, false)
56 }
57 ScalarExpr::Func {
58 args,
59 order_by,
60 filter,
61 ..
62 } => combine_bindings(
63 args.iter()
64 .chain(order_by.iter().map(|order| &order.expr))
65 .chain(filter.as_deref()),
66 schema,
67 ),
68 ScalarExpr::Array(items)
69 | ScalarExpr::Row(items)
70 | ScalarExpr::And(items)
71 | ScalarExpr::Or(items) => combine_bindings(items, schema),
72 ScalarExpr::Binary { lhs, rhs, .. } => {
73 combine_bindings([lhs.as_ref(), rhs.as_ref()], schema)
74 }
75 ScalarExpr::UnaryMinus(expr)
76 | ScalarExpr::Not(expr)
77 | ScalarExpr::IsNull { expr, .. }
78 | ScalarExpr::Cast { expr, .. }
79 | ScalarExpr::InSubquery { expr, .. } => expression_binding(expr, schema),
80 ScalarExpr::Between { expr, low, high } => {
81 combine_bindings([expr.as_ref(), low.as_ref(), high.as_ref()], schema)
82 }
83 ScalarExpr::InList { expr, list, .. } => {
84 combine_bindings(std::iter::once(expr.as_ref()).chain(list), schema)
85 }
86 ScalarExpr::Case {
87 base,
88 when,
89 else_branch,
90 } => combine_bindings(
91 base.as_deref()
92 .into_iter()
93 .chain(
94 when.iter()
95 .flat_map(|(condition, value)| [condition, value]),
96 )
97 .chain(else_branch.as_deref()),
98 schema,
99 ),
100 ScalarExpr::WindowCall { .. }
101 | ScalarExpr::Default
102 | ScalarExpr::Star
103 | ScalarExpr::QualifiedStar(_)
104 | ScalarExpr::ScalarSubquery(_)
105 | ScalarExpr::Exists { .. } => (false, false),
106 }
107}
108
109fn combine_bindings<'a>(
110 expressions: impl IntoIterator<Item = &'a ScalarExpr>,
111 schema: &RowSchema,
112) -> (bool, bool) {
113 expressions
114 .into_iter()
115 .map(|expression| expression_binding(expression, schema))
116 .fold((true, false), |(valid, has_column), binding| {
117 (valid && binding.0, has_column || binding.1)
118 })
119}
120
121#[cfg(test)]
122mod tests {
123 use super::*;
124 use crate::ColumnIdentity;
125
126 #[test]
127 fn join_side_binding_uses_structured_identities() {
128 let left = RowSchema::with_identities(
129 vec!["order.key".into()],
130 vec![ColumnIdentity::qualified("left.alias", "order.key")],
131 vec![None],
132 );
133 let right = RowSchema::with_qualified_types("right.alias", vec!["id".into()], vec![None]);
134 let lhs = ScalarExpr::qualified_column("left.alias", "order.key");
135 let rhs = ScalarExpr::qualified_column("right.alias", "id");
136 assert_eq!(
137 decide_join_sides(&left, &right, &lhs, &rhs),
138 Some((&lhs, &rhs))
139 );
140 }
141
142 #[test]
143 fn ambiguous_unqualified_join_key_is_not_assigned_arbitrarily() {
144 let left = RowSchema::with_identities(
145 vec!["id".into(), "id".into()],
146 vec![
147 ColumnIdentity::qualified("left", "id"),
148 ColumnIdentity::qualified("other", "id"),
149 ],
150 vec![None, None],
151 );
152 let right = RowSchema::with_qualified_types("right", vec!["id".into()], vec![None]);
153 let lhs = ScalarExpr::Column("id".into());
154 let rhs = ScalarExpr::qualified_column("right", "id");
155 assert_eq!(decide_join_sides(&left, &right, &lhs, &rhs), None);
156 }
157}