1use uqa_core::Value;
10use uqa_sql::ast::BinaryOp;
11use uqa_sql::expr::IntegerWidth;
12use uqa_sql::{SQLError, SQLParam};
13
14use crate::{RowSchema, ScalarExpr};
15
16mod compile;
17mod evaluate;
18
19pub struct ProjectedPredicate {
23 expression: ProjectedExpr,
24}
25
26pub(super) enum ProjectedIntPredicate {
27 Comparison {
28 field: usize,
29 op: BinaryOp,
30 literal: i64,
31 field_on_left: bool,
32 },
33 Between {
34 field: usize,
35 low: i64,
36 high: i64,
37 },
38}
39
40pub(super) enum ProjectedExpr {
41 Field(usize),
42 Literal(Value),
43 Binary {
44 op: BinaryOp,
45 lhs: Box<Self>,
46 rhs: Box<Self>,
47 integer_width: Option<IntegerWidth>,
48 },
49 UnaryMinus(Box<Self>),
50 IntFieldComparison {
51 field: usize,
52 op: BinaryOp,
53 literal: i64,
54 field_on_left: bool,
55 },
56 Not(Box<Self>),
57 And(Vec<Self>),
58 IntFieldConjunction(Vec<ProjectedIntPredicate>),
59 Or(Vec<Self>),
60 IsNull {
61 expression: Box<Self>,
62 negated: bool,
63 },
64 Between {
65 expression: Box<Self>,
66 low: Box<Self>,
67 high: Box<Self>,
68 },
69 IntFieldBetween {
70 field: usize,
71 low: i64,
72 high: i64,
73 },
74 InList {
75 expression: Box<Self>,
76 list: Vec<Self>,
77 negated: bool,
78 },
79 Like {
80 expression: Box<Self>,
81 pattern: uqa_sql::expr::CompiledLikePattern,
82 },
83 Cast {
84 expression: Box<Self>,
85 ty: String,
86 },
87}
88
89impl ProjectedPredicate {
90 pub fn compile(
91 expression: &ScalarExpr,
92 fields: &[String],
93 params: &[SQLParam],
94 ) -> Result<Option<Self>, SQLError> {
95 Self::compile_with_schema(expression, &RowSchema::new(fields.to_vec()), params)
96 }
97
98 pub fn compile_with_schema(
100 expression: &ScalarExpr,
101 schema: &RowSchema,
102 params: &[SQLParam],
103 ) -> Result<Option<Self>, SQLError> {
104 match compile::compile(expression, schema, params) {
105 Ok(expression) => Ok(expression.map(|expression| Self { expression })),
106 Err(SQLError::Unsupported(_)) => Ok(None),
107 Err(error) => Err(error),
108 }
109 }
110
111 #[inline]
112 pub fn keep(&self, values: &[&Value]) -> Result<bool, SQLError> {
113 evaluate::keep(&self.expression, values)
114 }
115
116 #[inline]
118 pub fn keep_indexed(&self, values: &[Value], projection: &[usize]) -> Result<bool, SQLError> {
119 evaluate::keep_indexed(&self.expression, values, projection)
120 }
121
122 #[inline]
126 pub fn keep_row(&self, row: &crate::PhysicalRowView<'_>) -> Result<bool, SQLError> {
127 evaluate::keep_row(&self.expression, row)
128 }
129}
130
131#[cfg(test)]
132mod tests {
133 use super::*;
134 use crate::{eval_scalar, ScalarEvalContext};
135 use uqa_sql::ast::ColumnType;
136 use uqa_sql::expr::truthy;
137
138 #[test]
139 fn positional_predicate_preserves_null_and_short_circuit_semantics() {
140 let expression = ScalarExpr::And(vec![
141 ScalarExpr::Between {
142 expr: Box::new(ScalarExpr::Column("x".into())),
143 low: Box::new(ScalarExpr::Literal(Value::Int(2))),
144 high: Box::new(ScalarExpr::Literal(Value::Int(4))),
145 },
146 ScalarExpr::IsNull {
147 expr: Box::new(ScalarExpr::Column("y".into())),
148 negated: false,
149 },
150 ]);
151 let predicate = ProjectedPredicate::compile(&expression, &["x".into(), "y".into()], &[])
152 .unwrap()
153 .unwrap();
154 assert!(predicate.keep(&[&Value::Int(3), &Value::Null]).unwrap());
155 assert!(!predicate.keep(&[&Value::Int(5), &Value::Null]).unwrap());
156 assert!(!predicate.keep(&[&Value::Null, &Value::Null]).unwrap());
157
158 let stored = [Value::Null, Value::Int(3)];
159 assert!(predicate.keep_indexed(&stored, &[1, 0]).unwrap());
160 assert!(!predicate.keep_indexed(&stored, &[usize::MAX, 0]).unwrap());
161 }
162
163 #[test]
164 fn projected_q1_and_q6_predicates_match_the_canonical_evaluator() {
165 let fields = vec!["discount".into(), "quantity".into(), "ship_day".into()];
166 let q6 = ScalarExpr::And(vec![
167 ScalarExpr::Between {
168 expr: Box::new(ScalarExpr::Column("ship_day".into())),
169 low: Box::new(ScalarExpr::Param(1)),
170 high: Box::new(ScalarExpr::Literal(Value::Int(2_190))),
171 },
172 ScalarExpr::Between {
173 expr: Box::new(ScalarExpr::Column("discount".into())),
174 low: Box::new(ScalarExpr::Literal(Value::Int(2))),
175 high: Box::new(ScalarExpr::Literal(Value::Int(8))),
176 },
177 ScalarExpr::Binary {
178 op: BinaryOp::Greater,
179 lhs: Box::new(ScalarExpr::Literal(Value::Int(40))),
180 rhs: Box::new(ScalarExpr::Column("quantity".into())),
181 },
182 ]);
183 let params = vec![SQLParam::typed_scalar(
184 Value::Int(365),
185 ColumnType::SmallInteger,
186 )];
187 let predicate = ProjectedPredicate::compile(&q6, &fields, ¶ms)
188 .unwrap()
189 .unwrap();
190 assert!(matches!(
191 &predicate.expression,
192 ProjectedExpr::IntFieldConjunction(items) if items.len() == 3
193 ));
194 let discounts = [Value::Null, Value::Int(1), Value::Int(2), Value::Int(8)];
195 let quantities = [
196 Value::Null,
197 Value::Int(39),
198 Value::Int(40),
199 Value::Float(39.5),
200 ];
201 let ship_days = [
202 Value::Null,
203 Value::Int(364),
204 Value::Int(365),
205 Value::Int(2_190),
206 Value::Int(2_191),
207 ];
208 for discount in &discounts {
209 for quantity in &quantities {
210 for ship_day in &ship_days {
211 assert_projected_parity(
212 &q6,
213 &predicate,
214 &fields,
215 &[discount.clone(), quantity.clone(), ship_day.clone()],
216 ¶ms,
217 );
218 }
219 }
220 }
221
222 let q1 = ScalarExpr::Binary {
223 op: BinaryOp::LessEqual,
224 lhs: Box::new(ScalarExpr::Column("ship_day".into())),
225 rhs: Box::new(ScalarExpr::Literal(Value::Int(2_449))),
226 };
227 let predicate = ProjectedPredicate::compile(&q1, &fields, &[])
228 .unwrap()
229 .unwrap();
230 for ship_day in [Value::Null, Value::Int(2_449), Value::Int(2_450)] {
231 assert_projected_parity(
232 &q1,
233 &predicate,
234 &fields,
235 &[Value::Int(0), Value::Int(0), ship_day],
236 &[],
237 );
238 }
239 }
240
241 #[test]
242 fn projected_integer_comparisons_match_every_canonical_operator_and_operand_order() {
243 let fields = vec!["value".into()];
244 for op in [
245 BinaryOp::Equal,
246 BinaryOp::NotEqual,
247 BinaryOp::Less,
248 BinaryOp::LessEqual,
249 BinaryOp::Greater,
250 BinaryOp::GreaterEqual,
251 ] {
252 for field_on_left in [true, false] {
253 let field = ScalarExpr::Column("value".into());
254 let literal = ScalarExpr::Literal(Value::Int(7));
255 let expression = ScalarExpr::Binary {
256 op,
257 lhs: Box::new(if field_on_left {
258 field.clone()
259 } else {
260 literal.clone()
261 }),
262 rhs: Box::new(if field_on_left {
263 literal.clone()
264 } else {
265 field.clone()
266 }),
267 };
268 let predicate = ProjectedPredicate::compile(&expression, &fields, &[])
269 .unwrap()
270 .unwrap();
271 for value in [
272 Value::Null,
273 Value::Int(6),
274 Value::Int(7),
275 Value::Int(8),
276 Value::Float(7.0),
277 ] {
278 assert_projected_parity(&expression, &predicate, &fields, &[value], &[]);
279 }
280 }
281 }
282 }
283
284 #[test]
285 fn projected_like_predicates_match_the_canonical_evaluator() {
286 for (name, pattern) in [
287 ("like", "%"),
288 ("like", "%green%"),
289 ("like", "%special%requests%"),
290 ("like", "a_c"),
291 ("ilike", "%GREEN%"),
292 ] {
293 let expression = ScalarExpr::Func {
294 name: name.into(),
295 binding: None,
296 args: vec![
297 ScalarExpr::Column("text".into()),
298 ScalarExpr::Literal(Value::Str(pattern.into())),
299 ],
300 distinct: false,
301 order_by: Vec::new(),
302 filter: None,
303 };
304 let fields = vec!["text".into()];
305 let predicate = ProjectedPredicate::compile(&expression, &fields, &[])
306 .unwrap()
307 .unwrap();
308 for value in [
309 Value::Str("forest green part".into()),
310 Value::FixedChar("GREEN ".into()),
311 Value::Str("a-c".into()),
312 Value::Str("special pending requests".into()),
313 Value::Null,
314 ] {
315 assert_projected_parity(&expression, &predicate, &fields, &[value], &[]);
316 }
317 }
318 }
319
320 #[test]
321 fn qualified_like_runs_directly_on_a_composite_physical_row() {
322 let expression = ScalarExpr::Not(Box::new(ScalarExpr::Func {
323 name: "like".into(),
324 binding: None,
325 args: vec![
326 ScalarExpr::qualified_column("o", "comment"),
327 ScalarExpr::Literal(Value::Str("%special%requests%".into())),
328 ],
329 distinct: false,
330 order_by: Vec::new(),
331 filter: None,
332 }));
333 let left_schema =
334 crate::RowSchema::with_qualified_types("c", vec!["id".into()], vec![None]);
335 let right_schema =
336 crate::RowSchema::with_qualified_types("o", vec!["comment".into()], vec![None]);
337 let schema = crate::RowSchema::join(&left_schema, &right_schema, std::iter::empty());
338 let predicate = ProjectedPredicate::compile_with_schema(&expression, &schema, &[])
339 .unwrap()
340 .unwrap();
341
342 let accepted = crate::PhysicalRow::concat(
343 &crate::PhysicalRow::from_values(vec![Value::Int(1)]),
344 &crate::PhysicalRow::from_values(vec![Value::Str("ordinary order".into())]),
345 );
346 let rejected = crate::PhysicalRow::concat(
347 &crate::PhysicalRow::from_values(vec![Value::Int(1)]),
348 &crate::PhysicalRow::from_values(vec![Value::Str("special pending requests".into())]),
349 );
350 assert!(predicate.keep_row(&schema.view(&accepted)).unwrap());
351 assert!(!predicate.keep_row(&schema.view(&rejected)).unwrap());
352 }
353
354 #[test]
355 fn projected_predicate_folds_typed_literals_once() {
356 let expression = ScalarExpr::Binary {
357 op: BinaryOp::Less,
358 lhs: Box::new(ScalarExpr::Column("day".into())),
359 rhs: Box::new(ScalarExpr::Cast {
360 expr: Box::new(ScalarExpr::Literal(Value::Str("1995-03-15".into()))),
361 ty: "date".into(),
362 }),
363 };
364 let predicate = ProjectedPredicate::compile(&expression, &["day".into()], &[])
365 .unwrap()
366 .unwrap();
367
368 let ProjectedExpr::Binary { rhs, .. } = &predicate.expression else {
369 panic!("expected a compiled comparison");
370 };
371 assert!(matches!(
372 rhs.as_ref(),
373 ProjectedExpr::Literal(Value::Temporal(_))
374 ));
375 }
376
377 #[test]
378 fn projected_predicate_preserves_unary_minus_integer_width() {
379 let expression = ScalarExpr::Binary {
380 op: BinaryOp::Equal,
381 lhs: Box::new(ScalarExpr::UnaryMinus(Box::new(ScalarExpr::Cast {
382 expr: Box::new(ScalarExpr::Column("x".into())),
383 ty: "smallint".into(),
384 }))),
385 rhs: Box::new(ScalarExpr::Literal(Value::Int(-1))),
386 };
387 let predicate = ProjectedPredicate::compile(&expression, &["x".into()], &[])
388 .unwrap()
389 .unwrap();
390
391 assert!(predicate.keep(&[&Value::Int(1)]).unwrap());
392 let error = predicate
393 .keep(&[&Value::Int(i64::from(i16::MIN))])
394 .expect_err("negating smallint minimum must overflow");
395 assert_eq!(error.sqlstate(), Some("22003"));
396 }
397
398 fn assert_projected_parity(
399 expression: &ScalarExpr,
400 predicate: &ProjectedPredicate,
401 fields: &[String],
402 values: &[Value],
403 params: &[SQLParam],
404 ) {
405 let row = fields
406 .iter()
407 .cloned()
408 .zip(values.iter().cloned())
409 .collect::<uqa_sql::ResultRow>();
410 let expected = eval_scalar(expression, &ScalarEvalContext::new(Some(&row), params))
411 .map(|value| truthy(&value));
412 let references = values.iter().collect::<Vec<_>>();
413 let actual = predicate.keep(&references);
414 match (expected, actual) {
415 (Ok(expected), Ok(actual)) => assert_eq!(actual, expected, "row: {row:?}"),
416 (Err(expected), Err(actual)) => assert_eq!(actual.to_string(), expected.to_string()),
417 (expected, actual) => panic!("projected result {actual:?} != canonical {expected:?}"),
418 }
419 }
420}