use rustc_hir::def_id::DefId;
use rustc_middle::ty::{Ty, TyCtxt, TyKind};
use rustc_abi::FieldIdx;
use crate::verify::{
contract::{ContractExpr, ContractPlace, Property, PropertyArg, PropertyKind},
def_use::{PlaceBaseKey, PlaceKey},
helpers::Checkpoint,
target::get_struct_invariants_for_adt,
verifier::{AbstractValue, ForwardVisitResult, StateFact},
};
const MAX_TRACE_STEPS: usize = 32;
pub(super) fn discharge_from_field_invariant<'tcx>(
tcx: TyCtxt<'tcx>,
caller: DefId,
target: &PlaceKey,
forward: &ForwardVisitResult<'tcx>,
kind: PropertyKind,
required_ty: Option<Ty<'tcx>>,
required_elements: Option<u64>,
) -> Option<String> {
let body = tcx.optimized_mir(caller);
let mut current = target.clone();
let mut visited: Vec<PlaceKey> = Vec::new();
for _ in 0..MAX_TRACE_STEPS {
if visited.contains(¤t) {
break;
}
visited.push(current.clone());
if let Some(reason) = field_invariant_matches(
tcx,
caller,
body,
¤t,
kind.clone(),
required_ty,
required_elements,
) {
return Some(reason);
}
let Some(next) = substitute_base(¤t, forward) else {
break;
};
current = next;
}
None
}
fn substitute_base<'tcx>(place: &PlaceKey, forward: &ForwardVisitResult<'tcx>) -> Option<PlaceKey> {
let local = place.local()?;
let source = match forward.values.get(&local) {
Some(AbstractValue::Place(source))
| Some(AbstractValue::Ref(source))
| Some(AbstractValue::RawPtr(source)) => Some(source.clone()),
Some(AbstractValue::Cast(inner, _)) => match inner.as_ref() {
AbstractValue::Place(source)
| AbstractValue::Ref(source)
| AbstractValue::RawPtr(source) => Some(source.clone()),
_ => None,
},
_ => None,
};
let source = source.or_else(|| {
forward.facts.iter().find_map(|fact| match fact {
StateFact::PointsTo { pointer, source }
if pointer.base == place.base && pointer.fields.is_empty() =>
{
Some(source.clone())
}
_ => None,
})
})?;
let mut fields = source.fields.clone();
fields.extend_from_slice(&place.fields);
Some(PlaceKey {
base: source.base,
fields,
})
}
fn field_invariant_matches<'tcx>(
tcx: TyCtxt<'tcx>,
caller: DefId,
body: &rustc_middle::mir::Body<'tcx>,
place: &PlaceKey,
kind: PropertyKind,
required_ty: Option<Ty<'tcx>>,
required_elements: Option<u64>,
) -> Option<String> {
if place.fields.is_empty() {
return None;
}
let local = place.local()?;
if local.as_usize() >= body.local_decls.len() {
return None;
}
let base_ty = body.local_decls[local].ty;
let (adt_def, substs) = match base_ty.kind() {
TyKind::Ref(_, pointee, _) | TyKind::RawPtr(pointee, _) => match pointee.kind() {
TyKind::Adt(adt, subs) => (*adt, *subs),
_ => return None,
},
TyKind::Adt(adt, subs) => (*adt, *subs),
_ => return None,
};
if !adt_def.is_struct() {
return None;
}
let struct_def_id = adt_def.did();
for invariant in get_struct_invariants_for_adt(tcx, struct_def_id) {
if !invariant_kind_implies(tcx, caller, &invariant.kind, &kind, required_ty) {
continue;
}
let Some(PropertyArg::Place(contract_place)) = invariant.args.first() else {
continue;
};
let invariant_key = PlaceKey::from_contract_place(contract_place);
if invariant_key.fields != place.fields {
continue;
}
if !invariant_args_cover(&invariant, required_ty, required_elements) {
continue;
}
let struct_name = tcx.def_path_str(struct_def_id);
return Some(format!(
"{kind:?} assumed from struct invariant on `{struct_name}` for pointee field path {:?}",
place.fields
));
}
if kind == PropertyKind::Alive
&& matches!(base_ty.kind(), TyKind::Ref(..))
&& place.fields.len() == 1
{
let field_idx = FieldIdx::from_usize(place.fields[0]);
let variant = adt_def.non_enum_variant();
if field_idx.as_usize() < variant.fields.len() {
#[cfg(not(rapx_rustc_ge_198))]
let field_ty = variant.fields[field_idx].ty(tcx, substs);
#[cfg(rapx_rustc_ge_198)]
let field_ty = variant.fields[field_idx].ty(tcx, substs).skip_norm_wip();
if matches!(field_ty.kind(), TyKind::Ref(..)) {
let struct_name = tcx.def_path_str(struct_def_id);
return Some(format!(
"Alive inferred from reference-typed field in `{struct_name}`"
));
}
}
}
None
}
fn invariant_kind_implies<'tcx>(
tcx: TyCtxt<'tcx>,
caller: DefId,
declared: &PropertyKind,
required: &PropertyKind,
required_ty: Option<Ty<'tcx>>,
) -> bool {
if crate::verify::contract::decomp::kind_implies(declared, required) {
if matches!(declared, PropertyKind::ValidPtr)
&& matches!(required, PropertyKind::Allocated | PropertyKind::InBound)
&& !required_ty.is_some_and(|ty| {
super::common::safe_type_layout(tcx, caller, ty).is_some_and(|(_, size)| size > 0)
})
{
return false;
}
return true;
}
false
}
fn invariant_args_cover<'tcx>(
invariant: &Property<'tcx>,
required_ty: Option<Ty<'tcx>>,
required_elements: Option<u64>,
) -> bool {
let declared_ty = invariant.args.iter().find_map(|arg| match arg {
PropertyArg::Ty(ty) => Some(*ty),
_ => None,
});
let declared_elements = invariant.args.iter().find_map(|arg| match arg {
PropertyArg::Expr(ContractExpr::Const(value)) => u64::try_from(*value).ok(),
_ => None,
});
let ty_ok = match (required_ty, declared_ty) {
(Some(required), Some(declared)) => {
required == declared || format!("{required:?}") == format!("{declared:?}")
}
(None, _) => true,
(Some(_), None) => false,
};
let elements_ok = match (required_elements, declared_elements) {
(Some(required), Some(declared)) => declared >= required,
(None, _) => true,
(Some(_), None) => false,
};
ty_ok && elements_ok
}
pub(super) fn discharge_from_contract_fact<'tcx>(
property: &Property<'tcx>,
forward: &ForwardVisitResult<'tcx>,
) -> Option<String> {
let target_key = contract_property_key(property)?;
for fact in &forward.facts {
let StateFact::Contract(contract) = fact else {
continue;
};
if !crate::verify::contract::decomp::kind_implies(&contract.kind, &property.kind) {
continue;
}
let Some(contract_key) = contract_property_key(contract) else {
continue;
};
if contract_key != target_key {
continue;
}
if !contract_args_cover(contract, property) {
continue;
}
return Some(format!(
"{:?} assumed from an entry contract covering the same place",
property.kind
));
}
None
}
fn contract_property_key<'tcx>(property: &Property<'tcx>) -> Option<PlaceKey> {
let arg = property.args.first()?;
let place = match arg {
PropertyArg::Place(place) => place,
PropertyArg::Expr(ContractExpr::Place(place)) => place,
_ => return None,
};
let mut key = PlaceKey::from_contract_place(place);
if let PlaceBaseKey::Arg(index) = key.base {
key.base = PlaceBaseKey::Local(index + 1);
}
Some(key)
}
fn contract_args_cover<'tcx>(contract: &Property<'tcx>, property: &Property<'tcx>) -> bool {
let ty_of = |candidate: &Property<'tcx>| {
candidate.args.iter().find_map(|arg| match arg {
PropertyArg::Ty(ty) => Some(*ty),
_ => None,
})
};
let elements_of = |candidate: &Property<'tcx>| {
candidate.args.iter().find_map(|arg| match arg {
PropertyArg::Expr(ContractExpr::Const(value)) => u64::try_from(*value).ok(),
_ => None,
})
};
let ty_ok = match (ty_of(property), ty_of(contract)) {
(Some(required), Some(declared)) => {
required == declared || format!("{required:?}") == format!("{declared:?}")
}
(None, _) => true,
(Some(_), None) => false,
};
let elements_ok = match (elements_of(property), elements_of(contract)) {
(Some(required), Some(declared)) => declared >= required,
(None, _) => true,
(Some(_), None) => false,
};
ty_ok && elements_ok
}
pub(super) fn discharge_from_contract_fact_with_checkpoint<'tcx>(
property: &Property<'tcx>,
forward: &ForwardVisitResult<'tcx>,
checkpoint: &Checkpoint<'tcx>,
) -> Option<String> {
let target_key = checkpoint_target_key(checkpoint, property)
.or_else(|| contract_property_key(property))?;
for fact in &forward.facts {
let StateFact::Contract(contract) = fact else {
continue;
};
if !crate::verify::contract::decomp::kind_implies(&contract.kind, &property.kind) {
continue;
}
let Some(contract_key) = contract_property_key(contract) else {
continue;
};
if contract_key != target_key
&& !provenance_chain_reaches(&contract_key, &target_key, forward)
{
continue;
}
if !contract_args_cover(contract, property) {
continue;
}
return Some(format!(
"{:?} assumed from an entry contract covering the same place",
property.kind
));
}
None
}
fn checkpoint_target_key<'tcx>(
checkpoint: &Checkpoint<'tcx>,
property: &Property<'tcx>,
) -> Option<PlaceKey> {
let arg = property.args.first()?;
let place = match arg {
PropertyArg::Place(place) => place,
PropertyArg::Expr(ContractExpr::Place(place)) => place,
_ => return None,
};
if let ContractPlace {
base: crate::verify::contract::PlaceBase::Arg(index),
..
} = place
{
let operand = checkpoint.args.get(*index)?;
match operand {
rustc_middle::mir::Operand::Copy(mir_place)
| rustc_middle::mir::Operand::Move(mir_place) => {
Some(PlaceKey::from_mir_place(mir_place))
}
_ => None,
}
} else {
None
}
}
fn provenance_chain_reaches<'tcx>(
contract: &PlaceKey,
target: &PlaceKey,
forward: &ForwardVisitResult<'tcx>,
) -> bool {
let mut seen: std::collections::HashSet<PlaceKey> = std::collections::HashSet::new();
let mut queue: Vec<PlaceKey> = vec![target.clone()];
while let Some(cur) = queue.pop() {
if &cur == contract {
return true;
}
if !seen.insert(cur.clone()) {
continue;
}
if cur.fields.is_empty() {
if let Some(local) = cur.local()
&& let Some(def) = forward
.latest_value_definition_before(local, forward.value_definitions.len())
{
match &def.value {
AbstractValue::Place(p)
| AbstractValue::Ref(p)
| AbstractValue::RawPtr(p) => queue.push(p.clone()),
_ => {}
}
}
}
for fact in &forward.facts {
let StateFact::Cast { target, source, .. } = fact else { continue; };
if target == &cur {
match source {
AbstractValue::Place(p)
| AbstractValue::Ref(p)
| AbstractValue::RawPtr(p) => queue.push(p.clone()),
_ => {}
}
}
}
}
false
}