Skip to content

Commit 610a68f

Browse files
authored
fix(qc): apply limit to selectors (#5461)
[ORM-1044](https://linear.app/prisma-company/issue/ORM-1044/fix-limit-not-being-applied-to-updatesdeletes) Generates an in-memory take for query selectors when the update/delete have a limit
1 parent 92dceb1 commit 610a68f

21 files changed

Lines changed: 214 additions & 78 deletions

File tree

Cargo.lock

Lines changed: 1 addition & 0 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

query-compiler/query-compiler/Cargo.toml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@ quaint.workspace = true
1414
thiserror.workspace = true
1515
serde.workspace = true
1616
itertools.workspace = true
17+
bon.workspace = true
1718
pretty = { workspace = true, features = ["termcolor"] }
1819
indexmap = { workspace = true, features = ["serde"] }
1920

query-compiler/query-compiler/src/binding.rs

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@ use query_structure::{ScalarField, SelectedField};
66
const JOIN_PARENT: &str = "@parent";
77
const DEFAULTS: &str = "@defaults";
88
const GENERATED: &str = "@generated";
9+
const SELECTOR: &str = "@selector";
910

1011
const FIELD_SEPARATOR: &str = "$";
1112

@@ -32,3 +33,7 @@ pub fn defaults() -> Cow<'static, str> {
3233
pub fn generated(row_idx: usize, field_name: &str) -> Cow<'static, str> {
3334
format!("{GENERATED}{FIELD_SEPARATOR}row{row_idx}{FIELD_SEPARATOR}{field_name}").into()
3435
}
36+
37+
pub fn selector(field: &SelectedField) -> Cow<'static, str> {
38+
format!("{SELECTOR}{FIELD_SEPARATOR}{}", field.prisma_name()).into()
39+
}

query-compiler/query-compiler/src/data_mapper.rs

Lines changed: 1 addition & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -110,10 +110,7 @@ fn get_result_node(
110110
for prisma_name in selection_order {
111111
match field_map.get(prisma_name.as_str()) {
112112
Some(sf @ SelectedField::Scalar(f)) => {
113-
node.add_field(
114-
prisma_name,
115-
builder.new_value(sf.db_name().into_owned(), f.result_type()),
116-
);
113+
node.add_field(prisma_name, builder.new_value(sf.db_name().into_owned(), f.type_info()));
117114
}
118115
Some(SelectedField::Composite(_)) => todo!("MongoDB specific"),
119116
Some(SelectedField::Relation(f)) => {

query-compiler/query-compiler/src/expression.rs

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@ use std::{
44
};
55

66
use crate::result_node::ResultNode;
7+
use bon::bon;
78
use query_builder::DbQuery;
89
use query_core::{DataExpectation, DataRule};
910
use query_structure::{InternalEnum, PrismaValue, PrismaValueType, ScalarWriteOperation, TaggedPrismaValue};
@@ -194,7 +195,9 @@ pub struct Pagination {
194195
linking_fields: Option<Vec<String>>,
195196
}
196197

198+
#[bon]
197199
impl Pagination {
200+
#[builder]
198201
pub fn new(cursor: Option<HashMap<String, PrismaValue>>, take: Option<i64>, skip: Option<i64>) -> Self {
199202
Self {
200203
cursor,

query-compiler/query-compiler/src/result_node.rs

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11
use indexmap::IndexMap;
2-
use query_structure::{PrismaValueType, ScalarFieldResultType, TypeIdentifier};
2+
use query_structure::{FieldTypeInformation, PrismaValueType, TypeIdentifier};
33
use serde::Serialize;
44

55
use crate::expression::EnumsMap;
@@ -67,7 +67,7 @@ impl<'a> ResultNodeBuilder<'a> {
6767
ObjectBuilder::new(ObjectKind::Flattened)
6868
}
6969

70-
pub fn new_value(&mut self, db_name: String, result_type: ScalarFieldResultType) -> ResultNode {
70+
pub fn new_value(&mut self, db_name: String, result_type: FieldTypeInformation) -> ResultNode {
7171
let prisma_type = result_type.to_prisma_type();
7272
if let TypeIdentifier::Enum(id) = result_type.typ.id {
7373
self.enums.add(result_type.typ.dm.zip(id));

query-compiler/query-compiler/src/translate.rs

Lines changed: 16 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,9 @@ use query_core::{
1313
Computation, EdgeRef, Flow, Node, NodeRef, Query, QueryGraph, QueryGraphBuilderError, QueryGraphDependency,
1414
QueryGraphError, RowCountSink, RowSink,
1515
};
16-
use query_structure::{FieldSelection, Placeholder, PrismaValue, PrismaValueType, SelectedField, SelectionResult};
16+
use query_structure::{
17+
FieldSelection, FieldTypeInformation, Placeholder, PrismaValue, PrismaValueType, SelectedField, SelectionResult,
18+
};
1719
use thiserror::Error;
1820

1921
#[derive(Debug, Error)]
@@ -502,10 +504,22 @@ impl<'a, 'b> NodeTranslator<'a, 'b> {
502504
selection: FieldSelection,
503505
) -> Vec<(SelectedField, PrismaValue)> {
504506
let bindings_refer_to_fields = matches!(node, Node::Query(_));
507+
let binding_is_unique = matches!(node, Node::Query(q) if q.is_unique());
505508

506509
selection
507510
.selections()
508511
.map(|field| {
512+
let r#type = field
513+
.type_info()
514+
.as_ref()
515+
.map(FieldTypeInformation::to_prisma_type)
516+
.unwrap_or(PrismaValueType::Any);
517+
let r#type = if binding_is_unique {
518+
r#type
519+
} else {
520+
PrismaValueType::Array(r#type.into())
521+
};
522+
509523
(
510524
field.clone(),
511525
PrismaValue::Placeholder(Placeholder {
@@ -514,7 +528,7 @@ impl<'a, 'b> NodeTranslator<'a, 'b> {
514528
} else {
515529
binding::node_result(self.graph.edge_source(edge))
516530
},
517-
r#type: PrismaValueType::Any,
531+
r#type,
518532
}),
519533
)
520534
})

query-compiler/query-compiler/src/translate/query/read.rs

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -173,7 +173,7 @@ pub(super) fn add_inmemory_join(
173173
.map(|(parent_scalar, child_scalar)| {
174174
let placeholder = PrismaValue::placeholder(
175175
binding::join_parent_field(parent_scalar),
176-
parent_scalar.result_type().to_prisma_type(),
176+
parent_scalar.type_info().to_prisma_type(),
177177
);
178178
let condition = if parent.r#type().is_list() {
179179
ScalarCondition::InTemplate(ConditionValue::value(placeholder))
@@ -374,7 +374,11 @@ fn extract_pagination(args: &mut QueryArguments) -> Pagination {
374374
.map(|(sf, val)| (sf.db_name().into_owned(), val.clone()))
375375
.collect()
376376
});
377-
Pagination::new(cursor, args.take.abs(), args.skip)
377+
Pagination::builder()
378+
.maybe_cursor(cursor)
379+
.maybe_take(args.take.abs())
380+
.maybe_skip(args.skip)
381+
.build()
378382
}
379383

380384
fn extract_distinct_by(args: &mut QueryArguments) -> Vec<String> {

query-compiler/query-compiler/src/translate/query/write.rs

Lines changed: 67 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -4,13 +4,13 @@ use query_core::{
44
ConnectRecords, DeleteManyRecords, DeleteRecord, DisconnectRecords, RawQuery, UpdateManyRecords, UpdateRecord,
55
UpdateRecordWithSelection, UpdateRecordWithoutSelection, WriteQuery,
66
};
7-
use query_structure::{PrismaValue, QueryArguments, RelationLoadStrategy, Take};
7+
use query_structure::{PrismaValue, PrismaValueType, QueryArguments, RecordFilter, RelationLoadStrategy, Take};
88
use sql_query_builder::write::split_write_args_by_shape;
9-
use std::{collections::BTreeMap, iter};
9+
use std::{collections::BTreeMap, iter, mem};
1010

1111
use crate::{
1212
TranslateError, binding,
13-
expression::{Binding, Expression, FieldInitializer, FieldOperation},
13+
expression::{Binding, Expression, FieldInitializer, FieldOperation, Pagination},
1414
translate::TranslateResult,
1515
};
1616

@@ -117,12 +117,17 @@ pub(crate) fn translate_write_query(query: WriteQuery, builder: &dyn QueryBuilde
117117
WriteQuery::UpdateManyRecords(UpdateManyRecords {
118118
name: _,
119119
model,
120-
record_filter,
120+
mut record_filter,
121121
args,
122122
selected_fields,
123123
limit,
124124
}) => {
125125
let projection = selected_fields.as_ref().map(|f| &f.fields);
126+
127+
let selector_bindings = limit
128+
.map(|limit| extract_selectors_that_require_limit(&mut record_filter, limit))
129+
.unwrap_or_default();
130+
126131
let updates = builder
127132
.build_updates(&model, record_filter, args, projection, limit)
128133
.map_err(TranslateError::QueryBuildFailure)?
@@ -134,15 +139,24 @@ pub(crate) fn translate_write_query(query: WriteQuery, builder: &dyn QueryBuilde
134139
})
135140
.collect::<Vec<_>>();
136141

137-
if let Some(selected_fields) = selected_fields {
142+
let mut expr = if let Some(selected_fields) = selected_fields {
138143
let mut expr = Expression::Concat(updates);
139144
if !selected_fields.nested.is_empty() {
140145
expr = add_inmemory_join(expr, selected_fields.nested, builder)?;
141146
}
142147
expr
143148
} else {
144149
Expression::Sum(updates)
150+
};
151+
152+
if !selector_bindings.is_empty() {
153+
expr = Expression::Let {
154+
bindings: selector_bindings,
155+
expr: expr.into(),
156+
};
145157
}
158+
159+
expr
146160
}
147161

148162
WriteQuery::UpdateRecord(UpdateRecord::WithSelection(UpdateRecordWithSelection {
@@ -282,16 +296,31 @@ pub(crate) fn translate_write_query(query: WriteQuery, builder: &dyn QueryBuilde
282296

283297
WriteQuery::DeleteManyRecords(DeleteManyRecords {
284298
model,
285-
record_filter,
299+
mut record_filter,
286300
limit,
287-
}) => Expression::Sum(
288-
builder
289-
.build_deletes(&model, record_filter, limit)
290-
.map_err(TranslateError::QueryBuildFailure)?
291-
.into_iter()
292-
.map(Expression::Execute)
293-
.collect::<Vec<_>>(),
294-
),
301+
}) => {
302+
let selector_bindings = limit
303+
.map(|limit| extract_selectors_that_require_limit(&mut record_filter, limit))
304+
.unwrap_or_default();
305+
306+
let mut expr = Expression::Sum(
307+
builder
308+
.build_deletes(&model, record_filter, limit)
309+
.map_err(TranslateError::QueryBuildFailure)?
310+
.into_iter()
311+
.map(Expression::Execute)
312+
.collect::<Vec<_>>(),
313+
);
314+
315+
if !selector_bindings.is_empty() {
316+
expr = Expression::Let {
317+
bindings: selector_bindings,
318+
expr: expr.into(),
319+
};
320+
}
321+
322+
expr
323+
}
295324

296325
WriteQuery::ConnectRecords(ConnectRecords {
297326
parent_id,
@@ -327,3 +356,27 @@ pub(crate) fn translate_write_query(query: WriteQuery, builder: &dyn QueryBuilde
327356
}
328357
})
329358
}
359+
360+
/// Extracts selectors from the filter that require an in-memory limit operation as bindings.
361+
/// Selectors in the [`RecordFilter`] are replaced with placeholders that refer to the
362+
/// returned bindings.
363+
fn extract_selectors_that_require_limit(record_filter: &mut RecordFilter, limit: usize) -> Vec<Binding> {
364+
record_filter
365+
.selectors
366+
.iter_mut()
367+
.flatten()
368+
.flat_map(|result| result.pairs.iter_mut())
369+
.filter_map(|(field, value)| {
370+
let typ = value.r#type();
371+
if !matches!(typ, PrismaValueType::Array(_)) {
372+
return None;
373+
}
374+
375+
let name = binding::selector(field);
376+
let value = mem::replace(value, PrismaValue::placeholder(name.clone(), typ));
377+
let pagination = Pagination::builder().take(limit as i64).build();
378+
let expr = Expression::Value(value).into();
379+
Some(Binding::new(name, Expression::Paginate { expr, pagination }))
380+
})
381+
.collect_vec()
382+
}

query-compiler/query-compiler/tests/snapshots/queries__queries@create-m2m.json.snap

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -44,7 +44,7 @@ transaction
4444
2$id = mapField id (get 2)
4545
in execute «INSERT INTO "public"."_CategoryToPost" ("B","A")
4646
VALUES [($1)] ON CONFLICT DO NOTHING»
47-
params [product(var(1$id as Any), var(2$id as Any))];
47+
params [product(var(1$id as Int[]), var(2$id as Int[]))];
4848
let 4 = let 0 = validate (get 0)
4949
[ rowCountNeq 0
5050
] orRaise "MISSING_RECORD";

0 commit comments

Comments
 (0)