@@ -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 } ;
88use sql_query_builder:: write:: split_write_args_by_shape;
9- use std:: { collections:: BTreeMap , iter} ;
9+ use std:: { collections:: BTreeMap , iter, mem } ;
1010
1111use 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+ }
0 commit comments