Skip to content
Merged
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
103 changes: 36 additions & 67 deletions compiler/rustc_mir_transform/src/dataflow_const_prop.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,6 @@
//! Currently, this pass only propagates scalar values.

use std::assert_matches;
use std::cell::RefCell;
use std::fmt::Formatter;

use rustc_abi::{BackendRepr, FIRST_VARIANT, FieldIdx, Size, VariantIdx};
Expand Down Expand Up @@ -70,7 +69,7 @@ impl<'tcx> crate::MirPass<'tcx> for DataflowConstProp {
.in_scope(|| ConstAnalysis::new(tcx, body, map).iterate_to_fixpoint(tcx, body, None));

// Collect results and patch the body afterwards.
let mut visitor = Collector::new(tcx, &body.local_decls);
let mut visitor = Collector::new(tcx, body);
debug_span!("collect").in_scope(|| visit_reachable_results(body, &const_, &mut visitor));
let mut patch = visitor.patch;
debug_span!("patch").in_scope(|| patch.visit_body_preserves_cfg(body));
Expand All @@ -89,7 +88,7 @@ struct ConstAnalysis<'a, 'tcx> {
map: Map<'tcx>,
tcx: TyCtxt<'tcx>,
local_decls: &'a LocalDecls<'tcx>,
ecx: RefCell<InterpCx<'tcx, DummyMachine>>,
ecx: InterpCx<'tcx, DummyMachine>,
typing_env: ty::TypingEnv<'tcx>,
}

Expand Down Expand Up @@ -157,7 +156,7 @@ impl<'a, 'tcx> ConstAnalysis<'a, 'tcx> {
map,
tcx,
local_decls: &body.local_decls,
ecx: RefCell::new(InterpCx::new(tcx, DUMMY_SP, typing_env, DummyMachine)),
ecx: InterpCx::new(tcx, DUMMY_SP, typing_env, DummyMachine),
typing_env,
}
}
Expand Down Expand Up @@ -412,7 +411,6 @@ impl<'a, 'tcx> ConstAnalysis<'a, 'tcx> {
match self.eval_operand(operand, state) {
FlatSet::Elem(op) => self
.ecx
.borrow()
.int_to_int_or_float(&op, layout)
.discard_err()
.map_or(FlatSet::Top, |result| self.wrap_immediate(*result)),
Expand All @@ -427,7 +425,6 @@ impl<'a, 'tcx> ConstAnalysis<'a, 'tcx> {
match self.eval_operand(operand, state) {
FlatSet::Elem(op) => self
.ecx
.borrow()
.float_to_float_or_int(&op, layout)
.discard_err()
.map_or(FlatSet::Top, |result| self.wrap_immediate(*result)),
Expand Down Expand Up @@ -458,7 +455,6 @@ impl<'a, 'tcx> ConstAnalysis<'a, 'tcx> {
match self.eval_operand(operand, state) {
FlatSet::Elem(value) => self
.ecx
.borrow()
.unary_op(*op, &value)
.discard_err()
.map_or(FlatSet::Top, |val| self.wrap_immediate(*val)),
Expand Down Expand Up @@ -547,11 +543,8 @@ impl<'a, 'tcx> ConstAnalysis<'a, 'tcx> {
}
}
Operand::Constant(constant) => {
if let Some(constant) = self
.ecx
.borrow()
.eval_mir_constant(&constant.const_, constant.span, None)
.discard_err()
if let Some(constant) =
self.ecx.eval_mir_constant(&constant.const_, constant.span, None).discard_err()
{
self.assign_constant(state, place, constant, &[]);
}
Expand Down Expand Up @@ -581,9 +574,7 @@ impl<'a, 'tcx> ConstAnalysis<'a, 'tcx> {
return;
}
}
operand = if let Some(operand) =
self.ecx.borrow().project(&operand, proj_elem).discard_err()
{
operand = if let Some(operand) = self.ecx.project(&operand, proj_elem).discard_err() {
operand
} else {
return;
Expand All @@ -594,22 +585,17 @@ impl<'a, 'tcx> ConstAnalysis<'a, 'tcx> {
place,
operand,
&mut |elem, op| match elem {
TrackElem::Field(idx) => self.ecx.borrow().project_field(op, idx).discard_err(),
TrackElem::Variant(idx) => {
self.ecx.borrow().project_downcast(op, idx).discard_err()
}
TrackElem::Field(idx) => self.ecx.project_field(op, idx).discard_err(),
TrackElem::Variant(idx) => self.ecx.project_downcast(op, idx).discard_err(),
TrackElem::Discriminant => {
let variant = self.ecx.borrow().read_discriminant(op).discard_err()?;
let discr_value = self
.ecx
.borrow()
.discriminant_for_variant(op.layout.ty, variant)
.discard_err()?;
let variant = self.ecx.read_discriminant(op).discard_err()?;
let discr_value =
self.ecx.discriminant_for_variant(op.layout.ty, variant).discard_err()?;
Some(discr_value.into())
}
TrackElem::DerefLen => {
let op: OpTy<'_> = self.ecx.borrow().deref_pointer(op).discard_err()?.into();
let len_usize = op.len(&self.ecx.borrow()).discard_err()?;
let op: OpTy<'_> = self.ecx.deref_pointer(op).discard_err()?.into();
let len_usize = op.len(&self.ecx).discard_err()?;
let layout = self
.tcx
.layout_of(self.typing_env.as_query_input(self.tcx.types.usize))
Expand All @@ -618,7 +604,7 @@ impl<'a, 'tcx> ConstAnalysis<'a, 'tcx> {
}
},
&mut |place, op| {
if let Some(imm) = self.ecx.borrow().read_immediate_raw(op).discard_err()
if let Some(imm) = self.ecx.read_immediate_raw(op).discard_err()
&& let Some(imm) = imm.right()
{
let elem = self.wrap_immediate(*imm);
Expand All @@ -642,7 +628,7 @@ impl<'a, 'tcx> ConstAnalysis<'a, 'tcx> {
(FlatSet::Bottom, _) | (_, FlatSet::Bottom) => (FlatSet::Bottom, FlatSet::Bottom),
// Both sides are known, do the actual computation.
(FlatSet::Elem(left), FlatSet::Elem(right)) => {
match self.ecx.borrow().binary_op(op, &left, &right).discard_err() {
match self.ecx.binary_op(op, &left, &right).discard_err() {
// Ideally this would return an Immediate, since it's sometimes
// a pair and sometimes not. But as a hack we always return a pair
// and just make the 2nd component `Bottom` when it does not exist.
Expand Down Expand Up @@ -715,11 +701,8 @@ impl<'a, 'tcx> ConstAnalysis<'a, 'tcx> {
return None;
}
let enum_ty_layout = self.tcx.layout_of(self.typing_env.as_query_input(enum_ty)).ok()?;
let discr_value = self
.ecx
.borrow()
.discriminant_for_variant(enum_ty_layout.ty, variant_index)
.discard_err()?;
let discr_value =
self.ecx.discriminant_for_variant(enum_ty_layout.ty, variant_index).discard_err()?;
Some(discr_value.to_scalar())
}

Expand Down Expand Up @@ -781,23 +764,27 @@ impl<'tcx> Patch<'tcx> {
struct Collector<'a, 'tcx> {
patch: Patch<'tcx>,
local_decls: &'a LocalDecls<'tcx>,
ecx: InterpCx<'tcx, DummyMachine>,
}

impl<'a, 'tcx> Collector<'a, 'tcx> {
pub(crate) fn new(tcx: TyCtxt<'tcx>, local_decls: &'a LocalDecls<'tcx>) -> Self {
Self { patch: Patch::new(tcx), local_decls }
pub(crate) fn new(tcx: TyCtxt<'tcx>, body: &'a Body<'tcx>) -> Self {
Self {
patch: Patch::new(tcx),
local_decls: &body.local_decls,
ecx: InterpCx::new(tcx, DUMMY_SP, body.typing_env(tcx), DummyMachine),
}
}

#[instrument(level = "trace", skip(self, ecx, map), ret)]
#[instrument(level = "trace", skip(self, map), ret)]
fn try_make_constant(
&self,
ecx: &mut InterpCx<'tcx, DummyMachine>,
&mut self,
place: Place<'tcx>,
state: &State<FlatSet<Scalar>>,
map: &Map<'tcx>,
) -> Option<Const<'tcx>> {
let ty = place.ty(self.local_decls, self.patch.tcx).ty;
let layout = ecx.layout_of(ty).ok()?;
let layout = self.ecx.layout_of(ty).ok()?;

if layout.is_zst() {
return Some(Const::zero_sized(ty));
Expand All @@ -815,7 +802,8 @@ impl<'a, 'tcx> Collector<'a, 'tcx> {
}

if matches!(layout.backend_repr, BackendRepr::Scalar(..) | BackendRepr::ScalarPair { .. }) {
let alloc_id = ecx
let alloc_id = self
.ecx
.intern_with_temp_alloc(layout, |ecx, dest| {
try_write_constant(ecx, dest, place, ty, state, map)
})
Expand Down Expand Up @@ -960,13 +948,8 @@ impl<'tcx> ResultsVisitor<'tcx, ConstAnalysis<'_, 'tcx>> for Collector<'_, 'tcx>
) {
match &statement.kind {
StatementKind::Assign((_, rvalue)) => {
OperandCollector {
state,
visitor: self,
ecx: &mut analysis.ecx.borrow_mut(),
map: &analysis.map,
}
.visit_rvalue(rvalue, location);
OperandCollector { state, visitor: self, map: &analysis.map }
.visit_rvalue(rvalue, location);
}
_ => (),
}
Expand All @@ -985,12 +968,7 @@ impl<'tcx> ResultsVisitor<'tcx, ConstAnalysis<'_, 'tcx>> for Collector<'_, 'tcx>
// Don't overwrite the assignment if it already uses a constant (to keep the span).
}
StatementKind::Assign((place, _)) => {
if let Some(value) = self.try_make_constant(
&mut analysis.ecx.borrow_mut(),
place,
state,
&analysis.map,
) {
if let Some(value) = self.try_make_constant(place, state, &analysis.map) {
self.patch.assignments.insert(location, value);
}
}
Expand All @@ -1005,13 +983,8 @@ impl<'tcx> ResultsVisitor<'tcx, ConstAnalysis<'_, 'tcx>> for Collector<'_, 'tcx>
terminator: &Terminator<'tcx>,
location: Location,
) {
OperandCollector {
state,
visitor: self,
ecx: &mut analysis.ecx.borrow_mut(),
map: &analysis.map,
}
.visit_terminator(terminator, location);
OperandCollector { state, visitor: self, map: &analysis.map }
.visit_terminator(terminator, location);
}
}

Expand Down Expand Up @@ -1070,7 +1043,6 @@ impl<'tcx> MutVisitor<'tcx> for Patch<'tcx> {
struct OperandCollector<'a, 'b, 'tcx> {
state: &'a State<FlatSet<Scalar>>,
visitor: &'a mut Collector<'b, 'tcx>,
ecx: &'a mut InterpCx<'tcx, DummyMachine>,
map: &'a Map<'tcx>,
}

Expand All @@ -1083,18 +1055,15 @@ impl<'tcx> Visitor<'tcx> for OperandCollector<'_, '_, 'tcx> {
location: Location,
) {
if let PlaceElem::Index(local) = elem
&& let Some(value) =
self.visitor.try_make_constant(self.ecx, local.into(), self.state, self.map)
&& let Some(value) = self.visitor.try_make_constant(local.into(), self.state, self.map)
{
self.visitor.patch.before_effect.insert((location, local.into()), value);
}
}

fn visit_operand(&mut self, operand: &Operand<'tcx>, location: Location) {
if let Some(place) = operand.place() {
if let Some(value) =
self.visitor.try_make_constant(self.ecx, place, self.state, self.map)
{
if let Some(value) = self.visitor.try_make_constant(place, self.state, self.map) {
self.visitor.patch.before_effect.insert((location, place), value);
} else if !place.projection.is_empty() {
// Try to propagate into `Index` projections.
Expand Down
Loading