Skip to main content

query/datafusion/
pg_oid_alias_expr_planner.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use arrow_schema::DataType;
16use datafusion_common::{DFSchema, ExprSchema, Result, ScalarValue};
17use datafusion_expr::expr::BinaryExpr;
18use datafusion_expr::planner::{ExprPlanner, PlannerResult, RawBinaryExpr};
19use datafusion_expr::{Expr, Operator};
20use datafusion_pg_catalog::pg_catalog::oid_field::{OID_ALIAS_KEY, kind};
21use sqlparser::ast::BinaryOperator;
22
23/// Rewrites PostgreSQL's regproc zero sentinel before DataFusion type coercion.
24#[derive(Debug)]
25pub(crate) struct PgOidAliasExprPlanner;
26
27impl ExprPlanner for PgOidAliasExprPlanner {
28    fn plan_binary_op(
29        &self,
30        expr: RawBinaryExpr,
31        schema: &DFSchema,
32    ) -> Result<PlannerResult<RawBinaryExpr>> {
33        let RawBinaryExpr {
34            op,
35            mut left,
36            mut right,
37        } = expr;
38
39        let operator = match op {
40            BinaryOperator::Eq => Operator::Eq,
41            BinaryOperator::NotEq => Operator::NotEq,
42            _ => return Ok(PlannerResult::Original(RawBinaryExpr { op, left, right })),
43        };
44
45        let (column, zero_on_left) = match (&left, &right) {
46            (Expr::Literal(value, _), Expr::Column(column)) if is_integral_zero(value) => {
47                (column, true)
48            }
49            (Expr::Column(column), Expr::Literal(value, _)) if is_integral_zero(value) => {
50                (column, false)
51            }
52            _ => return Ok(PlannerResult::Original(RawBinaryExpr { op, left, right })),
53        };
54
55        // A raw SQL column is resolved against the schema before the default
56        // coercion planner runs. Do not infer alias semantics from casts or any
57        // other expression shape.
58        let Ok(field) = schema.field_from_column(column) else {
59            return Ok(PlannerResult::Original(RawBinaryExpr { op, left, right }));
60        };
61        if field.metadata().get(OID_ALIAS_KEY).map(String::as_str) != Some(kind::REGPROC) {
62            return Ok(PlannerResult::Original(RawBinaryExpr { op, left, right }));
63        }
64
65        let Some(sentinel) = regproc_zero_sentinel(field.data_type()) else {
66            return Ok(PlannerResult::Original(RawBinaryExpr { op, left, right }));
67        };
68
69        let sentinel = Expr::Literal(sentinel, None);
70        if zero_on_left {
71            left = sentinel;
72        } else {
73            right = sentinel;
74        }
75
76        Ok(PlannerResult::Planned(Expr::BinaryExpr(BinaryExpr::new(
77            Box::new(left),
78            operator,
79            Box::new(right),
80        ))))
81    }
82}
83
84fn is_integral_zero(value: &ScalarValue) -> bool {
85    matches!(
86        value,
87        ScalarValue::Int8(Some(0))
88            | ScalarValue::Int16(Some(0))
89            | ScalarValue::Int32(Some(0))
90            | ScalarValue::Int64(Some(0))
91            | ScalarValue::UInt8(Some(0))
92            | ScalarValue::UInt16(Some(0))
93            | ScalarValue::UInt32(Some(0))
94            | ScalarValue::UInt64(Some(0))
95    )
96}
97
98fn regproc_zero_sentinel(data_type: &DataType) -> Option<ScalarValue> {
99    match data_type {
100        DataType::Utf8 => Some(ScalarValue::Utf8(Some("-".to_string()))),
101        DataType::LargeUtf8 => Some(ScalarValue::LargeUtf8(Some("-".to_string()))),
102        DataType::Utf8View => Some(ScalarValue::Utf8View(Some("-".to_string()))),
103        _ => None,
104    }
105}
106
107#[cfg(test)]
108mod tests {
109    use std::collections::HashMap;
110    use std::sync::Arc;
111
112    use arrow_schema::{Field, Fields};
113    use datafusion_common::Column;
114    use datafusion_expr::ExprSchemable;
115    use datafusion_expr::expr::Cast;
116    use datafusion_expr::simplify::SimplifyContext;
117    use datafusion_optimizer::simplify_expressions::ExprSimplifier;
118
119    use super::*;
120
121    fn schema(data_type: DataType, alias: Option<&str>) -> DFSchema {
122        let mut field = Field::new("typreceive", data_type, true);
123        if let Some(alias) = alias {
124            field = field.with_metadata(HashMap::from([(
125                OID_ALIAS_KEY.to_string(),
126                alias.to_string(),
127            )]));
128        }
129        DFSchema::from_unqualified_fields(Fields::from(vec![field]), HashMap::new()).unwrap()
130    }
131
132    fn column() -> Expr {
133        Expr::Column(Column::new_unqualified("typreceive"))
134    }
135
136    fn plan(expr: RawBinaryExpr, schema: &DFSchema) -> PlannerResult<RawBinaryExpr> {
137        PgOidAliasExprPlanner.plan_binary_op(expr, schema).unwrap()
138    }
139
140    fn assert_planned_sentinel(
141        planned: PlannerResult<RawBinaryExpr>,
142        operator: Operator,
143        zero_on_left: bool,
144        sentinel: ScalarValue,
145    ) -> Expr {
146        let PlannerResult::Planned(Expr::BinaryExpr(expr)) = planned else {
147            panic!("expected a planned binary expression");
148        };
149        assert_eq!(expr.op, operator);
150        let literal = Expr::Literal(sentinel, None);
151        if zero_on_left {
152            assert_eq!(expr.left.as_ref(), &literal);
153            assert_eq!(expr.right.as_ref(), &column());
154        } else {
155            assert_eq!(expr.left.as_ref(), &column());
156            assert_eq!(expr.right.as_ref(), &literal);
157        }
158        Expr::BinaryExpr(expr)
159    }
160
161    #[test]
162    fn rewrites_zero_regproc_comparisons_in_both_operand_orders() {
163        let schema = schema(DataType::Utf8, Some(kind::REGPROC));
164
165        for (sql_operator, operator) in [
166            (BinaryOperator::Eq, Operator::Eq),
167            (BinaryOperator::NotEq, Operator::NotEq),
168        ] {
169            for zero_on_left in [true, false] {
170                let zero = Expr::Literal(ScalarValue::Int64(Some(0)), None);
171                let (left, right) = if zero_on_left {
172                    (zero, column())
173                } else {
174                    (column(), zero)
175                };
176                let planned = assert_planned_sentinel(
177                    plan(
178                        RawBinaryExpr {
179                            op: sql_operator.clone(),
180                            left,
181                            right,
182                        },
183                        &schema,
184                    ),
185                    operator,
186                    zero_on_left,
187                    ScalarValue::Utf8(Some("-".to_string())),
188                );
189                assert!(planned.nullable(&schema).unwrap());
190            }
191        }
192    }
193
194    #[test]
195    fn rewrites_every_integral_zero_with_the_column_string_storage_type() {
196        let zero_literals = [
197            ScalarValue::Int8(Some(0)),
198            ScalarValue::Int16(Some(0)),
199            ScalarValue::Int32(Some(0)),
200            ScalarValue::Int64(Some(0)),
201            ScalarValue::UInt8(Some(0)),
202            ScalarValue::UInt16(Some(0)),
203            ScalarValue::UInt32(Some(0)),
204            ScalarValue::UInt64(Some(0)),
205        ];
206        let string_types = [
207            (DataType::Utf8, ScalarValue::Utf8(Some("-".to_string()))),
208            (
209                DataType::LargeUtf8,
210                ScalarValue::LargeUtf8(Some("-".to_string())),
211            ),
212            (
213                DataType::Utf8View,
214                ScalarValue::Utf8View(Some("-".to_string())),
215            ),
216        ];
217
218        for (data_type, sentinel) in string_types {
219            let schema = schema(data_type, Some(kind::REGPROC));
220            for zero in &zero_literals {
221                assert_planned_sentinel(
222                    plan(
223                        RawBinaryExpr {
224                            op: BinaryOperator::Eq,
225                            left: column(),
226                            right: Expr::Literal(zero.clone(), None),
227                        },
228                        &schema,
229                    ),
230                    Operator::Eq,
231                    false,
232                    sentinel.clone(),
233                );
234            }
235        }
236    }
237
238    #[test]
239    fn leaves_non_matching_comparisons_untouched() {
240        let regproc = schema(DataType::Utf8, Some(kind::REGPROC));
241        let int32_regproc = schema(DataType::Int32, Some(kind::REGPROC));
242        let untagged = schema(DataType::Utf8, None);
243        let regtype = schema(DataType::Utf8, Some(kind::REGTYPE));
244
245        let cases = [
246            (
247                RawBinaryExpr {
248                    op: BinaryOperator::Eq,
249                    left: column(),
250                    right: Expr::Literal(ScalarValue::Int64(Some(1)), None),
251                },
252                &regproc,
253            ),
254            (
255                RawBinaryExpr {
256                    op: BinaryOperator::Eq,
257                    left: column(),
258                    right: Expr::Literal(ScalarValue::Int64(None), None),
259                },
260                &regproc,
261            ),
262            (
263                RawBinaryExpr {
264                    op: BinaryOperator::Lt,
265                    left: column(),
266                    right: Expr::Literal(ScalarValue::Int64(Some(0)), None),
267                },
268                &regproc,
269            ),
270            (
271                RawBinaryExpr {
272                    op: BinaryOperator::Eq,
273                    left: column(),
274                    right: Expr::Literal(ScalarValue::Int64(Some(0)), None),
275                },
276                &int32_regproc,
277            ),
278            (
279                RawBinaryExpr {
280                    op: BinaryOperator::Eq,
281                    left: column(),
282                    right: Expr::Literal(ScalarValue::Int64(Some(0)), None),
283                },
284                &untagged,
285            ),
286            (
287                RawBinaryExpr {
288                    op: BinaryOperator::Eq,
289                    left: column(),
290                    right: Expr::Literal(ScalarValue::Int64(Some(0)), None),
291                },
292                &regtype,
293            ),
294        ];
295
296        for (expr, schema) in cases {
297            assert!(matches!(plan(expr, schema), PlannerResult::Original(_)));
298        }
299    }
300
301    #[test]
302    fn leaves_casts_and_the_adbc_array_receiver_predicate_untouched() {
303        let schema = schema(DataType::Utf8, Some(kind::REGPROC));
304        let cast = Expr::Cast(Cast::new(Box::new(column()), DataType::Utf8));
305        let expr = RawBinaryExpr {
306            op: BinaryOperator::NotEq,
307            left: cast,
308            right: Expr::Literal(ScalarValue::Utf8(Some("array_recv".to_string())), None),
309        };
310
311        assert!(matches!(plan(expr, &schema), PlannerResult::Original(_)));
312    }
313
314    #[test]
315    fn type_coercion_keeps_regproc_as_a_string_after_the_rewrite() {
316        let schema = Arc::new(schema(DataType::Utf8, Some(kind::REGPROC)));
317        let planned = assert_planned_sentinel(
318            plan(
319                RawBinaryExpr {
320                    op: BinaryOperator::NotEq,
321                    left: column(),
322                    right: Expr::Literal(ScalarValue::Int64(Some(0)), None),
323                },
324                &schema,
325            ),
326            Operator::NotEq,
327            false,
328            ScalarValue::Utf8(Some("-".to_string())),
329        );
330        let simplifier = ExprSimplifier::new(
331            SimplifyContext::builder()
332                .with_schema(schema.clone())
333                .build(),
334        );
335        let coerced = simplifier.coerce(planned, &schema).unwrap();
336
337        assert!(!format!("{coerced}").contains("CAST"));
338        assert!(!format!("{coerced:?}").contains("Int64"));
339    }
340}