diff --git a/src/analyze/reconstruct_slice_indexing.rs b/src/analyze/reconstruct_slice_indexing.rs index 0e92c911..dfe226df 100644 --- a/src/analyze/reconstruct_slice_indexing.rs +++ b/src/analyze/reconstruct_slice_indexing.rs @@ -79,6 +79,21 @@ impl<'tcx> mir::visit::Visitor<'tcx> for IndexedPlaceFinder<'_, 'tcx> { } } +struct ReadFinder { + local: Local, + found: bool, +} + +impl<'tcx> mir::visit::Visitor<'tcx> for ReadFinder { + fn visit_local(&mut self, local: Local, context: mir::visit::PlaceContext, _: mir::Location) { + let is_store = + context == mir::visit::PlaceContext::MutatingUse(mir::visit::MutatingUseContext::Store); + if local == self.local && context.is_use() && !is_store { + self.found = true; + } + } +} + /// Reconstructs the trait call erased by MIR's first-class slice indexing operation. /// /// For example, optimized MIR for `slice[index]` contains: @@ -269,7 +284,6 @@ fn reconstruct_access<'tcx>( result_local, ); let receiver = receiver_operand(tcx, body, &bounds_check, &access, region); - remove_bounds_check_setup(body, &bounds_check, access.receiver.local); let (lang_item, method_name) = if access.mutable { (LangItem::IndexMut, sym::index_mut) @@ -279,7 +293,7 @@ fn reconstruct_access<'tcx>( let method = lang_item_method(tcx, lang_item, method_name); let args = tcx.mk_args(&[access.slice_ty.into(), tcx.types.usize.into()]); let func = super::fn_operand(tcx, method, args, bounds_check.source_info.span); - let call_args = [receiver, bounds_check.index] + let call_args = [receiver, bounds_check.index.clone()] .into_iter() .map(|node| Spanned { node, @@ -301,6 +315,7 @@ fn reconstruct_access<'tcx>( }, }); tracing::trace!(?result_local, ?method, "slice indexing call inserted"); + remove_bounds_check_setup(body, &bounds_check); } /// Replaces every use of the indexed place in the target block with `*result_local`. @@ -358,11 +373,12 @@ fn receiver_operand<'tcx>( } /// Removes the MIR temporaries that only supported the now-replaced bounds check. -fn remove_bounds_check_setup<'tcx>( - body: &mut Body<'tcx>, - bounds_check: &BoundsCheck<'tcx>, - receiver_local: Local, -) { +/// +/// A temporary that is still read elsewhere is kept: with optimizations enabled, rustc reuses +/// a `len()` the program computes itself as the bounds check's length (`a[a.len() - 1]`). +fn remove_bounds_check_setup<'tcx>(body: &mut Body<'tcx>, bounds_check: &BoundsCheck<'tcx>) { + // Ordered so that each local's readers among these come before it: the condition reads the + // length, which reads the `PtrMetadata` operand. let mut lowered_locals: Vec<_> = bounds_check.condition_local.into_iter().collect(); if let Some(len_place) = bounds_check .len @@ -377,13 +393,7 @@ fn remove_bounds_check_setup<'tcx>( continue; }; if lhs.local == len_place.local { - // Only include the PtrMetadata operand if it is a distinct temporary — i.e. - // not the receiver that the reconstructed Index::index call will use. - // When PtrMetadata is applied directly to the slice reference (the receiver), - // that local must not be NOP'd: its assignment may be in this same block. - if let Some(raw_place) = operand - .place() - .filter(|place| place.projection.is_empty() && place.local != receiver_local) + if let Some(raw_place) = operand.place().filter(|place| place.projection.is_empty()) { lowered_locals.push(raw_place.local); } @@ -391,18 +401,36 @@ fn remove_bounds_check_setup<'tcx>( } } - tracing::trace!( - ?lowered_locals, - "removing replaced bounds-check temporaries" - ); - for statement in &mut body.basic_blocks.as_mut()[bounds_check.block].statements { - let Some((lhs, _)) = statement.kind.as_assign() else { + for local in lowered_locals { + if is_read(body, local) { + tracing::trace!(?local, "keeping bounds-check temporary that is still read"); continue; - }; - if lowered_locals.contains(&lhs.local) { - statement.kind = StatementKind::Nop; } + tracing::trace!(?local, "removing replaced bounds-check temporary"); + for statement in &mut body.basic_blocks.as_mut()[bounds_check.block].statements { + if statement + .kind + .as_assign() + .is_some_and(|(lhs, _)| lhs.local == local) + { + statement.kind = StatementKind::Nop; + } + } + } +} + +/// Whether any statement or terminator in `body` uses `local` other than by assigning to it. +fn is_read(body: &Body<'_>, local: Local) -> bool { + use mir::visit::Visitor as _; + + let mut finder = ReadFinder { + local, + found: false, + }; + for (block, data) in body.basic_blocks.iter_enumerated() { + finder.visit_basic_block_data(block, data); } + finder.found } fn lang_item_method(tcx: TyCtxt<'_>, item: LangItem, name: Symbol) -> DefId { diff --git a/tests/ui/fail/slice_index_len.rs b/tests/ui/fail/slice_index_len.rs new file mode 100644 index 00000000..63e8df2a --- /dev/null +++ b/tests/ui/fail/slice_index_len.rs @@ -0,0 +1,24 @@ +//@error-in-other-file: Unsat +//@compile-flags: -C debug-assertions=off -C opt-level=1 +//@rustc-env: THRUST_SOLVER=tests/thrust-pcsat-wrapper + +#[thrust::trusted] +#[thrust_macros::requires(true)] +#[thrust_macros::ensures( + (*result).len() == 3 + && (*result)[0] == 10 + && (*result)[1] == 20 + && (*result)[2] == 30 +)] +fn slice() -> &'static [i32] { + unimplemented!() +} + +fn last(slice: &[i32]) -> i32 { + slice[slice.len() - 1] +} + +fn main() { + let slice = slice(); + assert!(last(slice) == 20); +} diff --git a/tests/ui/pass/slice_index_len.rs b/tests/ui/pass/slice_index_len.rs new file mode 100644 index 00000000..b6c00dba --- /dev/null +++ b/tests/ui/pass/slice_index_len.rs @@ -0,0 +1,24 @@ +//@check-pass +//@compile-flags: -C debug-assertions=off -C opt-level=1 +//@rustc-env: THRUST_SOLVER=tests/thrust-pcsat-wrapper + +#[thrust::trusted] +#[thrust_macros::requires(true)] +#[thrust_macros::ensures( + (*result).len() == 3 + && (*result)[0] == 10 + && (*result)[1] == 20 + && (*result)[2] == 30 +)] +fn slice() -> &'static [i32] { + unimplemented!() +} + +fn last(slice: &[i32]) -> i32 { + slice[slice.len() - 1] +} + +fn main() { + let slice = slice(); + assert!(last(slice) == 30); +}