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::CompositeRow { items, .. }
71 | ScalarExpr::And(items)
72 | ScalarExpr::Or(items) => combine_bindings(items, schema),
73 ScalarExpr::Binary { lhs, rhs, .. } => {
74 combine_bindings([lhs.as_ref(), rhs.as_ref()], schema)
75 }
76 ScalarExpr::UnaryMinus(expr)
77 | ScalarExpr::Not(expr)
78 | ScalarExpr::IsNull { expr, .. }
79 | ScalarExpr::Cast { expr, .. }
80 | ScalarExpr::InSubquery { expr, .. } => expression_binding(expr, schema),
81 ScalarExpr::Between { expr, low, high } => {
82 combine_bindings([expr.as_ref(), low.as_ref(), high.as_ref()], schema)
83 }
84 ScalarExpr::InList { expr, list, .. } => {
85 combine_bindings(std::iter::once(expr.as_ref()).chain(list), schema)
86 }
87 ScalarExpr::Case {
88 base,
89 when,
90 else_branch,
91 } => combine_bindings(
92 base.as_deref()
93 .into_iter()
94 .chain(
95 when.iter()
96 .flat_map(|(condition, value)| [condition, value]),
97 )
98 .chain(else_branch.as_deref()),
99 schema,
100 ),
101 ScalarExpr::WindowCall { .. }
102 | ScalarExpr::Default
103 | ScalarExpr::Star
104 | ScalarExpr::QualifiedStar(_)
105 | ScalarExpr::ScalarSubquery(_)
106 | ScalarExpr::Exists { .. } => (false, false),
107 }
108}
109
110fn combine_bindings<'a>(
111 expressions: impl IntoIterator<Item = &'a ScalarExpr>,
112 schema: &RowSchema,
113) -> (bool, bool) {
114 expressions
115 .into_iter()
116 .map(|expression| expression_binding(expression, schema))
117 .fold((true, false), |(valid, has_column), binding| {
118 (valid && binding.0, has_column || binding.1)
119 })
120}
121
122#[cfg(test)]
123mod tests {
124 use super::*;
125 use crate::ColumnIdentity;
126
127 #[test]
128 fn join_side_binding_uses_structured_identities() {
129 let left = RowSchema::with_identities(
130 vec!["order.key".into()],
131 vec![ColumnIdentity::qualified("left.alias", "order.key")],
132 vec![None],
133 );
134 let right = RowSchema::with_qualified_types("right.alias", vec!["id".into()], vec![None]);
135 let lhs = ScalarExpr::qualified_column("left.alias", "order.key");
136 let rhs = ScalarExpr::qualified_column("right.alias", "id");
137 assert_eq!(
138 decide_join_sides(&left, &right, &lhs, &rhs),
139 Some((&lhs, &rhs))
140 );
141 }
142
143 #[test]
144 fn ambiguous_unqualified_join_key_is_not_assigned_arbitrarily() {
145 let left = RowSchema::with_identities(
146 vec!["id".into(), "id".into()],
147 vec![
148 ColumnIdentity::qualified("left", "id"),
149 ColumnIdentity::qualified("other", "id"),
150 ],
151 vec![None, None],
152 );
153 let right = RowSchema::with_qualified_types("right", vec!["id".into()], vec![None]);
154 let lhs = ScalarExpr::Column("id".into());
155 let rhs = ScalarExpr::qualified_column("right", "id");
156 assert_eq!(decide_join_sides(&left, &right, &lhs, &rhs), None);
157 }
158}