Skip to content

Commit e32e01e

Browse files
authored
fix: render tuple functions (#5783)
[TML-1922](https://linear.app/prisma-company/issue/TML-1922/fix-case-insensitive-in-regression) Enables the query compiler to render tuples of parameters wrapped in function calls. Needed for properly rendering `? IN (LOWER(?), LOWER(?))`. Client PR: prisma/prisma#29243
1 parent abbfa42 commit e32e01e

31 files changed

Lines changed: 340 additions & 61 deletions

libs/query-template/src/fragment.rs

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,12 @@ pub enum Fragment {
99
chunk: String,
1010
},
1111
Parameter,
12-
ParameterTuple,
12+
#[serde(rename_all = "camelCase")]
13+
ParameterTuple {
14+
item_prefix: Cow<'static, str>,
15+
item_separator: Cow<'static, str>,
16+
item_suffix: Cow<'static, str>,
17+
},
1318
#[serde(rename_all = "camelCase")]
1419
ParameterTupleList {
1520
item_prefix: Cow<'static, str>,

libs/query-template/src/template.rs

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -28,7 +28,7 @@ impl<P> QueryTemplate<P> {
2828
match fragment {
2929
Fragment::StringChunk { chunk } => sql.push_str(chunk),
3030
Fragment::Parameter => self.placeholder_format.write(&mut sql, &mut placeholder_number)?,
31-
Fragment::ParameterTuple | Fragment::ParameterTupleList { .. } => return Err(fmt::Error), // Unsupported in Query Engine
31+
Fragment::ParameterTuple { .. } | Fragment::ParameterTupleList { .. } => return Err(fmt::Error), // Unsupported in Query Engine
3232
};
3333
}
3434
Ok(sql)
@@ -46,7 +46,7 @@ impl<P> fmt::Display for QueryTemplate<P> {
4646
match fragment {
4747
Fragment::StringChunk { chunk } => write!(f, "{chunk}")?,
4848
Fragment::Parameter => self.placeholder_format.write(f, &mut placeholder_number)?,
49-
Fragment::ParameterTuple => {
49+
Fragment::ParameterTuple { .. } => {
5050
f.write_str("[")?;
5151
self.placeholder_format.write(f, &mut placeholder_number)?;
5252
f.write_str("]")?;

libs/query-template/tests/template.rs

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -47,7 +47,11 @@ fn new_query_template_with_parameter_tuple(pf: PlaceholderFormat) -> QueryTempla
4747
qt.fragments.push(Fragment::StringChunk {
4848
chunk: " AND status IN ".to_string(),
4949
});
50-
qt.fragments.push(Fragment::ParameterTuple);
50+
qt.fragments.push(Fragment::ParameterTuple {
51+
item_prefix: "".into(),
52+
item_separator: ", ".into(),
53+
item_suffix: "".into(),
54+
});
5155
qt
5256
}
5357

quaint/src/ast/function.rs

Lines changed: 65 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -42,6 +42,7 @@ pub use self::uuid::*;
4242

4343
use super::{Aliasable, Expression};
4444
use std::borrow::Cow;
45+
use std::slice;
4546

4647
/// A database function definition
4748
#[derive(Debug, Clone, PartialEq)]
@@ -92,6 +93,70 @@ pub(crate) enum FunctionType<'a> {
9293
Uuid,
9394
}
9495

96+
impl<'a> FunctionType<'a> {
97+
/// Returns the arguments of the function as a slice of expressions.
98+
/// Only returns a non-empty slice for functions that accept arbitrary expressions.
99+
pub fn arguments(&self) -> &[Expression<'a>] {
100+
match self {
101+
Self::Count(count) => &count.exprs,
102+
Self::AggregateToString(agg) => slice::from_ref(&agg.value),
103+
Self::Sum(avg) => slice::from_ref(&avg.expr),
104+
Self::Lower(f) => slice::from_ref(&f.expression),
105+
Self::Upper(f) => slice::from_ref(&f.expression),
106+
Self::Coalesce(f) => &f.exprs,
107+
Self::Concat(f) => &f.exprs,
108+
Self::JsonExtract(f) => slice::from_ref(&f.column),
109+
Self::JsonExtractLastArrayElem(f) => slice::from_ref(&f.expr),
110+
Self::JsonExtractFirstArrayElem(f) => slice::from_ref(&f.expr),
111+
Self::JsonUnquote(f) => slice::from_ref(&f.expr),
112+
Self::JsonArrayAgg(f) => slice::from_ref(&f.expr),
113+
Self::TextSearch(f) => &f.exprs,
114+
Self::TextSearchRelevance(f) => &f.exprs,
115+
Self::RowToJson(_)
116+
| Self::RowNumber(_)
117+
| Self::Average(_)
118+
| Self::Minimum(_)
119+
| Self::Maximum(_)
120+
| Self::JsonBuildObject(_)
121+
| Self::UuidToBin
122+
| Self::UuidToBinSwapped
123+
| Self::Uuid => &[],
124+
}
125+
}
126+
127+
/// Returns the name of the function, if it has an unambiguous name that can be used
128+
/// in all of the databases.
129+
pub fn name(&self) -> Option<&'static str> {
130+
// The list is based on the default `Visitor::visit_function`.
131+
let name = match self {
132+
Self::RowToJson(_) => "ROW_TO_JSON",
133+
Self::RowNumber(_) => "ROW_NUMBER",
134+
Self::Count(_) => "COUNT",
135+
Self::Sum(_) => "SUM",
136+
Self::Lower(_) => "LOWER",
137+
Self::Upper(_) => "UPPER",
138+
Self::Coalesce(_) => "COALESCE",
139+
Self::Concat(_)
140+
| Self::AggregateToString(_)
141+
| Self::Average(_)
142+
| Self::Minimum(_)
143+
| Self::Maximum(_)
144+
| Self::JsonExtract(_)
145+
| Self::JsonExtractLastArrayElem(_)
146+
| Self::JsonExtractFirstArrayElem(_)
147+
| Self::JsonUnquote(_)
148+
| Self::JsonArrayAgg(_)
149+
| Self::JsonBuildObject(_)
150+
| Self::TextSearch(_)
151+
| Self::TextSearchRelevance(_)
152+
| Self::UuidToBin
153+
| Self::UuidToBinSwapped
154+
| Self::Uuid => return None,
155+
};
156+
Some(name)
157+
}
158+
}
159+
95160
impl<'a> Aliasable<'a> for Function<'a> {
96161
type Target = Function<'a>;
97162

quaint/src/visitor.rs

Lines changed: 74 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -129,7 +129,13 @@ pub trait Visitor<'a> {
129129
fn parameter_substitution(&mut self) -> Result;
130130

131131
/// What to use to substitute a list of parameters of variable length
132-
fn visit_parameterized_row(&mut self, value: Value<'a>) -> Result;
132+
fn visit_parameterized_row(
133+
&mut self,
134+
value: Value<'a>,
135+
item_prefix: impl Into<Cow<'static, str>>,
136+
separator: impl Into<Cow<'static, str>>,
137+
item_suffix: impl Into<Cow<'static, str>>,
138+
) -> Result;
133139

134140
/// What to use to aggregate an array of values into a string
135141
fn visit_aggregate_to_string(&mut self, value: Expression<'a>) -> Result;
@@ -618,7 +624,7 @@ pub trait Visitor<'a> {
618624
ExpressionKind::ConditionTree(tree) => self.visit_conditions(tree)?,
619625
ExpressionKind::Compare(compare) => self.visit_compare(compare)?,
620626
ExpressionKind::Parameterized(val) => self.visit_parameterized(val)?,
621-
ExpressionKind::ParameterizedRow(val) => self.visit_parameterized_row(val)?,
627+
ExpressionKind::ParameterizedRow(val) => self.visit_parameterized_row(val, "", ",", "")?,
622628
ExpressionKind::RawValue(val) => self.visit_raw_value(val.0)?,
623629
ExpressionKind::Column(column) => self.visit_column(*column)?,
624630
ExpressionKind::Row(row) => self.visit_row(row)?,
@@ -877,15 +883,13 @@ pub trait Visitor<'a> {
877883
kind: ExpressionKind::Row(mut cols),
878884
..
879885
},
880-
Expression {
881-
kind: ExpressionKind::ParameterizedRow(value),
886+
rhs @ Expression {
887+
kind: ExpressionKind::ParameterizedRow(_),
882888
..
883889
},
884890
) if cols.len() == 1 => {
885891
let col = cols.pop().unwrap();
886-
self.visit_expression(col)?;
887-
self.write(" IN ")?;
888-
self.visit_parameterized_row(value)
892+
self.visit_compare(Compare::In(Box::new(col), Box::new(rhs)))
889893
}
890894

891895
// expr IN (?, ?, ..., ?)
@@ -898,7 +902,36 @@ pub trait Visitor<'a> {
898902
) => {
899903
self.visit_expression(left)?;
900904
self.write(" IN ")?;
901-
self.visit_parameterized_row(value)
905+
self.visit_parameterized_row(value, "", ",", "")
906+
}
907+
908+
// expr IN (CALL(?), CALL(?), ..., CALL(?))
909+
(
910+
left,
911+
Expression {
912+
kind: ExpressionKind::Function(value),
913+
..
914+
},
915+
) if value.typ_.arguments().len() == 1
916+
&& value
917+
.typ_
918+
.arguments()
919+
.iter()
920+
.all(|arg| matches!(arg.kind, ExpressionKind::ParameterizedRow(_))) =>
921+
{
922+
self.visit_expression(left)?;
923+
self.write(" IN ")?;
924+
925+
let Some(ExpressionKind::ParameterizedRow(val)) =
926+
value.typ_.arguments().first().map(|arg| &arg.kind)
927+
else {
928+
unreachable!()
929+
};
930+
let Some(function_name) = &value.typ_.name() else {
931+
panic!("function call against a row of expressions must have a name")
932+
};
933+
934+
self.visit_parameterized_row(val.clone(), format!("{function_name}("), ",", ")")
902935
}
903936

904937
(
@@ -979,15 +1012,13 @@ pub trait Visitor<'a> {
9791012
kind: ExpressionKind::Row(mut cols),
9801013
..
9811014
},
982-
Expression {
983-
kind: ExpressionKind::ParameterizedRow(value),
1015+
rhs @ Expression {
1016+
kind: ExpressionKind::ParameterizedRow(_),
9841017
..
9851018
},
9861019
) if cols.len() == 1 => {
9871020
let col = cols.pop().unwrap();
988-
self.visit_expression(col)?;
989-
self.write(" NOT IN ")?;
990-
self.visit_parameterized_row(value)
1021+
self.visit_compare(Compare::NotIn(Box::new(col), Box::new(rhs)))
9911022
}
9921023

9931024
// expr NOT IN (?, ?, ..., ?)
@@ -1000,7 +1031,7 @@ pub trait Visitor<'a> {
10001031
) => {
10011032
self.visit_expression(left)?;
10021033
self.write(" NOT IN ")?;
1003-
self.visit_parameterized_row(value)
1034+
self.visit_parameterized_row(value, "", ",", "")
10041035
}
10051036

10061037
(
@@ -1014,6 +1045,35 @@ pub trait Visitor<'a> {
10141045
},
10151046
) => self.visit_multiple_tuple_comparison(row, *values, true),
10161047

1048+
// expr NOT IN (CALL(?), CALL(?), ..., CALL(?))
1049+
(
1050+
left,
1051+
Expression {
1052+
kind: ExpressionKind::Function(value),
1053+
..
1054+
},
1055+
) if value.typ_.arguments().len() == 1
1056+
&& value
1057+
.typ_
1058+
.arguments()
1059+
.iter()
1060+
.all(|arg| matches!(arg.kind, ExpressionKind::ParameterizedRow(_))) =>
1061+
{
1062+
self.visit_expression(left)?;
1063+
self.write(" NOT IN ")?;
1064+
1065+
let Some(ExpressionKind::ParameterizedRow(val)) =
1066+
value.typ_.arguments().first().map(|arg| &arg.kind)
1067+
else {
1068+
unreachable!()
1069+
};
1070+
let Some(function_name) = &value.typ_.name() else {
1071+
panic!("function call against a row of expressions must have a name")
1072+
};
1073+
1074+
self.visit_parameterized_row(val.clone(), format!("{function_name}("), ",", ")")
1075+
}
1076+
10171077
// expr IN (..)
10181078
(left, right) => {
10191079
self.visit_expression(left)?;

quaint/src/visitor/mssql.rs

Lines changed: 9 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -231,8 +231,15 @@ impl<'a> Visitor<'a> for Mssql<'a> {
231231
Ok(())
232232
}
233233

234-
fn visit_parameterized_row(&mut self, value: Value<'a>) -> visitor::Result {
235-
self.query_template.write_parameter_tuple();
234+
fn visit_parameterized_row(
235+
&mut self,
236+
value: Value<'a>,
237+
item_prefix: impl Into<Cow<'static, str>>,
238+
separator: impl Into<Cow<'static, str>>,
239+
item_suffix: impl Into<Cow<'static, str>>,
240+
) -> visitor::Result {
241+
self.query_template
242+
.write_parameter_tuple(item_prefix, separator, item_suffix);
236243
self.query_template.parameters.push(value);
237244
Ok(())
238245
}

quaint/src/visitor/mysql.rs

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@ use crate::{
55
visitor::{self, Visitor},
66
};
77
use query_template::{PlaceholderFormat, QueryTemplate};
8+
use std::borrow::Cow;
89
use std::fmt;
910

1011
/// A visitor to generate queries for the MySQL database.
@@ -356,8 +357,15 @@ impl<'a> Visitor<'a> for Mysql<'a> {
356357
Ok(())
357358
}
358359

359-
fn visit_parameterized_row(&mut self, value: Value<'a>) -> visitor::Result {
360-
self.query_template.write_parameter_tuple();
360+
fn visit_parameterized_row(
361+
&mut self,
362+
value: Value<'a>,
363+
item_prefix: impl Into<Cow<'static, str>>,
364+
separator: impl Into<Cow<'static, str>>,
365+
item_suffix: impl Into<Cow<'static, str>>,
366+
) -> visitor::Result {
367+
self.query_template
368+
.write_parameter_tuple(item_prefix, separator, item_suffix);
361369
self.query_template.parameters.push(value);
362370
Ok(())
363371
}

quaint/src/visitor/postgres.rs

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@ use crate::{
66
};
77
use itertools::Itertools;
88
use query_template::{PlaceholderFormat, QueryTemplate};
9+
use std::borrow::Cow;
910
use std::{fmt, ops::Deref};
1011

1112
/// A visitor to generate queries for the PostgreSQL database.
@@ -92,8 +93,15 @@ impl<'a> Visitor<'a> for Postgres<'a> {
9293
Ok(())
9394
}
9495

95-
fn visit_parameterized_row(&mut self, value: Value<'a>) -> visitor::Result {
96-
self.query_template.write_parameter_tuple();
96+
fn visit_parameterized_row(
97+
&mut self,
98+
value: Value<'a>,
99+
item_prefix: impl Into<Cow<'static, str>>,
100+
separator: impl Into<Cow<'static, str>>,
101+
item_suffix: impl Into<Cow<'static, str>>,
102+
) -> visitor::Result {
103+
self.query_template
104+
.write_parameter_tuple(item_prefix, separator, item_suffix);
97105
self.query_template.parameters.push(value);
98106
Ok(())
99107
}

quaint/src/visitor/query_writer.rs

Lines changed: 17 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,12 @@ use query_template::{Fragment, QueryTemplate};
66
pub(crate) trait QueryWriter {
77
fn write_string_chunk(&mut self, value: String);
88
fn write_parameter(&mut self);
9-
fn write_parameter_tuple(&mut self);
9+
fn write_parameter_tuple(
10+
&mut self,
11+
item_prefix: impl Into<Cow<'static, str>>,
12+
item_separator: impl Into<Cow<'static, str>>,
13+
item_suffix: impl Into<Cow<'static, str>>,
14+
);
1015
fn write_parameter_tuple_list(
1116
&mut self,
1217
item_prefix: impl Into<Cow<'static, str>>,
@@ -32,8 +37,17 @@ impl QueryWriter for QueryTemplate<Value<'_>> {
3237
self.fragments.push(Fragment::Parameter);
3338
}
3439

35-
fn write_parameter_tuple(&mut self) {
36-
self.fragments.push(Fragment::ParameterTuple);
40+
fn write_parameter_tuple(
41+
&mut self,
42+
item_prefix: impl Into<Cow<'static, str>>,
43+
item_separator: impl Into<Cow<'static, str>>,
44+
item_suffix: impl Into<Cow<'static, str>>,
45+
) {
46+
self.fragments.push(Fragment::ParameterTuple {
47+
item_prefix: item_prefix.into(),
48+
item_separator: item_separator.into(),
49+
item_suffix: item_suffix.into(),
50+
});
3751
}
3852

3953
fn write_parameter_tuple_list(

quaint/src/visitor/sqlite.rs

Lines changed: 10 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,7 @@ use crate::{
66

77
use crate::visitor::query_writer::QueryWriter;
88
use query_template::{PlaceholderFormat, QueryTemplate};
9-
use std::fmt;
9+
use std::{borrow::Cow, fmt};
1010

1111
/// A visitor to generate queries for the SQLite database.
1212
///
@@ -272,8 +272,15 @@ impl<'a> Visitor<'a> for Sqlite<'a> {
272272
Ok(())
273273
}
274274

275-
fn visit_parameterized_row(&mut self, value: Value<'a>) -> visitor::Result {
276-
self.query_template.write_parameter_tuple();
275+
fn visit_parameterized_row(
276+
&mut self,
277+
value: Value<'a>,
278+
item_prefix: impl Into<Cow<'static, str>>,
279+
separator: impl Into<Cow<'static, str>>,
280+
item_suffix: impl Into<Cow<'static, str>>,
281+
) -> visitor::Result {
282+
self.query_template
283+
.write_parameter_tuple(item_prefix, separator, item_suffix);
277284
self.query_template.parameters.push(value);
278285
Ok(())
279286
}

0 commit comments

Comments
 (0)