1use 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#[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 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 ®proc,
253 ),
254 (
255 RawBinaryExpr {
256 op: BinaryOperator::Eq,
257 left: column(),
258 right: Expr::Literal(ScalarValue::Int64(None), None),
259 },
260 ®proc,
261 ),
262 (
263 RawBinaryExpr {
264 op: BinaryOperator::Lt,
265 left: column(),
266 right: Expr::Literal(ScalarValue::Int64(Some(0)), None),
267 },
268 ®proc,
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 ®type,
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}