Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
74 changes: 51 additions & 23 deletions src/analyze/reconstruct_slice_indexing.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand All @@ -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,
Expand All @@ -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`.
Expand Down Expand Up @@ -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
Expand All @@ -377,32 +393,44 @@ 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);
}
}
}
}

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 {
Expand Down
24 changes: 24 additions & 0 deletions tests/ui/fail/slice_index_len.rs
Original file line number Diff line number Diff line change
@@ -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);
}
24 changes: 24 additions & 0 deletions tests/ui/pass/slice_index_len.rs
Original file line number Diff line number Diff line change
@@ -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);
}
Loading