diff --git a/SIFExtendedTransformer.scala b/SIFExtendedTransformer.scala new file mode 100644 index 00000000000..b880f07fab6 --- /dev/null +++ b/SIFExtendedTransformer.scala @@ -0,0 +1,1727 @@ +// This Source Code Form is subject to the terms of the Mozilla Public +// License, v. 2.0. If a copy of the MPL was not distributed with this +// file, You can obtain one at http://mozilla.org/MPL/2.0/. +// +// Copyright (c) 2011-2020 ETH Zurich. + +package viper.silver.sif + +import viper.silver.ast._ +import viper.silver.ast.utility.Simplifier +import viper.silver.verifier.errors +import viper.silver.verifier.errors.{AssertFailed, ErrorNode} + +import scala.collection.immutable.HashSet +import scala.collection.mutable +import scala.collection.mutable.ListBuffer + +trait SIFExtendedTransformer { + object Config { + /** If true, don't generate all the control flow variables, just the ones needed for each method. */ + var optimizeControlFlow: Boolean = true + /** If true, try to bunch as many statements into a single if statement which was introduced for checking active + * executions, instead of having one if stmt per original statement. */ + var optimizeSequential: Boolean = true + /** If true, add an 'assume p1' at the beginning of each method, to cut down on redundant paths the + * verification could consider */ + var optimizeRestrictActVars: Boolean = true + /** Applications of the functions which have an entry here, will be replaced by the expression + * determined by the entry in the second execution. */ + var primedFuncAppReplacements: mutable.HashMap[String, (FuncApp, Exp, Exp) => Exp] = new mutable.HashMap + /** Set this to only transform methods that contain relational assertions somewhere in their spec or body. + * May lead to invalid programs when such a method calls another methods that does contain such specs. + */ + var onlyTransformMethodsWithRelationalSpecs: Boolean = false + var generateAllLowFuncs: Boolean = true + } + def optimizeControlFlow(v: Boolean): Unit = { + Config.optimizeControlFlow = v + } + def optimizeSequential(v: Boolean): Unit = { + Config.optimizeSequential = v + } + def optimizeRestrictActVars(v: Boolean): Unit = { + Config.optimizeRestrictActVars = v + } + def generateAllLowFuncs(v: Boolean): Unit = { + Config.generateAllLowFuncs = v + } + def onlyTransformMethodsWithRelationalSpecs(v: Boolean): Unit = { + Config.onlyTransformMethodsWithRelationalSpecs = v + } + def addPrimedFuncAppReplacement(name: String, strategy: String): Unit = { + strategy match { + case "first_arg" => Config.primedFuncAppReplacements.put(name, + (func, p1, p2) => translatePrime(func.args.head, p1, p2)) + case "true" => Config.primedFuncAppReplacements.put(name, (func, _, _) => + TrueLit()(func.pos, func.info, func.errT)) + case _ => new IllegalArgumentException( + s"""Unknown strategy "$strategy" for primed function application replacement.""") + } + } + def clearPrimedFuncAppReplacement(): Unit = { + Config.primedFuncAppReplacements.clear() + } + + val primedNamesPerMethod = new mutable.HashMap[String, Map[String, String]] + val primedNames = new mutable.HashMap[String, String] + val relationalPredicates = new mutable.HashSet[Predicate] + val predLowFuncs = new mutable.HashMap[String, Option[Function]] + val predLowFuncInfo = new mutable.HashMap[String, Option[(String, Seq[LocalVarDecl], Seq[LocalVarDecl])]] + val predAllLowFuncs = new mutable.HashMap[String, Option[Function]]() + val predAllLowFuncInfo = new mutable.HashMap[String, Option[(String, Seq[LocalVarDecl], Seq[LocalVarDecl])]] + val usedNames = new mutable.HashSet[String] + var newFields : List[Field] = Nil + var newPredicates: Seq[Predicate] = Nil + var _program : Program = null + var getArgFunc : DomainFunc = null + var getOldFunc : DomainFunc = null + var getArgPFunc : DomainFunc = null + var getOldPFunc : DomainFunc = null + + private var _allLowMethods: Set[String] = HashSet[String]() + def allLowMethods: Set[String] = _allLowMethods + def setAllLowMethods(value: Set[String]): Unit = _allLowMethods = value + + private var _preservesLowMethods: Set[String] = HashSet[String]() + def preservesLowMethods: Set[String] = _preservesLowMethods + def setPreservesLowMethods(value: Set[String]): Unit =_preservesLowMethods = value + + private val _domainFuncsToDuplicate = new mutable.HashSet[DomainFunc]() + def domainFuncsToDuplicate: mutable.Set[DomainFunc] = _domainFuncsToDuplicate + def addDomainFuncToDuplicate(funcs: DomainFunc*): Unit = _domainFuncsToDuplicate ++= funcs + + var timing = false + var time : Option[LocalVar] = None + + val skip: Seqn = Seqn(Seq(), Seq())() + + def transform(p: Program, enableTiming: Boolean) : Program = { + primedNames.clear() + predLowFuncs.clear() + predLowFuncInfo.clear() + predAllLowFuncInfo.clear() + usedNames.clear() + newPredicates = Nil + timing = enableTiming + _program = p + + val allNames = p.collect({ + case d: Declaration => d.name + }) + usedNames ++= allNames + + collectRelationalPredicates(_program, relationalPredicates) + createNewNames(p) + var newFunctions: Seq[Function] = p.functions //.flatMap(f => translateFunction(f)) + newPredicates = Seq() + for (pred <- p.predicates) { + val (newPs, predFuncs) = translatePredicate(pred) + newPredicates ++= newPs + newFunctions ++= predFuncs.collect{case f if f.isDefined => f.get} + } + + val newMethods: Seq[Method] = if (Config.onlyTransformMethodsWithRelationalSpecs) { + val relationalMethods = p.methods.filter(m => m.existsDefined({ + case e: Exp if !isUnary(e) => true + })) + val unaryMethods = p.methods.filter(m => !relationalMethods.contains(m)) + relationalMethods.map(m => translateMethod(m)) ++ unaryMethods + }else{ + p.methods.map(m => translateMethod(m)) + } + + p.copy(functions = newFunctions, predicates = newPredicates, + methods = newMethods)(p.pos, p.info, p.errT) + } + + def getName(orig: String) : String = { + if (usedNames.contains(orig)){ + var index = 0 + while (usedNames.contains(orig + "_" + index)){ + index += 1 + } + val result = orig + "_" + index + usedNames.add(result) + return result + }else{ + usedNames.add(orig) + return orig + } + } + + /** Create the new names for the programs variables, functions (those which depend on the heap), + * and predicates, which are used for the second execution. Updates [[primedNames]]. + * @param p The program being encoded. + */ + def createNewNames(p: Program): Unit = { + // duplicate names for predicates + for (pred <- p.predicates) { + val duplicatedArgs = pred.formalArgs.map{a => + val newName = getName(a.name) + a.copy(name = newName)(a.pos, a.info, a.errT) + } + predLowFuncInfo.update(pred.name, if (relationalPredicates.contains(pred) && pred.body.isDefined) + Some(getName(pred.name + "_low"), pred.formalArgs, duplicatedArgs) + else None + ) + predAllLowFuncInfo.update(pred.name, pred.body match { + case Some(_) => + Some(getName(pred.name + "_all_low"), pred.formalArgs, duplicatedArgs) + case None => None + }) + } + } + + def collectRelationalPredicates(p: Program, relPreds: mutable.HashSet[Predicate]): Unit = { + def directlyRelational(pred: Predicate): Boolean = { + pred.body.isDefined && !isUnary(pred.body.get) + } + + relPreds.clear() + relPreds ++= p.predicates.filter(pred => directlyRelational(pred)) + val dependencies = mutable.HashMap[String, Seq[String]]() + + relPreds.foreach(pred => + // put the name of this as depending on all referenced predicates + pred.body.collect{ + case pap: PredicateAccessPredicate => pap + }.foreach(pap => + dependencies.update(pap.loc.predicateName, dependencies.getOrElse(pap.loc.predicateName, Seq()) :+ pred.name) + ) + ) + // go through dependent predicates to add them to relationals + val queue = mutable.Queue[String](relPreds.toSeq.map(rp => rp.name): _*) + while (queue.nonEmpty) { + val head: String = queue.dequeue() + for (dep <- dependencies.getOrElse(head, Seq())) { + if (relPreds.add(p.findPredicate(dep))) queue.enqueue(dep) + } + } + } + + def translateMethod(m: Method) : Method = { + val primedBefore = primedNames.clone() + val (p1d, p1r) = getNewBool("p1") + val (p2d, p2r) = getNewBool("p2") + var toAdd = Seq(p1d, p2d) + val (t1d, t1r) = getNewVar("t1", Int) + val (t2d, t2r) = getNewVar("t2", Int) + if (timing){ + toAdd = Seq(p1d, p2d, t1d, t2d) + primedNames.update(t1d.name, t2d.name) + time = Some(t1r) + } + + val newArgs = toAdd ++ m.formalArgs.flatMap{a => + val newName = getName(a.name) + primedNames.update(a.name, newName) + val primedArg = a.copy(name = newName)(a.pos, a.info, a.errT) + Seq(a, primedArg) + } + var toAddRet : Seq[LocalVarDecl] = Seq() + val (t1dr, t1rr) = getNewVar("t1", Int) + val (t2dr, t2rr) = getNewVar("t2", Int) + if (timing){ + toAddRet = Seq(t1dr, t2dr) + primedNames.update(t1dr.name, t2dr.name) + time = Some(t1rr) + } + val newReturns = toAddRet ++ m.formalReturns.flatMap{r => + val newName = getName(r.name) + primedNames.update(r.name, newName) + val primedRet = r.copy(name = newName)(r.pos, r.info, r.errT) + Seq(r, primedRet) + } + + var (newPres, newPosts) = addLownessConditions(m, m.pres, m.posts) + val condCtx = TranslationContext(p1r, p2r, EmptyControlFlowVars(), m) + newPres = newPres.map{e => translateSIFAss(e, condCtx.copy(translatingPrecond = true))} + newPosts = newPosts.map{e => translateSIFAss(e, condCtx)} + // Termination channels + var terminates: Option[SIFTerminatesExp] = None + m.pres.foreach(pre => pre.visit{ + case t: SIFTerminatesExp => terminates = Some(t) + }) + if (terminates.isDefined) { + newPres = simplifyConditions(newPres ++ + terminationChannelsLowChecks(terminates.get, condCtx)) + newPosts :+= translateSIFAss(Old(terminates.get.cond)( + terminates.get.cond.pos, terminates.get.cond.info, + ErrTrafo({case _ => SIFTerminationChannelCheckFailed(terminates.get.cond, SIFTermCondNotTight(terminates.get))}) + ), condCtx) + } + + if (timing){ + val timeUnchanged1 = Implies(Not(p1r)(), EqCmp(t1r, t1rr)())() + val timeUnchanged2 = Implies(Not(p2r)(), EqCmp(t2r, t2rr)())() + newPosts ++= Seq(timeUnchanged1, timeUnchanged2) + } + + var firstStatements : Seq[Stmt] = Seq() + if (Config.optimizeRestrictActVars) firstStatements :+= Inhale(p1r)() + if (timing){ + val assignTime1 = LocalVarAssign(t1rr, t1r)() + val assignTime2 = LocalVarAssign(t2rr, t2r)() + firstStatements ++= Seq(assignTime1, assignTime2) + time = Some(t1rr) + } + + var newBody: Option[Seqn] = None + if (m.body.isDefined) { + var body: Seqn = m.body.get + if (allLowMethods.contains(m.name) || preservesLowMethods.contains(m.name)) { + body = inferLowLoopInvariants(m, preservesOnly = preservesLowMethods.contains(m.name)) + } + + // find which control variables the method requires + val ctrlVars: MethodControlFlowVars = createControlFlowVars(body) + val ctx: TranslationContext = TranslationContext(p1r, p2r, ctrlVars, m) + + newBody = Some(Seqn(firstStatements ++ ctrlVars.initAssigns() ++ + Seq(translateStatement(body, ctx).asInstanceOf[Seqn]), ctrlVars.declarations())()) + } + + time = None + primedNamesPerMethod.update(m.name, primedNames.toMap) + primedNames.clear() + primedNames ++= primedBefore + Method(m.name, newArgs, newReturns, newPres, newPosts, newBody)(m.pos) + } + + def _conjoinOptions(in: Seq[Option[Exp]]): Exp = { + val defined = in.filter(x => x.isDefined).map(x => x.get) + defined.size match { + case 0 => TrueLit()() + case _ => defined.reduceRight((a, b) => And(a, b)(a.pos)) + } + } + + /** Takes a sequence of LocalVarDecls and produces an expression saying each pair of adjacent variables are equal. + */ + def _varDeclsToAllLow(in: Seq[LocalVarDecl]): Exp = { + in.map(decl => decl.localVar) + .map(v => SIFLowExp(v, None)()) + .reduceRight[Exp]((v, e) => And(v, e)()) + } + + def allReachableStateLow(m: Method, old: Boolean, predicateSrc: Seq[Exp]): Option[Exp] = { + var lowExpressions: Seq[Exp] = Seq() + + // all accs to fields in preconditions + val allFieldAccesses: Seq[FieldAccess] = m.pres.flatMap(e => e.deepCollect({ + case FieldAccessPredicate(loc, _) => Seq(loc) + })).flatten.distinct + for (fieldAcc <- allFieldAccesses) { + val eq = SIFLowExp(fieldAcc, None)(m.pos) + if (old) + lowExpressions :+= Old(eq)(m.pos) + else + lowExpressions :+= eq + } + + // all accs we get via predicates + lowExpressions ++= predicateSrc.flatMap(e => e.deepCollect{ + case PredicateAccessPredicate(loc, _) => + val funcApp = FuncApp(predAllLowFuncs(loc.predicateName).get, + loc.args ++ loc.args.map(a => translatePrime(a, null, null)))() + if (old) + Old(funcApp)(m.pos) + else + funcApp + }) + + if (lowExpressions.nonEmpty) + Some(lowExpressions.reduceRight[Exp]((a, b) => And(a, b)(m.pos))) + else + None + } + + def allVarsAndStateLow(m: Method, vars: Seq[LocalVarDecl], old: Boolean, predicateSrc: Seq[Exp]): Exp = { + val nonObligationVars = vars.filterNot(isObligationVar) + val allArgsLow: Option[Exp] = if (nonObligationVars.isEmpty) None else Some(_varDeclsToAllLow(nonObligationVars)) + val allStateLow: Option[Exp] = allReachableStateLow(m, old, predicateSrc) + _conjoinOptions(Seq(allArgsLow, allStateLow)) + } + + def addLownessConditions(m: Method, pres: Seq[Exp], posts: Seq[Exp]): (Seq[Exp], Seq[Exp]) = { + var newPres = pres + if (allLowMethods.contains(m.name)) newPres :+= allVarsAndStateLow(m, m.formalArgs, old = false, predicateSrc = m.pres) + var newPosts = posts + if (allLowMethods.contains(m.name)) newPosts :+= allVarsAndStateLow(m, m.formalReturns, old = false, predicateSrc = m.posts) + if (preservesLowMethods.contains(m.name)) newPosts :+= Implies( + allVarsAndStateLow(m, m.formalArgs, old = true, predicateSrc = m.pres), + allVarsAndStateLow(m, m.formalReturns, old = false, predicateSrc = m.posts))(m.pos) + (newPres, newPosts) + } + + def inferLowLoopInvariants(m: Method, preservesOnly: Boolean): Seqn = { + val _obligationVars: Seq[String] = obligationVars(m) + m.body.get.transform{ + case w@While(cond, invs, body) => + val targets: Seq[LocalVar] = w.deepCollect({ + case LocalVarAssign(lhs, _) => Seq(lhs) + case MethodCall(_, _, ts) => ts + case NewStmt(t, _) => Seq(t) + }).flatten + .distinct + .filterNot(v => _obligationVars.contains(v.name)) + var additionalInvs: Seq[Exp] = targets.map(lv => SIFLowExp(lv, None)()) ++ + allReachableStateLow(m, old = false, predicateSrc = m.pres) + if (preservesOnly) { + additionalInvs = additionalInvs.map( + i => Implies(allVarsAndStateLow(m, m.formalArgs, old=true, predicateSrc = m.pres), i)()) + } + val newInvs: Seq[Exp] = invs ++ additionalInvs + While(cond, newInvs, body)(w.pos, w.info, w.errT) + } + } + + def isObligationVar(v: LocalVarDecl): Boolean = { + val vi = v.info.getUniqueInfo[SIFInfo] + vi match { + case Some(info) => info.obligationVar + case _ => false + } + } + + def obligationVars(m: Method): Seq[String] = { + m.deepCollect[LocalVarDecl]({ + case d: LocalVarDecl if isObligationVar(d) => d + }).map(v => v.name) + } + + def createControlFlowVars(methodBody: Seqn): MethodControlFlowVars = { + val gotos = methodBody.collect({case Goto(l) => l}).toSet + val labels = methodBody.collect({case l@Label(n, _) if gotos.contains(n) => l}).toSet + if (!Config.optimizeControlFlow) return new MethodControlFlowVars(true, true, true, true, labels) + var hasRet, hasBreak, hasCont, hasExcept: Boolean = false + methodBody.visit({ + case _: SIFReturnStmt => hasRet = true + case _: SIFBreakStmt => hasBreak = true + case _: SIFContinueStmt => hasCont = true + case _: SIFRaiseStmt => hasExcept = true + case _: SIFTryCatchStmt => hasExcept = true + }) + if (labels.nonEmpty && (hasRet || hasBreak || hasCont || hasExcept)) + throw new IllegalArgumentException + new MethodControlFlowVars(hasRet, hasBreak, hasCont, hasExcept, labels) + } + + def incrementTime(p1: Exp, p2: Exp) : Seqn = { + if (timing){ + val timeInc1 = If(p1, Seqn(Seq(LocalVarAssign(time.get, Add(time.get, IntLit(1)())())()), Seq())(), skip)() + val timeInc2 = If(p2, Seqn(Seq(LocalVarAssign(translatePrime(time.get, p1, p2), translatePrime(Add(time.get, IntLit(1)())(), p1, p2))()), Seq())(), skip)() + Seqn(Seq(timeInc1, timeInc2), Seq())() + }else{ + skip + } + + } + + private def terminationChannelsLowChecks(terminates: SIFTerminatesExp, + ctx: TranslationContext): Seq[Exp] = { + val condLowEventReasonTrafo = ErrTrafo({case _ => + SIFTerminationChannelCheckFailed(terminates, SIFTermCondLowEvent(terminates))}) + + ReTrafo({case _ => SIFTermCondLowEvent(terminates)}) + val condLowReasonTrafo = ErrTrafo({case _ => + SIFTerminationChannelCheckFailed(terminates, SIFTermCondNotLow(terminates))}) + + ReTrafo({case _ => SIFTermCondNotLow(terminates)}) + val pos = terminates.pos + val info = terminates.info + Seq( + // Note: cast here is not very elegant, maybe there's a better way to attach the error reason + // to result of translateSIFAss + translateSIFAss(Implies(Not(terminates.cond)(pos, info, condLowEventReasonTrafo), + SIFLowEventExp()(pos, info, condLowEventReasonTrafo))(pos, info, condLowEventReasonTrafo), ctx) + .asInstanceOf[Implies].copy()(pos, info, condLowEventReasonTrafo), + translateSIFAss(SIFLowExp(terminates.cond)(pos, info, condLowReasonTrafo), ctx) + .asInstanceOf[Implies].copy()(pos, info, condLowReasonTrafo) + ) + } + + /** Translate a predicate into MPP. Creates two copies of the predicate: one version for first execution, + * one for the second. The predicates only contain the unary expressions of the original predicate. + * Generates a function containing all the relational expressions in the predicates body. + * @param pred The predicate to translate. + * @return List of the new predicates, plus the low-function. + */ + def translatePredicate(pred: Predicate): (Seq[Predicate], Seq[Option[Function]]) = { + val unaryBody: Option[Exp] = translateToUnary(pred.body) + + val newPred = pred.copy(body = unaryBody)(pred.pos, pred.info, pred.errT) + + var lowF: Option[Function] = None + var allLowF: Option[Function] = None + if (pred.body.isDefined) { + val (allLowFName, formalArgs, duplicatedFormalArgs) = predAllLowFuncInfo(pred.name).get + + val access1 = PredicateAccess(formalArgs.map{a => a.localVar}, newPred.name)(pred.pos) + val access2 = PredicateAccess(duplicatedFormalArgs.map{a => a.localVar}, newPred.name)(pred.pos) + val fPres: Seq[Exp] = Seq(And(PredicateAccessPredicate(access1, WildcardPerm()())(), + PredicateAccessPredicate(access2, WildcardPerm()())())()) + val lowFFormalArgs = pred.formalArgs ++ duplicatedFormalArgs + val primedBefore = primedNames.clone() + formalArgs.zip(duplicatedFormalArgs).foreach(t => primedNames.update(t._1.name, t._2.name)) + + def unfoldingPredicates(body: Exp): Exp = { + Unfolding(PredicateAccessPredicate(access1, WildcardPerm()())(), + Unfolding(PredicateAccessPredicate(access2, WildcardPerm()())(), + body)())() + } + if (relationalPredicates.contains(pred)) { + val (lowFName, _, _) = predLowFuncInfo(pred.name).get + val fBody: Exp = unfoldingPredicates(translatePredLowFuncBody(pred.body.get)) + lowF = Some(Function(lowFName, lowFFormalArgs, Bool, fPres, Seq(), Some(fBody)) + (pred.pos, pred.info, pred.errT)) + } + + val allLowBody: Exp = unfoldingPredicates(translatePredAllLowFuncBody(pred.body.get)) + allLowF = Some(Function(allLowFName, lowFFormalArgs, Bool, fPres, Seq(), Some(allLowBody)) + (pred.pos, pred.info, pred.errT)) + + primedNames.clear() + primedNames ++= primedBefore + } + + predLowFuncs.update(pred.name, lowF) + predAllLowFuncs.update(pred.name, allLowF) + if (Config.generateAllLowFuncs) + (Seq(newPred), Seq(lowF, allLowF)) + else + (Seq(newPred), Seq(lowF, None)) + } + + private def bypassPreamble(p1: Exp, p2: Exp, ctrlVars: MethodControlFlowVars, + bypass1r: LocalVar, bypass2r: LocalVar): Seq[Stmt] = { + Seq(LocalVarAssign(bypass1r, + Not(ctrlVars.activeExecNormal(Some(p1)))())(), + LocalVarAssign(bypass2r, Not(ctrlVars.activeExecPrime(Some(p2)))())()) + } + + /** Translate a while statement into MPP. + */ + private def translateWhileStmt(w: While, ctx: TranslationContext): Seqn = { + val p1 = ctx.p1; val p2 = ctx.p2; val ctrlVars = ctx.ctrlVars + + // check if we need to do a reconstruction of this loop (iff it has ret/break/except stmt or we don't optimize) + val recNeeded: Boolean = { + var rn = !Config.optimizeControlFlow + w.body.visit({ + case _: SIFReturnStmt => rn = true + case _: SIFBreakStmt => rn = true + case _: SIFRaiseStmt => rn = true + case _: SIFTryCatchStmt => rn = true + }) + rn + } + + var newVarDecls = Seq[LocalVarDecl]() + //var targetValRefs = Seq[LocalVar]() + var targetValAssigns = Seq[Stmt]() + var targetValEqualities1 = Set[Exp]() + var targetValEqualities2 = Set[Exp]() + + val (bypass1d, bypass1r) = getNewBool("bypass1") + val (bypass2d, bypass2r) = getNewBool("bypass2") + newVarDecls ++= Seq(bypass1d, bypass2d) + var stmts: Seq[Stmt] = bypassPreamble(p1, p2, ctrlVars, bypass1r, bypass2r) + + def targetCollectF: PartialFunction[Node, Seq[LocalVar]] = { + case LocalVarAssign(lhs, _) => Seq(lhs) + case MethodCall(_, _, ts) => ts + case NewStmt(t, _) => Seq(t) + case SIFReturnStmt(_, _) => Seq(ctrlVars.ret1r.get) + case SIFRaiseStmt(_) => Seq(ctrlVars.except1r.get) + case SIFTryCatchStmt(body, handlers, elseBlock, finallyBlock) => + (body.deepCollect(targetCollectF).distinct.flatten ++ + handlers.map(h => h.body.deepCollect(targetCollectF).flatten).distinct.flatten ++ (elseBlock match { + case Some(eb) => eb.deepCollect(targetCollectF).distinct.flatten + case None => Seq() + }) ++ (finallyBlock match { + case Some(fb) => fb.deepCollect(targetCollectF).distinct.flatten + case None => Seq() + })).distinct + } + + // continue and break control variables will be assigned even if there is no continue in this loop -> add to targets + val targets = w.deepCollect(targetCollectF).flatten.distinct ++ (if (ctrlVars.cont1r.isDefined) + Seq(ctrlVars.cont1r.get, ctrlVars.cont2r.get) else Seq()) ++ (if (ctrlVars.break1r.isDefined) + Seq(ctrlVars.break1r.get, ctrlVars.break2r.get) else Seq()) + + var tmpAssigns1 = Seq[LocalVarAssign]() + var tmpAssigns2 = Seq[LocalVarAssign]() + for (t <- targets) { + // make sure the variable is defined outside the loop + if (primedNames.contains(t.name)){ + val (tmp1d, tmp1r) = getNewVar("tmp1", t.typ) + val (tmp2d, tmp2r) = getNewVar("tmp2", t.typ) + newVarDecls ++= Seq(tmp1d, tmp2d) + //targetValRefs ++= Seq(tmp1r, tmp2r) + tmpAssigns1 :+= LocalVarAssign(tmp1r, t)() + tmpAssigns2 :+= LocalVarAssign(tmp2r, translatePrime(t, p1, p2))() + val eq1 = EqCmp(tmp1r, t)() + val eq2 = EqCmp(tmp2r, translatePrime(t, p1, p2))() + targetValEqualities1 ++= Seq(eq1) + targetValEqualities2 ++= Seq(eq2) + } + } + if (tmpAssigns1.nonEmpty) targetValAssigns :+= If(bypass1r, Seqn(tmpAssigns1, Seq())(), skip)() + if (tmpAssigns2.nonEmpty) targetValAssigns :+= If(bypass2r, Seqn(tmpAssigns2, Seq())(), skip)() + + /*if (timing){ + val (tmp1d, tmp1r) = getNewVar("tmp1", Int) + val (tmp2d, tmp2r) = getNewVar("tmp2", Int) + newVarDecls ++= Seq(tmp1d, tmp2d) + val assign1 = LocalVarAssign(tmp1r, time.get)() + val assign2 = LocalVarAssign(tmp2r, translatePrime(time.get, p1, p2))() + targetValAssigns ++= Seq(assign1, assign2) + val eq1 = EqCmp(tmp1r, time.get)() + val eq2 = EqCmp(tmp2r, translatePrime(time.get, p1, p2))() + targetValEqualities1 ++= Seq(eq1) + targetValEqualities2 ++= Seq(eq2) + }*/ + stmts ++= targetValAssigns + + // tmp assigns for all ctrlVars + var ctrlVarToOldMap: Map[LocalVar, LocalVar] = Map() + if (recNeeded) { + var ctrlFlowTmpAssigns = Seq[Stmt]() + for (v <- ctrlVars.declarations().map(d => d.localVar)) { + val (tmp1d, tmp1r) = getNewBool("old" + v.name) + ctrlVarToOldMap += (v -> tmp1r) + newVarDecls :+= tmp1d + stmts :+= LocalVarAssign(tmp1r, v)() + ctrlFlowTmpAssigns :+= LocalVarAssign(v, tmp1r)() + } + } + + val newCond = Or(And(And(ctrlVars.activeExecNoContNormal(Some(p1)), Not(bypass1r)())(), w.cond)(), + And(And(ctrlVars.activeExecNoContPrime(Some(p2)), Not(bypass2r)())(), translatePrime(w.cond, p1, p2))())() + var bodyPreamble: Seq[Stmt] = Seq() + if (ctrlVars.cont1r.isDefined) bodyPreamble = bodyPreamble :+ + LocalVarAssign(ctrlVars.cont1r.get, FalseLit()())() :+ + LocalVarAssign(ctrlVars.cont2r.get, FalseLit()())() + val (p1d, p1r) = getNewBool("p1") + val (p2d, p2r) = getNewBool("p2") + newVarDecls ++= Seq(p1d, p2d) + val p1Assign = LocalVarAssign(p1r, And(ctrlVars.activeExecNormal(Some(p1)), w.cond)())() + val p2Assign = LocalVarAssign(p2r, And(ctrlVars.activeExecPrime(Some(p2)), translatePrime(w.cond, p1, p2))())() + bodyPreamble ++= Seq(p1Assign, p2Assign) + + var newStdInvs = w.invs + // check if there is an InhaleExhaleExp in invariants + if (w.invs.exists(inv => inv.contains[InhaleExhaleExp])) { + val (idle1d, idle1r) = getNewBool("idle1") + val (idle2d, idle2r) = getNewBool("idle2") + primedNames.update(idle1r.name, idle2r.name) + newVarDecls ++= Seq(idle1d, idle2d) + // assign false before the loop + stmts ++= Seq(LocalVarAssign(idle1r, FalseLit()())(), + LocalVarAssign(idle2r, FalseLit()())()) + // assign inside loop body that the execution is idling + val idle1Assign = LocalVarAssign(idle1r, + And(ctrlVars.activeExecNormal(Some(p1)), Not(w.cond)())())() + val idle2Assign = LocalVarAssign(idle2r, + And(ctrlVars.activeExecPrime(Some(p2)), Not(translatePrime(w.cond, p1, p2))())())() + bodyPreamble ++= Seq(idle1Assign, idle2Assign) + // make all exhales of InhaleExhales dependent on not idling + newStdInvs = newStdInvs.map(inv => inv.transform{ + case ie@InhaleExhaleExp(in, ex) => InhaleExhaleExp(in, + Implies(Not(idle1r)(), ex)())(ie.pos, ie.info, ie.errT) + }) + } + + // --- Terminates --- + var terminates: Option[SIFTerminatesExp] = None + w.invs.foreach(inv => inv.visit{ + case t: SIFTerminatesExp => terminates = Some(t) + }) + if (terminates.isDefined) { + stmts ++= terminationChannelsLowChecks(terminates.get, ctx) + .map(tc => Assert(tc)(tc.pos, tc.info, tc.errT)) + val (cond1d, cond1r) = getNewBool("cond") + val (cond2d, cond2r) = getNewBool("cond") + primedNames.update(cond1r.name, cond2r.name) + newVarDecls ++= Seq(cond1d, cond2d) + stmts ++= Seq(If(p1, Seqn(Seq(LocalVarAssign(cond1r, terminates.get.cond)()), Seq())(), skip)(), + If(p2, Seqn(Seq(LocalVarAssign(cond2r, translatePrime(terminates.get.cond, p1, p2))()), Seq())(), skip)()) + newStdInvs :+= Implies(Not(cond1r)(), w.cond)( + terminates.get.pos, terminates.get.info, ErrTrafo({ + case _ => SIFTerminationChannelCheckFailed(terminates.get, SIFTermCondNotTight(terminates.get)) + })) + } + + val invCtx = ctx.copy( + p1 = And(p1, Not(bypass1r)())(), + p2 = And(p2, Not(bypass2r)())() + ) + /* + val invCtx = ctx.copy( + p1 = ctrlVars.activeExecNoContNormal(Some(p1)), + p2 = ctrlVars.activeExecNoContPrime(Some(p2)) + ) + */ + newStdInvs = simplifyConditions(newStdInvs.map(e => translateSIFAss(e, invCtx, invCtx))) + val newInvs: Seq[Exp] = newStdInvs ++ + targetValEqualities1.map(e => Implies(bypass1r, e)()) ++ + targetValEqualities2.map(e => Implies(bypass2r, e)()) + + val bodyRes = translateStatement(w.body, ctx.copy(p1 = p1r, p2 = p2r)) + /*val bodyPostamble = Seq( + Inhale(Or(Not(p1)(), ctrlVars.activeExecNoContNormal(None))())(), + Inhale(Or(Not(p2)(), ctrlVars.activeExecNoContPrime(None))())() + )*/ + stmts :+= While(newCond, newInvs, Seqn(bodyPreamble ++ Seq(bodyRes), Seq())())() + + // loop reconstruction + if (recNeeded) { + val recCond1 = And(Not(bypass1r)(), Seq(ctrlVars.ret1r, ctrlVars.break1r, ctrlVars.except1r) + .collect({ case Some(x) => x }) + .reduceRight[Exp]((x, y) => Or(x, y)()) + )() + val recCond2 = And(Not(bypass2r)(), Seq(ctrlVars.ret2r, ctrlVars.break2r, ctrlVars.except2r) + .collect({ case Some(x) => x }) + .reduceRight[Exp]((x, y) => Or(x, y)()))() + val recCond = Or(recCond1, recCond2)() + val ctrlVarAssigns: Seq[Stmt] = ctrlVars.declarations() + .map(d => d.localVar) + .map(v => LocalVarAssign(v, ctrlVarToOldMap(v))()) + val recInhales: Seq[Stmt] = Seq() :+ //newStdInvs.map(i => Inhale(i)()) :+ + Inhale(Implies(ctrlVars.activeExecNoContNormal(Some(p1)), w.cond)())() :+ + Inhale(Implies(ctrlVars.activeExecNoContNormal(Some(p2)), translatePrime(w.cond, p1, p2))())() + val recKillInhales: Seq[Stmt] = Seq( + Inhale(Or(Not(p1r)(), Not(ctrlVars.activeExecNoContNormal(None))())())(), + Inhale(Or(Not(p2r)(), Not(ctrlVars.activeExecNoContPrime(None))())())() + ) + val recThn = Seqn( + (ctrlVarAssigns ++ recInhales ++ bodyPreamble ++ Seq(bodyRes) ++ recKillInhales), + Seq())() + stmts :+= If(recCond, recThn, skip)(info=SimpleInfo(Seq("Loop Reconstruction.\n "))) + } + if (Seq(ctrlVars.break1r, ctrlVars.cont1r).collect({case Some(x) => x}).nonEmpty) { + stmts ++= Seq( + If(Not(bypass1r)(), Seqn(Seq(ctrlVars.break1r, ctrlVars.cont1r) + .collect({ case Some(x) => x }) + .map(v => LocalVarAssign(v, FalseLit()())()), + Seq())(), skip)(), + If(Not(bypass2r)(), Seqn(Seq(ctrlVars.break2r, ctrlVars.cont2r) + .collect({ case Some(x) => x }) + .map(v => LocalVarAssign(v, FalseLit()())()), + Seq())(), skip)() + ) + } + Seqn(stmts, newVarDecls)() + } + + /** Translate a try/catch block into MPP. + */ + private def translateTryCatchStmt(tryStmt: SIFTryCatchStmt, ctx: TranslationContext): Seqn = { + val p1 = ctx.p1; val p2 = ctx.p2; val ctrlVars = ctx.ctrlVars + var stmts: Seq[Stmt] = Seq() + var newVarDecls: Seq[LocalVarDecl] = Seq() + val hasFinally: Boolean = tryStmt.finallyBlock.isDefined + // bypass preamble + val (bypass1d, bypass1r) = getNewBool("bypass1") + val (bypass2d, bypass2r) = getNewBool("bypass2") + newVarDecls ++= Seq(bypass1d, bypass2d) + val bypassAssigns: Seq[Stmt] = bypassPreamble(p1, p2, ctrlVars, bypass1r, bypass2r) + stmts ++= bypassAssigns + // assigning old values of ret and except flags + def oldAssign(ctrl: Option[LocalVar], name: String): Option[LocalVar] = { + if (ctrl.isEmpty) return None + val (decl, v) = getNewBool(name) + newVarDecls :+= decl + stmts :+= LocalVarAssign(v, ctrl.get)() + Some(v) + } + var oldret1r, oldret2r, oldbreak1r, oldbreak2r, oldcont1r, oldcont2r, + oldexcept1r, oldexcept2r: Option[LocalVar] = None + if (hasFinally) { + oldret1r = oldAssign(ctrlVars.ret1r, name = "oldret1") + oldret2r = oldAssign(ctrlVars.ret2r, name = "oldret2") + oldbreak1r = oldAssign(ctrlVars.break1r, name = "oldbreak1") + oldbreak2r = oldAssign(ctrlVars.break2r, name = "oldbreak2") + oldcont1r = oldAssign(ctrlVars.cont1r, name = "oldcont1") + oldcont2r = oldAssign(ctrlVars.cont2r, name = "oldcont2") + oldexcept1r = oldAssign(ctrlVars.except1r, name = "oldexcept1") + oldexcept2r = oldAssign(ctrlVars.except2r, name = "oldexcept2") + } + // translate try body + stmts :+= translateStatement(tryStmt.body, ctx) + // create variable 'thisexcept', to express that we had an exception in this tryblock + val (thisexcept1d, thisexcept1r) = getNewBool("thisexcept1") + val (thisexcept2d, thisexcept2r) = getNewBool("thisexcept2") + newVarDecls ++= Seq(thisexcept1d, thisexcept2d) + stmts :+= LocalVarAssign(thisexcept1r, And(ctrlVars.except1r.get, Not(bypass1r)())())() + stmts :+= LocalVarAssign(thisexcept2r, And(ctrlVars.except2r.get, Not(bypass2r)())())() + // translate the exception handlers + for (handler <- tryStmt.catchBlocks) { + val (p1d, p1r) = getNewBool("p1") + val (p2d, p2r) = getNewBool("p2") + newVarDecls ++= Seq(p1d, p2d) + stmts = stmts :+ + LocalVarAssign(p1r, And(p1, And(thisexcept1r, translateNormal(handler.exception, p1, p2))())())() :+ + LocalVarAssign(p2r, And(p2, And(thisexcept2r, translatePrime(handler.exception, p1, p2))())())() :+ + If(p1r, Seqn(Seq(LocalVarAssign(ctrlVars.except1r.get, FalseLit()())()), Seq())(), skip)() :+ + If(p2r, Seqn(Seq(LocalVarAssign(ctrlVars.except2r.get, FalseLit()())()), Seq())(), skip)() + stmts :+= translateStatement(handler.body, TranslationContext(p1r, p2r, ctrlVars, ctx.currentMethod)) + // assign null to error variable if exception was caught + stmts :+= translateStatement(LocalVarAssign(handler.errVar, NullLit()())(), ctx) + } + // translate the else block + if (tryStmt.elseBlock.isDefined) { + val (p1d2, p1r2) = getNewBool("p1") + val (p2d2, p2r2) = getNewBool("p2") + newVarDecls ++= Seq(p1d2, p2d2) + stmts :+= LocalVarAssign(p1r2, And(p1, Not(thisexcept1r)())())() + stmts :+= LocalVarAssign(p2r2, And(p2, Not(thisexcept2r)())())() + stmts :+= translateStatement(tryStmt.elseBlock.get, TranslationContext(p1r2, p2r2, ctrlVars, ctx.currentMethod)) + } + // translate the finally block + if (hasFinally) { + def tmpAssigns(tuples: (LocalVar, Option[LocalVar], Option[LocalVar])*): Seq[Stmt] = { + tuples + .filter(tuple => tuple._2.isDefined) + .flatMap{ + case (tmp: LocalVar, ctrl: Option[LocalVar], old: Option[LocalVar]) => + Seq(LocalVarAssign(tmp, ctrl.get)(), + LocalVarAssign(ctrl.get, old.get)()) + } + } + def tmpReAssigns(tuples: (LocalVar, Option[LocalVar], Option[Seq[Option[LocalVar]]])*): Seq[Stmt] = { + tuples + .filter(tuple => tuple._2.isDefined) + .flatMap{ + case (tmp: LocalVar, ctrl: Option[LocalVar], None) => + Seq(LocalVarAssign(ctrl.get, Or(ctrl.get, tmp)())()) + case (tmp: LocalVar, ctrl: Option[LocalVar], Some(unless)) => + val negatedUnless = unless.map{ + case Some(v) => Some(Not(v)()) + case None => None + } + Seq(LocalVarAssign(ctrl.get, Or(ctrl.get, And(tmp, _conjoinOptions(negatedUnless))())())()) + } + } + + // store ret and except in tmp variables + val (tmpret1d, tmpret1r) = getNewBool("tmp_ret1") + val (tmpret2d, tmpret2r) = getNewBool("tmp_ret2") + val (tmpbreak1d, tmpbreak1r) = getNewBool("tmp_break1") + val (tmpbreak2d, tmpbreak2r) = getNewBool("tmp_break2") + val (tmpcont1d, tmpcont1r) = getNewBool("tmp_cont1") + val (tmpcont2d, tmpcont2r) = getNewBool("tmp_cont2") + val (tmpexcept1d, tmpexcept1r) = getNewBool("tmp_except1") + val (tmpexcept2d, tmpexcept2r) = getNewBool("tmp_except2") + newVarDecls ++= Seq(tmpret1d, tmpret2d, tmpbreak1d, tmpbreak2d, tmpcont1d, tmpcont2d, + tmpexcept1d, tmpexcept2d) + val tmpAssigns1: Seq[Stmt] = tmpAssigns( + (tmpret1r, ctrlVars.ret1r, oldret1r), + (tmpbreak1r, ctrlVars.break1r, oldbreak1r), + (tmpcont1r, ctrlVars.cont1r, oldcont1r), + (tmpexcept1r, ctrlVars.except1r, oldexcept1r) + ) + val tmpAssigns2: Seq[Stmt] = tmpAssigns( + (tmpret2r, ctrlVars.ret2r, oldret2r), + (tmpbreak2r, ctrlVars.break2r, oldbreak2r), + (tmpcont2r, ctrlVars.cont2r, oldcont2r), + (tmpexcept2r, ctrlVars.except2r, oldexcept2r) + ) + stmts :+= If(p1, Seqn(tmpAssigns1, Seq())(), skip)() + stmts :+= If(p2, Seqn(tmpAssigns2, Seq())(), skip)() + stmts :+= translateStatement(tryStmt.finallyBlock.get, ctx) + val tmpReAssigns1 = tmpReAssigns( + (tmpexcept1r, ctrlVars.except1r, Some(Seq(ctrlVars.ret1r, ctrlVars.break1r))), + (tmpret1r, ctrlVars.ret1r, None), + (tmpbreak1r, ctrlVars.break1r, None), + (tmpcont1r, ctrlVars.cont1r, None) + ) + val tmpReAssigns2 = tmpReAssigns( + (tmpexcept2r, ctrlVars.except2r, Some(Seq(ctrlVars.ret2r, ctrlVars.break2r))), + (tmpret2r, ctrlVars.ret2r, None), + (tmpbreak2r, ctrlVars.break2r, None), + (tmpcont2r, ctrlVars.cont2r, None) + ) + stmts :+= If(p1, Seqn(tmpReAssigns1, Seq())(), skip)() + stmts :+= If(p2, Seqn(tmpReAssigns2, Seq())(), skip)() + } + Seqn(stmts, newVarDecls)(info = SimpleInfo(Seq("Try/catch block\n "))) + } + + def translateWandPackage(p: Package, ctx: TranslationContext): Stmt = { + val Package(w@MagicWand(left, right), proofScript) = p + // Package( + // MagicWand( + // translateSIFAss(left, ctx), + // translateSIFAss(right, ctx) + // )(w.pos, w.info, w.errT), + // Seqn(Seq(translateStatement(proofScript, ctx)), Seq())() + // )(p.pos, p.info, p.errT) + val p1 = ctx.p1 + val p2 = ctx.p2 + lazy val act1: Exp = ctx.ctrlVars.activeExecNormal(Some(p1)) + lazy val act2: Exp = ctx.ctrlVars.activeExecPrime(Some(p2)) + + val leftNormal = translateNormal(left, p1, p2) + val rightNormal = translateNormal(right, p1, p2) + val wandNormal = MagicWand(leftNormal, rightNormal)(w.pos, w.info, w.errT) + val proofScriptNormal = proofScript.transform{ + case e: Exp => translateNormal(e, act1, act2) + } + val packageNormal = Package(wandNormal, proofScriptNormal)(p.pos, p.info, p.errT) + val ifNormal = If(act1, Seqn(Seq(packageNormal), Seq())(), skip)() + + val leftPrime = translatePrime(left, p1, p2) + val rightPrime = translatePrime(right, p1, p2) + val wandPrime = MagicWand(leftPrime, rightPrime)(w.pos, w.info, w.errT) + val proofScriptPrime = proofScript.transform{ + case e: Exp => translatePrime(e, act1, act2) + } + val packagePrime = Package(wandPrime, proofScriptPrime)(p.pos, p.info, p.errT) + val ifPrime = If(act2, Seqn(Seq(packagePrime), Seq())(), skip)() + + Seqn(Seq(ifNormal, ifPrime), Seq())() + } + + def translateWandApply(a: Apply, ctx: TranslationContext): Stmt = { + val Apply(w@MagicWand(left, right)) = a + // Apply( + // MagicWand( + // translateSIFAss(left, ctx), + // translateSIFAss(right, ctx) + // )(w.pos, w.info, w.errT), + // )(a.pos, a.info, a.errT) + + val p1 = ctx.p1 + val p2 = ctx.p2 + lazy val act1: Exp = ctx.ctrlVars.activeExecNormal(Some(p1)) + lazy val act2: Exp = ctx.ctrlVars.activeExecPrime(Some(p2)) + + val leftNormal = translateNormal(left, p1, p2) + val rightNormal = translateNormal(right, p1, p2) + val wandNormal = MagicWand(leftNormal, rightNormal)(w.pos, w.info, w.errT) + val applyNormal = Apply(wandNormal)(a.pos, a.info, a.errT) + val ifNormal = If(act1, Seqn(Seq(applyNormal), Seq())(), skip)() + + val leftPrime = translatePrime(left, p1, p2) + val rightPrime = translatePrime(right, p1, p2) + val wandPrime = MagicWand(leftPrime, rightPrime)(w.pos, w.info, w.errT) + val applyPrime = Apply(wandPrime)(a.pos, a.info, a.errT) + val ifPrime = If(act2, Seqn(Seq(applyPrime), Seq())(), skip)() + + Seqn(Seq(ifNormal, ifPrime), Seq())() + } + + def translateStatement(s: Stmt, ctx: TranslationContext) : Stmt = { + val p1 = ctx.p1 + val p2 = ctx.p2 + val ctrlVars = ctx.ctrlVars + lazy val act1: Exp = ctx.ctrlVars.activeExecNormal(Some(p1)) + lazy val act2: Exp = ctx.ctrlVars.activeExecPrime(Some(p2)) + + def executeConditionally(ift1: Seqn, ift2: Seqn): Stmt = { + val a1 = If(act1, ift1, skip)() + val a2 = If(act2, ift2, skip)() + Seqn(Seq(a1, a2, incrementTime(p1, p2)), Seq())() + } + def translateAssignment(to1: LocalVar, v1: Exp, to2: LocalVar, v2: Exp, orig: Stmt) : Stmt = { + executeConditionally( + Seqn(Seq(LocalVarAssign(to1, v1)(orig.pos, orig.info, orig.errT)), Seq())(), + Seqn(Seq(LocalVarAssign(to2, v2)(orig.pos, orig.info, orig.errT)), Seq())() + ) + } + + /** Return true if the statement can be bunched together with others, without doing the interleaving of the + * two executions, thus saving on the number of if-statements generated. + */ + def isCompressible(s: Stmt): Boolean = { + s match { + case _: LocalVarAssign => true + case Inhale(e) => isUnary(e) + case Exhale(e) => isUnary(e) + case _: SIFReturnStmt => true + case _: SIFBreakStmt => true + case _: SIFContinueStmt => true + case _: Goto => false + case _: Label => false + case _ => false + } + } + + /** Do the translation of a statement without wrapping it in an if statement. Just returns the two versions in a + * tuple, to allow putting multiple statements in an if block. Requires [[isCompressible(s)]]. + */ + def translateStmtPartial(s: Stmt): (Stmt, Stmt) = { + assert(isCompressible(s)) + s match { + case l@LocalVarAssign(lhs, rhs) => (LocalVarAssign(translateNormal(lhs, p1, p2), + translateNormal(rhs, p1, p2))(l.pos, l.info, l.errT), + LocalVarAssign(translatePrime(lhs, p1, p2), translatePrime(rhs, p1, p2))(l.pos, l.info, l.errT)) + case a@FieldAssign(lhs, rhs) => (a, FieldAssign(translatePrime(lhs, p1, p2), + translatePrime(rhs, p1, p2))(a.pos, a.info, a.errT)) + case i@Inhale(e) => (Inhale(translateNormal(e, p1, p2))(i.pos, i.info, i.errT), + Inhale(translatePrime(e, p1, p2))(i.pos, i.info, i.errT)) + case ex@Exhale(e) => (Exhale(translateNormal(e, p1, p2))(ex.pos, ex.info, ex.errT), + Exhale(translatePrime(e, p1, p2))(ex.pos, ex.info, ex.errT)) + case b: SIFBreakStmt => { + // TODO + (LocalVarAssign(ctrlVars.break1r.get, TrueLit()())(b.pos, b.info, b.errT), + LocalVarAssign(ctrlVars.break2r.get, TrueLit()())(b.pos, b.info, b.errT)) + } + case c: SIFContinueStmt => (LocalVarAssign(ctrlVars.cont1r.get, TrueLit()())(c.pos, c.info, c.errT), + LocalVarAssign(ctrlVars.cont2r.get, TrueLit()())(c.pos, c.info, c.errT)) + case r@SIFReturnStmt(e, resVar) => + { + // TODO + val assign1 = resVar match { + case Some(rv) => Seq(LocalVarAssign(translateNormal(rv, p1, p2), + translateNormal(e.get, p1, p2))(r.pos, r.info, r.errT)) + case None => Seq() + } + val assign2 = resVar match { + case Some(rv) => Seq(LocalVarAssign(translatePrime(rv, p1, p2), + translatePrime(e.get, p1, p2))(r.pos, r.info, r.errT)) + case None => Seq() + } + (Seqn(assign1 :+ LocalVarAssign(ctrlVars.ret1r.get, TrueLit()())(), Seq())(), + Seqn(assign2 :+ LocalVarAssign(ctrlVars.ret2r.get, TrueLit()())(), Seq())()) + } + + case _ => throw new IllegalArgumentException(s"The statement $s can't be translated partially") + } + } + + def optimizeSequential(s: Seqn): Seq[Stmt] = { + def sequenceSplit(in: Seq[Stmt]): (Seq[Stmt], Seq[Stmt]) = { + // ensure that we make a split after return, break, continue, raise, because the ctrlVars will have changed + var stop: Boolean = false + def keepGoing(stmt: Stmt): Boolean = { + val oldStop = stop + stmt match { + case _: SIFReturnStmt => stop = true + case _: SIFBreakStmt => stop = true + case _: SIFContinueStmt => stop = true + case _: SIFRaiseStmt => stop = true + case _ => + } + !oldStop && isCompressible(stmt) + } + in.span(stmt => keepGoing(stmt)) + } + + var newStmts = Seq[Stmt]() + var (comp, rest) = sequenceSplit(s.ss) +// println("optimizing compressible statements:") + while (comp.nonEmpty || rest.nonEmpty) { +// println(s"split into comp: $comp and rest: $rest") + // collect all compressible statements we have until here + if (comp.nonEmpty) { + val (fstExComp, secExComp): (Seq[Stmt], Seq[Stmt]) = comp.map(stmt => translateStmtPartial(stmt)).unzip + newStmts :+= executeConditionally(Seqn(fstExComp, Seq())(), Seqn(secExComp, Seq())()) + } + // translate all non-compressible statements + var split = rest.span(stmt => !isCompressible(stmt)) + val nonComp = split._1 + rest = split._2 +// println(s"split into non-comp: $nonComp and rest: $rest") + nonComp.foreach(stmt => newStmts :+= translateStatement(stmt, ctx)) + // start anew + split = sequenceSplit(rest) + comp = split._1 + rest = split._2 + } + newStmts + } + + s match { + case l@LocalVarAssign(lhs, rhs) => translateAssignment(translateNormal(lhs, p1, p2), + translateNormal(rhs, p1, p2), translatePrime(lhs, p1, p2), translatePrime(rhs, p1, p2), l) + case a@FieldAssign(lhs, rhs) => executeConditionally(Seqn(Seq(a), Seq())(), + Seqn(Seq(FieldAssign(translatePrime(lhs, p1, p2), + translatePrime(rhs, p1, p2))(a.pos, a.info, a.errT)), Seq())()) + case NewStmt(lhs, fields) => { + val allFields = fields ++ fields.map{f => newFields.find(f2 => f2.name == primedNames(f.name)).get} + val (tmpd, tmpr) = getNewVar("tmp", Ref) + val newNew = NewStmt(tmpr, allFields)() + /*val allFieldAssigns = allFields.map { f => + val (hd, hr) = getNewVar("havoc", f.typ) + val fieldAcc = FieldAccess(tmpr, f)() + Seqn(Seq(FieldAssign(fieldAcc, hr)()), Seq(hd))() + }*/ + val assign1 = If(act1, Seqn(Seq(LocalVarAssign(lhs, tmpr)()), Seq())(), skip)() + val assign2 = If(act2, Seqn(Seq(LocalVarAssign(translatePrime(lhs, p1, p2), tmpr)()), + Seq())(), skip)() + Seqn(Seq(newNew) ++ /*allFieldAssigns ++*/ Seq(assign1, assign2, incrementTime(p1, p2)), Seq(tmpd))() + } + case i@If(cond, thn, els) => { + val (p1d, p1r) = getNewBool("p1") + val (p2d, p2r) = getNewBool("p2") + val (p3d, p3r) = getNewBool("p3") + val (p4d, p4r) = getNewBool("p4") + + val p1Assign = LocalVarAssign(p1r, And(act1, cond)())(i.pos) + val p2Assign = LocalVarAssign(p2r, And(act2, translatePrime(cond, p1, p2))())(i.pos) + val p3Assign = LocalVarAssign(p3r, And(act1, Not(cond)())())(i.pos) + val p4Assign = LocalVarAssign(p4r, And(act2, Not(translatePrime(cond, p1, p2))())())(i.pos) + + val thnRes = translateStatement(thn, TranslationContext(p1r, p2r, ctrlVars, ctx.currentMethod)) + val elsRes = translateStatement(els, TranslationContext(p3r, p4r, ctrlVars, ctx.currentMethod)) + Seqn(Seq(p1Assign, p2Assign, p3Assign, p4Assign, incrementTime(p1, p2), thnRes, elsRes), Seq(p1d, p2d, p3d, p4d))() + } + case w: While => translateWhileStmt(w, ctx) + case mc@MethodCall(name, args, targets) => { + var argDecls = Seq[LocalVarDecl]() + var newArgs = Seq[Exp](act1, act2) + var argAssigns1 = Seq[LocalVarAssign]() + var argAssigns2 = Seq[LocalVarAssign]() + + if (timing){ + newArgs ++= Seq(time.get, translatePrime(time.get, p1, p2)) + } + + for (a <- args){ + val (tmp1d, tmp1r) = getNewVar("tmp1", a.typ) + val (tmp2d, tmp2r) = getNewVar("tmp2", a.typ) + argDecls ++= Seq(tmp1d, tmp2d) + newArgs ++= Seq(tmp1r, tmp2r) + argAssigns1 :+= LocalVarAssign(tmp1r, a)() + argAssigns2 :+= LocalVarAssign(tmp2r, translatePrime(a, p1, p2))() + } + val argAssignsConditional: Seq[Stmt] = Seq( + If(act1, Seqn(argAssigns1, Seq())(), skip)(), + If(act2, Seqn(argAssigns2, Seq())(), skip)() + ) + + var targetDecls = Seq[LocalVarDecl]() + var newTargets = Seq[LocalVar]() + var targetAssigns1 = Seq[LocalVarAssign]() + var targetAssigns2 = Seq[LocalVarAssign]() + + if (timing){ + newTargets ++= Seq(time.get, translatePrime(time.get, p1, p2)) + } + + for (t <- targets){ + val (tmp1d, tmp1r) = getNewVar("tmp1", t.typ) + val (tmp2d, tmp2r) = getNewVar("tmp2", t.typ) + targetDecls ++= Seq(tmp1d, tmp2d) + newTargets ++= Seq(tmp1r, tmp2r) + targetAssigns1 :+= LocalVarAssign(t, tmp1r)() + targetAssigns2 :+= LocalVarAssign(translatePrime(t, p1, p2), tmp2r)() + } + val targetAssignsConditional = if (targets.nonEmpty) Seq[Stmt]( + If(act1, Seqn(targetAssigns1, Seq())(), skip)(), + If(act2, Seqn(targetAssigns2, Seq())(), skip)() + ) else Seq() + + val call = MethodCall(name, newArgs, newTargets)(mc.pos, mc.info, mc.errT) + If(Or(act1, act2)(), + Seqn(argAssignsConditional ++ Seq(call) ++ targetAssignsConditional, argDecls ++ targetDecls)(), + skip + )(info = SimpleInfo(Seq(s"Method call: $name\n "))) + } + case s: Seqn => { + val seq = if (Config.optimizeSequential) flattenSeqn(s) else s + var newDecls = Seq[Declaration]() + for (d <- seq.scopedDecls.filter(d => d.isInstanceOf[LocalVarDecl])){ + val newName = getName(d.name) + primedNames.update(d.name, newName) + val newD = LocalVarDecl(newName, d.asInstanceOf[LocalVarDecl].typ)() + newDecls ++= Seq(d, newD) + } + newDecls ++= seq.scopedDecls.filter(d => d.isInstanceOf[Label]) + val newStmts = if (Config.optimizeSequential) optimizeSequential(seq) + else seq.ss.map{stmt => translateStatement(stmt, ctx)} + Seqn(newStmts, newDecls)() + } + case a@Assert(e1) => + val newCtx = a.info.getUniqueInfo[SIFInfo] match { + case Some(info) if info.continueUnaware => ctx.copy(p1 = ctrlVars.activeExecNoContNormal(Some(p1)), + p2 = ctrlVars.activeExecNoContPrime(Some(p2))) + case _ => ctx.copy(p1 = act1, p2 = act2) + } + Assert(translateSIFAss(e1, newCtx))(s.pos, errT= fwTs(s, s)) + case i@Inhale(FalseLit()) => i + case Assume(e1) => Assume(translateSIFAss(e1, ctx.copy(p1 = act1, p2 = act2)))(s.pos, errT= fwTs(s, s)) + case Inhale(e1) => Inhale(translateSIFAss(e1, ctx.copy(p1 = act1, p2 = act2)))(s.pos, errT= fwTs(s, s)) + case Exhale(e1) => Exhale(translateSIFAss(e1, ctx.copy(p1 = act1, p2 = act2)))(s.pos, errT= fwTs(s, s)) + case d : LocalVarDeclStmt => d + case u@Unfold(acc) => + val predicate2 = PredicateAccess(acc.loc.args.map(a => translatePrime(a, p1, p2)), + acc.loc.predicateName)() + val (lowFunc, lhs) = getPredicateLowFuncExp(acc.loc.predicateName, ctx) + val assert = lowFunc match { + case Some(f) => + val et = ErrTrafo({case _: AssertFailed => errors.UnfoldFailed(u, SIFUnfoldNotLow(u))}) + Assert(Implies( + lhs, + Implies( + And( + PermGeCmp(CurrentPerm(acc.loc)(), FullPerm()())(), + PermGeCmp(CurrentPerm(predicate2)(), FullPerm()())() + )(), + FuncApp(f, acc.loc.args ++ acc.loc.args.map(a => translatePrime(a, p1, p2)))() + )())() + )(u.pos, u.info, errT = et) + case None => skip + } + val if1 = If(act1, Seqn(Seq(u), Seq())(), skip)() + val if2 = If(act2, Seqn(Seq( + Unfold(PredicateAccessPredicate(predicate2, translatePrime(acc.perm, p1, p2))())(u.pos, u.info, u.errT) + ), Seq())(), skip)() + Seqn(Seq(assert, if1, if2), Seq())() + case f@Fold(acc) => + val if1 = If(act1, Seqn(Seq(f), Seq())(), skip)() + val if2 = If(act2, Seqn(Seq( + Fold(PredicateAccessPredicate( + PredicateAccess(acc.loc.args.map(a => translatePrime(a, p1, p2)), + acc.loc.predicateName)(), translatePrime(acc.perm, p1, p2))())(f.pos, f.info, f.errT) + ), Seq())(), skip)() + val (lowFunc, lhs) = getPredicateLowFuncExp(acc.loc.predicateName, ctx) + val assert: Stmt = lowFunc match { + case Some(func) => + val et = ErrTrafo({case AssertFailed(_,_,_) => errors.FoldFailed(f, SIFFoldNotLow(f))}) + Assert(Implies( + lhs, + FuncApp(func.copy()(func.pos, func.info, errT = et), + acc.loc.args ++ acc.loc.args.map(a => translatePrime(a, p1, p2)))() + )())(f.pos, f.info, errT = et) + case None => skip + } + Seqn(Seq(if1, if2, assert), Seq())() + case r: SIFReturnStmt => { + // TODO + val (first, second) = translateStmtPartial(r).asInstanceOf[(Seqn, Seqn)] + val r1 = If(act1, first, skip)() + val r2 = If(act2, second, skip)() + Seqn(Seq(r1, r2, incrementTime(p1, p2)), Seq())() + } + case b@SIFBreakStmt() => { + // TODO + translateAssignment(ctrlVars.break1r.get, TrueLit()(), + ctrlVars.break2r.get, TrueLit()(), b) + } + case c@SIFContinueStmt() => translateAssignment(ctrlVars.cont1r.get, TrueLit()(), + ctrlVars.cont2r.get, TrueLit()(), c) + case tryCatch: SIFTryCatchStmt => translateTryCatchStmt(tryCatch, ctx) + case SIFRaiseStmt(assignment) => + var stmts = Seq[Stmt]() + val assign1 = assignment match { + case Some(a) => Some(LocalVarAssign(translateNormal(a.lhs, p1, p2), + translateNormal(a.rhs, p1, p2))()) + case None => None + } + val assign2 = assignment match { + case Some(a) => Some(LocalVarAssign(translatePrime(a.lhs, p1, p2), + translatePrime(a.rhs, p1, p2))()) + case None => None + } + stmts :+= If(act1, + Seqn(Seq(assign1, Some(LocalVarAssign(ctrlVars.except1r.get, TrueLit()())())) + .collect({case Some(x) => x}), Seq())(), + skip)() + stmts :+= If(act2, + Seqn(Seq(assign2, Some(LocalVarAssign(ctrlVars.except2r.get, TrueLit()())())) + .collect({case Some(x) => x}), Seq())(), + skip)() + Seqn(stmts, Seq())() + case d@SIFDeclassifyStmt(e) => Inhale(Implies( + And(act1, act2)(), EqCmp(translateNormal(e, p1, p2), translatePrime(e, p1, p2))() + )())(d.pos, d.info, d.errT) + case SIFAssertNoException() => + val exp: Exp = if (ctrlVars.except1r.isDefined) And( + Implies(p1, Not(ctrlVars.except1r.get)())(s.pos, s.info, s.errT), + Implies(p2, Not(ctrlVars.except2r.get)())(s.pos, s.info, s.errT) + )(s.pos, s.info, s.errT) + else TrueLit()() + Exhale(exp)(s.pos, s.info, s.errT) + case SIFInlinedCallStmt(stmts) => + val newCtrlVars = createControlFlowVars(stmts) + val inlinedCtx = TranslationContext(p1, p2, newCtrlVars, ctx.currentMethod) + Seqn(newCtrlVars.initAssigns() :+ translateStatement(stmts, inlinedCtx), newCtrlVars.declarations())() + case Goto(l) => { + val varName1 = ctrlVars.labelNames.get(l).get + val varName2 = primedNames.get(varName1).get + val assign1 = If(act1, Seqn(Seq(LocalVarAssign(LocalVar(varName1, Bool)(), TrueLit()())()), Seq())(), Seqn(Seq(), Seq())())() + val assign2 = If(act2, Seqn(Seq(LocalVarAssign(LocalVar(varName2, Bool)(), TrueLit()())()), Seq())(), Seqn(Seq(), Seq())())() + Seqn(Seq(assign1, assign2), Seq())() + } + case lb@Label(l, _) if ctrlVars.labelNames.contains(l) => { + val varName1 = ctrlVars.labelNames.get(l).get + val varName2 = primedNames.get(varName1).get + val thisAct1 = ctx.ctrlVars.activeExecNormalExceptLabel(Some(p1), varName1) + val thisAct2 = ctx.ctrlVars.activeExecPrimeExceptLabel(Some(p2), varName2) + val assign1 = If(thisAct1, Seqn(Seq(LocalVarAssign(LocalVar(varName1, Bool)(), FalseLit()())()), Seq())(), Seqn(Seq(), Seq())())() + val assign2 = If(thisAct2, Seqn(Seq(LocalVarAssign(LocalVar(varName2, Bool)(), FalseLit()())()), Seq())(), Seqn(Seq(), Seq())())() + Seqn(Seq(lb, assign1, assign2), Seq())() + } + case lb : Label => lb + case p: Package => translateWandPackage(p, ctx) + case a: Apply => translateWandApply(a, ctx) + case other => throw new IllegalArgumentException("unexpected: " + other) + } + } + + def getNewVar(name: String, typ: Type) : (LocalVarDecl, LocalVar) = { + val newName = getName(name) + (LocalVarDecl(newName, typ)(), LocalVar(newName, typ)()) + } + + def getNewBool(name: String) : (LocalVarDecl, LocalVar) = { + getNewVar(name, Bool) + } + + def isRelational(e: Exp): Boolean = { + !isUnary(e) + } + + def isUnary(e: Exp): Boolean = { + val relVars = e.filter{ + case _: SIFLowExp => true + case _: SIFLowEventExp => true + case DomainFuncApp("Low", _, _) => true + case _ => false + } + relVars.isEmpty + } + + def translateSIFExp1(e: Exp, p1: Exp, p2: Exp): Exp = { + e match { + case re if isRelational(e) => Implies(And(p1, p2)(), translateNormal(e, p1, p2))(e.pos, errT = fwTs(re, re)) + case _ if isUnary(e) => translateNormal(e, p1, p2) + } + } + + def translateSIFExp2(e: Exp, p1: Exp, p2: Exp): Exp = { + e match { + case re if isRelational(e) => TrueLit()(e. pos, errT = fwTs(re, re)) + case _ => translatePrime(e, p1, p2) + } + } + + def translateSIFLowExpComparison(l: SIFLowExp, p1: Exp, p2: Exp): Exp = { + val primedExp = translatePrime(l.exp, p1, p2) + l.comparator match { + case None => EqCmp(l.exp, primedExp)(l.pos, errT = fwTs(l, l)) + case Some(str) => + _program.findDomainFunctionOptionally(str) match { + case Some(df) => DomainFuncApp(df, Seq(l.exp, primedExp), l.typVarMap)(l.pos, l.info, errT = fwTs(l, l)) + case None => _program.findFunctionOptionally(str) match { + case Some(f) => FuncApp(f, Seq(l.exp, primedExp))(l.pos, l.info, errT = fwTs(l, l)) + case None => sys.error(s"Unknown comparator $str.") + } + } + } + } + + def translateAssDefault(e: Exp, p1: Exp, p2: Exp): And = { + And(Implies(p1, translateSIFExp1(e, p1, p2))(e.pos, errT = fwTs(e, e)), + Implies(p2, translateSIFExp2(e, p1, p2))(e.pos, errT = fwTs(e, e)))(e.pos, errT = fwTs(e, e)) + } + + def fwTs(t: TransformableErrors, node: ErrorNode) = { + Trafos(t.errT.eTransformations, t.errT.rTransformations, Some(node)) + } + + def translateSIFAss(e: Exp, ctx: TranslationContext, relAssertCtx: TranslationContext = null): Exp = { + val p1 = ctx.p1 + val p2 = ctx.p2 + val relCtx = if (relAssertCtx == null) ctx else relAssertCtx + def bothExecutions(e: Exp, pos: Position = NoPosition, info: Info = NoInfo, + errT: ErrorTrafo = NoTrafos): Exp = { + Implies(And(relCtx.p1, relCtx.p2)(), e)(pos, info, errT) + } + + e match { + case And(e1, e2) => And(translateSIFAss(e1, ctx, relAssertCtx), translateSIFAss(e2, ctx, relAssertCtx))(e.pos, errT = fwTs(e, e)) + case i@Implies(e1, e2) if !isUnary(i) => { + Implies(translateSIFAss(e1, ctx, relAssertCtx), translateSIFAss(e2, ctx, relAssertCtx))(e.pos, errT = fwTs(e, e)) + } + case Implies(e1, e2) if e2.exists({ + case PredicateAccess(_, name) => predLowFuncs(name).isDefined + case _ => false + }) => + And(translateAssDefault(e, p1, p2), Implies( + translateSIFAss(e1, ctx, relAssertCtx), + translatePredLowFuncOnly(e2, p1, p2) + )())(e.pos, e.info, e.errT) + case fa@Forall(vars, triggers, exp) => { + if (fa.isPure){ + for (v <- vars){ + if (primedNames.contains(v.name)){ + primedNames.remove(v.name) + } + } + /*var varEqs: Exp = TrueLit()() + val pvars = vars.map { v => + val primeName = getName(v.name) + primedNames.update(v.name, primeName) + varEqs = And(varEqs, EqCmp(v.localVar, LocalVar(primeName)(v.typ))())() + LocalVarDecl(primedNames.get(v.name).get, v.typ)() + }*/ + val newTriggers = triggers.map{t => Trigger(t.exps.map{e => translatePrime(e, p1, p2)})()} + //val res = Forall(vars ++ pvars, newTriggers, Implies(varEqs, translateSIFAss(exp, p1, p2))(e.pos, errT = NodeTrafo(e)))(e.pos, errT = NodeTrafo(e)) + val res = Forall(vars, triggers ++ newTriggers, + translateSIFAss(exp, ctx, relAssertCtx))(e.pos, errT = fwTs(e, e)).autoTrigger + res + } else { + val normal = translateNormal(fa, p1, p2) + val prime = translatePrime(fa, p1, p2) + And(Implies(p1, normal)(e.pos, errT = fwTs(e, e)), Implies(p2, prime)(e.pos, errT = fwTs(e, e)))(e.pos, errT = fwTs(e, e)) + } + } + case l: SIFLowEventExp => + val act1 = ctx.ctrlVars.activeExecNormal(Some(p1)) + val act2 = ctx.ctrlVars.activeExecPrime(Some(p2)) + val dynCheckInfo = l.info.getUniqueInfo[SIFDynCheckInfo] + dynCheckInfo match { + case None => EqCmp(act1, act2)(e.pos, errT = fwTs(e, e)) + case Some(dci) => And(EqCmp(act1, act2)(e.pos, errT = fwTs(e, e)), Implies(act1, + EqCmp(translateNormal(dci.dynCheck, p1, p2), + translatePrime(dci.dynCheck, p1, p2))())())(e.pos, errT = fwTs(e, e)) + } + case _: SIFLowExitExp => + val act1 = ctx.ctrlVars.activeExecNoContNormal(None) + val act2 = ctx.ctrlVars.activeExecNoContPrime(None) + Implies(And(relCtx.p1, relCtx.p2)(), EqCmp(act1, act2)(e.pos, errT = fwTs(e, e)))(e.pos, errT = fwTs(e, e)) + case l@SIFLowExp(_, _, _) => + val comparison = translateSIFLowExpComparison(l, relCtx.p1, relCtx.p2) + val dynCheckInfo = l.info.getUniqueInfo[SIFDynCheckInfo] + dynCheckInfo match { + case None => bothExecutions(comparison, e.pos) + case Some(dci) => + val inhalePart = bothExecutions(Implies( + EqCmp(translateNormal(dci.dynCheck, relCtx.p1, relCtx.p2), translatePrime(dci.dynCheck, relCtx.p1, relCtx.p2))(), comparison + )(), e.pos) + if (dci.onlyDynVersion) { + inhalePart + } else { + InhaleExhaleExp(inhalePart, bothExecutions(comparison, e.pos) + )(l.pos, l.info, errT = fwTs(l, l)) + } + } + // for the domain method low, used e.g. for list resource + case f@DomainFuncApp("Low", args, _) => translateSIFAss( + SIFLowExp(args.head, None)(f.pos, f.info, f.errT), ctx, relAssertCtx) + case pap@PredicateAccessPredicate(pred, _) => + val (lowFunc, lhs) = getPredicateLowFuncExp(pred.predicateName, ctx) + lowFunc match { + case Some(f: Function) => + val lowFuncApp = Implies(lhs, + FuncApp(f, pred.args ++ pred.args.map(a => translatePrime(a, p1, p2)))(pap.pos, pap.info, pap.errT) + )() + val dynCheckInfo = pap.info.getUniqueInfo[SIFDynCheckInfo] + val lowPart: Exp = dynCheckInfo match { + case Some(dfi) => + val inhalePart = Implies( + EqCmp(translateNormal(dfi.dynCheck, p1, p2), translatePrime(dfi.dynCheck, p1, p2))(), + lowFuncApp + )(e.pos, e.info, e.errT) + if (dfi.onlyDynVersion) inhalePart + else InhaleExhaleExp(inhalePart, lowFuncApp)(e.pos, e.info, e.errT) + case None => lowFuncApp + } + And(translateAssDefault(pap, p1, p2), lowPart)(e.pos, errT = fwTs(e, e)) + case None => translateAssDefault(pap, p1, p2) + } + case o@Old(oldExp) => Old(translateSIFAss(oldExp, ctx, relAssertCtx))(o.pos, o.info, o.errT) + case FuncApp(name, _) if predAllLowFuncs.values.exists(v => v.isDefined && v.get.name == name) => + bothExecutions(e) + case Unfolding(predAcc, body) if !isUnary(body) => bothExecutions(Unfolding( + translateNormal(predAcc, relCtx.p1, relCtx.p2), + Unfolding(translatePrime(predAcc, relCtx.p1, relCtx.p2), + translateSIFAss(body, ctx, relAssertCtx))(e.pos, e.info, e.errT))(e.pos, e.info, e.errT), + e.pos, e.info, e.errT) + // case w@MagicWand(left, right) => + // case Let(decl, e1, e2) => ??? + case _: SIFTerminatesExp => TrueLit()() + case _ => translateAssDefault(e, p1, p2) + } + } + + def translatePrime[T <: Exp](e: T, p1: Exp, p2: Exp) : T = { + e.transform{ + case d: LocalVarDecl if primedNames.contains(d.name) => + d.copy(name = primedNames(d.name))(d.pos, d.info, d.errT) + case l: LocalVar if primedNames.contains(l.name) => + l.copy(name = primedNames(l.name))(l.pos, l.info, l.errT) + case l: LocalVar if !primedNames.contains(l.name) => l + case FieldAccess(rcv, field) => + FieldAccess(translatePrime(rcv, p1, p2), field)(e.pos) + case f@FuncApp(name, _) if Config.primedFuncAppReplacements.keySet.contains(name) => + Config.primedFuncAppReplacements(name)(f, p1, p2) + case f@FuncApp(name, args) => FuncApp(name, + args.map(a => translatePrime(a, p1, p2)))(f.pos, f.info, f.typ, f.errT) + case DomainFuncApp("Low", _, _) => TrueLit()() + case df@DomainFuncApp(name, args, typVarMap) => + DomainFuncApp( + name, args.map(a => translatePrime(a, p1, p2)), typVarMap)( + df.pos, df.info, df.typ, df.domainName, df.errT) + case pa@PredicateAccess(args, "MayJoin") => PredicateAccess(args.map(a => translatePrime(a, p1, p2)), "MayJoinP")(pa.pos, pa.info, pa.errT) + case pa@PredicateAccess(args, name) => PredicateAccess(args.map(a => translatePrime(a, p1, p2)), + name)(pa.pos, pa.info, pa.errT) + case l: SIFLowExp => Implies(And(p1, p2)(), translateSIFLowExpComparison(l, p1, p2))() + case f@ForPerm(vars, location, body) => ForPerm(vars, + translateResourceAccess(location), + translatePrime(body, p1, p2))(f.pos, f.info, f.errT) + } + } + + def translateNormal[T <: Exp](e: T, p1: Exp, p2: Exp): T = { + e.transform{ + case l: SIFLowExp => Implies(And(p1, p2)(), translateSIFLowExpComparison(l, p1, p2))() + case DomainFuncApp("Low", args, _) => Implies(And(p1, p2)(), translateSIFLowExpComparison(SIFLowExp(args.head)(), p1, p2))() + } + } + + def translateToUnary(e: Option[Exp]): Option[Exp] = { + e match { + case Some(x) => Some(translateToUnary(x)) + case None => None + } + } + + /** Translate an expression getting rid of all the relational parts. */ + def translateToUnary(e: Exp): Exp = { + val transformed = e.transform{ + case _: SIFLowExp => TrueLit()() + case DomainFuncApp("Low", _, _) => TrueLit()() + case Implies(_: SIFLowExp, _: SIFLowExp) => TrueLit()() + case i@Implies(lhs, rhs) => Implies(lhs, translateToUnary(rhs))(i.pos, i.info, i.errT) + } + Simplifier.simplify(transformed) + } + + /** Translate only the relational parts of an expression. All unary parts are translated to True. + * @param e The expression to translate. + * @return The translation of the relational parts of `e`. + */ + def translatePredLowFuncBody(e: Exp): Exp = { + val translated = e match { + case l: SIFLowExp => translateSIFLowExpComparison(l, null, null) + case p@PredicateAccessPredicate(loc, _) => + val (lowFName, _, _) = predLowFuncInfo(loc.predicateName).get + FuncApp(lowFName, + loc.args ++ loc.args.map(a => translatePrime(a, null, null)))( + p.pos, NoInfo, Bool, p.errT) + case a@And(left, right) => And(translatePredLowFuncBody(left), + translatePredLowFuncBody(right))(a.pos, a.info, a.errT) + case o@Or(left, right) => Or(translatePredLowFuncBody(left), + translatePredLowFuncBody(right))(o.pos, o.info, o.errT) + case i@Implies(left, right) => Implies(And(translateNormal(left, null, null), + translatePrime(left, null, null))(), + translatePredLowFuncBody(right) + )(i.pos, i.info, i.errT) + case _ => TrueLit()() + } + Simplifier.simplify(translated) + } + + def translatePredAllLowFuncBody(e: Exp): Exp = { + val translated = e match { + case FieldAccessPredicate(loc, _) => EqCmp(loc, translatePrime(loc, null, null))() + case p@PredicateAccessPredicate(loc, _) => + if (predAllLowFuncInfo(loc.predicateName).isDefined){ + val (lowFName, _, _) = predAllLowFuncInfo(loc.predicateName).get + FuncApp(lowFName, + loc.args ++ loc.args.map(a => translatePrime(a, null, null)))( + p.pos, NoInfo, Bool, p.errT) + }else{ + TrueLit()() + } + + case a@And(left, right) => And(translatePredAllLowFuncBody(left), + translatePredAllLowFuncBody(right))(a.pos, a.info, a.errT) + case o@Or(left, right) => Or(translatePredAllLowFuncBody(left), + translatePredAllLowFuncBody(right))(o.pos, o.info, o.errT) + case i@Implies(left, right) => Implies(And(translateNormal(left, null, null), + translatePrime(left, null, null))(), + translatePredAllLowFuncBody(right) + )(i.pos, i.info, i.errT) + case _ => TrueLit()() + } + Simplifier.simplify(translated) + } + + def translatePredLowFuncOnly(e: Exp, p1: Exp, p2: Exp): Exp = { + val translated: Exp = e match { + case PredicateAccessPredicate(pred, _) => + predLowFuncs(pred.predicateName) match { + case Some(f: Function) => Implies(And(p1, p2)(), + FuncApp(f, pred.args ++ pred.args.map(a => translatePrime(a, p1, p2)))(e.pos, e.info, e.errT) + )() + case None => TrueLit()() + } + case a@And(left, right) => And(translatePredLowFuncOnly(left, p1, p2), + translatePredLowFuncOnly(right, p1, p2))(a.pos, a.info, a.errT) + case o@Or(left, right) => Or(translatePredLowFuncOnly(left, p1, p2), + translatePredLowFuncOnly(right, p1, p2))(o.pos, o.info, o.errT) + case i@Implies(left, right) => Implies(And(translateNormal(left, null, null), + translatePrime(left, null, null))(), + translatePredLowFuncOnly(right, p1, p2) + )(i.pos, i.info, i.errT) + case _ => TrueLit()() + } + Simplifier.simplify(translated) + } + + def translateResourceAccess(ra: ResourceAccess): ResourceAccess = { + ra match { + case FieldAccess(rcv, field) => { + FieldAccess(rcv, field)(ra.pos, ra.info, ra.errT) + } + case PredicateAccess(args, name) => { + PredicateAccess(args, name)(ra.pos, ra.info, ra.errT) + } + case _ => sys.error("Unsupported") + } + } + + def getPredicateLowFunction(predName: String, m: Method): Option[Function] = { + if (allLowMethods.contains(m.name) || preservesLowMethods.contains(m.name)) + predAllLowFuncs(predName) + else + predLowFuncs(predName) + } + + def getPredicateLowFuncExp(predName: String, ctx: TranslationContext): (Option[Function], Exp) = { + val lowFunc = getPredicateLowFunction(predName, ctx.currentMethod) + lazy val act1: Exp = ctx.ctrlVars.activeExecNormal(Some(ctx.p1)) + lazy val act2: Exp = ctx.ctrlVars.activeExecPrime(Some(ctx.p2)) + lazy val isInPreservesLow: Boolean = preservesLowMethods.contains(ctx.currentMethod.name) + lowFunc match { + case Some(f) => + var lhs = And(act1, act2)() + if (isInPreservesLow) { + val allStateLow = translateSIFAss(allVarsAndStateLow(ctx.currentMethod, ctx.currentMethod.formalArgs, + old = !ctx.translatingPrecond, predicateSrc = ctx.currentMethod.pres), ctx) + lhs = And(lhs, allStateLow)() + } + (Some(f), lhs) + case None => (None, TrueLit()()) + } + } + + def simplifyConditions(in: Seq[Exp]): Seq[Exp] = { + val simplified = in.map(e => e.transform{ + case x: Exp if x.isPure => Simplifier.simplify(x) + }) + simplified.filter(e => !e.isInstanceOf[TrueLit]) + } + + def flattenSeqn(in: Seqn): Seqn = { + var newDecls: Seq[Declaration] = in.scopedDecls + val newSS: Seq[Stmt] = in.ss.flatMap({ + case s: Seqn => + val innerFlat = flattenSeqn(s) + newDecls ++= innerFlat.scopedDecls + innerFlat.ss + case x: Stmt => Seq(x) + }) + Seqn(newSS, newDecls)(in.pos, in.info, in.errT) + } + + case class TranslationContext(p1: Exp, p2: Exp, + ctrlVars: MethodControlFlowVars, + currentMethod: Method, + translatingPrecond: Boolean = false + ) {} + + class MethodControlFlowVars(hasRet: Boolean, hasBreak: Boolean, hasCont: Boolean, hasExcept: Boolean, labels: Set[Label]) { + var ret1d, ret2d, break1d, break2d, cont1d, cont2d, except1d, except2d: Option[LocalVarDecl] = None + var ret1r, ret2r, break1r, break2r, cont1r, cont2r, except1r, except2r: Option[LocalVar] = None + val labelRefs1 : ListBuffer[LocalVar] = new ListBuffer[LocalVar]() + val labelDecls1 : ListBuffer[LocalVarDecl] = new ListBuffer[LocalVarDecl]() + val labelRefs2 : ListBuffer[LocalVar] = new ListBuffer[LocalVar]() + val labelDecls2 : ListBuffer[LocalVarDecl] = new ListBuffer[LocalVarDecl]() + val labelNames : mutable.HashMap[String, String] = new mutable.HashMap[String, String]() + + if (hasRet) {val t = getNewBool("ret1"); ret1d = Some(t._1); ret1r = Some(t._2)} + if (hasRet) {val t = getNewBool("ret2"); ret2d = Some(t._1); ret2r = Some(t._2)} + if (hasBreak) {val t = getNewBool("break1"); break1d = Some(t._1); break1r = Some(t._2)} + if (hasBreak) {val t = getNewBool("break2"); break2d = Some(t._1); break2r = Some(t._2)} + if (hasCont) {val t = getNewBool("cont1"); cont1d = Some(t._1); cont1r = Some(t._2)} + if (hasCont) {val t = getNewBool("cont2"); cont2d = Some(t._1); cont2r = Some(t._2)} + if (hasExcept) {val t = getNewBool("except1"); except1d = Some(t._1); except1r = Some(t._2)} + if (hasExcept) {val t = getNewBool("except2"); except2d = Some(t._1); except2r = Some(t._2)} + if (hasRet) primedNames.update(ret1r.get.name, ret2r.get.name) + if (hasBreak) primedNames.update(break1r.get.name, break2r.get.name) + if (hasCont) primedNames.update(cont1r.get.name, cont2r.get.name) + if (hasExcept) primedNames.update(except1r.get.name, except2r.get.name) + for (label <- labels){ + val t1 = getNewBool(label.name + "1") + labelRefs1.append(t1._2) + labelDecls1.append(t1._1) + labelNames.update(label.name, t1._2.name) + val t2 = getNewBool(label.name + "2") + labelRefs2.append(t2._2) + labelDecls2.append(t2._1) + primedNames.update(t1._2.name, t2._2.name) + } + + def declarations(): Seq[LocalVarDecl] = { + Seq(ret1d, ret2d, break1d, break2d, cont1d, cont2d, except1d, except2d) + .filter(v => v.isDefined) + .map(v => v.get) ++ labelDecls1 ++ labelDecls2 + } + + def initAssigns(): Seq[Stmt] = { + (Seq(ret1r, ret2r, break1r, break2r, cont1r, cont2r, except1r, except2r) + .filter(v => v.isDefined) + .map(v => LocalVarAssign(v.get, FalseLit()())())) ++ + (labelRefs1 ++ labelRefs2).map(v => LocalVarAssign(v, FalseLit()())()) + } + + private def conjoinVars(s: Seq[Exp]): Option[Exp] = { + val negations: Seq[Exp] = s.map(x => Not(x)()) + if (negations.nonEmpty) + Some(negations.reduceRight[Exp]((x, y) => And(x, y)())) + else + None + } + + private def activeExecHelper(p: Option[Exp], s: Seq[Exp]): Exp = { + val controls: Option[Exp] = conjoinVars(s) + val expressions: Seq[Exp] = Seq(p, controls).collect({case Some(x) => x}) + expressions match { + case Nil => TrueLit()() + case Seq(head) => head + case _ => expressions.reduceRight[Exp]((x, y) => And(x, y)()) + } + } + + def removeOptions(s: Seq[Option[Exp]]) : Seq[Exp] = { + s.filter(v => v.isDefined).map(v => v.get) + } + + def activeExecNormalExceptLabel(p1: Option[Exp], labelName: String): Exp = { + activeExecHelper(p1, removeOptions(Seq(ret1r, break1r, cont1r, except1r)) ++ labelRefs1.filter(v => !v.name.equals(labelName))) + } + + def activeExecPrimeExceptLabel(p2: Option[Exp], labelNamePrime: String): Exp = { + activeExecHelper(p2, removeOptions(Seq(ret2r, break2r, cont2r, except2r)) ++ labelRefs2.filter(v => !v.name.equals(labelNamePrime))) + } + + def activeExecNormal(p1: Option[Exp]): Exp = { + activeExecHelper(p1, removeOptions(Seq(ret1r, break1r, cont1r, except1r)) ++ labelRefs1) + } + def activeExecPrime(p2: Option[Exp]): Exp = { + activeExecHelper(p2, removeOptions(Seq(ret2r, break2r, cont2r, except2r)) ++ labelRefs2) + } + + def activeExecNoContNormal(p1: Option[Exp]): Exp = { + activeExecHelper(p1, removeOptions(Seq(ret1r, break1r, except1r)) ++ labelRefs1) + } + def activeExecNoContPrime(p2: Option[Exp]): Exp = { + activeExecHelper(p2, removeOptions(Seq(ret2r, break2r, except2r)) ++ labelRefs2) + } + } + + object EmptyControlFlowVars { + def apply(): MethodControlFlowVars = { + new MethodControlFlowVars(false, false, false, false, Set()) + } + } +} + +object SIFExtendedTransformer extends SIFExtendedTransformer diff --git a/docs/dev-guide/src/config/flags.md b/docs/dev-guide/src/config/flags.md index b3ef1913dc9..f58005cdcdb 100644 --- a/docs/dev-guide/src/config/flags.md +++ b/docs/dev-guide/src/config/flags.md @@ -59,6 +59,7 @@ | [`SERVER_ADDRESS`](#server_address) | `Option` | `None` | A | | [`SERVER_MAX_CONCURRENCY`](#server_max_concurrency) | `Option` | `None` | A | | [`SERVER_MAX_STORED_VERIFIERS`](#server_max_stored_verifiers) | `Option` | `None` | A | +| [`SIF`](#sif) | `bool` | `false`| A | | [`SIMPLIFY_ENCODING`](#simplify_encoding) | `bool` | `true` | A | | [`SKIP_UNSUPPORTED_FEATURES`](#skip_unsupported_features) | `bool` | `false` | A | | [`SMT_QI_BOUND_GLOBAL`](#smt_qi_bound_global) | `Option` | `None` | A | @@ -373,6 +374,10 @@ Maximum amount of instantiated Viper verifiers the server will keep around for r > **Note:** This does _not_ limit how many verification requests the server handles concurrently, only the size of what is essentially its verifier cache. +## `SIF` + +When enabled, check the program for secure information flow. + ## `SIMPLIFY_ENCODING` When enabled, the encoded program is simplified before it is passed to the Viper backend. diff --git a/prusti-common/src/vir/optimizations/folding/expressions.rs b/prusti-common/src/vir/optimizations/folding/expressions.rs index 169a6c01773..d0215550d3a 100644 --- a/prusti-common/src/vir/optimizations/folding/expressions.rs +++ b/prusti-common/src/vir/optimizations/folding/expressions.rs @@ -399,6 +399,13 @@ impl ast::FallibleExprFolder for ExprOptimizer { Err(()) } + fn fallible_fold_low( + &mut self, + _expr: vir::polymorphic::Low, + ) -> Result { + Err(()) + } + fn fallible_fold_predicate_access_predicate( &mut self, _predicate_access_predicate: ast::PredicateAccessPredicate, diff --git a/prusti-common/src/vir/optimizations/functions/simplifier.rs b/prusti-common/src/vir/optimizations/functions/simplifier.rs index e2406e87059..41b5a6f8626 100644 --- a/prusti-common/src/vir/optimizations/functions/simplifier.rs +++ b/prusti-common/src/vir/optimizations/functions/simplifier.rs @@ -19,8 +19,9 @@ impl Simplifier for ast::Function { /// #[tracing::instrument(level = "debug", skip(self), fields(self = %self), ret(Display))] fn simplify(mut self) -> Self { - let new_body = self.body.map(|b| b.simplify()); - self.body = new_body; + self.body = self.body.map(|b| b.simplify()); + self.posts = self.posts.into_iter().map(|p| p.simplify()).collect(); + self.pres = self.pres.into_iter().map(|p| p.simplify()).collect(); self } } @@ -29,7 +30,8 @@ impl Simplifier for ast::Expr { #[must_use] fn simplify(self) -> Self { let mut folder = ExprSimplifier {}; - folder.fold(self) + let res = folder.fold(self); + res } } @@ -44,13 +46,14 @@ impl ExprSimplifier { argument: box ast::Expr::Const(ast::ConstExpr { value: ast::Const::Bool(b), - .. + position: inner_pos, }), position: pos, }) => ast::Expr::Const(ast::ConstExpr { value: ast::Const::Bool(!b), position: pos, - }), + }) + .set_default_pos(inner_pos), ast::Expr::UnaryOp(ast::UnaryOp { op_kind: ast::UnaryOpKind::Not, argument: @@ -58,41 +61,25 @@ impl ExprSimplifier { op_kind: ast::BinaryOpKind::EqCmp, box left, box right, - .. + position: inner_pos, }), position: pos, - }) => ast::Expr::BinOp(ast::BinOp { + }) if !matches!(left.get_type(), ast::Type::Float(_)) => ast::Expr::BinOp(ast::BinOp { op_kind: ast::BinaryOpKind::NeCmp, left: box left, right: box right, position: pos, - }), - ast::Expr::BinOp(ast::BinOp { - op_kind: ast::BinaryOpKind::And, - left: - box ast::Expr::Const(ast::ConstExpr { - value: ast::Const::Bool(b1), - .. - }), - right: - box ast::Expr::Const(ast::ConstExpr { - value: ast::Const::Bool(b2), - .. - }), - position: pos, - }) => ast::Expr::Const(ast::ConstExpr { - value: ast::Const::Bool(b1 && b2), - position: pos, - }), + }) + .set_default_pos(inner_pos), ast::Expr::BinOp(ast::BinOp { op_kind: ast::BinaryOpKind::And, left: box ast::Expr::Const(ast::ConstExpr { value: ast::Const::Bool(b), - .. + position: inner_pos, }), right: box conjunct, - .. + position: pos, }) | ast::Expr::BinOp(ast::BinOp { op_kind: ast::BinaryOpKind::And, @@ -100,25 +87,24 @@ impl ExprSimplifier { right: box ast::Expr::Const(ast::ConstExpr { value: ast::Const::Bool(b), - .. + position: inner_pos, }), - .. - }) => { - if b { - conjunct - } else { - false.into() - } + position: pos, + }) => if b { + conjunct + } else { + Into::::into(false).set_pos(inner_pos) } + .set_default_pos(pos), ast::Expr::BinOp(ast::BinOp { op_kind: ast::BinaryOpKind::Or, left: box ast::Expr::Const(ast::ConstExpr { value: ast::Const::Bool(b), - .. + position: inner_pos, }), right: box disjunct, - .. + position: pos, }) | ast::Expr::BinOp(ast::BinOp { op_kind: ast::BinaryOpKind::Or, @@ -126,52 +112,49 @@ impl ExprSimplifier { right: box ast::Expr::Const(ast::ConstExpr { value: ast::Const::Bool(b), - .. + position: inner_pos, }), - .. - }) => { - if b { - true.into() - } else { - disjunct - } + position: pos, + }) => if b { + Into::::into(true).set_pos(inner_pos) + } else { + disjunct } + .set_default_pos(pos), ast::Expr::BinOp(ast::BinOp { op_kind: ast::BinaryOpKind::Implies, left: guard, right: box ast::Expr::Const(ast::ConstExpr { value: ast::Const::Bool(b), - .. + position: inner_pos, }), position: pos, - }) => { - if b { - true.into() - } else { - ast::Expr::UnaryOp(ast::UnaryOp { - op_kind: ast::UnaryOpKind::Not, - argument: guard, - position: pos, - }) - } + }) => if b { + Into::::into(true).set_pos(pos) + } else { + ast::Expr::UnaryOp(ast::UnaryOp { + op_kind: ast::UnaryOpKind::Not, + argument: guard, + position: pos, + }) } + .set_default_pos(inner_pos), ast::Expr::BinOp(ast::BinOp { op_kind: ast::BinaryOpKind::Implies, left: box ast::Expr::Const(ast::ConstExpr { value: ast::Const::Bool(b), - .. + position: inner_pos, }), right: box body, - .. - }) => { - if b { - body - } else { - true.into() - } + position: pos, + }) => if b { + body + } else { + Into::::into(true).set_pos(inner_pos) } + .set_default_pos(pos), ast::Expr::BinOp(ast::BinOp { op_kind: ast::BinaryOpKind::And, left: box op1, diff --git a/prusti-common/src/vir/optimizations/methods/mod.rs b/prusti-common/src/vir/optimizations/methods/mod.rs index bb18d270283..ede9b9f1c28 100644 --- a/prusti-common/src/vir/optimizations/methods/mod.rs +++ b/prusti-common/src/vir/optimizations/methods/mod.rs @@ -13,6 +13,7 @@ mod purifier; mod quantifier_fixer; mod unfolding_fixer; mod var_remover; +mod simplifier; use super::log_method; use crate::{config::Optimizations, vir::polymorphic_vir::cfg::CfgMethod}; @@ -20,7 +21,7 @@ use crate::{config::Optimizations, vir::polymorphic_vir::cfg::CfgMethod}; use self::{ assert_remover::remove_trivial_assertions, cfg_cleaner::clean_cfg, empty_if_remover::remove_empty_if, purifier::purify_vars, quantifier_fixer::fix_quantifiers, - unfolding_fixer::fix_unfoldings, var_remover::remove_unused_vars, + simplifier::simplify_exprs, unfolding_fixer::fix_unfoldings, var_remover::remove_unused_vars, }; #[allow(clippy::let_and_return)] @@ -58,6 +59,7 @@ pub fn optimize_method_encoding( }; let cfg = apply!(fix_unfoldings, cfg); let cfg = apply!(fix_quantifiers, cfg); + let cfg = apply!(simplify_exprs, cfg); let cfg = apply!(remove_empty_if, cfg); let cfg = apply!(remove_unused_vars, cfg); let cfg = apply!(remove_trivial_assertions, cfg); diff --git a/prusti-common/src/vir/optimizations/methods/simplifier.rs b/prusti-common/src/vir/optimizations/methods/simplifier.rs new file mode 100644 index 00000000000..16c3ab4cbc6 --- /dev/null +++ b/prusti-common/src/vir/optimizations/methods/simplifier.rs @@ -0,0 +1,45 @@ +// © 2022, ETH Zurich +// +// This Source Code Form is subject to the terms of the Mozilla Public +// License, v. 2.0. If a copy of the MPL was not distributed with this +// file, You can obtain one at http://mozilla.org/MPL/2.0/. + +use log::debug; +use std::mem; +use vir::polymorphic::*; + +use crate::vir::optimizations::functions::Simplifier; + +/// This optimization simplifies all expressions in a method. +/// It is required when resource access predicates appear on the RHS of +/// implications as implications are transformed into ors which need to be +/// transformed back into implications otherwise, we would have an impure +/// expression in ors which is disallowed in Viper. + +pub fn simplify_exprs(mut cfg: CfgMethod) -> CfgMethod { + debug!("Simplifying exprs in {}", cfg.name()); + let mut sentinel_stmt = Stmt::comment("moved out stmt"); + for block in &mut cfg.basic_blocks { + for stmt in &mut block.stmts { + mem::swap(&mut sentinel_stmt, stmt); + sentinel_stmt = sentinel_stmt.simplify(); + mem::swap(&mut sentinel_stmt, stmt); + } + } + cfg +} + +struct StmtSimplifier; + +impl Simplifier for Stmt { + fn simplify(self) -> Self { + let mut folder = StmtSimplifier; + folder.fold(self) + } +} + +impl StmtFolder for StmtSimplifier { + fn fold_expr(&mut self, expr: Expr) -> Expr { + expr.simplify() + } +} diff --git a/prusti-common/src/vir/to_viper.rs b/prusti-common/src/vir/to_viper.rs index 034628c5696..22b156e4d4f 100644 --- a/prusti-common/src/vir/to_viper.rs +++ b/prusti-common/src/vir/to_viper.rs @@ -750,6 +750,10 @@ impl<'v> ToViper<'v, viper::Expr<'v>> for Expr { ast.int_to_backend_bv(size, base.to_viper(context, ast)) } }, + Expr::Low(expr, position) => { + ast.low_with_pos(expr.to_viper(context, ast), position.to_viper(context, ast)) + } + Expr::LowEvent => ast.low_event(), }; if config::simplify_encoding() { ast.simplified_expression(expr) @@ -1055,9 +1059,11 @@ impl<'a, 'v> ToViper<'v, viper::Method<'v>> for &'a CfgMethod { // self.convert_basic_block_path(path, ast, &mut blocks_ast, &mut declarations); } else { // Sort blocks by label, except for the first block - let mut blocks: Vec<_> = self.basic_blocks.iter().enumerate().skip(1).collect(); - blocks.sort_by_key(|(index, _)| index_to_label(self.basic_blocks_labels(), *index)); - blocks.insert(0, (0, &self.basic_blocks[0])); + // let mut blocks: Vec<_> = self.basic_blocks.iter().enumerate().skip(1).collect(); + // blocks.sort_by_key(|(index, _)| index_to_label(self.basic_blocks_labels(), *index)); + // blocks.insert(0, (0, &self.basic_blocks[0])); + + let blocks = topologicaly_sort_blocks(&self.basic_blocks); for (index, block) in blocks.into_iter() { blocks_ast.push(block_to_viper( @@ -1094,6 +1100,55 @@ impl<'a, 'v> ToViper<'v, viper::Method<'v>> for &'a CfgMethod { } } +fn topologicaly_sort_blocks(blocks: &[CfgBlock]) -> Vec<(usize, &CfgBlock)> { + let mut added_blocks: Vec = (0..blocks.len()).map(|_| false).collect(); + let mut traversed_blocks: Vec = (0..blocks.len()).map(|_| false).collect(); + let mut res = Vec::new(); + + topo_rec( + blocks, + 0, + &mut res, + &mut added_blocks, + &mut traversed_blocks, + ); + + res.into_iter() + .rev() + .map(|idx| (idx, &blocks[idx])) + .collect() +} + +fn topo_rec( + blocks: &[CfgBlock], + curr: usize, + res: &mut Vec, + added_blocks: &mut [bool], + traversed_blocks: &mut [bool], +) { + debug_assert!(!traversed_blocks[curr], "Found a back edge!"); + if added_blocks[curr] { + return; + } + + traversed_blocks[curr] = true; + match &blocks[curr].successor { + Successor::Undefined => unreachable!("Undefined edge!"), + Successor::Return => (), + Successor::Goto(to) => topo_rec(blocks, to.index(), res, added_blocks, traversed_blocks), + Successor::GotoSwitch(tos, default) => { + for (_, to) in tos.iter() { + topo_rec(blocks, to.index(), res, added_blocks, traversed_blocks); + } + topo_rec(blocks, default.index(), res, added_blocks, traversed_blocks); + } + } + traversed_blocks[curr] = false; + + res.push(curr); + added_blocks[curr] = true; +} + fn cfg_method_convert_basic_block_path<'v>( cfg_method: &CfgMethod, mut path: Vec, diff --git a/prusti-contracts/prusti-contracts/src/lib.rs b/prusti-contracts/prusti-contracts/src/lib.rs index b7e40364743..427f5a69878 100644 --- a/prusti-contracts/prusti-contracts/src/lib.rs +++ b/prusti-contracts/prusti-contracts/src/lib.rs @@ -322,6 +322,14 @@ mod private { panic!() } } + + /// Declassify the given expression, making it low + #[macro_export] + macro_rules! declassify { + ($base:expr) => { + $crate::prusti_assume!($crate::low($base)) + }; + } } /// This function is used to evaluate an expression in the context just @@ -365,4 +373,14 @@ pub fn snapshot_equality(_l: T, _r: T) -> bool { true } +/// Mark the given expression as low security level +pub fn low(_t: T) -> bool { + true +} + +/// Check that all execusion must either all reach this point or none +pub fn low_event() -> bool { + true +} + pub use private::*; diff --git a/prusti-server/src/backend.rs b/prusti-server/src/backend.rs index 5a36ad6c399..d8867df7456 100644 --- a/prusti-server/src/backend.rs +++ b/prusti-server/src/backend.rs @@ -21,7 +21,22 @@ impl<'a> Backend<'a> { ast_utils.with_local_frame(16, || { let ast_factory = context.new_ast_factory(); - let viper_program = program.to_viper(LoweringContext::default(), &ast_factory); + let mut viper_program = + program.to_viper(LoweringContext::default(), &ast_factory); + + if config::sif() { + if config::dump_viper_program() { + stopwatch.start_next("dumping viper program before sif transformation"); + dump_viper_program( + &ast_utils, + viper_program, + &format!("{}_before_sif", program.get_name_with_check_mode()), + ); + } + let sif_transformer = context.new_sif_transformer(); + stopwatch.start_next("sif translation"); + viper_program = sif_transformer.sif_transformation(viper_program); + } if config::dump_viper_program() { stopwatch.start_next("dumping viper program"); diff --git a/prusti-tests/tests/parse/pass/sif/declassify.rs b/prusti-tests/tests/parse/pass/sif/declassify.rs new file mode 100644 index 00000000000..67b9669d09c --- /dev/null +++ b/prusti-tests/tests/parse/pass/sif/declassify.rs @@ -0,0 +1,17 @@ +// compile-flags: -Psif=true + +use prusti_contracts::*; + +fn declassify_int(i: i32) { + declassify!(i); +} + +fn declassify_bool(b: bool) { + declassify!(b); +} + +fn declassify_mod(i: i32) { + declassify!(i % 2); +} + +fn main() {} diff --git a/prusti-tests/tests/parse/pass/sif/low.rs b/prusti-tests/tests/parse/pass/sif/low.rs new file mode 100644 index 00000000000..5743fc4c597 --- /dev/null +++ b/prusti-tests/tests/parse/pass/sif/low.rs @@ -0,0 +1,30 @@ +// compile-flags: -Psif=true + +use prusti_contracts::*; + +#[requires(low(n))] +fn foo(n: i32) -> i32 { + n +} + +#[ensures(low(result))] +fn bar() -> i32 { + 42 +} + +fn baz() -> i32 { + let x = 1; + let y = 2; + prusti_assert!(low(x + y)); + x + y +} + +fn l(n: u32) { + let mut i = 0; + while i < n { + body_invariant!(low_event()); + i += 1; + } +} + +fn main() {} diff --git a/prusti-tests/tests/verify/fail/sif/arrays.rs b/prusti-tests/tests/verify/fail/sif/arrays.rs new file mode 100644 index 00000000000..5ad947a74ca --- /dev/null +++ b/prusti-tests/tests/verify/fail/sif/arrays.rs @@ -0,0 +1,69 @@ +// compile-flags: -Psif=true -Pserver_address=MOCK +// The sif flag is used in the server which, during the compiletest is only spawned with the default config. +// So we need to start a new server with this test config to make it work. + +use prusti_contracts::*; + +#[requires(forall(|i: usize| 0 <= i && i < ts.len() ==> low(ts[i])))] +#[ensures(low(result))] +fn sum_array_low(ts: &[i32]) -> i32 { + let mut i = 0; + let mut res = 0; + + while i < ts.len() { + res += ts[i]; + i += 1; + } + + res +} + +#[ensures(forall(|i: usize| 0 <= i && i < ts.len() ==> low(ts[i])) ==> low(result))] +fn sum_array(ts: &[i32]) -> i32 { + let mut i = 0; + let mut res = 0; + + while i < ts.len() { + res += ts[i]; + i += 1; + } + + res +} + +#[ensures(low(result))] +fn produce_low() -> i32 { + 42 +} + +#[ensures(forall(|i: usize| 0 <= i && i < array.len() ==> low(array[i])))] +fn fill_with_low(array: &mut [i32]) { + let mut i = 0; + while i < array.len() { + array[i] = produce_low(); + i += 1; + } +} + +fn produce_high() -> i32 { + 12 +} + +fn foo() { + let mut array = [1, 2, 3, 4, 5]; + fill_with_low(&mut array); + + array[2] = produce_high(); + + let res2 = sum_array(&array); + prusti_assert!(low(res2)); //~ERROR the asserted expression might not hold +} + +fn main() { + let mut array = [1, 2, 3, 4, 5]; + fill_with_low(&mut array); + + array[2] = produce_high(); + + let res1 = sum_array_low(&array); //~ERROR precondition might not hold +} diff --git a/prusti-tests/tests/verify/fail/sif/functions.rs b/prusti-tests/tests/verify/fail/sif/functions.rs new file mode 100644 index 00000000000..60062136692 --- /dev/null +++ b/prusti-tests/tests/verify/fail/sif/functions.rs @@ -0,0 +1,24 @@ +// compile-flags: -Psif=true -Pserver_address=MOCK +// The sif flag is used in the server which, during the compiletest is only spawned with the default config. +// So we need to start a new server with this test config to make it work. + +use prusti_contracts::*; + +#[requires(low(i))] +#[ensures(low(result))] +fn foo(i: i32) -> i32 { + i + 42 +} + +#[requires(low(n))] +fn bar(n: i32) {} + +fn baz_safe() -> i32 { + 42 +} + +fn main() { + let i = baz_safe(); + let j = foo(i); //~ERROR precondition might not hold + bar(j); +} diff --git a/prusti-tests/tests/verify/fail/sif/ifs.rs b/prusti-tests/tests/verify/fail/sif/ifs.rs new file mode 100644 index 00000000000..37feaf61a6d --- /dev/null +++ b/prusti-tests/tests/verify/fail/sif/ifs.rs @@ -0,0 +1,42 @@ +// compile-flags: -Psif=true -Pserver_address=MOCK +// The sif flag is used in the server which, during the compiletest is only spawned with the default config. +// So we need to start a new server with this test config to make it work. + +use prusti_contracts::*; + +#[requires(low(l))] +#[ensures(!b ==> low(result))] +fn choose(b: bool, h: i32, l: i32) -> i32 { + if b { + h + } else { + l + } +} + +#[requires(low(l1) && low(l2))] +#[ensures(low(result))] //~EROOR postcondition might not hold. +fn choose2(b: bool, l1: i32, l2: i32) -> i32 { + if b { + l1 + } else { + l2 + } +} + +fn produce_high() -> i32 { + 42 +} + +#[ensures(low(result))] +fn produce_low() -> i32 { + 12 +} + +fn main() { + let l = produce_low(); + let h = produce_high(); + + let res = choose(true, h, l); + prusti_assert!(low(res)); //~ERROR the asserted expression might not hold +} diff --git a/prusti-tests/tests/verify/fail/sif/impl.rs b/prusti-tests/tests/verify/fail/sif/impl.rs new file mode 100644 index 00000000000..764a19a291a --- /dev/null +++ b/prusti-tests/tests/verify/fail/sif/impl.rs @@ -0,0 +1,58 @@ +// compile-flags: -Psif=true -Pserver_address=MOCK +// The sif flag is used in the server which, during the compiletest is only spawned with the default config. +// So we need to start a new server with this test config to make it work. + +use prusti_contracts::*; + +#[ensures(low(i1) && low(i2) ==> low(result))] +fn add(i1: i32, i2: i32) -> i32 { + i1 + i2 +} + +#[ensures(low(result))] +fn produce_low() -> i32 { + 40 +} + +fn produce_high() -> i32 { + 2 +} + +#[requires(low(i))] +fn requires_low(i: i32) {} + +fn requires_high(_i: i32) {} + +fn add_low_high() { + let i1_low = produce_low(); + let i2_high = produce_high(); + + requires_low(add(i1_low, i2_high)); //~ERROR precondition might not hold. +} + +fn add_high_low() { + let i1_high = produce_high(); + let i2_low = produce_low(); + + requires_low(add(i1_high, i2_low)); //~ERROR precondition might not hold. +} + +fn add_high_high() { + let i1_high = produce_high(); + let i2_high = produce_high(); + + requires_low(add(i1_high, i2_high)); //~ERROR precondition might not hold. +} + +fn main() { + let i1_low = produce_low(); + let i2_low = produce_low(); + let i1_high = produce_high(); + let i2_high = produce_high(); + + requires_low(add(i1_low, i2_low)); + requires_high(add(i1_low, i2_low)); // low values can be used in as high + requires_high(add(i1_low, i2_high)); + requires_high(add(i1_high, i2_low)); + requires_high(add(i1_high, i2_high)); +} diff --git a/prusti-tests/tests/verify/fail/sif/loops.rs b/prusti-tests/tests/verify/fail/sif/loops.rs new file mode 100644 index 00000000000..d6c7bea8c1d --- /dev/null +++ b/prusti-tests/tests/verify/fail/sif/loops.rs @@ -0,0 +1,48 @@ +// compile-flags: -Psif=true -Pserver_address=MOCK +// The sif flag is used in the server which, during the compiletest is only spawned with the default config. +// So we need to start a new server with this test config to make it work. + +use prusti_contracts::*; + +#[requires(low(x))] +#[ensures(low(result))] //~ERROR postcondition might not hold. +fn loop_max(x: i32, y: i32) -> i32 { + let mut res = x; + while res < y { + body_invariant!(low(res)); + res += 1; + } + res +} + +#[requires(low(y))] +#[ensures(low(result))] //~ERROR postcondition might not hold. +fn loop_max2(x: i32, y: i32) -> i32 { + let mut res = x; + while res < y { + res += 1; + } + res +} + +#[requires(low(y))] +#[ensures(low(result))] +fn loop_max3(x: i32, y: i32) -> i32 { + let mut res = x; + while res < y { + body_invariant!(low(res)); //~ERROR loop invariant might not hold in the first loop iteration. + res += 1; + } + res +} + +#[ensures(low(result))] //~ERROR postcondition might not hold. +fn test(n: u32) -> u32 { + let mut res = 0; + while res < n { + res += 1; + } + res +} + +fn main() {} diff --git a/prusti-tests/tests/verify/fail/sif/loops_encoding.rs b/prusti-tests/tests/verify/fail/sif/loops_encoding.rs new file mode 100644 index 00000000000..07a0c802739 --- /dev/null +++ b/prusti-tests/tests/verify/fail/sif/loops_encoding.rs @@ -0,0 +1,23 @@ +// compile-flags: -Psif=true -Pserver_address=MOCK +// The sif flag is used in the server which, during the compiletest is only spawned with the default config. +// So we need to start a new server with this test config to make it work. + +use prusti_contracts::*; + +#[trusted] +#[requires(low(i))] +#[requires(low_event())] +fn print(i: u32) {} + +#[requires(low_event())] +fn foo(y: u32) { + let mut i = 0; + while i < y + 1 { + body_invariant!(low(i)); //~ERROR loop invariant might not hold after a loop iteration + i += 1; + } + print(i); + print(y); +} + +fn main() {} diff --git a/prusti-tests/tests/verify/fail/sif/low_event.rs b/prusti-tests/tests/verify/fail/sif/low_event.rs new file mode 100644 index 00000000000..586d85caebc --- /dev/null +++ b/prusti-tests/tests/verify/fail/sif/low_event.rs @@ -0,0 +1,19 @@ +// compile-flags: -Psif=true -Pserver_address=MOCK +// The sif flag is used in the server which, during the compiletest is only spawned with the default config. +// So we need to start a new server with this test config to make it work. + +use prusti_contracts::*; + +#[trusted] +#[requires(low(i) && low_event())] +fn print_int(i: i32) {} + +fn foo(b: bool) { + if b { + print_int(0); + } else { + print_int(1); + } +} + +fn main() {} diff --git a/prusti-tests/tests/verify/pass/sif/arrays.rs b/prusti-tests/tests/verify/pass/sif/arrays.rs new file mode 100644 index 00000000000..44afae72bf7 --- /dev/null +++ b/prusti-tests/tests/verify/pass/sif/arrays.rs @@ -0,0 +1,88 @@ +// compile-flags: -Psif=true -Pserver_address=MOCK +// The sif flag is used in the server which, during the compiletest is only spawned with the default config. +// So we need to start a new server with this test config to make it work. + +use prusti_contracts::*; + +// #[requires(low(ts))] +// #[ensures(low(result))] +// fn sum_array_low(ts: &[i32]) -> i32 { +// let mut i = 0; +// let mut res = 0; + +// while i < ts.len() { +// body_invariant!(low(i)); +// body_invariant!(low(res)); +// res += ts[i]; +// i += 1; +// } + +// res +// } + +// #[requires(forall(|i: usize| low(ts[i])))] +// #[ensures(low(result))] +// fn sum_array_low2(ts: &[i32]) -> i32 { +// let mut i = 0; +// let mut res = 0; + +// while i < ts.len() { +// body_invariant!(low(i)); +// body_invariant!(low(res)); +// res += ts[i]; +// i += 1; +// } + +// res +// } + +// #[ensures(low(ts) ==> low(result))] +// fn sum_array(ts: &[i32]) -> i32 { +// let mut i = 0; +// let mut res = 0; + +// while i < ts.len() { +// body_invariant!(low(i)); +// body_invariant!(low(ts) ==> low(res)); +// res += ts[i]; +// i += 1; +// } + +// res +// } + +#[trusted] +#[ensures(low(result))] +fn produce_low() -> i32 { + todo!() +} + +// #[requires(0 <= i && i < array.len())] +// #[ensures(low(array[i]))] +// fn replace_low(array: &mut [i32], i: usize) { +// array[i] = produce_low(); +// } + +#[requires(low(array.len()))] +#[ensures(forall(|i: usize| i < array.len() ==> low(array[i])))] +fn fill_with_low(array: &mut [i32]) { + let mut i = 0; + while i < array.len() { + body_invariant!(low(i)); + body_invariant!(forall(|j: usize| 0 <= j && j < i ==> low(array[j]))); + array[i] = produce_low(); + // replace_low(array, i); + i += 1; + } +} + +fn main() { + // let mut low_array = [1, 2, 3, 4, 5]; + // fill_with_low(&mut low_array); + + // let res1 = sum_array_low(&low_array); + // prusti_assert!(low(res1)); + + // let res2 = sum_array(&low_array); + // prusti_assert!(low(res2)); +} diff --git a/prusti-tests/tests/verify/pass/sif/examples/banerjee.rs b/prusti-tests/tests/verify/pass/sif/examples/banerjee.rs new file mode 100644 index 00000000000..c6eea9b21d1 --- /dev/null +++ b/prusti-tests/tests/verify/pass/sif/examples/banerjee.rs @@ -0,0 +1,66 @@ +// compile-flags: -Psif=true -Pserver_address=MOCK + +// Example from "Secure Information Flow and Pointer Confinement in a Java-like Language" +// A. Banerjee and D. A. Naumann +// CSFW 2002 + +use prusti_contracts::*; + +// ignore-test: rust's borrow checker makes this example hard to translate without too much modification + +// HIV is high, name is low +struct Patient { + name: String, + hiv: String, +} + +impl Patient { + fn get_name(&self) -> &str { + &self.name + } + + #[requires(low(name))] + #[ensures(low(&self.name))] + fn set_name(&mut self, name: String) { + self.name = name; + } + + fn get_hiv(&self) -> &str { + &self.hiv + } + + fn set_hiv(&mut self, hiv: String) { + self.hiv = hiv; + } +} + +#[trusted] +#[ensures(low(result))] +#[ensures(low(&result.name))] +fn read_file() -> Patient { + Patient { + name: String::new(), + hiv: String::new(), + } +} + +// the result is high +#[trusted] +fn read_from_trusted_channel() -> String { + String::new() +} + +fn main() { + let lbuf: &str; + let mut hbuf: &str; + let mut lp = read_file(); + let mut xp = Patient { + name: String::new(), + hiv: String::new(), + }; + lbuf = lp.get_name(); + hbuf = lp.get_name(); + xp.set_name(lbuf); + hbuf = read_from_trusted_channel(); + xp.set_hiv(hbuf); +} diff --git a/prusti-tests/tests/verify/pass/sif/examples/constanzo.rs b/prusti-tests/tests/verify/pass/sif/examples/constanzo.rs new file mode 100644 index 00000000000..62fd2bc0f8a --- /dev/null +++ b/prusti-tests/tests/verify/pass/sif/examples/constanzo.rs @@ -0,0 +1,25 @@ +// compile-flags: -Psif=true -Pserver_address=MOCK + +// Example from "A Separation Logic for Enforcing Declarative Information Flow Control Policies" +// D. Costanzo and Z. Shao +// POST 2014 + +use prusti_contracts::*; + +#[trusted] +#[requires(low(n))] +fn print(n: usize) {} + +#[requires(forall(|i: usize| 0 <= i && i < schedule.len() ==> (schedule[i] == 0 ==> low(i))))] +fn print_availibility(schedule: &[u32]) { + let mut i = 0; + while i < schedule.len() { + let x = schedule[i]; + if x == 0 { + print(i); + } + i += 1; + } +} + +fn main() {} diff --git a/prusti-tests/tests/verify/pass/sif/functions.rs b/prusti-tests/tests/verify/pass/sif/functions.rs new file mode 100644 index 00000000000..986bfc83033 --- /dev/null +++ b/prusti-tests/tests/verify/pass/sif/functions.rs @@ -0,0 +1,25 @@ +// compile-flags: -Psif=true -Pserver_address=MOCK +// The sif flag is used in the server which, during the compiletest is only spawned with the default config. +// So we need to start a new server with this test config to make it work. + +use prusti_contracts::*; + +#[requires(low(i))] +#[ensures(low(result))] +fn foo(i: i32) -> i32 { + i + 42 +} + +#[requires(low(n))] +fn bar(n: i32) {} + +#[ensures(low(result))] +fn baz() -> i32 { + 42 +} + +fn main() { + let i = baz(); + let j = foo(i); + bar(j); +} diff --git a/prusti-tests/tests/verify/pass/sif/ifs.rs b/prusti-tests/tests/verify/pass/sif/ifs.rs new file mode 100644 index 00000000000..7411fa8f978 --- /dev/null +++ b/prusti-tests/tests/verify/pass/sif/ifs.rs @@ -0,0 +1,48 @@ +// compile-flags: -Psif=true -Pserver_address=MOCK +// The sif flag is used in the server which, during the compiletest is only spawned with the default config. +// So we need to start a new server with this test config to make it work. + +use prusti_contracts::*; + +#[requires(low(l))] +#[ensures(!b ==> low(result))] +fn choose(b: bool, h: i32, l: i32) -> i32 { + if b { + h + } else { + l + } +} + +#[requires(low(l))] +#[ensures(low(result))] +fn foo(b: bool, l: i32) -> i32 { + let res; + if b { + res = l + 3; + } else { + res = l - 3; + } + if b { + res - 3 + } else { + res + 3 + } +} + +fn produce_high() -> i32 { + 42 +} + +#[ensures(low(result))] +fn produce_low() -> i32 { + 12 +} + +fn main() { + let l = produce_low(); + let h = produce_high(); + + let res = choose(false, h, l); + prusti_assert!(low(res)); +} diff --git a/prusti-tests/tests/verify/pass/sif/impl.rs b/prusti-tests/tests/verify/pass/sif/impl.rs new file mode 100644 index 00000000000..5c3648a2ea7 --- /dev/null +++ b/prusti-tests/tests/verify/pass/sif/impl.rs @@ -0,0 +1,37 @@ +// compile-flags: -Psif=true -Pserver_address=MOCK +// The sif flag is used in the server which, during the compiletest is only spawned with the default config. +// So we need to start a new server with this test config to make it work. + +use prusti_contracts::*; + +#[ensures(low(i1) && low(i2) ==> low(result))] +fn add(i1: i32, i2: i32) -> i32 { + i1 + i2 +} + +#[ensures(low(result))] +fn produce_low() -> i32 { + 40 +} + +fn produce_high() -> i32 { + 2 +} + +#[requires(low(i))] +fn requires_low(i: i32) {} + +fn requires_high(_i: i32) {} + +fn main() { + let i1_low = produce_low(); + let i2_low = produce_low(); + let i1_high = produce_high(); + let i2_high = produce_high(); + + requires_low(add(i1_low, i2_low)); + requires_high(add(i1_low, i2_low)); // low values can be used in as high + requires_high(add(i1_low, i2_high)); + requires_high(add(i1_high, i2_low)); + requires_high(add(i1_high, i2_high)); +} diff --git a/prusti-tests/tests/verify/pass/sif/loop_encoding.rs b/prusti-tests/tests/verify/pass/sif/loop_encoding.rs new file mode 100644 index 00000000000..6142e741f2c --- /dev/null +++ b/prusti-tests/tests/verify/pass/sif/loop_encoding.rs @@ -0,0 +1,30 @@ +// compile-flags: -Psif=true -Pserver_address=MOCK +// The sif flag is used in the server which, during the compiletest is only spawned with the default config. +// So we need to start a new server with this test config to make it work. + +use prusti_contracts::*; + +#[trusted] +#[requires(low(i))] +fn print(i: u32) {} + +#[requires(low(y))] +#[requires(low_event())] +fn foo(y: u32) { + let mut i = 0; + while i < y { + body_invariant!(low(i)); + i += 1; + } + print(i); +} + +fn bar(x: u32) -> u32 { + let mut i = 0; + while i < x { + i += 1; + } + i +} + +fn main() {} diff --git a/prusti-tests/tests/verify/pass/sif/loops.rs b/prusti-tests/tests/verify/pass/sif/loops.rs new file mode 100644 index 00000000000..afaaa4166a3 --- /dev/null +++ b/prusti-tests/tests/verify/pass/sif/loops.rs @@ -0,0 +1,38 @@ +// compile-flags: -Psif=true -Pserver_address=MOCK +// The sif flag is used in the server which, during the compiletest is only spawned with the default config. +// So we need to start a new server with this test config to make it work. + +use prusti_contracts::*; + +#[requires(low(x) && low(y))] +#[ensures(result == x || result == y)] +#[ensures(low(result))] +fn loop_max(x: i32, y: i32) -> i32 { + let mut res = x; + while res < y { + res += 1; + } + res +} + +#[ensures(low(result))] +fn test(n: u32) -> u32 { + let mut res = 0; + while res < n { + res += 1; + } + declassify!(res); + res +} + +#[ensures(low(n))] +fn l(n: u32) { + let mut i = 0; + while i < n { + body_invariant!(low_event()); + body_invariant!(low(i)); + i += 1; + } +} + +fn main() {} diff --git a/prusti-tests/tests/verify/pass/sif/reborrow.rs b/prusti-tests/tests/verify/pass/sif/reborrow.rs new file mode 100644 index 00000000000..0964ae29c91 --- /dev/null +++ b/prusti-tests/tests/verify/pass/sif/reborrow.rs @@ -0,0 +1,39 @@ +// compile-flags: -Psif=true -Pserver_address=MOCK +// The sif flag is used in the server which, during the compiletest is only spawned with the default config. +// So we need to start a new server with this test config to make it work. + +use prusti_contracts::*; + +struct Point { + x: i32, + y: i32, +} + +#[ensures(low(old(pt.x)) ==> low(*result))] +#[after_expiry(result => pt.x == *before_expiry(result))] +fn get_mut_x<'a>(pt: &'a mut Point) -> &'a mut i32 { + &mut pt.x +} + +#[trusted] +#[ensures(low(result))] +fn produces_low() -> i32 { + 12 +} + +#[trusted] +#[requires(low(n))] +fn requires_low(n: i32) {} + +fn main() { + let mut p = Point { + x: produces_low(), + y: 42, + }; + { + let x = get_mut_x(&mut p); + requires_low(*x); + *x = 19; + } + prusti_assert!(p.x == 19); +} diff --git a/prusti-tests/tests/verify/pass/sif/wand.rs b/prusti-tests/tests/verify/pass/sif/wand.rs new file mode 100644 index 00000000000..cdf30f4ca05 --- /dev/null +++ b/prusti-tests/tests/verify/pass/sif/wand.rs @@ -0,0 +1,33 @@ +// compile-flags: -Psif=true -Pserver_address=MOCK +// The sif flag is used in the server which, during the compiletest is only spawned with the default config. +// So we need to start a new server with this test config to make it work. + +use prusti_contracts::*; + +struct Wrapper(T); + +impl Wrapper { + #[pure] + fn get(&self) -> &T { + &self.0 + } + + // #[after_expiry(low(before_expiry(*result)) ==> low(self.get()))] + #[after_expiry(before_expiry(*result) == *self.get())] + fn get_mut(&mut self) -> &mut T { + &mut self.0 + } +} + +#[trusted] +#[ensures(low(result))] +fn produce_low() -> u32 { + todo!() +} + +fn main() { + let mut m = Wrapper(42); + + *m.get_mut() = produce_low(); + prusti_assert!(low(m.get())); +} diff --git a/prusti-utils/src/config.rs b/prusti-utils/src/config.rs index 65a6bc9e624..7d885fe2c1d 100644 --- a/prusti-utils/src/config.rs +++ b/prusti-utils/src/config.rs @@ -22,6 +22,7 @@ pub struct Optimizations { pub optimize_folding: bool, pub remove_empty_if: bool, pub purify_vars: bool, + pub simplify_exprs: bool, pub fix_quantifiers: bool, pub fix_unfoldings: bool, pub remove_unused_vars: bool, @@ -37,6 +38,7 @@ impl Optimizations { optimize_folding: false, remove_empty_if: false, purify_vars: false, + simplify_exprs: false, fix_quantifiers: false, fix_unfoldings: false, remove_unused_vars: false, @@ -52,6 +54,7 @@ impl Optimizations { optimize_folding: true, remove_empty_if: true, purify_vars: true, + simplify_exprs: true, fix_quantifiers: true, // Disabled because https://github.com/viperproject/prusti-dev/issues/892 has been fixed fix_unfoldings: false, @@ -115,6 +118,7 @@ lazy_static::lazy_static! { settings.set_default("optimizations", "all").unwrap(); settings.set_default("intern_names", true).unwrap(); settings.set_default("enable_purification_optimization", false).unwrap(); + settings.set_default("sif", false).unwrap(); // settings.set_default("enable_manual_axiomatization", false).unwrap(); settings.set_default("unsafe_core_proof", false).unwrap(); settings.set_default("verify_core_proof", true).unwrap(); @@ -701,6 +705,7 @@ pub fn verify_only_basic_block_path() -> Vec { /// - `"optimize_folding"` /// - `"remove_empty_if"` /// - `"purify_vars"` +/// - `"simplify_exprs"` /// - `"fix_quantifiers"` /// - `"fix_unfoldings"` /// - `"remove_unused_vars"` @@ -720,6 +725,7 @@ pub fn optimizations() -> Optimizations { "optimize_folding" => opt.optimize_folding = true, "remove_empty_if" => opt.remove_empty_if = true, "purify_vars" => opt.purify_vars = true, + "simplify_exprs" => opt.simplify_exprs = true, "fix_quantifiers" => opt.fix_quantifiers = true, "fix_unfoldings" => opt.fix_unfoldings = true, "remove_unused_vars" => opt.remove_unused_vars = true, @@ -870,6 +876,11 @@ pub fn unsafe_core_proof() -> bool { read_setting("unsafe_core_proof") } +/// When enabled, checks the program for secure information flow +pub fn sif() -> bool { + read_setting("sif") +} + /// Whether the core proof (memory safety) should be verified. /// /// **Note:** This option is taken into account only when `unsafe_core_proof` is diff --git a/prusti-viper/src/encoder/foldunfold/footprint.rs b/prusti-viper/src/encoder/foldunfold/footprint.rs index cc576b22d0d..a774eee7557 100644 --- a/prusti-viper/src/encoder/foldunfold/footprint.rs +++ b/prusti-viper/src/encoder/foldunfold/footprint.rs @@ -176,6 +176,8 @@ impl ExprFootprintGetter for vir::Expr { vir::Expr::SnapApp(vir::SnapApp { ref base, .. }) => base.get_footprint(predicates), vir::Expr::Cast(vir::Cast { ref base, .. }) => base.get_footprint(predicates), + vir::Expr::Low(vir::Low { ref base, .. }) => base.get_footprint(predicates), + vir::Expr::LowEvent => FxHashSet::default(), } } } diff --git a/prusti-viper/src/encoder/mir/pure/interpreter/interpreter_poly.rs b/prusti-viper/src/encoder/mir/pure/interpreter/interpreter_poly.rs index d5f2c6fa481..ab70a90bfcc 100644 --- a/prusti-viper/src/encoder/mir/pure/interpreter/interpreter_poly.rs +++ b/prusti-viper/src/encoder/mir/pure/interpreter/interpreter_poly.rs @@ -583,6 +583,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> BackwardMirInterpreter<'tcx> | "prusti_contracts::specification_entailment" | "prusti_contracts::call_description" | "prusti_contracts::snap" + | "prusti_contracts::low" + | "prusti_contracts::low_event" | "prusti_contracts::snapshot_equality" => { let expr = self.encoder.encode_prusti_operation( full_func_proc_name, diff --git a/prusti-viper/src/encoder/mir/pure/specifications/interface.rs b/prusti-viper/src/encoder/mir/pure/specifications/interface.rs index ff05151e2dd..56d8fd2d945 100644 --- a/prusti-viper/src/encoder/mir/pure/specifications/interface.rs +++ b/prusti-viper/src/encoder/mir/pure/specifications/interface.rs @@ -5,7 +5,7 @@ // file, You can obtain one at http://mozilla.org/MPL/2.0/. use crate::encoder::{ - errors::{SpannedEncodingResult, WithSpan}, + errors::{EncodingError, SpannedEncodingResult, WithSpan}, mir::{ places::PlacesEncoderInterface, pure::{ @@ -27,6 +27,7 @@ use crate::encoder::{ mir_encoder::{MirEncoder, PlaceEncoder, PRECONDITION_LABEL}, snapshot::interface::SnapshotEncoderInterface, }; +use prusti_common::config; use prusti_rustc_interface::{ hir::def_id::DefId, middle::{mir, ty::subst::SubstsRef}, @@ -255,6 +256,26 @@ impl<'v, 'tcx: 'v> SpecificationEncoderInterface<'tcx> for crate::encoder::Encod vir_poly::Expr::snap_app(encoded_args[0].clone()), vir_poly::Expr::snap_app(encoded_args[1].clone()), )), + "prusti_contracts::low" => { + if config::sif() { + Ok(vir_poly::Expr::low(encoded_args[0].clone())) + } else { + Err(EncodingError::incorrect( + "Found low specification when sif option is disabled!", + )) + .with_span(span) + } + } + "prusti_contracts::low_event" => { + if config::sif() { + Ok(vir_poly::Expr::LowEvent) + } else { + Err(EncodingError::incorrect( + "Found low specification when sif option is disabled!", + )) + .with_span(span) + } + } _ => unimplemented!(), } } diff --git a/prusti-viper/src/encoder/procedure_encoder.rs b/prusti-viper/src/encoder/procedure_encoder.rs index c671915b6cc..92a5ce011d3 100644 --- a/prusti-viper/src/encoder/procedure_encoder.rs +++ b/prusti-viper/src/encoder/procedure_encoder.rs @@ -4,85 +4,82 @@ // License, v. 2.0. If a copy of the MPL was not distributed with this // file, You can obtain one at http://mozilla.org/MPL/2.0/. -use crate::encoder::mir::spans::interface::SpanInterface; -use crate::encoder::builtin_encoder::{BuiltinMethodKind}; -use crate::encoder::errors::{ - SpannedEncodingError, ErrorCtxt, EncodingError, WithSpan, - EncodingResult, SpannedEncodingResult +use super::{ + counterexamples::DiscriminantsStateInterface, high::generics::HighGenericsEncoderInterface, }; -use crate::encoder::errors::error_manager::PanicCause; -use crate::encoder::foldunfold; -use crate::encoder::high::types::HighTypeEncoderInterface; -use crate::encoder::initialisation::InitInfo; -use crate::encoder::loop_encoder::{LoopEncoder, LoopEncoderError}; -use crate::encoder::mir_encoder::{MirEncoder, FakeMirEncoder, PlaceEncoder, PlaceEncoding, ExprOrArrayBase}; -use crate::encoder::mir_encoder::PRECONDITION_LABEL; -use crate::encoder::mir_successor::MirSuccessor; -use crate::encoder::places::{Local, LocalVariableManager, Place}; -use crate::encoder::Encoder; -use crate::encoder::snapshot::interface::SnapshotEncoderInterface; -use crate::encoder::mir::procedures::encoder::specification_blocks::SpecificationBlocks; -use crate::error_unsupported; +use crate::{ + encoder::{ + builtin_encoder::BuiltinMethodKind, + errors::{ + error_manager::PanicCause, EncodingError, EncodingErrorKind, EncodingResult, ErrorCtxt, + SpannedEncodingError, SpannedEncodingResult, WithSpan, + }, + foldunfold, + high::types::HighTypeEncoderInterface, + initialisation::InitInfo, + loop_encoder::{LoopEncoder, LoopEncoderError}, + mir::{ + contracts::{ContractsEncoderInterface, ProcedureContract}, + procedures::encoder::specification_blocks::SpecificationBlocks, + pure::{PureFunctionEncoderInterface, SpecificationEncoderInterface}, + sequences::MirSequencesEncoderInterface, + spans::interface::SpanInterface, + specifications::SpecificationsInterface, + type_invariants::TypeInvariantEncoderInterface, + types::MirTypeEncoderInterface, + }, + mir_encoder::{ + ExprOrArrayBase, FakeMirEncoder, MirEncoder, PlaceEncoder, PlaceEncoding, + PRECONDITION_LABEL, + }, + mir_successor::MirSuccessor, + places::{Local, LocalVariableManager, Place}, + snapshot::interface::SnapshotEncoderInterface, + Encoder, + }, + error_unsupported, + utils::is_reference, +}; +use ::log::{debug, trace}; use prusti_common::{ config, utils::to_string::ToString, - vir::{ToGraphViz, fixes::fix_ghost_vars}, - vir_local, vir_expr, vir_stmt -}; -use vir_crate::{ - polymorphic::{ - self as vir, - compute_identifier, - borrows::Borrow, - collect_assigned_vars, - CfgBlockIndex, ExprIterator, Successor, Type}, + vir::{fixes::fix_ghost_vars, ToGraphViz}, + vir_expr, vir_local, vir_stmt, }; use prusti_interface::{ data::ProcedureDefId, environment::{ - borrowck::facts, + borrowck::{facts, regions::PlaceRegionsError}, + mir_utils::SliceOrArrayRef, polonius_info::{ LoanPlaces, PoloniusInfo, PoloniusInfoError, ReborrowingDAG, ReborrowingDAGNode, ReborrowingKind, ReborrowingZombity, }, BasicBlockIndex, LoopAnalysisError, PermissionKind, Procedure, }, - PrustiError, + specs::{ + typed, + typed::{Pledge, SpecificationItem}, + }, + utils, PrustiError, }; -use std::collections::{BTreeMap}; -use std::fmt::Debug; -use prusti_interface::utils; -use prusti_rustc_interface::middle::mir::Mutability; -use prusti_rustc_interface::middle::mir; -use prusti_rustc_interface::middle::mir::{TerminatorKind}; -use prusti_rustc_interface::middle::ty::{self, subst::SubstsRef}; -use prusti_rustc_interface::target::abi::Integer; -use rustc_hash::{FxHashMap, FxHashSet}; -use prusti_rustc_interface::span::Span; -use prusti_rustc_interface::errors::MultiSpan; -use prusti_interface::specs::typed; -use ::log::{trace, debug}; -use prusti_interface::environment::borrowck::regions::PlaceRegionsError; -use crate::encoder::errors::EncodingErrorKind; -use std::convert::TryInto; -use prusti_interface::specs::typed::{Pledge, SpecificationItem}; -use vir_crate::polymorphic::Float; -use crate::utils::is_reference; -use crate::encoder::mir::{ - sequences::MirSequencesEncoderInterface, - contracts::{ - ContractsEncoderInterface, - ProcedureContract, +use prusti_rustc_interface::{ + errors::MultiSpan, + middle::{ + mir, + mir::{Mutability, TerminatorKind}, + ty::{self, subst::SubstsRef}, }, - pure::PureFunctionEncoderInterface, - types::MirTypeEncoderInterface, - pure::SpecificationEncoderInterface, - specifications::SpecificationsInterface, - type_invariants::TypeInvariantEncoderInterface, + span::Span, + target::abi::Integer, +}; +use rustc_hash::{FxHashMap, FxHashSet}; +use std::{collections::BTreeMap, convert::TryInto, fmt::Debug}; +use vir_crate::polymorphic::{ + self as vir, borrows::Borrow, collect_assigned_vars, compute_identifier, CfgBlockIndex, + ExprIterator, Float, Successor, Type, }; -use super::high::generics::HighGenericsEncoderInterface; -use super::counterexamples::DiscriminantsStateInterface; -use prusti_interface::environment::mir_utils::SliceOrArrayRef; pub struct ProcedureEncoder<'p, 'v: 'p, 'tcx: 'v> { encoder: &'p Encoder<'v, 'tcx>, @@ -121,7 +118,8 @@ pub struct ProcedureEncoder<'p, 'v: 'p, 'tcx: 'v> { FxHashMap, FxHashMap)>, /// A map that stores local variables used to preserve the value of a place accross the loop /// when we cannot do that by using permissions. - pure_var_for_preserving_value_map: FxHashMap>, + pure_var_for_preserving_value_map: + FxHashMap>, /// Information about which places are definitely initialised. init_info: InitInfo, /// Mapping from old expressions to ghost variables with which they were replaced. @@ -139,7 +137,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { #[tracing::instrument(name = "ProcedureEncoder::new", level = "debug", skip_all, fields(proc_def_id = ?procedure.get_id()))] pub fn new( encoder: &'p Encoder<'v, 'tcx>, - procedure: &'p Procedure<'tcx> + procedure: &'p Procedure<'tcx>, ) -> SpannedEncodingResult { let mir = procedure.get_mir(); let proc_def_id = procedure.get_id(); @@ -149,7 +147,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let init_info = InitInfo::new(mir, tcx, proc_def_id, &mir_encoder) .with_default_span(procedure.get_span())?; - let specification_blocks = SpecificationBlocks::build(encoder.env().query, mir, procedure, false); + let specification_blocks = + SpecificationBlocks::build(encoder.env().query, mir, procedure, false); let cfg_method = vir::CfgMethod::new( // method name @@ -222,7 +221,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ) -> SpannedEncodingResult<()> { let block = &self.mir[bb]; let _ = self.try_encode_assert(bb, block, encoded_statements)? - || self.try_encode_assume(bb, block, encoded_statements)?; + || self.try_encode_assume(bb, block, encoded_statements)?; Ok(()) } @@ -241,13 +240,11 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { if self.encoder.get_prusti_assumption(cl_def_id).is_none() { return Ok(false); } - let assume_expr = self.encoder.encode_invariant(self.mir, bb, self.proc_def_id, cl_substs)?; + let assume_expr = + self.encoder + .encode_invariant(self.mir, bb, self.proc_def_id, cl_substs)?; - let assume_stmt = vir::Stmt::Inhale( - vir::Inhale { - expr: assume_expr - } - ); + let assume_stmt = vir::Stmt::Inhale(vir::Inhale { expr: assume_expr }); encoded_statements.push(assume_stmt); @@ -278,14 +275,14 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { .encoder .get_definition_span(assertion.assertion.to_def_id()); - let assert_expr = self.encoder.encode_invariant(self.mir, bb, self.proc_def_id, cl_substs)?; + let assert_expr = + self.encoder + .encode_invariant(self.mir, bb, self.proc_def_id, cl_substs)?; - let assert_stmt = vir::Stmt::Assert( - vir::Assert { - expr: assert_expr, - position: self.register_error(span, ErrorCtxt::Panic(PanicCause::Assert)) - } - ); + let assert_stmt = vir::Stmt::Assert(vir::Assert { + expr: assert_expr, + position: self.register_error(span, ErrorCtxt::Panic(PanicCause::Assert)), + }); encoded_statements.push(assert_stmt); @@ -302,11 +299,15 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { variable, } => { let msg = if self.mir.local_decls[variable].is_user_variable() { - "creation of loan 'FIXME: extract variable name' in loop is unsupported".to_string() + "creation of loan 'FIXME: extract variable name' in loop is unsupported" + .to_string() } else { "creation of temporary loan in loop is unsupported".to_string() }; - SpannedEncodingError::unsupported(msg, self.mir_encoder.get_span_of_basic_block(loop_head)) + SpannedEncodingError::unsupported( + msg, + self.mir_encoder.get_span_of_basic_block(loop_head), + ) } PoloniusInfoError::LoansInNestedLoops(location1, _loop1, _location2, _loop2) => { @@ -324,11 +325,13 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ) } - PoloniusInfoError::MultipleMagicWandsPerLoop(location) => SpannedEncodingError::unsupported( - "the creation of loans in this loop is not supported \ + PoloniusInfoError::MultipleMagicWandsPerLoop(location) => { + SpannedEncodingError::unsupported( + "the creation of loans in this loop is not supported \ (MultipleMagicWandsPerLoop)", - self.mir.source_info(location).span, - ), + self.mir.source_info(location).span, + ) + } PoloniusInfoError::MagicWandHasNoRepresentativeLoan(location) => { SpannedEncodingError::unsupported( @@ -338,10 +341,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ) } - PoloniusInfoError::PlaceRegionsError( - PlaceRegionsError::Unsupported(msg), - span, - ) => { + PoloniusInfoError::PlaceRegionsError(PlaceRegionsError::Unsupported(msg), span) => { SpannedEncodingError::unsupported(msg, span) } @@ -364,7 +364,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let mir_span = self.mir.span; // Retrieve the contract - let procedure_contract = self.encoder + let procedure_contract = self + .encoder .get_procedure_contract_for_def(self.proc_def_id, self.substs) .with_span(mir_span)?; assert_one_magic_wand(procedure_contract.borrow_infos.len()).with_span(mir_span)?; @@ -375,9 +376,9 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let name = self.mir_encoder.encode_local_var_name(local); let typ = self .encoder - .encode_type(self.mir_encoder.get_local_ty(local)).unwrap(); // will panic if attempting to encode unsupported type - self.cfg_method - .add_formal_return(&name, typ) + .encode_type(self.mir_encoder.get_local_ty(local)) + .unwrap(); // will panic if attempting to encode unsupported type + self.cfg_method.add_formal_return(&name, typ) } // Preprocess loops @@ -399,8 +400,12 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // Load Polonius info self.polonius_info = Some( - PoloniusInfo::new(self.encoder.env(), self.procedure, &self.cached_loop_invariant_block) - .map_err(|err| self.translate_polonius_error(err))?, + PoloniusInfo::new( + self.encoder.env(), + self.procedure, + &self.cached_loop_invariant_block, + ) + .map_err(|err| self.translate_polonius_error(err))?, ); // Initialize CFG blocks @@ -432,7 +437,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { .register_span(self.mir_encoder.get_span_of_basic_block(bbi)); self.cfg_method.add_stmt( start_cfg_block, - vir::Stmt::Assign( vir::Assign { + vir::Stmt::Assign(vir::Assign { target: vir::Expr::local(executed_flag_var.clone()).set_pos(bb_pos), source: false.into(), kind: vir::AssignKind::Copy, @@ -453,9 +458,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { )?; if !unresolved_edges.is_empty() { return Err(SpannedEncodingError::internal( - format!( - "there are unresolved CFG edges in the encoding: {unresolved_edges:?}" - ), + format!("there are unresolved CFG edges in the encoding: {unresolved_edges:?}"), mir_span, )); } @@ -467,8 +470,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ); // Prepare assertions to check specification refinement - let (precondition_weakening, postcondition_strengthening) - = self.encode_spec_refinement(PRECONDITION_LABEL)?; + let (precondition_weakening, postcondition_strengthening) = + self.encode_spec_refinement(PRECONDITION_LABEL)?; // Encode preconditions self.encode_preconditions(start_cfg_block, precondition_weakening)?; @@ -485,8 +488,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let local_ty = self.locals.get_type(*local); let typ = self.encoder.encode_type(local_ty).unwrap(); // will panic if attempting to encode unsupported type let var_name = self.locals.get_name(*local); - self.cfg_method - .add_local_var(&var_name, typ); + self.cfg_method.add_local_var(&var_name, typ); } self.check_vir()?; @@ -506,7 +508,9 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } // Patch snapshots - self.cfg_method = self.encoder.patch_snapshots_method(self.cfg_method) + self.cfg_method = self + .encoder + .patch_snapshots_method(self.cfg_method) .with_span(mir_span)?; // Add fold/unfold @@ -516,10 +520,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { .iter() .map(|(loan, location)| (loan.index().into(), *location)) .collect(); - let method_pos = self.register_error( - self.mir.span, - ErrorCtxt::Unexpected - ); + let method_pos = self.register_error(self.mir.span, ErrorCtxt::Unexpected); let method_with_fold_unfold = foldunfold::add_fold_unfold( self.encoder, self.cfg_method, @@ -527,19 +528,15 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { &self.cfg_blocks_map, method_pos, ) - .map_err(|foldunfold_error| { - match foldunfold_error { - foldunfold::FoldUnfoldError::Unsupported(msg) => { - SpannedEncodingError::unsupported(msg, mir_span) - } - - _ => SpannedEncodingError::internal( - format!( - "cannot generate fold-unfold Viper statements. {foldunfold_error}", - ), - mir_span, - ), + .map_err(|foldunfold_error| match foldunfold_error { + foldunfold::FoldUnfoldError::Unsupported(msg) => { + SpannedEncodingError::unsupported(msg, mir_span) } + + _ => SpannedEncodingError::internal( + format!("cannot generate fold-unfold Viper statements. {foldunfold_error}",), + mir_span, + ), })?; // Fix variable declarations. @@ -643,7 +640,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { matches!( expr, vir::Expr::Const(vir::ConstExpr { - value: vir::Const::Bool(true), .. + value: vir::Const::Bool(true), + .. }) ) } @@ -665,14 +663,28 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } impl<'a, 'b, 'c> ExprFolder for FindFnApps<'a, 'b, 'c> { fn fold_func_app(&mut self, expr: vir::FuncApp) -> vir::Expr { - let vir::FuncApp { function_name, type_arguments, arguments, formal_arguments, return_type, .. } = expr; + let vir::FuncApp { + function_name, + type_arguments, + arguments, + formal_arguments, + return_type, + .. + } = expr; if self.recurse { - let identifier: vir::FunctionIdentifier = - compute_identifier(&function_name, &type_arguments, &formal_arguments, &return_type).into(); + let identifier: vir::FunctionIdentifier = compute_identifier( + &function_name, + &type_arguments, + &formal_arguments, + &return_type, + ) + .into(); let pres = self.encoder.get_function(&identifier).unwrap().pres.clone(); // Avoid recursively collecting preconditions: they should be self-framing anyway self.recurse = false; - let pres: vir::Expr = pres.into_iter().fold(true.into(), |acc, expr| vir_expr! { [acc] && [expr] }); + let pres: vir::Expr = pres + .into_iter() + .fold(true.into(), |acc, expr| vir_expr! { [acc] && [expr] }); self.push_precond(pres, Some(function_name)); self.recurse = true; @@ -686,16 +698,27 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // We only care about the exhale part (e.g. for `walk_exhale`) self.fold(*expr.exhale_expr) } - fn fold_predicate_access_predicate(&mut self, expr: vir::PredicateAccessPredicate) -> vir::Expr { + fn fold_predicate_access_predicate( + &mut self, + expr: vir::PredicateAccessPredicate, + ) -> vir::Expr { self.fold(*expr.argument); true.into() } - fn fold_field_access_predicate(&mut self, expr: vir::FieldAccessPredicate) -> vir::Expr { + fn fold_field_access_predicate( + &mut self, + expr: vir::FieldAccessPredicate, + ) -> vir::Expr { self.fold(*expr.base); true.into() } fn fold_bin_op(&mut self, expr: vir::BinOp) -> vir::Expr { - let vir::BinOp { op_kind, left, right, position } = expr; + let vir::BinOp { + op_kind, + left, + right, + position, + } = expr; let left = self.fold_boxed(left); let right = self.fold_boxed(right); match op_kind { @@ -704,27 +727,38 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { vir::BinaryOpKind::Or if is_const_true(&left) => *left, vir::BinaryOpKind::Or if is_const_true(&right) => *right, _ => vir::Expr::BinOp(vir::BinOp { - op_kind, left, right, position, + op_kind, + left, + right, + position, }), } } } let mut walker = FindFnApps { - recurse: true, preconds: Vec::new(), encoder: self.encoder + recurse: true, + preconds: Vec::new(), + encoder: self.encoder, }; walker.walk(stmt); walker.preconds } fn block_preconditions(&self, block: &CfgBlockIndex) -> Vec<(Option, vir::Expr)> { let bb = &self.cfg_method.basic_blocks[block.block_index]; - let mut preconds: Vec<_> = bb.stmts.iter().flat_map(|stmt| self.stmt_preconditions(stmt)).collect(); + let mut preconds: Vec<_> = bb + .stmts + .iter() + .flat_map(|stmt| self.stmt_preconditions(stmt)) + .collect(); match &bb.successor { Successor::Undefined => (), Successor::Return => (), Successor::Goto(cbi) => preconds.extend(self.block_preconditions(cbi)), Successor::GotoSwitch(succs, def) => { preconds.extend( - succs.iter().flat_map(|cond_bb| self.block_preconditions(&cond_bb.1)) + succs + .iter() + .flat_map(|cond_bb| self.block_preconditions(&cond_bb.1)), ); preconds.extend(self.block_preconditions(def)); } @@ -796,9 +830,9 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { .get_loop_body(loop_head) .iter() .copied() - .filter( - |&bb| self.procedure.is_reachable_block(bb) && !self.procedure.is_spec_block(bb) - ) + .filter(|&bb| { + self.procedure.is_reachable_block(bb) && !self.procedure.is_spec_block(bb) + }) .collect(); // Identify important blocks @@ -806,13 +840,15 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let loop_exit_blocks_set: FxHashSet<_> = loop_exit_blocks.iter().cloned().collect(); let before_invariant_block: BasicBlockIndex = self.cached_loop_invariant_block[&loop_head]; let before_inv_block_pos = loop_body - .iter().copied() + .iter() + .copied() .position(|bb| bb == before_invariant_block) .unwrap(); let after_inv_block_pos = 1 + before_inv_block_pos; // Find boolean switch exit blocks before the invariant let boolean_exit_blocks_before_inv: Vec<_> = loop_body[0..after_inv_block_pos] - .iter().copied() + .iter() + .copied() .filter(|bb| loop_exit_blocks_set.contains(bb)) .filter(|&bb| self.procedure.successors(bb).len() == 2) .collect(); @@ -894,8 +930,14 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { heads.push(first_b1_head); // Calculate if `G` and `B1` are safe to include as the first `body_invariant!(...)` - let mut preconds = first_g_head.map(|cbi| self.block_preconditions(&cbi)).unwrap_or_default(); - preconds.extend(first_b1_head.map(|cbi| self.block_preconditions(&cbi)).unwrap_or_default()); + let mut preconds = first_g_head + .map(|cbi| self.block_preconditions(&cbi)) + .unwrap_or_default(); + preconds.extend( + first_b1_head + .map(|cbi| self.block_preconditions(&cbi)) + .unwrap_or_default(), + ); // Build the "invariant" CFG block (start - G - B1 - *invariant* - B2 - G - B1 - end) // (1) checks the loop invariant on entry @@ -920,7 +962,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ); heads.push(Some(inv_pre_block)); self.cfg_method - .set_successor(inv_pre_block, vir::Successor::Goto(inv_post_block_perms)); + .set_successor(inv_pre_block, vir::Successor::Goto(inv_post_block_perms)); { let stmts = self.encode_loop_invariant_exhale_stmts(loop_head, before_invariant_block, false)?; @@ -928,13 +970,18 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } // We'll add later more statements at the end of inv_pre_block, to havoc local variables let fnspec_span = { - let (stmts, fnspec_span) = - self.encode_loop_invariant_inhale_fnspec_stmts(loop_head, before_invariant_block, false)?; - self.cfg_method.add_stmts(inv_post_block_fnspc, stmts); fnspec_span + let (stmts, fnspec_span) = self.encode_loop_invariant_inhale_fnspec_stmts( + loop_head, + before_invariant_block, + false, + )?; + self.cfg_method.add_stmts(inv_post_block_fnspc, stmts); + fnspec_span }; { - let stmts = - self.encode_loop_invariant_inhale_perm_stmts(loop_head, before_invariant_block, false).with_span(fnspec_span)?; + let stmts = self + .encode_loop_invariant_inhale_perm_stmts(loop_head, before_invariant_block, false) + .with_span(fnspec_span)?; self.cfg_method.add_stmts(inv_post_block_perms, stmts); } @@ -964,24 +1011,45 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { Some((mid_g, mid_b1)) } else { // Cannot add loop guard to loop invariant - let fn_names: Vec<_> = preconds.iter().filter_map(|(name, _)| name.as_ref()).map(|name| { - name.strip_prefix("m_").unwrap_or(name) - }).collect(); + let fn_names: Vec<_> = preconds + .iter() + .filter_map(|(name, _)| name.as_ref()) + .map(|name| name.strip_prefix("m_").unwrap_or(name)) + .collect(); let warning_msg = if fn_names.is_empty() { "the loop guard was not automatically added as a `body_invariant!(...)`, consider doing this manually".to_string() } else { "the loop guard was not automatically added as a `body_invariant!(...)`, \ - due to the following pure functions with preconditions: [".to_string() + &fn_names.join(", ") + "]" + due to the following pure functions with preconditions: [" + .to_string() + + &fn_names.join(", ") + + "]" }; // Span of loop guard - let span_lg = loop_guard_evaluation.last().map(|bb| self.mir_encoder.get_span_of_basic_block(*bb)); + let span_lg = loop_guard_evaluation + .last() + .map(|bb| self.mir_encoder.get_span_of_basic_block(*bb)); // Span of entire loop body before inv; the last elem (if any) will be the `body_invariant!(...)` bb - let span_lb = loop_body_before_inv.iter().rev().skip(1).map(|bb| self.mir_encoder.get_span_of_basic_block(*bb)) - .fold(None, |acc: Option, span| Some(acc.map(|other| - if span.ctxt() == other.ctxt() { span.to(other) } else { other } - ).unwrap_or(span)) ); + let span_lb = loop_body_before_inv + .iter() + .rev() + .skip(1) + .map(|bb| self.mir_encoder.get_span_of_basic_block(*bb)) + .fold(None, |acc: Option, span| { + Some( + acc.map(|other| { + if span.ctxt() == other.ctxt() { + span.to(other) + } else { + other + } + }) + .unwrap_or(span), + ) + }); // Multispan to highlight both - let span = MultiSpan::from_spans(vec![span_lg, span_lb].into_iter().flatten().collect()); + let span = + MultiSpan::from_spans(vec![span_lg, span_lb].into_iter().flatten().collect()); // It's only a warning so we might as well emit it straight away PrustiError::warning(warning_msg, span).emit(&self.encoder.env().diagnostic); None @@ -1024,16 +1092,13 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ))], ); { - let stmts = self.encode_loop_invariant_exhale_stmts( - loop_head, - before_invariant_block, - true - )?; + let stmts = + self.encode_loop_invariant_exhale_stmts(loop_head, before_invariant_block, true)?; self.cfg_method.add_stmts(end_body_block, stmts); } self.cfg_method.add_stmt( end_body_block, - vir::Stmt::Inhale( vir::Inhale {expr: false.into()} ), + vir::Stmt::Inhale(vir::Inhale { expr: false.into() }), ); heads.push(Some(end_body_block)); @@ -1068,20 +1133,28 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { if let Some(((mid_g_head, mid_g_edges), (mid_b1_head, mid_b1_edges))) = mid_groups { let mid_b1_block = mid_b1_head.unwrap_or(inv_post_block_fnspc); // Link edges of "invariant_perm" (start - G - B1 - *invariant_perm* - G - B1 - invariant_fnspec - B2 - G - B1 - end) - self.cfg_method - .set_successor(inv_post_block_perms, vir::Successor::Goto(mid_g_head.unwrap_or(mid_b1_block))); + self.cfg_method.set_successor( + inv_post_block_perms, + vir::Successor::Goto(mid_g_head.unwrap_or(mid_b1_block)), + ); // We know that there is at least one loop iteration of B2, thus it is safe to assume // that anything beforehand couldn't have exited the loop. Here we use // `self.cfg_method.set_successor(..., Successor::Return)` as the same as `assume false` for (curr_block, target) in &mid_g_edges { if self.loop_encoder.get_loop_depth(*target) < loop_depth { - self.cfg_method.set_successor(*curr_block, Successor::Return); + self.cfg_method.add_stmt( + *curr_block, + vir::Stmt::Inhale(vir::Inhale { expr: false.into() }), + ); + self.cfg_method + .set_successor(*curr_block, Successor::Return); } } - let mid_g_edges = mid_g_edges.into_iter().filter(|(_, target)| - self.loop_encoder.get_loop_depth(*target) >= loop_depth - ).collect::>(); + let mid_g_edges = mid_g_edges + .into_iter() + .filter(|(_, target)| self.loop_encoder.get_loop_depth(*target) >= loop_depth) + .collect::>(); // Link edges from the mid G group (start - G - B1 - invariant_perm - *G* - B1 - invariant_fnspec - B2 - G - B1 - end) still_unresolved_edges.extend(self.encode_unresolved_edges(mid_g_edges, |bb| { if bb == after_guard_block { @@ -1095,12 +1168,18 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // that anything beforehand couldn't have exited the loop. for (curr_block, target) in &mid_b1_edges { if self.loop_encoder.get_loop_depth(*target) < loop_depth { - self.cfg_method.set_successor(*curr_block, Successor::Return); + self.cfg_method.add_stmt( + *curr_block, + vir::Stmt::Inhale(vir::Inhale { expr: false.into() }), + ); + self.cfg_method + .set_successor(*curr_block, Successor::Return); } } - let mid_b1_edges = mid_b1_edges.into_iter().filter(|(_, target)| - self.loop_encoder.get_loop_depth(*target) >= loop_depth - ).collect::>(); + let mid_b1_edges = mid_b1_edges + .into_iter() + .filter(|(_, target)| self.loop_encoder.get_loop_depth(*target) >= loop_depth) + .collect::>(); // Link edges from the mid B1 group (start - G - B1 - invariant_perm - G - *B1* - invariant_fnspec - B2 - G - B1 - end) still_unresolved_edges.extend(self.encode_unresolved_edges(mid_b1_edges, |bb| { if bb == after_inv_block { @@ -1111,8 +1190,10 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { })?); } else { // Link edges of "invariant_perm" (start - G - B1 - *invariant_perm* - invariant_fnspec - B2 - G - B1 - end) - self.cfg_method - .set_successor(inv_post_block_perms, vir::Successor::Goto(inv_post_block_fnspc)); + self.cfg_method.set_successor( + inv_post_block_perms, + vir::Successor::Goto(inv_post_block_fnspc), + ); } // Link edges of "invariant" (start - G - B1 - *invariant* - B2 - G - B1 - end) @@ -1169,13 +1250,18 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { vir::Type::Snapshot(_) => BuiltinMethodKind::HavocRef, vir::Type::Seq(_) => BuiltinMethodKind::HavocRef, vir::Type::Map(_) => BuiltinMethodKind::HavocRef, - vir::Type::Ref => return Err(SpannedEncodingError::internal( - format!("unexpected type of local variable {var:?}"), - self.mir.span, - )), + vir::Type::Ref => { + return Err(SpannedEncodingError::internal( + format!("unexpected type of local variable {var:?}"), + self.mir.span, + )) + } }; - let stmt = vir::Stmt::MethodCall( vir::MethodCall { - method_name: self.encoder.encode_builtin_method_use(builtin_method).with_span(loop_head_span)?, + let stmt = vir::Stmt::MethodCall(vir::MethodCall { + method_name: self + .encoder + .encode_builtin_method_use(builtin_method) + .with_span(loop_head_span)?, arguments: vec![], targets: vec![var], }); @@ -1212,10 +1298,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { .insert(curr_block); if self.loop_encoder.is_loop_head(bbi) { - self.cfg_method.add_stmt( - curr_block, - vir::Stmt::comment("This is a loop head"), - ); + self.cfg_method + .add_stmt(curr_block, vir::Stmt::comment("This is a loop head")); } self.encode_execution_flag(bbi, curr_block)?; @@ -1278,7 +1362,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let executed_flag_var = self.cfg_block_has_been_executed[&bbi].clone(); self.cfg_method.add_stmt( cfg_block, - vir::Stmt::Assign( vir::Assign { + vir::Stmt::Assign(vir::Assign { target: vir::Expr::local(executed_flag_var).set_pos(pos), source: true.into(), kind: vir::AssignKind::Copy, @@ -1366,9 +1450,10 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { Err(err) => { let unsupported_msg = match err.kind() { EncodingErrorKind::Unsupported(msg) - if config::allow_unreachable_unsupported_code() => { + if config::allow_unreachable_unsupported_code() => + { msg.to_string() - }, + } _ => { // Propagate the error return Err(err); @@ -1384,13 +1469,11 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { }; let stmts = vec![ vir::Stmt::comment(head_stmt), - vir::Stmt::comment( - format!("Unsupported feature: {unsupported_msg}") - ), - vir::Stmt::Assert( vir::Assert { + vir::Stmt::comment(format!("Unsupported feature: {unsupported_msg}")), + vir::Stmt::Assert(vir::Assert { expr: false.into(), position: pos, - }) + }), ]; (stmts, Some(MirSuccessor::Kill)) } @@ -1419,48 +1502,54 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { mir::StatementKind::Assign(box (lhs, ref rhs)) => { // Array access on the LHS should always be mutable (idx is always calculated // before, and just a separate local variable here) - let (lhs_place_encoding, ty, _) = self.mir_encoder.encode_place(lhs).with_span(span)?; + let (lhs_place_encoding, ty, _) = + self.mir_encoder.encode_place(lhs).with_span(span)?; match lhs_place_encoding { - PlaceEncoding::SliceAccess { box base, index, rust_slice_ty: rust_ty, .. } | - PlaceEncoding::ArrayAccess { box base, index, rust_array_ty: rust_ty, .. } => { + PlaceEncoding::SliceAccess { + box base, + index, + rust_slice_ty: rust_ty, + .. + } + | PlaceEncoding::ArrayAccess { + box base, + index, + rust_array_ty: rust_ty, + .. + } => { // Current stmt is of the form `arr[idx] = val`. This does not have an expiring // temporary variable, so we encode it differently from indexing into an array. - self.encode_array_direct_assign( - base, - index, - rust_ty, - rhs, - location, - )? + self.encode_array_direct_assign(base, index, rust_ty, rhs, location)? } _ => { - let (encoded_lhs, pre_stmts) = self.postprocess_place_encoding(lhs_place_encoding, ArrayAccessKind::Mutable(None, location)) + let (encoded_lhs, pre_stmts) = self + .postprocess_place_encoding( + lhs_place_encoding, + ArrayAccessKind::Mutable(None, location), + ) .with_span(span)?; stmts.extend(pre_stmts); - self.encode_assign( - encoded_lhs, - rhs, - ty, - location, - )? + self.encode_assign(encoded_lhs, rhs, ty, location)? } } } - ref x => return Err(SpannedEncodingError::unsupported( - format!("unsupported statement kind: {x:?}"), - span, - )), + ref x => { + return Err(SpannedEncodingError::unsupported( + format!("unsupported statement kind: {x:?}"), + span, + )) + } }; stmts.extend(encoding_stmts); Ok(self.set_stmts_default_pos(stmts, stmt.source_info.span)) } fn set_stmts_default_pos(&self, stmts: Vec, default_span: Span) -> Vec { - let pos = self.encoder.error_manager().register_span(self.proc_def_id, default_span); - stmts - .into_iter() - .map(|s| s.set_default_pos(pos)) - .collect() + let pos = self + .encoder + .error_manager() + .register_span(self.proc_def_id, default_span); + stmts.into_iter().map(|s| s.set_default_pos(pos)).collect() } /// Encode assignment of RHS to LHS, depending on what kind of thing the RHS is @@ -1478,106 +1567,51 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { self.encode_assign_operand(&encoded_lhs, operand, location)? } mir::Rvalue::Aggregate(ref aggregate, ref operands) => { - self.encode_assign_aggregate( - &encoded_lhs, - ty, - aggregate, - operands, - location - )? - }, + self.encode_assign_aggregate(&encoded_lhs, ty, aggregate, operands, location)? + } mir::Rvalue::BinaryOp(op, box (ref left, ref right)) => { - self.encode_assign_binary_op( - op, - left, - right, - encoded_lhs, - ty, - location - )? + self.encode_assign_binary_op(op, left, right, encoded_lhs, ty, location)? + } + mir::Rvalue::CheckedBinaryOp(op, box (ref left, ref right)) => { + self.encode_assign_checked_binary_op(op, left, right, encoded_lhs, ty, location)? } - mir::Rvalue::CheckedBinaryOp(op, box (ref left, ref right)) => self - .encode_assign_checked_binary_op( - op, - left, - right, - encoded_lhs, - ty, - location, - )?, mir::Rvalue::UnaryOp(op, ref operand) => { - self.encode_assign_unary_op( - op, - operand, - encoded_lhs, - ty, - location - )? + self.encode_assign_unary_op(op, operand, encoded_lhs, ty, location)? } mir::Rvalue::NullaryOp(op, op_ty) => { - self.encode_assign_nullary_op( - op, - op_ty, - encoded_lhs, - ty, - location - )? + self.encode_assign_nullary_op(op, op_ty, encoded_lhs, ty, location)? } mir::Rvalue::Discriminant(src) => { - self.encode_assign_discriminant( - src, - location, - encoded_lhs, - ty - )? + self.encode_assign_discriminant(src, location, encoded_lhs, ty)? } mir::Rvalue::Ref(_region, mir_borrow_kind, place) => { - self.encode_assign_ref( - mir_borrow_kind, - place, - location, - encoded_lhs, - ty - )? + self.encode_assign_ref(mir_borrow_kind, place, location, encoded_lhs, ty)? } - mir::Rvalue::Cast(mir::CastKind::PointerExposeAddress, ref operand, dst_ty) | - mir::Rvalue::Cast(mir::CastKind::PointerFromExposedAddress, ref operand, dst_ty) | - mir::Rvalue::Cast(mir::CastKind::IntToInt, ref operand, dst_ty) => { - self.encode_cast( - operand, - dst_ty, - encoded_lhs, - ty, - location, - )? + mir::Rvalue::Cast(mir::CastKind::PointerExposeAddress, ref operand, dst_ty) + | mir::Rvalue::Cast(mir::CastKind::PointerFromExposedAddress, ref operand, dst_ty) + | mir::Rvalue::Cast(mir::CastKind::IntToInt, ref operand, dst_ty) => { + self.encode_cast(operand, dst_ty, encoded_lhs, ty, location)? } mir::Rvalue::Len(place) => { - self.encode_assign_sequence_len( - encoded_lhs, - place, - ty, - location, - )? + self.encode_assign_sequence_len(encoded_lhs, place, ty, location)? } - mir::Rvalue::Repeat(ref operand, times) => { - self.encode_assign_array_repeat_initializer( + mir::Rvalue::Repeat(ref operand, times) => self + .encode_assign_array_repeat_initializer( encoded_lhs, operand, times, ty, location, - )? - } - mir::Rvalue::Cast(mir::CastKind::Pointer(ty::adjustment::PointerCast::Unsize), ref operand, cast_ty) => { + )?, + mir::Rvalue::Cast( + mir::CastKind::Pointer(ty::adjustment::PointerCast::Unsize), + ref operand, + cast_ty, + ) => { let rhs_ty = self.mir_encoder.get_operand_ty(operand); if rhs_ty.is_array_ref() && cast_ty.is_slice_ref() { trace!("slice: operand={:?}, ty={:?}", operand, cast_ty); - self.encode_assign_slice( - encoded_lhs, - operand, - cast_ty, - location, - )? + self.encode_assign_slice(encoded_lhs, operand, cast_ty, location)? } else { return Err(SpannedEncodingError::unsupported( format!("unsizing a {rhs_ty} into a {cast_ty} is not supported"), @@ -1585,15 +1619,17 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { )); } } - mir::Rvalue::Cast(mir::CastKind::Pointer(_), _, _) | - mir::Rvalue::Cast(mir::CastKind::DynStar, _, _) => { + mir::Rvalue::Cast(mir::CastKind::Pointer(_), _, _) + | mir::Rvalue::Cast(mir::CastKind::DynStar, _, _) => { return Err(SpannedEncodingError::unsupported( - "raw pointers are not supported", span + "raw pointers are not supported", + span, )); } mir::Rvalue::Cast(cast_kind, _, _) => { return Err(SpannedEncodingError::unsupported( - format!("casts {cast_kind:?} are not supported"), span + format!("casts {cast_kind:?} are not supported"), + span, )); } mir::Rvalue::AddressOf(_, _) => { @@ -1603,16 +1639,12 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } mir::Rvalue::ThreadLocalRef(_) => { return Err(SpannedEncodingError::unsupported( - "references to thread-local storage are not supported", span + "references to thread-local storage are not supported", + span, )); } mir::Rvalue::ShallowInitBox(_, op_ty) => { - self.encode_assign_box( - op_ty, - encoded_lhs, - ty, - location, - )? + self.encode_assign_box(op_ty, encoded_lhs, ty, location)? } mir::Rvalue::CopyForDeref(ref place) => { self.encode_assign_operand(&encoded_lhs, &mir::Operand::Copy(*place), location)? @@ -1639,12 +1671,22 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { /// Encode the lhs and the rhs of the assignment that create the loan #[tracing::instrument(level = "debug", skip(self))] - fn encode_loan_places(&mut self, loan_places: &LoanPlaces<'tcx>) -> SpannedEncodingResult<(vir::Expr, Option, bool, Vec)> { + fn encode_loan_places( + &mut self, + loan_places: &LoanPlaces<'tcx>, + ) -> SpannedEncodingResult<(vir::Expr, Option, bool, Vec)> { let location = loan_places.location; let span = self.mir_encoder.get_span_of_location(location); - let (expiring_base, mut stmts, expiring_ty, _) = self.encode_place(loan_places.dest, ArrayAccessKind::Mutable(None, location), location)?; - trace!("expiring_base: {:?}", (&expiring_base, &stmts, &expiring_ty)); + let (expiring_base, mut stmts, expiring_ty, _) = self.encode_place( + loan_places.dest, + ArrayAccessKind::Mutable(None, location), + location, + )?; + trace!( + "expiring_base: {:?}", + (&expiring_base, &stmts, &expiring_ty) + ); // the original encoding of arrays is with a sort-of magic temporary variable, so // `postprocess_place_encoding` will return `i32` here instead of Array$3$i32. so here @@ -1654,12 +1696,19 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let (rhs_place_encoding, ..) = self.mir_encoder.encode_place(rhs_place).unwrap(); if let PlaceEncoding::ArrayAccess { .. } = rhs_place_encoding { // encode expiry of the array borrow - let (expired_expr, regained_array, wand_rhs) = self.array_magic_wand_at[&location].clone(); + let (expired_expr, regained_array, wand_rhs) = + self.array_magic_wand_at[&location].clone(); // expiring base is something like ref$i32, so we need .val_ref.val_int - let deref = self.encoder.encode_value_expr(expiring_base.clone(), expiring_ty).with_span(span)?; + let deref = self + .encoder + .encode_value_expr(expiring_base.clone(), expiring_ty) + .with_span(span)?; let target_ty = expiring_ty.peel_refs(); - let expiring_base_value = self.encoder.encode_value_expr(deref, target_ty).with_span(span)?; + let expiring_base_value = self + .encoder + .encode_value_expr(deref, target_ty) + .with_span(span)?; // the original magic wand refered to the temporary variable that we created for array // encoding. the expiry here refers to the non-temporary rust variable, so we need @@ -1670,25 +1719,31 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // so there's no LHS. instead, we add a label just before the inhale, and refer to // that let new_lhs_label = self.cfg_method.get_fresh_label_name(); - stmts.push( - vir::Stmt::label(new_lhs_label.clone()) - ); + stmts.push(vir::Stmt::label(new_lhs_label.clone())); - let wand_rhs_patched_lhs = wand_rhs - .map_old_expr_label(|label| if label == "lhs" { + let wand_rhs_patched_lhs = wand_rhs.map_old_expr_label(|label| { + if label == "lhs" { trace!("{} -> {}", label, new_lhs_label); new_lhs_label.clone() } else { label } - ); + }); - let expiring_base_pred = self.mir_encoder.encode_place_predicate_permission(expiring_base.clone(), vir::PermAmount::Write).unwrap(); - stmts.push(vir_stmt!{ exhale [expiring_base_pred] }); + let expiring_base_pred = self + .mir_encoder + .encode_place_predicate_permission( + expiring_base.clone(), + vir::PermAmount::Write, + ) + .unwrap(); + stmts.push(vir_stmt! { exhale [expiring_base_pred] }); - let arr_pred = vir::Expr::pred_permission(regained_array.clone(), vir::PermAmount::Write).unwrap(); - stmts.push(vir_stmt!{ inhale [arr_pred] }); - stmts.push(vir_stmt!{ inhale [wand_rhs_patched_lhs] }); + let arr_pred = + vir::Expr::pred_permission(regained_array.clone(), vir::PermAmount::Write) + .unwrap(); + stmts.push(vir_stmt! { inhale [arr_pred] }); + stmts.push(vir_stmt! { inhale [wand_rhs_patched_lhs] }); // NOTE: the third tuple elem is called is_mut in // `construct_vir_reborrowing_node_for_assignment`, and is currently only used to @@ -1700,10 +1755,17 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } } - let mut encode = |rhs_place, stmts: &mut Vec, array_encode_kind| -> SpannedEncodingResult<_> { - let (restored, pre_stmts, _, _) = self.encode_place(rhs_place, array_encode_kind, location)?; + let mut encode = |rhs_place, + stmts: &mut Vec, + array_encode_kind| + -> SpannedEncodingResult<_> { + let (restored, pre_stmts, _, _) = + self.encode_place(rhs_place, array_encode_kind, location)?; stmts.extend(pre_stmts); - let ref_field = self.encoder.encode_value_field(expiring_ty).with_span(span)?; + let ref_field = self + .encoder + .encode_value_field(expiring_ty) + .with_span(span)?; let expiring = expiring_base.clone().field(ref_field.clone()); Ok((expiring, restored, ref_field)) }; @@ -1714,32 +1776,43 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { mir::BorrowKind::Mut { .. } => true, _ => return Err(Self::unsupported_borrow_kind(mir_borrow_kind).with_span(span)), }; - let array_encode_kind = if is_mut { ArrayAccessKind::Mutable(None, location) } else { ArrayAccessKind::Shared }; + let array_encode_kind = if is_mut { + ArrayAccessKind::Mutable(None, location) + } else { + ArrayAccessKind::Shared + }; let (expiring, restored, _) = encode(rhs_place, &mut stmts, array_encode_kind)?; assert_eq!(expiring.get_type(), restored.get_type()); (expiring, Some(restored), is_mut, stmts) } mir::Rvalue::Use(mir::Operand::Move(rhs_place)) => { - let (expiring, restored_base, ref_field) = encode(rhs_place, &mut stmts, ArrayAccessKind::Shared)?; + let (expiring, restored_base, ref_field) = + encode(rhs_place, &mut stmts, ArrayAccessKind::Shared)?; let restored = restored_base.field(ref_field); assert_eq!(expiring.get_type(), restored.get_type()); (expiring, Some(restored), true, stmts) } mir::Rvalue::Use(mir::Operand::Copy(rhs_place)) => { - let (expiring, restored_base, ref_field) = encode(rhs_place, &mut stmts, ArrayAccessKind::Shared)?; + let (expiring, restored_base, ref_field) = + encode(rhs_place, &mut stmts, ArrayAccessKind::Shared)?; let restored = restored_base.field(ref_field); assert_eq!(expiring.get_type(), restored.get_type()); (expiring, Some(restored), false, stmts) } - mir::Rvalue::Cast(mir::CastKind::Pointer(ty::adjustment::PointerCast::Unsize), ref operand, ty) => { + mir::Rvalue::Cast( + mir::CastKind::Pointer(ty::adjustment::PointerCast::Unsize), + ref operand, + ty, + ) => { trace!("cast: operand={:?}, ty={:?}", operand, ty); let place = match *operand { mir::Operand::Move(place) => place, mir::Operand::Copy(place) => place, _ => unreachable!("operand: {:?}", operand), }; - let (restored, r_stmts, ..) = self.encode_place(place, ArrayAccessKind::Shared, location)?; + let (restored, r_stmts, ..) = + self.encode_place(place, ArrayAccessKind::Shared, location)?; stmts.extend(r_stmts); (expiring_base, Some(restored), false, stmts) @@ -1748,12 +1821,15 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { mir::Rvalue::Use(mir::Operand::Constant(ref expr)) => { // TODO: Encoding of string literals is not yet supported, so // do not return an expression in restored here. - let restored: Option = - if is_str(expr.ty()) { - None - } else { - Some(self.encoder.encode_const_expr(expr.ty(), expr.literal).with_span(span)?) - }; + let restored: Option = if is_str(expr.ty()) { + None + } else { + Some( + self.encoder + .encode_const_expr(expr.ty(), expr.literal) + .with_span(span)?, + ) + }; (expiring_base, restored, false, stmts) } @@ -1770,13 +1846,13 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { is_in_package_stmt: bool, ) -> Vec { let mut stmts = if let Some(var) = self.old_to_ghost_var.get(&rhs) { - vec![vir::Stmt::Assign( vir::Assign { + vec![vir::Stmt::Assign(vir::Assign { target: var.clone(), source: lhs.clone(), kind: vir::AssignKind::Move, })] } else { - vec![vir::Stmt::TransferPerm( vir::TransferPerm { + vec![vir::Stmt::TransferPerm(vir::TransferPerm { left: lhs.clone(), right: rhs.clone(), unchecked: false, @@ -1784,8 +1860,9 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { }; if self.check_foldunfold_state && !is_in_package_stmt { - let pos = self.register_error(self.mir.source_info(location).span, ErrorCtxt::Unexpected); - stmts.push(vir::Stmt::Assert( vir::Assert { + let pos = + self.register_error(self.mir.source_info(location).span, ErrorCtxt::Unexpected); + stmts.push(vir::Stmt::Assert(vir::Assert { expr: vir::Expr::eq_cmp(lhs, rhs), position: pos, })); @@ -1795,17 +1872,17 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } fn encode_obtain(&mut self, expr: vir::Expr, pos: vir::Position) -> Vec { - let mut stmts = vec![ - vir::Stmt::Obtain( vir::Obtain { - expr: expr.clone(), - position: pos, - }) - ]; + let mut stmts = vec![vir::Stmt::Obtain(vir::Obtain { + expr: expr.clone(), + position: pos, + })]; if self.check_foldunfold_state { let new_pos = self.encoder.error_manager().duplicate_position(pos); - self.encoder.error_manager().set_error(new_pos, ErrorCtxt::Unexpected); - stmts.push(vir::Stmt::Assert( vir::Assert { + self.encoder + .error_manager() + .set_error(new_pos, ErrorCtxt::Unexpected); + stmts.push(vir::Stmt::Assert(vir::Assert { expr, position: new_pos, })); @@ -1816,17 +1893,19 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { /// A borrow is mutable if it was a MIR unique borrow, a move of /// a borrow, or a argument of a function. - fn is_mutable_borrow(&self, loan: facts::Loan) - -> EncodingResult - { + fn is_mutable_borrow(&self, loan: facts::Loan) -> EncodingResult { if let Some(stmt) = self.polonius_info().get_assignment_for_loan(loan)? { Ok(match stmt.kind { mir::StatementKind::Assign(box (_, ref rhs)) => match rhs { - &mir::Rvalue::Ref(_, mir::BorrowKind::Shared, _) | - &mir::Rvalue::Use(mir::Operand::Copy(_)) => false, - &mir::Rvalue::Ref(_, mir::BorrowKind::Mut { .. }, _) | - &mir::Rvalue::Use(mir::Operand::Move(_)) => true, - &mir::Rvalue::Cast(mir::CastKind::Pointer(ty::adjustment::PointerCast::Unsize), _, _ty) => false, + &mir::Rvalue::Ref(_, mir::BorrowKind::Shared, _) + | &mir::Rvalue::Use(mir::Operand::Copy(_)) => false, + &mir::Rvalue::Ref(_, mir::BorrowKind::Mut { .. }, _) + | &mir::Rvalue::Use(mir::Operand::Move(_)) => true, + &mir::Rvalue::Cast( + mir::CastKind::Pointer(ty::adjustment::PointerCast::Unsize), + _, + _ty, + ) => false, &mir::Rvalue::Use(mir::Operand::Constant(_)) => false, x => unreachable!("{:?}", x), }, @@ -1868,11 +1947,9 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { is_in_package_stmt, )?, ReborrowingKind::Call { loan, .. } => { - if let Some(slice_expiry_node) = self.construct_vir_reborrowing_node_for_slice( - loan, - node, - location, - )? { + if let Some(slice_expiry_node) = + self.construct_vir_reborrowing_node_for_slice(loan, node, location)? + { slice_expiry_node } else { self.construct_vir_reborrowing_node_for_call( @@ -1880,7 +1957,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { loan, node, location, - is_in_package_stmt + is_in_package_stmt, )? } } @@ -1931,33 +2008,33 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let loan_location = self.polonius_info().get_loan_location(&loan); trace!("loan_location: {:?}", loan_location); - let loan_places = self.polonius_info().get_loan_places(&loan) + let loan_places = self + .polonius_info() + .get_loan_places(&loan) .map_err(EncodingError::from) .with_span(span)?; trace!("loan_places: {:?}", loan_places); - Ok(if let Some(regained) = self.slice_created_at.get(&loan_location) { - let guard = self.construct_location_guard(loan_location); - trace!("guard: {:?}", guard); - trace!("regained: {:?}", regained); - Some( - vir::borrows::Node::new( + Ok( + if let Some(regained) = self.slice_created_at.get(&loan_location) { + let guard = self.construct_location_guard(loan_location); + trace!("guard: {:?}", guard); + trace!("regained: {:?}", regained); + Some(vir::borrows::Node::new( guard, node.loan.index().into(), convert_loans_to_borrows(&node.reborrowing_loans), convert_loans_to_borrows(&node.reborrowed_loans), - vec![ - vir::Stmt::comment("hi"), - ], + vec![vir::Stmt::comment("hi")], vec![regained.clone()], Vec::new(), Vec::new(), None, - ) - ) - } else { - None - }) + )) + } else { + None + }, + ) } #[tracing::instrument(level = "trace", skip(self, _mir_dag))] @@ -1974,11 +2051,14 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let node_is_leaf = node.reborrowed_loans.is_empty(); let loan_location = self.polonius_info().get_loan_location(&loan); - let loan_places = self.polonius_info().get_loan_places(&loan) + let loan_places = self + .polonius_info() + .get_loan_places(&loan) .map_err(EncodingError::from) - .with_span(span)?.unwrap(); - let (expiring, restored, is_mut, mut stmts) = self.encode_loan_places(&loan_places) - .with_span(span)?; + .with_span(span)? + .unwrap(); + let (expiring, restored, is_mut, mut stmts) = + self.encode_loan_places(&loan_places).with_span(span)?; let borrowed_places = restored.clone().into_iter().collect(); trace!("construct_vir_reborrowing_node_for_assignment(loan={:?}, loan_places={:?}, expiring={:?}, restored={:?}, stmts={:?}", loan, loan_places, expiring, restored, stmts); @@ -1986,12 +2066,18 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // Move the permissions from the "in loans" ("reborrowing loans") to the current loan if node.incoming_zombies && restored.is_some() { - let lhs_label = self.get_label_after_location(loan_location)?.unwrap().to_string(); + let lhs_label = self + .get_label_after_location(loan_location)? + .unwrap() + .to_string(); for &in_loan in node.reborrowing_loans.iter() { // TODO: Is this the correct span? if self.is_mutable_borrow(in_loan).with_span(span)? { let in_location = self.polonius_info().get_loan_location(&in_loan); - let in_label = self.get_label_after_location(in_location)?.unwrap().to_string(); + let in_label = self + .get_label_after_location(in_location)? + .unwrap() + .to_string(); used_lhs_label = true; stmts.extend(self.encode_transfer_permissions( expiring.clone().old(in_label), @@ -2069,9 +2155,11 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // Get the borrow information. if !self.procedure_contracts.contains_key(&loan_location) { return Err(SpannedEncodingError::internal( - format!("There is no procedure contract for loan {loan:?}. This could happen if you \ - are chaining pure functions, which is not fully supported."), - span + format!( + "There is no procedure contract for loan {loan:?}. This could happen if you \ + are chaining pure functions, which is not fully supported." + ), + span, )); } let (contract, fake_exprs) = self.procedure_contracts[&loan_location].clone(); @@ -2084,10 +2172,15 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let borrow_infos = &contract.borrow_infos; match borrow_infos.len().cmp(&1) { std::cmp::Ordering::Less => (), - std::cmp::Ordering::Greater => return Err(SpannedEncodingError::internal( - format!("We require at most one magic wand in the postcondition. But we have {:?}", borrow_infos.len()), - span, - )), + std::cmp::Ordering::Greater => { + return Err(SpannedEncodingError::internal( + format!( + "We require at most one magic wand in the postcondition. But we have {:?}", + borrow_infos.len() + ), + span, + )) + } std::cmp::Ordering::Equal => { let borrow_info = &borrow_infos[0]; @@ -2105,9 +2198,9 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // Obtain the LHS permission. for (path, _) in &borrow_info.blocking_paths { - let (encoded_place, _, _) = self.encode_generic_place( - contract.def_id, Some(loan_location), *path - ).with_span(span)?; + let (encoded_place, _, _) = self + .encode_generic_place(contract.def_id, Some(loan_location), *path) + .with_span(span)?; let encoded_place = replace_fake_exprs(encoded_place); // Move the permissions from the "in loans" ("reborrowing loans") to the current loan @@ -2142,15 +2235,13 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ErrorCtxt::ApplyMagicWandOnExpiry, ); // Inhale the magic wand. - let magic_wand = vir::Expr::MagicWand( vir::MagicWand { + let magic_wand = vir::Expr::MagicWand(vir::MagicWand { left: box lhs.clone(), right: box rhs.clone(), borrow: Some(loan.index().into()), position: pos, }); - stmts.push(vir::Stmt::Inhale( vir::Inhale { - expr: magic_wand - })); + stmts.push(vir::Stmt::Inhale(vir::Inhale { expr: magic_wand })); // Emit the apply statement. let statement = vir::Stmt::apply_magic_wand(lhs, rhs, loan.index().into(), pos); debug!("{:?} at {:?}", statement, loan_location); @@ -2183,9 +2274,14 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ) -> SpannedEncodingResult> { let mut stmts: Vec = vec![]; if !loans.is_empty() { - let vir_reborrowing_dag = - self.construct_vir_reborrowing_dag(&loans, zombie_loans, location, end_location, is_in_package_stmt)?; - stmts.push(vir::Stmt::ExpireBorrows( vir::ExpireBorrows { + let vir_reborrowing_dag = self.construct_vir_reborrowing_dag( + &loans, + zombie_loans, + location, + end_location, + is_in_package_stmt, + )?; + stmts.push(vir::Stmt::ExpireBorrows(vir::ExpireBorrows { dag: vir_reborrowing_dag, })); } @@ -2202,15 +2298,24 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { .polonius_info() .get_all_loans_dying_between(begin_loc, end_loc); // FIXME: is 'end_loc' correct here? What about 'begin_loc'? - self.encode_expiration_of_loans(all_dying_loans, &zombie_loans, begin_loc, Some(end_loc), false) + self.encode_expiration_of_loans( + all_dying_loans, + &zombie_loans, + begin_loc, + Some(end_loc), + false, + ) } #[tracing::instrument(level = "debug", skip(self))] - fn encode_expiring_borrows_at(&mut self, location: mir::Location) - -> SpannedEncodingResult> - { + fn encode_expiring_borrows_at( + &mut self, + location: mir::Location, + ) -> SpannedEncodingResult> { + debug!("encode_expiring_borrows_at '{:?}'", location); let (all_dying_loans, zombie_loans) = self.polonius_info().get_all_loans_dying_at(location); - let stmts = self.encode_expiration_of_loans(all_dying_loans, &zombie_loans, location, None, false)?; + let stmts = + self.encode_expiration_of_loans(all_dying_loans, &zombie_loans, location, None, false)?; Ok(self.set_stmts_default_pos(stmts, self.mir_encoder.get_span_of_location(location))) } @@ -2282,29 +2387,27 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // Use a local variable for the discriminant. // See Silicon issue https://github.com/viperproject/silicon/issues/356 let discr_var = match switch_ty.kind() { - ty::TyKind::Bool => { - self.cfg_method.add_fresh_local_var(vir::Type::Bool) - } + ty::TyKind::Bool => self.cfg_method.add_fresh_local_var(vir::Type::Bool), - ty::TyKind::Int(_) - | ty::TyKind::Uint(_) - | ty::TyKind::Char => { + ty::TyKind::Int(_) | ty::TyKind::Uint(_) | ty::TyKind::Char => { self.cfg_method.add_fresh_local_var(vir::Type::Int) } - ty::TyKind::Float(ty::FloatTy::F32) => { - self.cfg_method.add_fresh_local_var(vir::Type::Float(Float::F32)) - } + ty::TyKind::Float(ty::FloatTy::F32) => self + .cfg_method + .add_fresh_local_var(vir::Type::Float(Float::F32)), - ty::TyKind::Float(ty::FloatTy::F64) => { - self.cfg_method.add_fresh_local_var(vir::Type::Float(Float::F64)) - } + ty::TyKind::Float(ty::FloatTy::F64) => self + .cfg_method + .add_fresh_local_var(vir::Type::Float(Float::F64)), ref x => unreachable!("{:?}", x), }; - let encoded_discr = self.mir_encoder.encode_operand_expr(discr) + let encoded_discr = self + .mir_encoder + .encode_operand_expr(discr) .with_span(span)?; - stmts.push(vir::Stmt::Assign( vir::Assign { + stmts.push(vir::Stmt::Assign(vir::Assign { target: discr_var.clone().into(), source: encoded_discr, kind: vir::AssignKind::Copy, @@ -2325,12 +2428,12 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } } - ty::TyKind::Int(_) - | ty::TyKind::Uint(_) - | ty::TyKind::Char => vir::Expr::eq_cmp( - discr_var.clone().into(), - self.encoder.encode_int_cast(value, switch_ty), - ), + ty::TyKind::Int(_) | ty::TyKind::Uint(_) | ty::TyKind::Char => { + vir::Expr::eq_cmp( + discr_var.clone().into(), + self.encoder.encode_int_cast(value, switch_ty), + ) + } ref x => unreachable!("{:?}", x), }; @@ -2367,13 +2470,12 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { if guard_is_bool && cfg_targets.len() == 1 { let (target_guard, target) = cfg_targets.pop().unwrap(); let target_span = self.mir_encoder.get_span_of_basic_block(target); - let default_target_span = self.mir_encoder.get_span_of_basic_block(default_target); + let default_target_span = + self.mir_encoder.get_span_of_basic_block(default_target); if target_span > default_target_span { let guard_pos = target_guard.pos(); - cfg_targets = vec![( - target_guard.negate().set_pos(guard_pos), - default_target, - )]; + cfg_targets = + vec![(target_guard.negate().set_pos(guard_pos), default_target)]; default_target = target; } else { // Undo the pop @@ -2398,7 +2500,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { TerminatorKind::Abort => { let pos = self.register_error(term.source_info.span, ErrorCtxt::AbortTerminator); - stmts.push(vir::Stmt::Assert( vir::Assert { + stmts.push(vir::Stmt::Assert(vir::Assert { expr: false.into(), position: pos, })); @@ -2421,11 +2523,10 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ref value, .. } => { - let (encoded_lhs, pre_stmts, _, _) = self.encode_place(lhs, ArrayAccessKind::Mutable(None, location), location)?; + let (encoded_lhs, pre_stmts, _, _) = + self.encode_place(lhs, ArrayAccessKind::Mutable(None, location), location)?; stmts.extend(pre_stmts); - stmts.extend( - self.encode_assign_operand(&encoded_lhs, value, location)? - ); + stmts.extend(self.encode_assign_operand(&encoded_lhs, value, location)?); (stmts, MirSuccessor::Goto(target)) } @@ -2433,21 +2534,23 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ref args, destination, target, - func: - mir::Operand::Constant(box mir::Constant { - literal, - .. - }), + func: mir::Operand::Constant(box mir::Constant { literal, .. }), .. } => { let ty = literal.ty(); let func_const_val = literal.try_to_value(self.encoder.env().tcx()); if let ty::TyKind::FnDef(called_def_id, call_substs) = ty.kind() { let called_def_id = *called_def_id; - debug!("Encode function call {:?} with substs {:?}", called_def_id, call_substs); + debug!( + "Encode function call {:?} with substs {:?}", + called_def_id, call_substs + ); - let full_func_proc_name: &str = - &self.encoder.env().name.get_absolute_item_name(called_def_id); + let full_func_proc_name: &str = &self + .encoder + .env() + .name + .get_absolute_item_name(called_def_id); match full_func_proc_name { "std::rt::begin_panic" @@ -2460,19 +2563,18 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // Example of args[0]: 'const "internal error: entered unreachable code"' let panic_message = format!("{:?}", args[0]); - let panic_cause = self.mir_encoder.encode_panic_cause( - term.source_info.span - ); + let panic_cause = + self.mir_encoder.encode_panic_cause(term.source_info.span); let pos = self.register_error( - term.source_info.span, - ErrorCtxt::Panic(panic_cause), - ); + term.source_info.span, + ErrorCtxt::Panic(panic_cause), + ); if self.check_panics { stmts.push(vir::Stmt::comment(format!( "Rust panic - {panic_message}" ))); - stmts.push(vir::Stmt::Assert( vir::Assert { + stmts.push(vir::Stmt::Assert(vir::Assert { expr: false.into(), position: pos, })); @@ -2481,94 +2583,91 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } } - "std::boxed::Box::::new" - | "alloc::boxed::Box::::new" => { + "std::boxed::Box::::new" | "alloc::boxed::Box::::new" => { // This is the initialization of a box // args[0]: value to put in the box assert_eq!(args.len(), 1); - let (dst, pre_stmts, dest_ty, _) = self.encode_place(destination, ArrayAccessKind::Shared, location)?; + let (dst, pre_stmts, dest_ty, _) = + self.encode_place(destination, ArrayAccessKind::Shared, location)?; stmts.extend(pre_stmts); let boxed_ty = dest_ty.boxed_ty(); - let ref_field = self.encoder.encode_dereference_field(boxed_ty) + let ref_field = self + .encoder + .encode_dereference_field(boxed_ty) .with_span(span)?; let box_content = dst.clone().field(ref_field.clone()); - stmts.extend( - self.prepare_assign_target( - dst, - ref_field, - location, - vir::AssignKind::Move, - false - )? - ); + stmts.extend(self.prepare_assign_target( + dst, + ref_field, + location, + vir::AssignKind::Move, + false, + )?); // Allocate `box_content` - stmts.extend(self.encode_havoc_and_initialization(&box_content).with_span(span)?); - - // Initialize `box_content` stmts.extend( - self.encode_assign_operand( - &box_content, - &args[0], - location - )? + self.encode_havoc_and_initialization(&box_content) + .with_span(span)?, ); + + // Initialize `box_content` + stmts.extend(self.encode_assign_operand( + &box_content, + &args[0], + location, + )?); } - "std::cmp::PartialEq::eq" | - "core::cmp::PartialEq::eq" - if args.len() == 2 && - self.encoder.has_structural_eq_impl( - self.mir_encoder.get_operand_ty(&args[0]) - ) - => { + "std::cmp::PartialEq::eq" | "core::cmp::PartialEq::eq" + if args.len() == 2 + && self.encoder.has_structural_eq_impl( + self.mir_encoder.get_operand_ty(&args[0]), + ) => + { debug!("Encoding call of PartialEq::eq"); - stmts.extend( - self.encode_cmp_function_call( - called_def_id, - location, - term.source_info.span, - args, - destination, - target, - vir::BinaryOpKind::EqCmp, - call_substs, - )? - ); + stmts.extend(self.encode_cmp_function_call( + called_def_id, + location, + term.source_info.span, + args, + destination, + target, + vir::BinaryOpKind::EqCmp, + call_substs, + )?); } - "std::cmp::PartialEq::ne" | - "core::cmp::PartialEq::ne" - if args.len() == 2 && - self.encoder.has_structural_eq_impl( - self.mir_encoder.get_operand_ty(&args[0]) - ) - => { + "std::cmp::PartialEq::ne" | "core::cmp::PartialEq::ne" + if args.len() == 2 + && self.encoder.has_structural_eq_impl( + self.mir_encoder.get_operand_ty(&args[0]), + ) => + { debug!("Encoding call of PartialEq::ne"); - stmts.extend( - self.encode_cmp_function_call( - called_def_id, - location, - term.source_info.span, - args, - destination, - target, - vir::BinaryOpKind::NeCmp, - call_substs, - )? - ); + stmts.extend(self.encode_cmp_function_call( + called_def_id, + location, + term.source_info.span, + args, + destination, + target, + vir::BinaryOpKind::NeCmp, + call_substs, + )?); } - "std::ops::Fn::call" - | "core::ops::Fn::call" => { + "std::ops::Fn::call" | "core::ops::Fn::call" => { let cl_type: ty::Ty = call_substs[0].expect_ty(); match cl_type.kind() { ty::TyKind::Closure(cl_def_id, _) => { - debug!("Encoding call to closure {:?} with func {:?}", cl_def_id, func_const_val); + debug!( + "Encoding call to closure {:?} with func {:?}", + cl_def_id, func_const_val + ); stmts.extend(self.encode_impure_function_call( location, term.source_info.span, @@ -2590,18 +2689,15 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } "core::slice::::len" => { - stmts.extend( - self.encode_slice_len_call( - destination, - args, - location, - span, - )? - ); + stmts.extend(self.encode_slice_len_call( + destination, + args, + location, + span, + )?); } - "std::iter::Iterator::next" | - "core::iter::Iterator::next" => { + "std::iter::Iterator::next" | "core::iter::Iterator::next" => { return Err(SpannedEncodingError::unsupported( "iterators are not fully supported yet", term.source_info.span, @@ -2609,23 +2705,22 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } // TODO: use extern_spec - "core::ops::IndexMut::index_mut" | - "std::ops::IndexMut::index_mut" => { + "core::ops::IndexMut::index_mut" | "std::ops::IndexMut::index_mut" => { return Err(SpannedEncodingError::unsupported( "mutably slicing is not fully supported yet", term.source_info.span, )); } - "core::ops::Index::index" | - "std::ops::Index::index" => { + "core::ops::Index::index" | "std::ops::Index::index" => { stmts.extend( self.encode_sequence_index_call( destination, args, location, term.source_info.span, - ).with_span(span)? + ) + .with_span(span)?, ); } @@ -2633,7 +2728,10 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // The called method might be a trait method. // We try to resolve it to the concrete implementation // and type substitutions. - let (called_def_id, call_substs) = self.encoder.env().query + let (called_def_id, call_substs) = self + .encoder + .env() + .query .resolve_method_call(self.proc_def_id, called_def_id, call_substs); let is_pure_function = self.encoder.is_pure(called_def_id, Some(call_substs)) && @@ -2643,7 +2741,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { self.proc_def_id != called_def_id; if is_pure_function { let def_id = called_def_id; - let (function_name, _) = self.encoder + let (function_name, _) = self + .encoder .encode_pure_function_use(def_id, self.proc_def_id, call_substs) .with_default_span(term.source_info.span)?; debug!("Encoding pure function call '{}'", function_name); @@ -2653,29 +2752,25 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let _ = self.mir_encoder.encode_operand_expr(operand); } - stmts.extend( - self.encode_pure_function_call( - location, - term.source_info.span, - args, - destination, - target, - def_id, - call_substs, - )? - ); + stmts.extend(self.encode_pure_function_call( + location, + term.source_info.span, + args, + destination, + target, + def_id, + call_substs, + )?); } else { - stmts.extend( - self.encode_impure_function_call( - location, - term.source_info.span, - args, - destination, - target, - called_def_id, - call_substs, - )? - ); + stmts.extend(self.encode_impure_function_call( + location, + term.source_info.span, + args, + destination, + target, + called_def_id, + call_substs, + )?); } } } @@ -2712,12 +2807,12 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // Use local variables in the switch/if. // See Silicon issue https://github.com/viperproject/silicon/issues/356 let cond_var = self.cfg_method.add_fresh_local_var(vir::Type::Bool); - stmts.push(vir::Stmt::Assign( vir::Assign { + stmts.push(vir::Stmt::Assign(vir::Assign { target: cond_var.clone().into(), - source: self.mir_encoder.encode_operand_expr(cond) - .with_span( - self.mir_encoder.get_span_of_location(location) - )?, + source: self + .mir_encoder + .encode_operand_expr(cond) + .with_span(self.mir_encoder.get_span_of_location(location))?, kind: vir::AssignKind::Copy, })); @@ -2739,18 +2834,13 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { stmts.push(vir::Stmt::comment(format!("Rust assertion: {assert_msg}"))); if self.check_panics { - stmts.push(vir::Stmt::Assert( vir::Assert { + stmts.push(vir::Stmt::Assert(vir::Assert { expr: viper_guard, - position: self.register_error( - term.source_info.span, - error_ctxt, - ), + position: self.register_error(term.source_info.span, error_ctxt), })); } else { stmts.push(vir::Stmt::comment("This assertion will not be checked")); - stmts.push(vir::Stmt::Inhale( vir::Inhale { - expr: viper_guard, - })); + stmts.push(vir::Stmt::Inhale(vir::Inhale { expr: viper_guard })); }; (stmts, MirSuccessor::Goto(target)) @@ -2773,7 +2863,9 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { span: Span, ) -> SpannedEncodingResult> { assert!(args.len() == 1, "unexpected args to slice::len(): {args:?}"); - let slice_operand = self.mir_encoder.encode_operand_expr(&args[0]) + let slice_operand = self + .mir_encoder + .encode_operand_expr(&args[0]) .with_span(span)?; let mut stmts = vec![]; @@ -2783,28 +2875,30 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { stmts.push(vir::Stmt::label(label.clone())); let slice_ty_ref = self.mir_encoder.get_operand_ty(&args[0]); - let slice_ty = if let ty::TyKind::Ref(_, slice_ty, _) = slice_ty_ref.kind() { slice_ty } else { unreachable!() }; - let slice_types = self.encoder.encode_sequence_types(*slice_ty).with_span(span)?; + let slice_ty = if let ty::TyKind::Ref(_, slice_ty, _) = slice_ty_ref.kind() { + slice_ty + } else { + unreachable!() + }; + let slice_types = self + .encoder + .encode_sequence_types(*slice_ty) + .with_span(span)?; let rhs = slice_types.len(self.encoder, slice_operand); - let (encoded_lhs, encode_stmts, ty, _) = self.encode_place( - destination, - ArrayAccessKind::Mutable(None, location), - location - ).with_span(span)?; + let (encoded_lhs, encode_stmts, ty, _) = self + .encode_place( + destination, + ArrayAccessKind::Mutable(None, location), + location, + ) + .with_span(span)?; stmts.extend(encode_stmts); - stmts.extend( - self.encode_copy_value_assign( - encoded_lhs, - rhs, - ty, - location, - )? - ); + stmts.extend(self.encode_copy_value_assign(encoded_lhs, rhs, ty, location)?); - self.encode_transfer_args_permissions(location, args, &mut stmts, &label, false)?; + self.encode_transfer_args_permissions(location, args, &mut stmts, &label, false)?; // Store a label for permissions got back from the call debug!( @@ -2822,12 +2916,15 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { destination: mir::Place<'tcx>, args: &[mir::Operand<'tcx>], location: mir::Location, - error_span: Span + error_span: Span, ) -> EncodingResult> { // args[0] is the base array/slice, args[1] is the index // index is not specified exactly, as std::ops::Index::index is a trait method, so all we // know is there needs to be an impl Index for T - assert!(args.len() == 2, "unexpected args to sequence index call: {args:?}"); + assert!( + args.len() == 2, + "unexpected args to sequence index call: {args:?}" + ); let mut stmts = vec![]; @@ -2845,7 +2942,11 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { if !lhs_ty.is_slice_or_ref() && !lhs_ty.is_array_or_ref() { error_unsupported!("Non-slice LHS type '{:?}' not supported yet", lhs_ty); } - let mutability = if let ty::TyKind::Ref(_, _, mutability) = lhs_ty.kind() { mutability } else { unreachable!() }; + let mutability = if let ty::TyKind::Ref(_, _, mutability) = lhs_ty.kind() { + mutability + } else { + unreachable!() + }; let perm_amount = match mutability { Mutability::Mut => vir::PermAmount::Write, Mutability::Not => vir::PermAmount::Read, @@ -2855,23 +2956,32 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { stmts.push(vir_stmt!{ inhale [vir::Expr::pred_permission(encoded_lhs.clone(), perm_amount).unwrap()] }); let lhs_slice_ty = lhs_ty.peel_refs(); - let lhs_slice_expr = self.encoder.encode_value_expr(encoded_lhs.clone(), lhs_ty)?; + let lhs_slice_expr = self + .encoder + .encode_value_expr(encoded_lhs.clone(), lhs_ty)?; let base_seq = self.mir_encoder.encode_operand_place(&args[0])?.unwrap(); let base_seq_ty = self.mir_encoder.get_operand_ty(&args[0]); if !base_seq_ty.is_slice_or_ref() && !base_seq_ty.is_array_or_ref() { - error_unsupported!("Slicing is only supported for arrays/slices currently, not '{:?}'", base_seq_ty); + error_unsupported!( + "Slicing is only supported for arrays/slices currently, not '{:?}'", + base_seq_ty + ); } // base_seq is expected to be ref$Array$.. or ref$Slice$.., but lookup_pure wants the // contained Array$../Slice$.. let base_seq_expr = self.encoder.encode_value_expr(base_seq, base_seq_ty)?; - let enc_sequence_types = self.encoder.encode_sequence_types(base_seq_ty.peel_refs())?; + let enc_sequence_types = self + .encoder + .encode_sequence_types(base_seq_ty.peel_refs())?; - let j = vir_local!{ j: Int }; - let elem_snap_ty = self.encoder.encode_snapshot_type(enc_sequence_types.elem_ty_rs)?; + let j = vir_local! { j: Int }; + let elem_snap_ty = self + .encoder + .encode_snapshot_type(enc_sequence_types.elem_ty_rs)?; let rhs_lookup_j = enc_sequence_types.encode_lookup_pure_call( self.encoder, base_seq_expr.clone(), @@ -2882,7 +2992,11 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let encoded_idx = self.mir_encoder.encode_operand_place(&args[1])?.unwrap(); trace!("idx: {:?}", encoded_idx); let idx_ty = self.mir_encoder.get_operand_ty(&args[1]); - let idx_ident = self.encoder.env().name.get_absolute_item_name(idx_ty.ty_adt_def().unwrap().did()); + let idx_ident = self + .encoder + .env() + .name + .get_absolute_item_name(idx_ty.ty_adt_def().unwrap().did()); trace!("ident: {}", idx_ident); self.slice_created_at.insert(location, encoded_lhs); @@ -2892,16 +3006,32 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // TODO: there's fields like _5.f$start.val_int on `encoded_idx`, it just feels hacky to // manually re-do and hardcode them here when we probably just encoded the type // and the construction of the fields. - let usize_ty = self.encoder.env().tcx().mk_ty_from_kind(ty::TyKind::Uint(ty::UintTy::Usize)); + let usize_ty = self + .encoder + .env() + .tcx() + .mk_ty_from_kind(ty::TyKind::Uint(ty::UintTy::Usize)); let start = match &*idx_ident { - "std::ops::Range" | "core::ops::Range" | - "std::ops::RangeFrom" | "core::ops::RangeFrom" => { - let start_expr = self.encoder.encode_struct_field_value(encoded_idx.clone(), "start", usize_ty)?; + "std::ops::Range" + | "core::ops::Range" + | "std::ops::RangeFrom" + | "core::ops::RangeFrom" => { + let start_expr = self.encoder.encode_struct_field_value( + encoded_idx.clone(), + "start", + usize_ty, + )?; if self.check_panics { // Check indexing in bounds - stmts.push(vir::Stmt::Assert( vir::Assert { - expr: vir_expr!{ [start_expr] >= [vir::Expr::from(0usize)] }, - position: self.register_error(error_span, ErrorCtxt::SliceRangeBoundsCheckAssert("the range start value may be smaller than 0 when slicing".to_string())), + stmts.push(vir::Stmt::Assert(vir::Assert { + expr: vir_expr! { [start_expr] >= [vir::Expr::from(0usize)] }, + position: self.register_error( + error_span, + ErrorCtxt::SliceRangeBoundsCheckAssert( + "the range start value may be smaller than 0 when slicing" + .to_string(), + ), + ), })); } start_expr @@ -2909,71 +3039,102 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // RangeInclusive is wierdly differnet to all of the other Range*s in that the struct fields are private // and it is created with a new() fn and start/end are accessed with getter fns // See https://github.com/rust-lang/rust/issues/67371 for why this is the case... - "std::ops::RangeInclusive" | "core::ops::RangeInclusive" => return Err( - EncodingError::unsupported("slicing with RangeInclusive (e.g. [x..=y]) currently not supported".to_string()) - ), - "std::ops::RangeTo" | "core::ops::RangeTo" | - "std::ops::RangeFull" | "core::ops::RangeFull" | - "std::ops::RangeToInclusive" | "core::ops::RangeToInclusive" => vir::Expr::from(0usize), - _ => unreachable!("{}", idx_ident) + "std::ops::RangeInclusive" | "core::ops::RangeInclusive" => { + return Err(EncodingError::unsupported( + "slicing with RangeInclusive (e.g. [x..=y]) currently not supported" + .to_string(), + )) + } + "std::ops::RangeTo" + | "core::ops::RangeTo" + | "std::ops::RangeFull" + | "core::ops::RangeFull" + | "std::ops::RangeToInclusive" + | "core::ops::RangeToInclusive" => vir::Expr::from(0usize), + _ => unreachable!("{}", idx_ident), }; let end = match &*idx_ident { - "std::ops::Range" | "core::ops::Range" | - "std::ops::RangeTo" | "core::ops::RangeTo" => { - let end_expr = self.encoder.encode_struct_field_value(encoded_idx, "end", usize_ty)?; + "std::ops::Range" | "core::ops::Range" | "std::ops::RangeTo" | "core::ops::RangeTo" => { + let end_expr = + self.encoder + .encode_struct_field_value(encoded_idx, "end", usize_ty)?; if self.check_panics { // Check indexing in bounds - stmts.push(vir::Stmt::Assert( vir::Assert { - expr: vir_expr!{ [end_expr] <= [original_len] }, - position: self.register_error(error_span, ErrorCtxt::SliceRangeBoundsCheckAssert("the range end value may be out of bounds when slicing".to_string())), + stmts.push(vir::Stmt::Assert(vir::Assert { + expr: vir_expr! { [end_expr] <= [original_len] }, + position: self.register_error( + error_span, + ErrorCtxt::SliceRangeBoundsCheckAssert( + "the range end value may be out of bounds when slicing".to_string(), + ), + ), })); } end_expr } - "std::ops::RangeInclusive" | "core::ops::RangeInclusive" => return Err( - EncodingError::unsupported("slicing with RangeInclusive (e.g. [x..=y]) currently not supported".to_string()) - ), + "std::ops::RangeInclusive" | "core::ops::RangeInclusive" => { + return Err(EncodingError::unsupported( + "slicing with RangeInclusive (e.g. [x..=y]) currently not supported" + .to_string(), + )) + } "std::ops::RangeToInclusive" | "core::ops::RangeToInclusive" => { - let end_expr = self.encoder.encode_struct_field_value(encoded_idx, "end", usize_ty)?; - let end_expr = vir_expr!{ [end_expr] + [vir::Expr::from(1usize)] }; + let end_expr = + self.encoder + .encode_struct_field_value(encoded_idx, "end", usize_ty)?; + let end_expr = vir_expr! { [end_expr] + [vir::Expr::from(1usize)] }; if self.check_panics { // Check indexing in bounds - stmts.push(vir::Stmt::Assert( vir::Assert { - expr: vir_expr!{ [end_expr] <= [original_len] }, - position: self.register_error(error_span, ErrorCtxt::SliceRangeBoundsCheckAssert("the range end value may be out of bounds when slicing".to_string())), + stmts.push(vir::Stmt::Assert(vir::Assert { + expr: vir_expr! { [end_expr] <= [original_len] }, + position: self.register_error( + error_span, + ErrorCtxt::SliceRangeBoundsCheckAssert( + "the range end value may be out of bounds when slicing".to_string(), + ), + ), })); } end_expr } - "std::ops::RangeFrom" | "core::ops::RangeFrom" | - "std::ops::RangeFull" | "core::ops::RangeFull" => original_len, - _ => unreachable!("{}", idx_ident) + "std::ops::RangeFrom" + | "core::ops::RangeFrom" + | "std::ops::RangeFull" + | "core::ops::RangeFull" => original_len, + _ => unreachable!("{}", idx_ident), }; trace!("start: {}, end: {}", start, end); let slice_types_lhs = self.encoder.encode_sequence_types(lhs_slice_ty)?; - let elem_snap_ty = self.encoder.encode_snapshot_type(slice_types_lhs.elem_ty_rs)?; + let elem_snap_ty = self + .encoder + .encode_snapshot_type(slice_types_lhs.elem_ty_rs)?; // length - let length = vir_expr!{ [end] - [start] }; + let length = vir_expr! { [end] - [start] }; if self.check_panics { // start must be leq than end if idx_ident != "std::ops::RangeFull" && idx_ident != "core::ops::RangeFull" { - stmts.push(vir::Stmt::Assert( vir::Assert { - expr: vir_expr!{ [start] <= [end] }, - position: self.register_error(error_span, ErrorCtxt::SliceRangeBoundsCheckAssert("the range end may be smaller than the start when slicing".to_string())), + stmts.push(vir::Stmt::Assert(vir::Assert { + expr: vir_expr! { [start] <= [end] }, + position: self.register_error( + error_span, + ErrorCtxt::SliceRangeBoundsCheckAssert( + "the range end may be smaller than the start when slicing".to_string(), + ), + ), })); } } let slice_len_call = slice_types_lhs.len(self.encoder, lhs_slice_expr.clone()); - stmts.push(vir_stmt!{ + stmts.push(vir_stmt! { inhale [vir_expr!{ [slice_len_call] == [length] }] }); // lookup_pure: contents - let i = vir_local!{ i: Int }; + let i = vir_local! { i: Int }; let i_var: vir::Expr = i.clone().into(); let j_var: vir::Expr = j.clone().into(); @@ -2989,8 +3150,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // NOTE: the lhs_ and rhs_ here refer to the moment the slice is created, so lhs_lookup is // lookup on the slice currently being created, and rhs_lookup is on the array or slice // being sliced - let lookup_eq = vir_expr!{ [lhs_lookup_i] == [rhs_lookup_j] }; - let indices = vir_expr!{ + let lookup_eq = vir_expr! { [lhs_lookup_i] == [rhs_lookup_j] }; + let indices = vir_expr! { [vir_expr!{ [vir::Expr::from(0usize)] <= [i_var] }] && ( [vir_expr!{ [i_var] < [slice_len_call] }] && @@ -3005,7 +3166,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { }; // forall i: Int, j: Int :: { lhs_lookup(i), rhs_lookup(j) } 0 <= i && i < slice$len && j == i + start && start <= j && j < end ==> lhs_lookup(i) == rhs_lookup(j) - stmts.push(vir_stmt!{ + stmts.push(vir_stmt! { inhale [ vir::Expr::forall( vec![i, j], @@ -3015,7 +3176,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ] }); - self.encode_transfer_args_permissions(location, args, &mut stmts, &label, false)?; + self.encode_transfer_args_permissions(location, args, &mut stmts, &label, false)?; // Store a label for permissions got back from the call debug!( "Pure function call location {:?} has label {}", @@ -3046,37 +3207,40 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ) -> SpannedEncodingResult> { let arg_ty = self.mir_encoder.get_operand_ty(&args[0]); - if self.encoder.supports_snapshot_equality(arg_ty).with_span(call_site_span)? { - let lhs = self.mir_encoder.encode_operand_expr(&args[0]) + if self + .encoder + .supports_snapshot_equality(arg_ty) + .with_span(call_site_span)? + { + let lhs = self + .mir_encoder + .encode_operand_expr(&args[0]) .with_span(call_site_span)?; - let rhs = self.mir_encoder.encode_operand_expr(&args[1]) + let rhs = self + .mir_encoder + .encode_operand_expr(&args[1]) .with_span(call_site_span)?; let expr = match bin_op { - vir::BinaryOpKind::EqCmp => vir::Expr::eq_cmp( - vir::Expr::snap_app(lhs), - vir::Expr::snap_app(rhs), - ), - vir::BinaryOpKind::NeCmp => vir::Expr::ne_cmp( - vir::Expr::snap_app(lhs), - vir::Expr::snap_app(rhs), - ), - _ => unreachable!() + vir::BinaryOpKind::EqCmp => { + vir::Expr::eq_cmp(vir::Expr::snap_app(lhs), vir::Expr::snap_app(rhs)) + } + vir::BinaryOpKind::NeCmp => { + vir::Expr::ne_cmp(vir::Expr::snap_app(lhs), vir::Expr::snap_app(rhs)) + } + _ => unreachable!(), }; - let (target_value, mut stmts) = self.encode_pure_function_call_lhs_value(destination, target, location) + let (target_value, mut stmts) = self + .encode_pure_function_call_lhs_value(destination, target, location) .with_span(call_site_span)?; let inhaled_expr = vir::Expr::eq_cmp(target_value, expr); - let (call_stmts, label) = self.encode_pure_function_call_site( - location, - destination, - target, - inhaled_expr, - )?; + let (call_stmts, label) = + self.encode_pure_function_call_site(location, destination, target, inhaled_expr)?; stmts.extend(call_stmts); - self.encode_transfer_args_permissions(location, args, &mut stmts, &label, false)?; + self.encode_transfer_args_permissions(location, args, &mut stmts, &label, false)?; Ok(stmts) } else { @@ -3148,7 +3312,10 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { .env() .name .get_absolute_item_name(called_def_id); - debug!("Encoding non-pure function call '{}' with args {:?} and substs {:?}", full_func_proc_name, mir_args, substs); + debug!( + "Encoding non-pure function call '{}' with args {:?} and substs {:?}", + full_func_proc_name, mir_args, substs + ); // Spans for fake exprs that cannot be encoded in viper let mut fake_expr_spans: FxHashMap = FxHashMap::default(); @@ -3160,8 +3327,10 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // - the VIR Local that will hold hold the argument before the call // - the type of the argument // - if not constant, the VIR expression for the argument - let mut operands: Vec<(&mir::Operand<'tcx>, Local, ty::Ty<'tcx>, Option)> = vec![]; - let mut encoded_operands = mir_args.iter() + let mut operands: Vec<(&mir::Operand<'tcx>, Local, ty::Ty<'tcx>, Option)> = + vec![]; + let mut encoded_operands = mir_args + .iter() .map(|arg| self.mir_encoder.encode_operand_place(arg)) .collect::>, _>>() .with_span(call_site_span)?; @@ -3173,12 +3342,10 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let cl_ty = self.mir_encoder.get_operand_ty(&mir_args[0]); operands.push(( &mir_args[0], - mir_args[0].place() + mir_args[0] + .place() .and_then(|place| place.as_local()) - .map_or_else( - || self.locals.get_fresh(cl_ty), - |local| local.into() - ), + .map_or_else(|| self.locals.get_fresh(cl_ty), |local| local.into()), cl_ty, encoded_operands[0].take(), )); @@ -3205,7 +3372,11 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // TODO: weird fix for closure call substitutions, we need to // prepend the identity substs of the containing method ... - substs = self.encoder.env().tcx().mk_substs_from_iter(self.substs.iter().chain(substs)); + substs = self + .encoder + .env() + .tcx() + .mk_substs_from_iter(self.substs.iter().chain(substs)); } else { for (arg, encoded_operand) in mir_args.iter().zip(encoded_operands.iter_mut()) { let arg_ty = self.mir_encoder.get_operand_ty(arg); @@ -3213,10 +3384,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { arg, arg.place() .and_then(|place| place.as_local()) - .map_or_else( - || self.locals.get_fresh(arg_ty), - |local| local.into() - ), + .map_or_else(|| self.locals.get_fresh(arg_ty), |local| local.into()), arg_ty, encoded_operand.take(), )); @@ -3252,10 +3420,12 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { debug!("arg: {} {}", arg_place, place); if !self.encoder.is_pure(called_def_id, Some(substs)) { type_invs.push( - self.encoder.encode_invariant_func_app( - arg_ty, - vir::Expr::snap_app(place.clone()), - ).with_span(call_site_span)?, + self.encoder + .encode_invariant_func_app( + arg_ty, + vir::Expr::snap_app(place.clone()), + ) + .with_span(call_site_span)?, ); } fake_exprs.insert(arg_place, place); @@ -3265,28 +3435,28 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // We have a constant. constant_args.push(arg_place.clone()); - let val_field = self.encoder.encode_value_field(arg_ty).with_span(call_site_span)?; + let val_field = self + .encoder + .encode_value_field(arg_ty) + .with_span(call_site_span)?; // TODO: String constants are not being encoded currently, // instead just inhale the permission to the value. if !is_str(arg_ty) { - let arg_val_expr = self.mir_encoder.encode_operand_expr(mir_arg) + let arg_val_expr = self + .mir_encoder + .encode_operand_expr(mir_arg) .with_span(call_site_span)?; debug!("arg_val_expr: {} {}", arg_place, arg_val_expr); fake_exprs.insert(arg_place.clone().field(val_field), arg_val_expr); } else { - fake_expr_spans.insert( - arg, - call_site_span - ); - stmts.push(vir::Stmt::Inhale ( - vir::Inhale { - expr: vir::Expr::acc_permission( - arg_place.clone().field(val_field), - vir::PermAmount::Read - ) - } - )); + fake_expr_spans.insert(arg, call_site_span); + stmts.push(vir::Stmt::Inhale(vir::Inhale { + expr: vir::Expr::acc_permission( + arg_place.clone().field(val_field), + vir::PermAmount::Read, + ), + })); } let in_loop = self.loop_encoder.get_loop_depth(location.block) > 0; if in_loop { @@ -3305,7 +3475,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let (target_local, encoded_target) = { if target.is_some() { - let (encoded_target, pre_stmts, ty, _) = self.encode_place(destination, ArrayAccessKind::Shared, location)?; + let (encoded_target, pre_stmts, ty, _) = + self.encode_place(destination, ArrayAccessKind::Shared, location)?; stmts.extend(pre_stmts); let target_local = if let Some(target_local) = destination.as_local() { @@ -3322,9 +3493,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // The return type is Never // This means that the function call never returns // So, we `assume false` after the function call - stmts_after.push(vir::Stmt::Inhale( vir::Inhale { - expr: false.into() - })); + stmts_after.push(vir::Stmt::Inhale(vir::Inhale { expr: false.into() })); // Return a dummy local variable let never_ty = self.encoder.env().tcx().mk_ty_from_kind(ty::TyKind::Never); (self.locals.get_fresh(never_ty), None) @@ -3352,7 +3521,9 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } } */ - vir::Expr::PredicateAccessPredicate( vir::PredicateAccessPredicate {ref argument, ..} ) => { + vir::Expr::PredicateAccessPredicate( + vir::PredicateAccessPredicate { ref argument, .. }, + ) => { if argument.is_local() && const_arg_vars.contains(argument) { // Skip predicate permission true.into() @@ -3370,13 +3541,15 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { }; let procedure_contract = { - self.encoder.get_procedure_contract_for_call( - self.proc_def_id, - called_def_id, - &arguments, - target_local, - substs, - ).with_span(call_site_span)? + self.encoder + .get_procedure_contract_for_call( + self.proc_def_id, + called_def_id, + &arguments, + target_local, + substs, + ) + .with_span(call_site_span)? }; assert_one_magic_wand(procedure_contract.borrow_infos.len()).with_span(call_site_span)?; @@ -3386,28 +3559,27 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // Havoc and inhale variables that store constants for constant_arg in &constant_args { - stmts.extend(self.encode_havoc_and_initialization(constant_arg).with_span(call_site_span)?); + stmts.extend( + self.encode_havoc_and_initialization(constant_arg) + .with_span(call_site_span)?, + ); } // Encode precondition. - let ( - pre_type_spec, - pre_mandatory_type_spec, - pre_invs_spec, - pre_func_spec, - ) = self.encode_precondition_expr(&procedure_contract, substs, fake_expr_spans)?; + let (pre_type_spec, pre_mandatory_type_spec, pre_invs_spec, pre_func_spec) = + self.encode_precondition_expr(&procedure_contract, substs, fake_expr_spans)?; let pos = self.register_error(call_site_span, ErrorCtxt::ExhaleMethodPrecondition); - stmts.push(vir::Stmt::Assert( vir::Assert { + stmts.push(vir::Stmt::Assert(vir::Assert { expr: replace_fake_exprs(pre_func_spec), position: pos, })); - stmts.push(vir::Stmt::Assert( vir::Assert { + stmts.push(vir::Stmt::Assert(vir::Assert { expr: replace_fake_exprs(pre_invs_spec), position: pos, })); let pre_perm_spec = replace_fake_exprs(pre_type_spec); assert!(!pos.is_default()); - stmts.push(vir::Stmt::Exhale( vir::Exhale { + stmts.push(vir::Stmt::Exhale(vir::Exhale { expr: pre_perm_spec.remove_read_permissions(), position: pos, })); @@ -3425,10 +3597,10 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let from_place = perm.get_place().unwrap().clone(); let to_place = from_place.clone().old(pre_label.clone()); let old_perm = perm.replace_place(&from_place, &to_place); - stmts.push(vir::Stmt::TransferPerm( vir::TransferPerm { + stmts.push(vir::Stmt::TransferPerm(vir::TransferPerm { left: from_place, right: to_place, - unchecked: true + unchecked: true, })); pre_mandatory_perms_old.push(old_perm); } @@ -3479,31 +3651,31 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { .collect(); let post_perm_spec = replace_fake_exprs(post_type_spec); - stmts.push(vir::Stmt::Inhale( vir::Inhale { + stmts.push(vir::Stmt::Inhale(vir::Inhale { expr: post_perm_spec.remove_read_permissions(), })); if let Some(access) = return_type_spec { - stmts.push(vir::Stmt::Inhale( vir::Inhale { + stmts.push(vir::Stmt::Inhale(vir::Inhale { expr: replace_fake_exprs(access), })); } for (from_place, to_place) in read_transfer { - stmts.push(vir::Stmt::TransferPerm( vir::TransferPerm { + stmts.push(vir::Stmt::TransferPerm(vir::TransferPerm { left: replace_fake_exprs(from_place), right: replace_fake_exprs(to_place), unchecked: true, })); } - stmts.push(vir::Stmt::Inhale( vir::Inhale { + stmts.push(vir::Stmt::Inhale(vir::Inhale { expr: replace_fake_exprs(post_invs_spec), })); - stmts.push(vir::Stmt::Inhale( vir::Inhale { + stmts.push(vir::Stmt::Inhale(vir::Inhale { expr: replace_fake_exprs(post_func_spec), })); // Exhale the permissions that were moved into magic wands. assert!(!pos.is_default()); - stmts.push(vir::Stmt::Exhale( vir::Exhale { + stmts.push(vir::Stmt::Exhale(vir::Exhale { expr: pre_mandatory_perm_spec, position: pos, })); @@ -3531,14 +3703,18 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { called_def_id: ProcedureDefId, call_substs: SubstsRef<'tcx>, ) -> SpannedEncodingResult> { - let (function_name, return_type) = self.encoder.encode_pure_function_use(called_def_id, self.proc_def_id, call_substs) + let (function_name, return_type) = self + .encoder + .encode_pure_function_use(called_def_id, self.proc_def_id, call_substs) .with_span(call_site_span)?; debug!("Encoding pure function call '{}'", function_name); assert!(target.is_some()); let mut arg_exprs = vec![]; for operand in args.iter() { - let arg_expr = self.mir_encoder.encode_operand_expr(operand) + let arg_expr = self + .mir_encoder + .encode_operand_expr(operand) .with_span(call_site_span)?; arg_exprs.push(arg_expr); } @@ -3575,7 +3751,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { .iter() .enumerate() .map(|(i, arg)| { - self.mir_encoder.encode_operand_expr_type(arg) + self.mir_encoder + .encode_operand_expr_type(arg) .map(|ty| vir::LocalVar::new(format!("x{i}"), ty)) }) .collect::>() @@ -3583,7 +3760,10 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let pos = self.register_error(call_site_span, ErrorCtxt::PureFunctionCall); - let type_arguments = self.encoder.encode_generic_arguments(called_def_id, call_substs).with_span(call_site_span)?; + let type_arguments = self + .encoder + .encode_generic_arguments(called_def_id, call_substs) + .with_span(call_site_span)?; let func_call = vir::Expr::func_app( function_name, @@ -3591,32 +3771,27 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { arg_exprs, formal_args, return_type.clone(), - pos + pos, ); - let (target_value, mut stmts) = self.encode_pure_function_call_lhs_value(destination, target, location) + let (target_value, mut stmts) = self + .encode_pure_function_call_lhs_value(destination, target, location) .with_span(call_site_span)?; let inhaled_expr = if return_type.is_domain() || return_type.is_snapshot() { - let (target_place, pre_stmts) = self.encode_pure_function_call_lhs_place(destination, target, location)?; + let (target_place, pre_stmts) = + self.encode_pure_function_call_lhs_place(destination, target, location)?; stmts.extend(pre_stmts); - vir::Expr::eq_cmp( - vir::Expr::snap_app(target_place), - func_call, - ) + vir::Expr::eq_cmp(vir::Expr::snap_app(target_place), func_call) } else { vir::Expr::eq_cmp(target_value, func_call) }; - let (call_stmts, label) = self.encode_pure_function_call_site( - location, - destination, - target, - inhaled_expr - )?; + let (call_stmts, label) = + self.encode_pure_function_call_site(location, destination, target, inhaled_expr)?; stmts.extend(call_stmts); - self.encode_transfer_args_permissions(location, args, &mut stmts, &label, false)?; + self.encode_transfer_args_permissions(location, args, &mut stmts, &label, false)?; Ok(stmts) } @@ -3628,8 +3803,12 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ) -> SpannedEncodingResult<(vir::Expr, Vec)> { let span = self.mir_encoder.get_span_of_location(location); assert!(target.is_some()); - let (encoded_place, pre_stmts, ty, _) = self.encode_place(destination, ArrayAccessKind::Shared, location)?; - let encoded_lhs_value = self.encoder.encode_value_expr(encoded_place, ty).with_span(span)?; + let (encoded_place, pre_stmts, ty, _) = + self.encode_place(destination, ArrayAccessKind::Shared, location)?; + let encoded_lhs_value = self + .encoder + .encode_value_expr(encoded_place, ty) + .with_span(span)?; Ok((encoded_lhs_value, pre_stmts)) } @@ -3640,7 +3819,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { location: mir::Location, ) -> SpannedEncodingResult<(vir::Expr, Vec)> { assert!(target.is_some()); - let (encoded, pre_stmts, _, _) = self.encode_place(destination, ArrayAccessKind::Shared, location)?; + let (encoded, pre_stmts, _, _) = + self.encode_place(destination, ArrayAccessKind::Shared, location)?; Ok((encoded, pre_stmts)) } @@ -3659,7 +3839,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { stmts.push(vir::Stmt::label(label.clone())); // Havoc the content of the lhs - let (target_place, pre_stmts) = self.encode_pure_function_call_lhs_place(destination, target, location)?; + let (target_place, pre_stmts) = + self.encode_pure_function_call_lhs_place(destination, target, location)?; stmts.extend(pre_stmts); stmts.extend(self.encode_havoc(&target_place).with_span(span)?); let type_predicate = self @@ -3667,16 +3848,12 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { .encode_place_predicate_permission(target_place.clone(), vir::PermAmount::Write) .unwrap(); - stmts.push(vir::Stmt::Inhale( vir::Inhale { + stmts.push(vir::Stmt::Inhale(vir::Inhale { expr: type_predicate, })); // Initialize the lhs - stmts.push( - vir::Stmt::Inhale( vir::Inhale { - expr: call_result, - }) - ); + stmts.push(vir::Stmt::Inhale(vir::Inhale { expr: call_result })); // Store a label for permissions got back from the call debug!( @@ -3700,17 +3877,15 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let span = self.mir_encoder.get_span_of_location(location); for operand in args.iter() { let operand_ty = self.mir_encoder.get_operand_ty(operand); - let operand_place = self.mir_encoder.encode_operand_place(operand) + let operand_place = self + .mir_encoder + .encode_operand_place(operand) .with_span(span)?; match (operand_place, &operand_ty.kind()) { - ( - Some(ref place), - ty::TyKind::RawPtr(ty::TypeAndMut { - ty: inner_ty, .. - }), - ) + (Some(ref place), ty::TyKind::RawPtr(ty::TypeAndMut { ty: inner_ty, .. })) | (Some(ref place), ty::TyKind::Ref(_, inner_ty, _)) => { - let ref_field = self.encoder + let ref_field = self + .encoder .encode_dereference_field(*inner_ty) .with_span(span)?; let ref_place = place.clone().field(ref_field); @@ -3756,9 +3931,11 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { /// Encode permissions that are implicitly carried by the given local variable. /// `override_span` is used for local vars for fake expressions - fn encode_local_variable_permission(&self, local: Local, override_span: Option) - -> SpannedEncodingResult - { + fn encode_local_variable_permission( + &self, + local: Local, + override_span: Option, + ) -> SpannedEncodingResult { Ok(match self.locals.get_type(local).kind() { ty::TyKind::Ref(_, ty, mutability) => { // Use unfolded references. @@ -3772,13 +3949,12 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { "In ProcedureEncoder::encode_local_variable_permission the local \ {local:?} is fake but override_span is None", ), - self.mir.span + self.mir.span, )); } self.mir_encoder.get_local_span(local.into()) }; - let field = self.encoder.encode_dereference_field(*ty) - .with_span(span)?; + let field = self.encoder.encode_dereference_field(*ty).with_span(span)?; let place = vir::Expr::from(encoded_local).field(field); let perm_amount = match mutability { Mutability::Mut => vir::PermAmount::Write, @@ -3809,13 +3985,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { &self, contract: &ProcedureContract<'tcx>, substs: SubstsRef<'tcx>, - override_spans: FxHashMap // spans for fake locals - ) -> SpannedEncodingResult<( - vir::Expr, - Vec, - vir::Expr, - vir::Expr, - )> { + override_spans: FxHashMap, // spans for fake locals + ) -> SpannedEncodingResult<(vir::Expr, Vec, vir::Expr, vir::Expr)> { let borrow_infos = &contract.borrow_infos; let maybe_blocked_paths = if !borrow_infos.is_empty() { assert_eq!( @@ -3853,12 +4024,15 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { type_spec.push(access); } }; - let access = self.encode_local_variable_permission( - *local, - override_spans.get(local).copied() - )?; + let access = + self.encode_local_variable_permission(*local, override_spans.get(local).copied())?; match access { - vir::Expr::BinOp( vir::BinOp {op_kind: vir::BinaryOpKind::And, left: box access1, right: box access2, ..} ) => { + vir::Expr::BinOp(vir::BinOp { + op_kind: vir::BinaryOpKind::And, + left: box access1, + right: box access2, + .. + }) => { add(access1); add(access2); } @@ -3873,24 +4047,26 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { .map(|local| self.encode_prusti_local(*local).into()) .collect(); - let func_spec: Vec = contract.functional_precondition( - self.encoder.env(), - substs, - ).iter() - .map(|(assertion, assertion_substs)| self.encoder.encode_assertion( - assertion, - None, - &encoded_args, - None, - false, - self.proc_def_id, - assertion_substs, - )) + let func_spec: Vec = contract + .functional_precondition(self.encoder.env(), substs) + .iter() + .map(|(assertion, assertion_substs)| { + self.encoder.encode_assertion( + assertion, + None, + &encoded_args, + None, + false, + self.proc_def_id, + assertion_substs, + ) + }) .collect::, _>>()?; // TODO(tymap): do this with the previous step ... let precondition_spans = MultiSpan::from_spans( - contract.functional_precondition(self.encoder.env(), substs) + contract + .functional_precondition(self.encoder.env(), substs) .iter() .map(|(ts, _)| self.encoder.env().query.get_def_span(ts)) .collect(), @@ -3903,10 +4079,9 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let ty = self.locals.get_type(*arg); if !ty.is_unsafe_ptr() && !self.encoder.is_pure(contract.def_id, Some(substs)) { invs_spec.push( - self.encoder.encode_invariant_func_app( - ty, - self.encode_prusti_local(*arg).into(), - ).with_span(precondition_spans.clone())? + self.encoder + .encode_invariant_func_app(ty, self.encode_prusti_local(*arg).into()) + .with_span(precondition_spans.clone())?, ); } } @@ -3927,13 +4102,15 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { Option, )> { // Encode arguments and return - let encoded_args = self.procedure_contract() + let encoded_args = self + .procedure_contract() .args .iter() .map(|local| self.encode_prusti_local(*local).into()) .collect::>(); let encoded_return = self - .encode_prusti_local(self.procedure_contract().returned_value).into(); + .encode_prusti_local(self.procedure_contract().returned_value) + .into(); debug!("procedure_contract: {:?}", self.procedure_contract()); @@ -3944,41 +4121,53 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { if let SpecificationItem::Refined(from, to) = &procedure_spec.pres { // See comment in `ProcedureContractGeneric::functional_precondition`. - let trait_substs = self.encoder.env().query.find_trait_method_substs( - self.proc_def_id, - self.substs, - ).unwrap().1; + let trait_substs = self + .encoder + .env() + .query + .find_trait_method_substs(self.proc_def_id, self.substs) + .unwrap() + .1; - let from_pre = from.iter() - .map(|spec| self.encoder.encode_assertion( - spec, - None, - &encoded_args, - None, - false, - self.proc_def_id, - trait_substs, - )) + let from_pre = from + .iter() + .map(|spec| { + self.encoder.encode_assertion( + spec, + None, + &encoded_args, + None, + false, + self.proc_def_id, + trait_substs, + ) + }) .collect::, _>>()? .into_iter() .conjoin(); - let to_pre = to.iter() - .map(|spec| self.encoder.encode_assertion( - spec, - None, - &encoded_args, - None, - false, - self.proc_def_id, - self.substs, - )) + let to_pre = to + .iter() + .map(|spec| { + self.encoder.encode_assertion( + spec, + None, + &encoded_args, + None, + false, + self.proc_def_id, + self.substs, + ) + }) .collect::, _>>()? .into_iter() .conjoin(); // The spans are used for error reporting - let spec_functions_span = MultiSpan::from_spans(from.iter().chain(to.iter()) - .map(|spec_def_id| self.encoder.env().query.get_def_span(spec_def_id)).collect() + let spec_functions_span = MultiSpan::from_spans( + from.iter() + .chain(to.iter()) + .map(|spec_def_id| self.encoder.env().query.get_def_span(spec_def_id)) + .collect(), ); weakening = Some(RefinementCheckExpr { @@ -3989,43 +4178,53 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { if let SpecificationItem::Refined(from, to) = &procedure_spec.posts { // See comment in `ProcedureContractGeneric::functional_precondition`. - let trait_substs = self.encoder.env().query.find_trait_method_substs( - self.proc_def_id, - self.substs, - ).unwrap().1; + let trait_substs = self + .encoder + .env() + .query + .find_trait_method_substs(self.proc_def_id, self.substs) + .unwrap() + .1; let from_post = from .iter() - .map(|spec| self.encoder.encode_assertion( - spec, - Some(pre_label), - &encoded_args, - Some(&encoded_return), - false, - self.proc_def_id, - trait_substs, - )) + .map(|spec| { + self.encoder.encode_assertion( + spec, + Some(pre_label), + &encoded_args, + Some(&encoded_return), + false, + self.proc_def_id, + trait_substs, + ) + }) .collect::, _>>()? .into_iter() .conjoin(); let to_post = to .iter() - .map(|spec| self.encoder.encode_assertion( - spec, - Some(pre_label), - &encoded_args, - Some(&encoded_return), - false, - self.proc_def_id, - self.substs, - )) + .map(|spec| { + self.encoder.encode_assertion( + spec, + Some(pre_label), + &encoded_args, + Some(&encoded_return), + false, + self.proc_def_id, + self.substs, + ) + }) .collect::, _>>()? .into_iter() .conjoin(); // The spans are used for error reporting - let spec_functions_span = MultiSpan::from_spans(from.iter().chain(to.iter()) - .map(|spec_def_id| self.encoder.env().query.get_def_span(spec_def_id)).collect() + let spec_functions_span = MultiSpan::from_spans( + from.iter() + .chain(to.iter()) + .map(|spec_def_id| self.encoder.env().query.get_def_span(spec_def_id)) + .collect(), ); let strengthening_expr = self.wrap_arguments_into_old( @@ -4063,51 +4262,47 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ) -> SpannedEncodingResult<()> { self.cfg_method .add_stmt(start_cfg_block, vir::Stmt::comment("Preconditions:")); - let (type_spec, mandatory_type_spec, invs_spec, func_spec) = - self.encode_precondition_expr( + let (type_spec, mandatory_type_spec, invs_spec, func_spec) = self + .encode_precondition_expr( self.procedure_contract(), self.substs, - FxHashMap::default() + FxHashMap::default(), )?; self.cfg_method.add_stmt( start_cfg_block, - vir::Stmt::Inhale( vir::Inhale { - expr: type_spec}), + vir::Stmt::Inhale(vir::Inhale { expr: type_spec }), ); self.cfg_method.add_stmt( start_cfg_block, - vir::Stmt::Inhale( vir::Inhale { + vir::Stmt::Inhale(vir::Inhale { expr: mandatory_type_spec.into_iter().conjoin(), }), ); self.cfg_method.add_stmt( start_cfg_block, - vir::Stmt::Inhale( vir::Inhale { - expr: invs_spec, - }), + vir::Stmt::Inhale(vir::Inhale { expr: invs_spec }), ); // Weakening assertion must be put before inhaling the precondition, otherwise the weakening // soundness check becomes trivially satisfied. if let Some(weakening_spec) = weakening_spec { - let pos = self.register_error(weakening_spec.spec_functions_span, ErrorCtxt::AssertMethodPreconditionWeakening); + let pos = self.register_error( + weakening_spec.spec_functions_span, + ErrorCtxt::AssertMethodPreconditionWeakening, + ); self.cfg_method.add_stmt( start_cfg_block, - vir::Stmt::Assert( vir::Assert { + vir::Stmt::Assert(vir::Assert { expr: weakening_spec.refinement_check_expr, - position: pos + position: pos, }), ); } self.cfg_method.add_stmt( start_cfg_block, - vir::Stmt::Inhale( vir::Inhale { - expr: func_spec - }), - ); - self.cfg_method.add_stmt( - start_cfg_block, - vir::Stmt::label(PRECONDITION_LABEL), + vir::Stmt::Inhale(vir::Inhale { expr: func_spec }), ); + self.cfg_method + .add_stmt(start_cfg_block, vir::Stmt::label(PRECONDITION_LABEL)); Ok(()) } @@ -4156,9 +4351,9 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { Mutability::Not => vir::PermAmount::Read, Mutability::Mut => vir::PermAmount::Write, }; - let (place_expr, place_ty, _) = self.encode_generic_place( - contract.def_id, location, place - ).with_span(span)?; + let (place_expr, place_ty, _) = self + .encode_generic_place(contract.def_id, location, place) + .with_span(span)?; let vir_access = vir::Expr::pred_permission(place_expr.clone().old(label), perm_amount).unwrap(); if !self.encoder.is_pure(contract.def_id, Some(substs)) { @@ -4181,7 +4376,12 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { .iter() .map(|(place, mutability)| encode_place_perm(*place, *mutability, pre_label)) .collect::>()?; - if let Some(typed::Pledge { reference, lhs: body_lhs, rhs: body_rhs}) = pledges.first() { + if let Some(typed::Pledge { + reference, + lhs: body_lhs, + rhs: body_rhs, + }) = pledges.first() + { debug!( "pledge reference={:?} lhs={:?} rhs={:?}", reference, body_lhs, body_rhs @@ -4225,9 +4425,9 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { &encoded_args, )?; let ty = self.locals.get_type(contract.returned_value); - let return_span = self.mir_encoder.get_local_span( - contract.returned_value.into() - ); + let return_span = self + .mir_encoder + .get_local_span(contract.returned_value.into()); let (encoded_deref, ..) = self .mir_encoder .encode_deref(encoded_return, ty) @@ -4243,12 +4443,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { lhs.push(assertion_lhs); rhs.push(assertion_rhs); } - let lhs = lhs - .into_iter() - .conjoin(); - let rhs = rhs - .into_iter() - .conjoin(); + let lhs = lhs.into_iter().conjoin(); + let rhs = rhs.into_iter().conjoin(); Ok(Some((lhs, rhs))) } else { Ok(None) @@ -4281,7 +4477,12 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } else { // If the argument is not a reference, we wrap entire path into old. assertion = assertion.fold_expr(|e| { - if let vir::Expr::FuncApp(vir::FuncApp { function_name, arguments, .. }) = &e { + if let vir::Expr::FuncApp(vir::FuncApp { + function_name, + arguments, + .. + }) = &e + { // If `assertion` is e.g. a `foo(snap(_1))` with an `_1: T` and `fn foo(x: &T)` we cannot // wrap `foo(snap(old[pre](_1)))`, but should instead wrap as `foo(old[pre](snap(_1)))` // TODO: this could all be fixed by making all arguments into fields (e.g. `_1.local`). @@ -4345,9 +4546,9 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { place, mutability ); // TODO: Use a better span - let (place_expr, place_ty, _) = self.encode_generic_place( - contract.def_id, location, *place - ).with_span(self.mir.span)?; + let (place_expr, place_ty, _) = self + .encode_generic_place(contract.def_id, location, *place) + .with_span(self.mir.span)?; let old_place_expr = place_expr.clone().old(pre_label); let mut add_type_spec = |perm_amount| { let permissions = @@ -4381,9 +4582,16 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { .iter() .map(|local| self.encode_prusti_local(*local).into()) .collect(); - trace!("encode_postcondition_expr: encoded_args {:?} ({:?}) as {:?}", contract.args, - contract.args.iter().map(|a| self.locals.get_type(*a)).collect::>(), - encoded_args); + trace!( + "encode_postcondition_expr: encoded_args {:?} ({:?}) as {:?}", + contract.args, + contract + .args + .iter() + .map(|a| self.locals.get_type(*a)) + .collect::>(), + encoded_args + ); let encoded_return: vir::Expr = self.encode_prusti_local(contract.returned_value).into(); @@ -4393,13 +4601,10 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { self.mir.span }; let mut magic_wands = Vec::new(); - if let Some((mut lhs, mut rhs)) = self.encode_postcondition_magic_wand( - location, - contract, - pre_label, - post_label, - substs, - ).with_span(span)? { + if let Some((mut lhs, mut rhs)) = self + .encode_postcondition_magic_wand(location, contract, pre_label, post_label, substs) + .with_span(span)? + { if let Some((location, fake_exprs)) = magic_wand_store_info { let replace_fake_exprs = |mut expr: vir::Expr| -> vir::Expr { for (fake_arg, arg_expr) in fake_exprs.iter() { @@ -4409,25 +4614,23 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { }; lhs = replace_fake_exprs(lhs); rhs = replace_fake_exprs(rhs); - lhs = self.encoder.patch_snapshots(lhs) - .with_span(self.mir.span)?; - rhs = self.encoder.patch_snapshots(rhs) - .with_span(self.mir.span)?; + lhs = self.encoder.patch_snapshots(lhs).with_span(self.mir.span)?; + rhs = self.encoder.patch_snapshots(rhs).with_span(self.mir.span)?; debug!("Insert ({:?} {:?}) at {:?}", lhs, rhs, location); self.magic_wand_at_location .insert(location, (post_label.to_string(), lhs.clone(), rhs.clone())); } - magic_wands.push(vir::Expr::magic_wand(lhs, rhs, loan.map(|l| l.index().into()))); + magic_wands.push(vir::Expr::magic_wand( + lhs, + rhs, + loan.map(|l| l.index().into()), + )); } // Encode permissions for return type // TODO: Clean-up: remove unnecessary Option. - let return_perm = Some( - self.encode_local_variable_permission( - contract.returned_value, - None - )? - ); + let return_perm = + Some(self.encode_local_variable_permission(contract.returned_value, None)?); // Encode functional specification let mut func_spec = vec![]; @@ -4446,12 +4649,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let assertion_span = self.encoder.env().query.get_def_span(typed_assertion); func_spec_spans.push(assertion_span); let assertion_pos = self.mir_encoder.register_span(assertion_span); - assertion = self.wrap_arguments_into_old( - assertion, - pre_label, - contract, - &encoded_args, - )?; + assertion = + self.wrap_arguments_into_old(assertion, pre_label, contract, &encoded_args)?; func_spec.push(assertion.set_default_pos(assertion_pos)); } let postcondition_span = MultiSpan::from_spans(func_spec_spans); @@ -4460,14 +4659,18 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // Encode invariant for return value if !self.encoder.is_pure(contract.def_id, Some(substs)) { invs_spec.push( - self.encoder.encode_invariant_func_app( - self.locals.get_type(contract.returned_value), - encoded_return, - ).with_span(postcondition_span)? + self.encoder + .encode_invariant_func_app( + self.locals.get_type(contract.returned_value), + encoded_return, + ) + .with_span(postcondition_span)?, ); } - let full_func_spec = func_spec.into_iter().conjoin() + let full_func_spec = func_spec + .into_iter() + .conjoin() .set_default_pos(func_spec_pos); Ok(( @@ -4502,9 +4705,16 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } impl<'a> vir::ExprFolder for OldReplacer<'a> { #[allow(clippy::map_entry)] - fn fold_labelled_old(&mut self, vir::LabelledOld {label, base, position}: vir::LabelledOld) -> vir::Expr { + fn fold_labelled_old( + &mut self, + vir::LabelledOld { + label, + base, + position, + }: vir::LabelledOld, + ) -> vir::Expr { let base = self.fold_boxed(base); - let expr = vir::Expr::LabelledOld( vir::LabelledOld { + let expr = vir::Expr::LabelledOld(vir::LabelledOld { label: label.clone(), base, position, @@ -4556,14 +4766,18 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let span = self.mir.source_info(location).span; // Package magic wand(s) - if let Some((lhs, rhs)) = self.encode_postcondition_magic_wand( - None, - self.procedure_contract(), - pre_label, - post_label, - self.substs, - ).with_span(span)? { - let pos = self.register_error(self.mir.span, ErrorCtxt::PackageMagicWandForPostcondition); + if let Some((lhs, rhs)) = self + .encode_postcondition_magic_wand( + None, + self.procedure_contract(), + pre_label, + post_label, + self.substs, + ) + .with_span(span)? + { + let pos = + self.register_error(self.mir.span, ErrorCtxt::PackageMagicWandForPostcondition); let blocker = mir::RETURN_PLACE; // TODO: Check if it really is always start and not the mid point. @@ -4571,28 +4785,26 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { .polonius_info() .get_point(location, facts::PointType::Start); - let opt_region = self.polonius_info() - .place_regions - .for_local(blocker); + let opt_region = self.polonius_info().place_regions.for_local(blocker); let mut package_stmts = if let Some(region) = opt_region { - let (all_loans, zombie_loans) = self - .polonius_info() - .get_all_loans_kept_alive_by(start_point, region); - self.encode_expiration_of_loans(all_loans, &zombie_loans, location, None, true)? - } else { - // This happens when encoding the following function - // ``` - // struct MyStruct<'tcx>(TyCtxt<'tcx>); - // fn foo(tcx: TyCtxt) -> MyStruct { - // MyStruct(tcx) - // } - // ``` - return Err(SpannedEncodingError::unsupported( - "the encoding of pledges does not supporte this \ + let (all_loans, zombie_loans) = self + .polonius_info() + .get_all_loans_kept_alive_by(start_point, region); + self.encode_expiration_of_loans(all_loans, &zombie_loans, location, None, true)? + } else { + // This happens when encoding the following function + // ``` + // struct MyStruct<'tcx>(TyCtxt<'tcx>); + // fn foo(tcx: TyCtxt) -> MyStruct { + // MyStruct(tcx) + // } + // ``` + return Err(SpannedEncodingError::unsupported( + "the encoding of pledges does not supporte this \ kind of reborrowing", - self.mir_encoder.get_span_of_location(location), - )); - }; + self.mir_encoder.get_span_of_location(location), + )); + }; // We need to make sure that the lhs of the magic wand is // fully folded before the label. @@ -4620,10 +4832,10 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let arg_span = self.mir_encoder.get_local_span(arg_index); if is_reference(arg_ty) { let encoded_arg = self.mir_encoder.encode_local(arg_index)?; - let (deref_place, ..) = - self.mir_encoder - .encode_deref(encoded_arg.into(), arg_ty) - .with_span(arg_span)?; + let (deref_place, ..) = self + .mir_encoder + .encode_deref(encoded_arg.into(), arg_ty) + .with_span(arg_span)?; let old_deref_place = deref_place.clone().old(pre_label); package_stmts.extend(self.encode_transfer_permissions( deref_place, @@ -4661,11 +4873,16 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { "We can have at most one magic wand in the postcondition." ); for (path, _) in borrow_infos[0].blocking_paths.clone().iter() { - let (encoded_place, _, _) = self.encode_generic_place( - self.procedure_contract().def_id, None, *path - ).with_span(span)?; + let (encoded_place, _, _) = self + .encode_generic_place(self.procedure_contract().def_id, None, *path) + .with_span(span)?; let old_place = encoded_place.clone().old(post_label.clone()); - stmts.extend(self.encode_transfer_permissions(old_place, encoded_place, location, true)); + stmts.extend(self.encode_transfer_permissions( + old_place, + encoded_place, + location, + true, + )); } } @@ -4682,29 +4899,27 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // This clone is only due to borrow checker restrictions let contract = self.procedure_contract().clone(); - self.cfg_method.add_stmt(return_cfg_block, vir::Stmt::comment("Exhale postcondition")); + self.cfg_method + .add_stmt(return_cfg_block, vir::Stmt::comment("Exhale postcondition")); let postcondition_label = self.cfg_method.get_fresh_label_name(); - self.cfg_method.add_stmt(return_cfg_block, vir::Stmt::label(postcondition_label.clone())); + self.cfg_method.add_stmt( + return_cfg_block, + vir::Stmt::label(postcondition_label.clone()), + ); - let ( - type_spec, - return_type_spec, - invs_spec, - func_spec, - magic_wands, - _, - ) = self.encode_postcondition_expr( - None, - &contract, - PRECONDITION_LABEL, - &postcondition_label, - None, - false, - None, - true, - self.substs, - )?; + let (type_spec, return_type_spec, invs_spec, func_spec, magic_wands, _) = self + .encode_postcondition_expr( + None, + &contract, + PRECONDITION_LABEL, + &postcondition_label, + None, + false, + None, + true, + self.substs, + )?; let type_inv_pos = self.register_error( self.mir.span, @@ -4765,10 +4980,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { vir::PermAmount::Write, ) .unwrap(); - for stmt in self - .encode_obtain(deref_pred, type_inv_pos) - .drain(..) - { + for stmt in self.encode_obtain(deref_pred, type_inv_pos).drain(..) { self.cfg_method.add_stmt(return_cfg_block, stmt); } @@ -4796,25 +5008,26 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { self.cfg_method.add_stmt( return_cfg_block, - vir::Stmt::Assign( vir::Assign { + vir::Stmt::Assign(vir::Assign { target: var, source: encoded_deref, - kind: vir::AssignKind::Move + kind: vir::AssignKind::Move, }), ); } } // Fold the result. - self.cfg_method.add_stmt( - return_cfg_block, - vir::Stmt::comment("Fold the result"), - ); + self.cfg_method + .add_stmt(return_cfg_block, vir::Stmt::comment("Fold the result")); let ty = self.locals.get_type(contract.returned_value); let encoded_return: vir::Expr = self.encode_prusti_local(contract.returned_value).into(); - let return_span = self.mir_encoder.get_local_span(contract.returned_value.into()); + let return_span = self + .mir_encoder + .get_local_span(contract.returned_value.into()); let encoded_return_expr = if is_reference(ty) { - let (encoded_deref, ..) = self.mir_encoder + let (encoded_deref, ..) = self + .mir_encoder .encode_deref(encoded_return, ty) .with_span(return_span)?; encoded_deref @@ -4825,7 +5038,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { .mir_encoder .encode_place_predicate_permission(encoded_return_expr, vir::PermAmount::Write) .unwrap(); - let obtain_return_stmt = vir::Stmt::Obtain( vir::Obtain { + let obtain_return_stmt = vir::Stmt::Obtain(vir::Obtain { expr: return_pred, position: type_inv_pos, }); @@ -4838,12 +5051,15 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { vir::Stmt::comment("Assert possible strengthening"), ); if let Some(strengthening_spec) = strengthening_spec { - let patched_strengthening_spec = - self.replace_old_places_with_ghost_vars(None, strengthening_spec.refinement_check_expr); - let pos = self.register_error(strengthening_spec.spec_functions_span, ErrorCtxt::AssertMethodPostconditionStrengthening); + let patched_strengthening_spec = self + .replace_old_places_with_ghost_vars(None, strengthening_spec.refinement_check_expr); + let pos = self.register_error( + strengthening_spec.spec_functions_span, + ErrorCtxt::AssertMethodPostconditionStrengthening, + ); self.cfg_method.add_stmt( return_cfg_block, - vir::Stmt::Assert( vir::Assert { + vir::Stmt::Assert(vir::Assert { expr: patched_strengthening_spec, position: pos, }), @@ -4859,7 +5075,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let patched_func_spec = self.replace_old_places_with_ghost_vars(None, func_spec); self.cfg_method.add_stmt( return_cfg_block, - vir::Stmt::Assert( vir::Assert { + vir::Stmt::Assert(vir::Assert { expr: patched_func_spec, position: func_pos, }), @@ -4873,7 +5089,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let patched_invs_spec = self.replace_old_places_with_ghost_vars(None, invs_spec); self.cfg_method.add_stmt( return_cfg_block, - vir::Stmt::Assert( vir::Assert { + vir::Stmt::Assert(vir::Assert { expr: patched_invs_spec, position: type_inv_pos, }), @@ -4889,7 +5105,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { debug_assert!(!perm_pos.is_default()); self.cfg_method.add_stmt( return_cfg_block, - vir::Stmt::Exhale( vir::Exhale { + vir::Stmt::Exhale(vir::Exhale { expr: patched_type_spec, position: perm_pos, }), @@ -4902,7 +5118,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { if let Some(access) = return_type_spec { self.cfg_method.add_stmt( return_cfg_block, - vir::Stmt::Exhale( vir::Exhale { + vir::Stmt::Exhale(vir::Exhale { expr: access, position: perm_pos, }), @@ -4916,7 +5132,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { for magic_wand in magic_wands { self.cfg_method.add_stmt( return_cfg_block, - vir::Stmt::Exhale( vir::Exhale { + vir::Stmt::Exhale(vir::Exhale { expr: magic_wand, position: perm_pos, }), @@ -4962,7 +5178,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { place: &vir::Expr, ) -> vir::Expr { let tmp_var = self.get_pure_var_for_preserving_value(loop_head, place); - vir::Expr::BinOp ( vir::BinOp { + vir::Expr::BinOp(vir::BinOp { op_kind: vir::BinaryOpKind::EqCmp, left: box tmp_var.into(), right: box place.clone(), @@ -4985,7 +5201,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let snap_array = vir::Expr::snap_app(array_base); let old_snap_array = vir::Expr::old(snap_array.clone(), old_label); - vir_expr!{ [snap_array] == [old_snap_array] } + vir_expr! { [snap_array] == [old_snap_array] } } /// Arguments: @@ -5022,17 +5238,21 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let enclosing_permission_forest = if loops.len() > 1 { let next_to_last = loops.len() - 2; let enclosing_loop_head = loops[next_to_last]; - Some(self.loop_encoder.compute_loop_invariant( - enclosing_loop_head, - self.cached_loop_invariant_block[&enclosing_loop_head], - ).map_err(|err| match err { - LoopAnalysisError::UnsupportedPlaceContext(place_ctxt, loc) => { - SpannedEncodingError::internal( - format!("loop uses the unexpected PlaceContext '{place_ctxt:?}'"), - self.mir_encoder.get_span_of_location(loc), + Some( + self.loop_encoder + .compute_loop_invariant( + enclosing_loop_head, + self.cached_loop_invariant_block[&enclosing_loop_head], ) - } - })?) + .map_err(|err| match err { + LoopAnalysisError::UnsupportedPlaceContext(place_ctxt, loc) => { + SpannedEncodingError::internal( + format!("loop uses the unexpected PlaceContext '{place_ctxt:?}'"), + self.mir_encoder.get_span_of_location(loc), + ) + } + })?, + ) } else { None }; @@ -5058,16 +5278,21 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ExprOrArrayBase::ArrayBase(b) | ExprOrArrayBase::SliceBase(b) => (b, true), }; - debug!("kind={:?} mir_place={:?} encoded_place={:?} ty={:?}", kind, mir_place, encoded_place, ty); + debug!( + "kind={:?} mir_place={:?} encoded_place={:?} ty={:?}", + kind, mir_place, encoded_place, ty + ); match kind { // Gives read permission to this node. It must not be a leaf node. PermissionKind::ReadNode => { if is_array_access { // if it's already there, it's at least read -> nothing to do here - array_pred_perms.entry(encoded_place) + array_pred_perms + .entry(encoded_place) .or_insert(vir::PermAmount::Read); } else { - let perm = vir::Expr::acc_permission(encoded_place, vir::PermAmount::Read); + let perm = + vir::Expr::acc_permission(encoded_place, vir::PermAmount::Read); permissions.push(perm); } } @@ -5077,7 +5302,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { if is_array_access { array_pred_perms.insert(encoded_place, vir::PermAmount::Write); } else { - let perm = vir::Expr::acc_permission(encoded_place, vir::PermAmount::Write); + let perm = + vir::Expr::acc_permission(encoded_place, vir::PermAmount::Write); permissions.push(perm); } } @@ -5095,7 +5321,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { .loop_encoder .is_definitely_initialised(mir_place, loop_head); debug!(" perm_amount={} def_init={}", perm_amount, def_init); - if let Some(base) = utils::try_pop_deref(self.encoder.env().tcx(), mir_place) + if let Some(base) = + utils::try_pop_deref(self.encoder.env().tcx(), mir_place) { // will panic if attempting to encode unsupported type let ref_ty = self.mir_encoder.encode_place(base).unwrap().1; @@ -5136,8 +5363,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { &field_place, )); } - if def_init - && !(mutbl == &Mutability::Not && drop_read_references) + if def_init && !(mutbl == &Mutability::Not && drop_read_references) { permissions.push( vir::Expr::pred_permission(field_place, perm_amount) @@ -5147,12 +5373,18 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } _ => { if is_array_access { - array_pred_perms.entry(encoded_place) - .and_modify(|e| if perm_amount == vir::PermAmount::Write { *e = vir::PermAmount::Write; }) + array_pred_perms + .entry(encoded_place) + .and_modify(|e| { + if perm_amount == vir::PermAmount::Write { + *e = vir::PermAmount::Write; + } + }) .or_insert(perm_amount); } else { permissions.push( - vir::Expr::pred_permission(encoded_place, perm_amount).unwrap(), + vir::Expr::pred_permission(encoded_place, perm_amount) + .unwrap(), ); } @@ -5163,23 +5395,30 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // that we will lose information about the children of that // place after the loop and we need to preserve it via local // variables. - let encoded_child = self.mir_encoder.encode_place(child_place)?.0; + let encoded_child = + self.mir_encoder.encode_place(child_place)?.0; match encoded_child.into_array_base() { ExprOrArrayBase::Expr(e) => { - equalities.push(self.construct_value_preserving_equality( - loop_head, - &e, - )); + equalities.push( + self.construct_value_preserving_equality( + loop_head, &e, + ), + ); } ExprOrArrayBase::ArrayBase(b) => { - let eq = self.construct_value_preserving_array_equality(loop_head, b); + let eq = self + .construct_value_preserving_array_equality( + loop_head, b, + ); // arrays can be mentioned multiple times, so we // need to check here if !equalities.contains(&eq) { equalities.push(eq); } } - ExprOrArrayBase::SliceBase(_) => unimplemented!("slices in loops not yet implemented"), + ExprOrArrayBase::SliceBase(_) => unimplemented!( + "slices in loops not yet implemented" + ), } } } @@ -5202,21 +5441,24 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { trace!("array_pred_perms: {:?}", array_pred_perms); for (place, perm) in array_pred_perms.into_iter() { permissions.push( - vir::Expr::pred_permission(place, perm) - .expect("invalid place in array_pred_perms") + vir::Expr::pred_permission(place, perm).expect("invalid place in array_pred_perms"), ); } // encode type invariants let mut invs_spec = Vec::new(); for permission in &permissions { - if let vir::Expr::PredicateAccessPredicate( vir::PredicateAccessPredicate {predicate_type, argument, ..}) = permission { + if let vir::Expr::PredicateAccessPredicate(vir::PredicateAccessPredicate { + predicate_type, + argument, + .. + }) = permission + { let ty = self.encoder.decode_type_predicate_type(predicate_type)?; if !self.encoder.is_pure(self.proc_def_id, Some(self.substs)) { - let inv_func_app = self.encoder.encode_invariant_func_app( - ty, - (**argument).clone(), - )?; + let inv_func_app = self + .encoder + .encode_invariant_func_app(ty, (**argument).clone())?; invs_spec.push(inv_func_app); } } @@ -5284,8 +5526,12 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { for stmt in &self.mir.basic_blocks[bbi].statements { if let mir::StatementKind::Assign(box ( _, - mir::Rvalue::Aggregate(box mir::AggregateKind::Closure(cl_def_id, cl_substs), _), - )) = stmt.kind { + mir::Rvalue::Aggregate( + box mir::AggregateKind::Closure(cl_def_id, cl_substs), + _, + ), + )) = stmt.kind + { if let Some(spec) = self.encoder.get_loop_specs(cl_def_id) { encoded_specs.push(self.encoder.encode_invariant( self.mir, @@ -5294,7 +5540,9 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { cl_substs, )?); let invariant = match spec { - prusti_interface::specs::typed::LoopSpecification::Invariant(inv) => inv, + prusti_interface::specs::typed::LoopSpecification::Invariant(inv) => { + inv + } _ => continue, }; encoded_spec_spans.push(self.encoder.env().tcx().def_span(invariant)); @@ -5318,9 +5566,9 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } let (func_spec, func_spec_span) = self.encode_loop_invariant_specs(loop_head, loop_inv_block)?; - let (permissions, equalities, invs_spec) = - self.encode_loop_invariant_permissions(loop_head, loop_inv_block, true) - .with_span(func_spec_span.clone())?; + let (permissions, equalities, invs_spec) = self + .encode_loop_invariant_permissions(loop_head, loop_inv_block, true) + .with_span(func_spec_span.clone())?; // TODO: use different positions, and generate different error messages, for the exhale // before the loop and after the loop body @@ -5353,7 +5601,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { stmts.push(vir::Stmt::label(label)); } for (place, field) in &self.pure_var_for_preserving_value_map[&loop_head] { - stmts.push(vir::Stmt::Assign( vir::Assign { + stmts.push(vir::Stmt::Assign(vir::Assign { target: field.into(), source: place.clone(), kind: vir::AssignKind::Ghost, @@ -5362,28 +5610,28 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } assert!(!assert_pos.is_default()); let obtain_predicates = permissions.iter().map(|p| { - vir::Stmt::Obtain( vir::Obtain { + vir::Stmt::Obtain(vir::Obtain { expr: p.clone(), position: assert_pos, }) // TODO: Use a better position. }); stmts.extend(obtain_predicates); - stmts.push(vir::Stmt::Assert( vir::Assert { + stmts.push(vir::Stmt::Assert(vir::Assert { expr: func_spec.into_iter().conjoin(), position: assert_pos, })); - stmts.push(vir::Stmt::Assert( vir::Assert { + stmts.push(vir::Stmt::Assert(vir::Assert { expr: invs_spec.into_iter().conjoin(), position: exhale_pos, })); let equalities_expr = equalities.into_iter().conjoin(); - stmts.push(vir::Stmt::Assert( vir::Assert { + stmts.push(vir::Stmt::Assert(vir::Assert { expr: equalities_expr, position: exhale_pos, })); let permission_expr = permissions.into_iter().conjoin(); - stmts.push(vir::Stmt::Exhale( vir::Exhale { + stmts.push(vir::Stmt::Exhale(vir::Exhale { expr: permission_expr, position: exhale_pos, })); @@ -5406,13 +5654,13 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let mut stmts = vec![vir::Stmt::comment(format!( "Inhale the loop permissions invariant of block {loop_head:?}" ))]; - stmts.push(vir::Stmt::Inhale( vir::Inhale { + stmts.push(vir::Stmt::Inhale(vir::Inhale { expr: permission_expr, })); - stmts.push(vir::Stmt::Inhale( vir::Inhale { + stmts.push(vir::Stmt::Inhale(vir::Inhale { expr: equality_expr, })); - stmts.push(vir::Stmt::Inhale( vir::Inhale { + stmts.push(vir::Stmt::Inhale(vir::Inhale { expr: invs_spec.into_iter().conjoin(), })); Ok(stmts) @@ -5431,7 +5679,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let mut stmts = vec![vir::Stmt::comment(format!( "Inhale the loop fnspec invariant of block {loop_head:?}" ))]; - stmts.push(vir::Stmt::Inhale( vir::Inhale { + stmts.push(vir::Stmt::Inhale(vir::Inhale { expr: func_spec.into_iter().conjoin(), })); Ok((stmts, func_spec_span)) @@ -5441,7 +5689,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let var_name = self.locals.get_name(local); let typ = self .encoder - .encode_type(self.locals.get_type(local)).unwrap(); // will panic if attempting to encode unsupported type + .encode_type(self.locals.get_type(local)) + .unwrap(); // will panic if attempting to encode unsupported type vir::LocalVar::new(var_name, typ) } @@ -5472,9 +5721,15 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ) -> EncodingResult<(vir::Expr, ty::Ty<'tcx>, Option)> { let mir_encoder = if let Some(location) = location { let block = &self.mir.basic_blocks[location.block]; - assert_eq!(block.statements.len(), location.statement_index, "expected terminator location"); + assert_eq!( + block.statements.len(), + location.statement_index, + "expected terminator location" + ); match &block.terminator().kind { - mir::TerminatorKind::Call{ args, destination, .. } => { + mir::TerminatorKind::Call { + args, destination, .. + } => { let tcx = self.encoder.env().tcx(); let arg_tys = args.iter().map(|arg| arg.ty(self.mir, tcx)).collect(); let return_ty = destination.ty(self.mir, tcx).ty; @@ -5500,7 +5755,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } Place::SubstitutedPlace { substituted_root, - place + place, } => { let (place_encoding, ty, variant) = mir_encoder.encode_place(place)?; let expr = place_encoding.try_into_expr()?; @@ -5510,7 +5765,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } use vir::ExprFolder; impl ExprFolder for RootReplacer { - fn fold_local(&mut self, vir::Local {position, ..}: vir::Local) -> vir::Expr { + fn fold_local(&mut self, vir::Local { position, .. }: vir::Local) -> vir::Expr { vir::Expr::local_with_pos(self.new_root.clone(), position) } } @@ -5595,15 +5850,20 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // inhale lookup_pure(array, index) == encoded_rhs // // now we have all the contents as before, just one item updated - let (encoded_array, mut stmts) = self.postprocess_place_encoding( - base, - ArrayAccessKind::Shared, // shouldn't be nested, so doesn't matter[tm] - ).with_span(span)?; + let (encoded_array, mut stmts) = self + .postprocess_place_encoding( + base, + ArrayAccessKind::Shared, // shouldn't be nested, so doesn't matter[tm] + ) + .with_span(span)?; let label = self.cfg_method.get_fresh_label_name(); stmts.push(vir::Stmt::label(&label)); - let sequence_types = self.encoder.encode_sequence_types(array_ty).with_span(span)?; + let sequence_types = self + .encoder + .encode_sequence_types(array_ty) + .with_span(span)?; let sequence_len = sequence_types.len(self.encoder, encoded_array.clone()); let array_acc_expr = vir::Expr::predicate_access_predicate( @@ -5613,54 +5873,78 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ); // exhale and re-inhale to havoc - stmts.push(vir_stmt!{ exhale [array_acc_expr] }); - stmts.push(vir_stmt!{ inhale [array_acc_expr] }); + stmts.push(vir_stmt! { exhale [array_acc_expr] }); + stmts.push(vir_stmt! { inhale [array_acc_expr] }); - let old = |e| { vir::Expr::labelled_old(&label, e) }; + let old = |e| vir::Expr::labelled_old(&label, e); // For sequences with fixed len (e.g. arrays) this will be Some if sequence_types.sequence_len.is_none() { // inhale that len unchanged - let len_eq = vir_expr!{ [ sequence_len ] == [ old(sequence_len.clone()) ] }; - stmts.push(vir_stmt!{ inhale [len_eq] }); + let len_eq = vir_expr! { [ sequence_len ] == [ old(sequence_len.clone()) ] }; + stmts.push(vir_stmt! { inhale [len_eq] }); } - let idx_val_int = self.encoder.patch_snapshots(vir::Expr::snap_app(index)).with_span(span)?; + let idx_val_int = self + .encoder + .patch_snapshots(vir::Expr::snap_app(index)) + .with_span(span)?; // inhale infos about array contents back - let i_var: vir::Expr = vir_local!{ i: Int }.into(); - let zero_le_i = vir_expr!{ [vir::Expr::from(0usize)] <= [ i_var ] }; - let i_lt_len = vir_expr!{ [ i_var ] < [ sequence_len ] }; - let i_ne_idx = vir_expr!{ [ i_var ] != [ old(idx_val_int.clone()) ] }; - let idx_conditions = vir_expr!{ [zero_le_i] && ([i_lt_len] && [i_ne_idx]) }; - let lookup_ret_ty = self.encoder.encode_snapshot_type(sequence_types.elem_ty_rs).with_span(span)?; - let lookup_array_i = sequence_types.encode_lookup_pure_call(self.encoder, encoded_array.clone(), i_var.clone(), lookup_ret_ty.clone()); + let i_var: vir::Expr = vir_local! { i: Int }.into(); + let zero_le_i = vir_expr! { [vir::Expr::from(0usize)] <= [ i_var ] }; + let i_lt_len = vir_expr! { [ i_var ] < [ sequence_len ] }; + let i_ne_idx = vir_expr! { [ i_var ] != [ old(idx_val_int.clone()) ] }; + let idx_conditions = vir_expr! { [zero_le_i] && ([i_lt_len] && [i_ne_idx]) }; + let lookup_ret_ty = self + .encoder + .encode_snapshot_type(sequence_types.elem_ty_rs) + .with_span(span)?; + let lookup_array_i = sequence_types.encode_lookup_pure_call( + self.encoder, + encoded_array.clone(), + i_var.clone(), + lookup_ret_ty.clone(), + ); // FIXME: using `old` here is a work-around for https://github.com/viperproject/prusti-dev/issues/877 // FIXME: which is due to an issue in Silicon, see https://github.com/viperproject/silicon/issues/603 - let lookup_array_i_old = sequence_types.encode_lookup_pure_call(self.encoder, old(encoded_array.clone()), i_var, lookup_ret_ty.clone()); - let lookup_same_as_old = vir_expr!{ [lookup_array_i] == [old(lookup_array_i)] }; - let forall_body = vir_expr!{ [idx_conditions] ==> [lookup_same_as_old] }; - let all_others_unchanged = vir_expr!{ forall i: Int :: { [lookup_array_i_old] } :: [ forall_body ] }; + let lookup_array_i_old = sequence_types.encode_lookup_pure_call( + self.encoder, + old(encoded_array.clone()), + i_var, + lookup_ret_ty.clone(), + ); + let lookup_same_as_old = vir_expr! { [lookup_array_i] == [old(lookup_array_i)] }; + let forall_body = vir_expr! { [idx_conditions] ==> [lookup_same_as_old] }; + let all_others_unchanged = + vir_expr! { forall i: Int :: { [lookup_array_i_old] } :: [ forall_body ] }; - stmts.push(vir_stmt!{ inhale [ all_others_unchanged ]}); + stmts.push(vir_stmt! { inhale [ all_others_unchanged ]}); - let tmp = vir::Expr::from(self.cfg_method.add_fresh_local_var(sequence_types.elem_pred_type.clone())); + let tmp = vir::Expr::from( + self.cfg_method + .add_fresh_local_var(sequence_types.elem_pred_type.clone()), + ); stmts.extend( - self.encode_assign( - tmp.clone(), - rhs, - sequence_types.elem_ty_rs, - location, - ).with_span(span)? + self.encode_assign(tmp.clone(), rhs, sequence_types.elem_ty_rs, location) + .with_span(span)?, ); - let tmp_val_field = self.encoder.encode_value_expr(tmp, sequence_types.elem_ty_rs).with_span(span)?; + let tmp_val_field = self + .encoder + .encode_value_expr(tmp, sequence_types.elem_ty_rs) + .with_span(span)?; - let indexed_lookup_pure_call = sequence_types - .encode_lookup_pure_call(self.encoder, encoded_array, old(idx_val_int), lookup_ret_ty); - let indexed_updated = vir_expr!{ [ indexed_lookup_pure_call ] == [ vir::Expr::snap_app(tmp_val_field) ] }; + let indexed_lookup_pure_call = sequence_types.encode_lookup_pure_call( + self.encoder, + encoded_array, + old(idx_val_int), + lookup_ret_ty, + ); + let indexed_updated = + vir_expr! { [ indexed_lookup_pure_call ] == [ vir::Expr::snap_app(tmp_val_field) ] }; - stmts.push(vir_stmt!{ inhale [ indexed_updated ] }); + stmts.push(vir_stmt! { inhale [ indexed_updated ] }); Ok(stmts) } @@ -5677,7 +5961,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let span = self.mir_encoder.get_span_of_location(location); let stmts = match operand { mir::Operand::Move(place) => { - let (src, mut stmts, ty, _) = self.encode_place(*place, ArrayAccessKind::Shared, location)?; + let (src, mut stmts, ty, _) = + self.encode_place(*place, ArrayAccessKind::Shared, location)?; let encode_stmts = match ty.kind() { ty::TyKind::RawPtr(..) | ty::TyKind::Ref(..) => { // Reborrow. @@ -5687,26 +5972,28 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { field.clone(), location, vir::AssignKind::Move, - false + false, )?; - alloc_stmts.push(vir::Stmt::Assign( vir::Assign { + alloc_stmts.push(vir::Stmt::Assign(vir::Assign { target: lhs.clone().field(field.clone()), source: src.field(field), kind: vir::AssignKind::Move, })); alloc_stmts } - _ if config::enable_purification_optimization() && - prusti_common::vir::optimizations::purification::is_purifiable_type(lhs.get_type()) => { + _ if config::enable_purification_optimization() + && prusti_common::vir::optimizations::purification::is_purifiable_type( + lhs.get_type(), + ) => + { self.encode_copy2(src, lhs.clone(), ty, location)? } _ => { // Just move. - let move_assign = - vir::Stmt::Assign( vir::Assign { - target: lhs.clone(), - source: src, - kind: vir::AssignKind::Move + let move_assign = vir::Stmt::Assign(vir::Assign { + target: lhs.clone(), + source: src, + kind: vir::AssignKind::Move, }); vec![move_assign] } @@ -5724,7 +6011,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } mir::Operand::Copy(place) => { - let (src, mut stmts, ty, _) = self.encode_place(*place, ArrayAccessKind::Shared, location)?; + let (src, mut stmts, ty, _) = + self.encode_place(*place, ArrayAccessKind::Shared, location)?; let encode_stmts = match ty.kind() { ty::TyKind::RawPtr(..) => { return Err(SpannedEncodingError::unsupported( @@ -5740,9 +6028,9 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ref_field.clone(), location, vir::AssignKind::SharedBorrow(loan.index().into()), - false + false, )?; - stmts.push(vir::Stmt::Assign( vir::Assign { + stmts.push(vir::Stmt::Assign(vir::Assign { target: lhs.clone().field(ref_field.clone()), source: src.field(ref_field), kind: vir::AssignKind::SharedBorrow(loan.index().into()), @@ -5774,17 +6062,18 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { field.clone(), location, vir::AssignKind::Copy, - true + true, )?; // TODO Encoding of string literals is not yet supported, // so do not encode an assignment if the RHS is a string if !is_str(ty) { // Initialize the constant - let const_val = self.encoder + let const_val = self + .encoder .encode_const_expr(ty, expr.literal) .with_span(span)?; // Initialize value of lhs - stmts.push(vir::Stmt::Assign( vir::Assign { + stmts.push(vir::Stmt::Assign(vir::Assign { target: lhs.clone().field(field), source: const_val, kind: vir::AssignKind::Copy, @@ -5815,13 +6104,15 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { location: mir::Location, ) -> SpannedEncodingResult> { let span = self.mir_encoder.get_span_of_location(location); - let encoded_left = self.mir_encoder.encode_operand_expr(left) + let encoded_left = self.mir_encoder.encode_operand_expr(left).with_span(span)?; + let encoded_right = self + .mir_encoder + .encode_operand_expr(right) .with_span(span)?; - let encoded_right = self.mir_encoder.encode_operand_expr(right) + let encoded_value = self + .mir_encoder + .encode_bin_op_expr(op, encoded_left, encoded_right, ty) .with_span(span)?; - let encoded_value = - self.mir_encoder.encode_bin_op_expr(op, encoded_left, encoded_right, ty) - .with_span(span)?; self.encode_copy_value_assign(encoded_lhs, encoded_value, ty, location) } @@ -5834,12 +6125,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ) -> SpannedEncodingResult> { let span = self.mir_encoder.get_span_of_location(location); let field = self.encoder.encode_value_field(ty).with_span(span)?; - self.encode_copy_value_assign2( - encoded_lhs, - encoded_rhs, - field, - location - ) + self.encode_copy_value_assign2(encoded_lhs, encoded_rhs, field, location) } /// Assignment with a(n overflow-)checked binary operation on the RHS. @@ -5860,20 +6146,19 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } else { unreachable!() }; - let encoded_left = self.mir_encoder.encode_operand_expr(left) + let encoded_left = self.mir_encoder.encode_operand_expr(left).with_span(span)?; + let encoded_right = self + .mir_encoder + .encode_operand_expr(right) + .with_span(span)?; + let encoded_value = self + .mir_encoder + .encode_bin_op_expr(op, encoded_left.clone(), encoded_right.clone(), operand_ty) .with_span(span)?; - let encoded_right = self.mir_encoder.encode_operand_expr(right) + let encoded_check = self + .mir_encoder + .encode_bin_op_check(op, encoded_left, encoded_right, operand_ty) .with_span(span)?; - let encoded_value = self.mir_encoder.encode_bin_op_expr( - op, - encoded_left.clone(), - encoded_right.clone(), - operand_ty, - ).with_span(span)?; - let encoded_check = - self.mir_encoder - .encode_bin_op_check(op, encoded_left, encoded_right, operand_ty) - .with_span(span)?; let field_types = if let ty::TyKind::Tuple(ref x) = ty.kind() { x } else { @@ -5883,19 +6168,25 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { .encoder .encode_raw_ref_field("tuple_0".to_string(), field_types[0]) .with_span(span)?; - let value_field_value = self.encoder.encode_value_field(field_types[0]).with_span(span)?; + let value_field_value = self + .encoder + .encode_value_field(field_types[0]) + .with_span(span)?; let check_field = self .encoder .encode_raw_ref_field("tuple_1".to_string(), field_types[1]) .with_span(span)?; - let check_field_value = self.encoder.encode_value_field(field_types[1]).with_span(span)?; + let check_field_value = self + .encoder + .encode_value_field(field_types[1]) + .with_span(span)?; let mut stmts = if !self .init_info .is_vir_place_accessible(&encoded_lhs, location) { let mut alloc_stmts = self.encode_havoc(&encoded_lhs).with_span(span)?; let mut inhale_acc = |place| { - alloc_stmts.push(vir::Stmt::Inhale( vir::Inhale { + alloc_stmts.push(vir::Stmt::Inhale(vir::Inhale { expr: vir::Expr::acc_permission(place, vir::PermAmount::Write), })); }; @@ -5918,7 +6209,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { Vec::with_capacity(2) }; // Initialize lhs.field - stmts.push(vir::Stmt::Assign( vir::Assign { + stmts.push(vir::Stmt::Assign(vir::Assign { target: encoded_lhs .clone() .field(value_field) @@ -5926,7 +6217,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { source: encoded_value, kind: vir::AssignKind::Copy, })); - stmts.push(vir::Stmt::Assign( vir::Assign { + stmts.push(vir::Stmt::Assign(vir::Assign { target: encoded_lhs.field(check_field).field(check_field_value), source: encoded_check, kind: vir::AssignKind::Copy, @@ -5946,10 +6237,10 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ty: ty::Ty<'tcx>, location: mir::Location, ) -> SpannedEncodingResult> { - let encoded_val = self.mir_encoder.encode_operand_expr(operand) - .with_span( - self.mir_encoder.get_span_of_location(location) - )?; + let encoded_val = self + .mir_encoder + .encode_operand_expr(operand) + .with_span(self.mir_encoder.get_span_of_location(location))?; let encoded_value = self.mir_encoder.encode_unary_op_expr(op, encoded_val); // Initialize `lhs.field` self.encode_copy_value_assign(encoded_lhs, encoded_value, ty, location) @@ -5971,10 +6262,12 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let param_env = tcx.param_env(self.proc_def_id); let layout = match tcx.layout_of(param_env.and(op_ty)) { Ok(layout_of) => layout_of.layout, - Err(_) => return Err(SpannedEncodingError::internal( - format!("could not fetch layout of type `{op_ty}` while translating `{op:?}`"), - self.mir.source_info(location).span, - )), + Err(_) => { + return Err(SpannedEncodingError::internal( + format!("could not fetch layout of type `{op_ty}` while translating `{op:?}`"), + self.mir.source_info(location).span, + )) + } }; let bytes = match op { mir::NullOp::SizeOf => layout.size().bytes(), @@ -5982,12 +6275,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { mir::NullOp::AlignOf => layout.align().abi.bytes(), }; let bytes_vir = vir::Expr::from(bytes); - self.encode_copy_value_assign( - encoded_lhs, - bytes_vir, - ty, - location, - ) + self.encode_copy_value_assign(encoded_lhs, bytes_vir, ty, location) } fn encode_assign_box( @@ -5999,10 +6287,10 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ) -> SpannedEncodingResult> { assert_eq!(op_ty, ty.boxed_ty()); let span = self.mir_encoder.get_span_of_location(location); - let ref_field = self.encoder.encode_dereference_field(op_ty) - .with_span( - self.mir_encoder.get_span_of_location(location) - )?; + let ref_field = self + .encoder + .encode_dereference_field(op_ty) + .with_span(self.mir_encoder.get_span_of_location(location))?; let box_content = encoded_lhs.clone().field(ref_field.clone()); let mut stmts = self.prepare_assign_target( @@ -6010,14 +6298,17 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ref_field, location, vir::AssignKind::Move, - false + false, )?; // Allocate `box_content` // TODO: This is also encoding initialization, which should be avoided because both // `Rvalue::NullaryOp(NullOp::Box)` and `Rvalue::ShallowInitBox` don't initialize the // content of the box. - stmts.extend(self.encode_havoc_and_initialization(&box_content).with_span(span)?); + stmts.extend( + self.encode_havoc_and_initialization(&box_content) + .with_span(span)?, + ); // Leave `box_content` uninitialized Ok(stmts) @@ -6034,7 +6325,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ty: ty::Ty<'tcx>, ) -> SpannedEncodingResult> { let span = self.mir_encoder.get_span_of_location(location); - let (encoded_src, mut stmts, src_ty, _) = self.encode_place(src, ArrayAccessKind::Shared, location)?; + let (encoded_src, mut stmts, src_ty, _) = + self.encode_place(src, ArrayAccessKind::Shared, location)?; let encode_stmts = match src_ty.kind() { ty::TyKind::Adt(adt_def, _) if !adt_def.is_box() => { let num_variants = adt_def.variants().len(); @@ -6051,16 +6343,10 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { self.proc_def_id, ); } - let encoded_rhs = self.encoder.encode_discriminant_func_app( - encoded_src, - *adt_def, - )?; - self.encode_copy_value_assign( - encoded_lhs, - encoded_rhs, - ty, - location - )? + let encoded_rhs = self + .encoder + .encode_discriminant_func_app(encoded_src, *adt_def)?; + self.encode_copy_value_assign(encoded_lhs, encoded_rhs, ty, location)? } else { vec![] } @@ -6069,12 +6355,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ty::TyKind::Int(_) | ty::TyKind::Uint(_) | ty::TyKind::Float(_) => { let value_field = self.encoder.encode_value_field(src_ty).with_span(span)?; let discr_value = encoded_src.field(value_field); - self.encode_copy_value_assign( - encoded_lhs, - discr_value, - ty, - location - )? + self.encode_copy_value_assign(encoded_lhs, discr_value, ty, location)? } ref x => { @@ -6100,26 +6381,28 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let span = self.mir_encoder.get_span_of_location(location); let loan = self.polonius_info().get_loan_at_location(location); let (vir_assign_kind, array_encode_kind) = match mir_borrow_kind { - mir::BorrowKind::Shared => - (vir::AssignKind::SharedBorrow(loan.index().into()), ArrayAccessKind::Shared), - mir::BorrowKind::Mut { .. } => - (vir::AssignKind::MutableBorrow(loan.index().into()), - ArrayAccessKind::Mutable(Some(loan.index().into()), location)), + mir::BorrowKind::Shared => ( + vir::AssignKind::SharedBorrow(loan.index().into()), + ArrayAccessKind::Shared, + ), + mir::BorrowKind::Mut { .. } => ( + vir::AssignKind::MutableBorrow(loan.index().into()), + ArrayAccessKind::Mutable(Some(loan.index().into()), location), + ), _ => return Err(Self::unsupported_borrow_kind(mir_borrow_kind).with_span(span)), }; - let (encoded_value, mut stmts, _, _) = self.encode_place(place, array_encode_kind, location)?; + let (encoded_value, mut stmts, _, _) = + self.encode_place(place, array_encode_kind, location)?; // Initialize ref_var.ref_field let field = self.encoder.encode_value_field(ty).with_span(span)?; - stmts.extend( - self.prepare_assign_target( - encoded_lhs.clone(), - field.clone(), - location, - vir_assign_kind, - false - )? - ); - stmts.push(vir::Stmt::Assign( vir::Assign { + stmts.extend(self.prepare_assign_target( + encoded_lhs.clone(), + field.clone(), + location, + vir_assign_kind, + false, + )?); + stmts.push(vir::Stmt::Assign(vir::Assign { target: encoded_lhs.field(field), source: encoded_value, kind: vir_assign_kind, @@ -6170,24 +6453,33 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } else { unreachable!("encode_assign_slice on a non-ref?!") }; - let slice_types = self.encoder.encode_sequence_types(*slice_ty).with_span(span)?; + let slice_types = self + .encoder + .encode_sequence_types(*slice_ty) + .with_span(span)?; stmts.extend(self.encode_havoc(&encoded_lhs).with_span(span)?); let val_ref_field = self.encoder.encode_value_field(ty).with_span(span)?; let slice_expr = encoded_lhs.field(val_ref_field); - stmts.push(vir_stmt!{ inhale [vir::Expr::FieldAccessPredicate( vir::FieldAccessPredicate { - base: box slice_expr.clone(), - permission: vir::PermAmount::Write, - position: vir::Position::default(), - })]}); + stmts.push( + vir_stmt! { inhale [vir::Expr::FieldAccessPredicate( vir::FieldAccessPredicate { + base: box slice_expr.clone(), + permission: vir::PermAmount::Write, + position: vir::Position::default(), + })]}, + ); - let slice_perm = vir::Expr::PredicateAccessPredicate( vir::PredicateAccessPredicate { + let slice_perm = vir::Expr::PredicateAccessPredicate(vir::PredicateAccessPredicate { predicate_type: slice_types.sequence_pred_type.clone(), argument: box slice_expr.clone(), - permission: if is_mut { vir::PermAmount::Write } else { vir::PermAmount::Read }, + permission: if is_mut { + vir::PermAmount::Write + } else { + vir::PermAmount::Read + }, position: vir::Position::default(), }); - stmts.push(vir_stmt!{ inhale [slice_perm] }); + stmts.push(vir_stmt! { inhale [slice_perm] }); let (rhs_place, rhs_ty) = if let mir::Operand::Move(place) = operand { let (rhs_place, rhs_ty, ..) = self.mir_encoder.encode_place(*place).with_span(span)?; @@ -6204,7 +6496,10 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let val_ref_field = self.encoder.encode_value_field(rhs_ty).with_span(span)?; let rhs_expr = rhs_place.field(val_ref_field); - let sequence_types = self.encoder.encode_sequence_types(*rhs_array_ty).with_span(span)?; + let sequence_types = self + .encoder + .encode_sequence_types(*rhs_array_ty) + .with_span(span)?; let slice_len_call = slice_types.len(self.encoder, slice_expr.clone()); @@ -6212,7 +6507,10 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { expr: vir_expr!{ [slice_len_call] == [sequence_types.len(self.encoder, rhs_expr.clone())] } })); - let elem_snap_ty = self.encoder.encode_snapshot_type(sequence_types.elem_ty_rs).with_span(span)?; + let elem_snap_ty = self + .encoder + .encode_snapshot_type(sequence_types.elem_ty_rs) + .with_span(span)?; let i: vir::Expr = vir_local! { i: Int }.into(); @@ -6223,20 +6521,16 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { elem_snap_ty.clone(), ); - let slice_lookup_call = slice_types.encode_lookup_pure_call( - self.encoder, - slice_expr, - i.clone(), - elem_snap_ty, - ); + let slice_lookup_call = + slice_types.encode_lookup_pure_call(self.encoder, slice_expr, i.clone(), elem_snap_ty); let indices = vir_expr! { ([vir::Expr::from(0usize)] <= [i]) && ([i] < [sequence_types.len(self.encoder, rhs_expr)]) }; - stmts.push(vir::Stmt::Inhale( vir::Inhale { + stmts.push(vir::Stmt::Inhale(vir::Inhale { expr: vir_expr! { - forall i: Int :: - { [array_lookup_call] } { [slice_lookup_call] } :: - ([indices] ==> ([array_lookup_call] == [slice_lookup_call])) } + forall i: Int :: + { [array_lookup_call] } { [slice_lookup_call] } :: + ([indices] ==> ([array_lookup_call] == [slice_lookup_call])) }, })); // Store a label for permissions got back from the call @@ -6258,18 +6552,16 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { location: mir::Location, ) -> SpannedEncodingResult> { let span = self.mir_encoder.get_span_of_location(location); - let (encoded_place, mut stmts, place_ty, ..) = self.encode_place( - place, - ArrayAccessKind::Mutable(None, location), - location - )?; + let (encoded_place, mut stmts, place_ty, ..) = + self.encode_place(place, ArrayAccessKind::Mutable(None, location), location)?; match place_ty.kind() { - ty::TyKind::Array(..) | - ty::TyKind::Slice(..) => { - let slice_types = self.encoder.encode_sequence_types(place_ty) - .with_span(span)?; + ty::TyKind::Array(..) | ty::TyKind::Slice(..) => { + let slice_types = self + .encoder + .encode_sequence_types(place_ty) + .with_span(span)?; - stmts.push(vir::Stmt::Assert( vir::Assert { + stmts.push(vir::Stmt::Assert(vir::Assert { expr: vir::Expr::predicate_access_predicate( slice_types.sequence_pred_type.clone(), encoded_place.clone(), @@ -6280,19 +6572,14 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let rhs = slice_types.len(self.encoder, encoded_place); - stmts.extend( - self.encode_copy_value_assign( - encoded_lhs, - rhs, - dst_ty, - location, - )? - ); - }, - other => return Err( - EncodingError::unsupported(format!("length operation on unsupported type '{other:?}'")) - .with_span(span) - ), + stmts.extend(self.encode_copy_value_assign(encoded_lhs, rhs, dst_ty, location)?); + } + other => { + return Err(EncodingError::unsupported(format!( + "length operation on unsupported type '{other:?}'" + )) + .with_span(span)) + } } Ok(stmts) @@ -6309,11 +6596,21 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let span = self.mir_encoder.get_span_of_location(location); let sequence_types = self.encoder.encode_sequence_types(ty).with_span(span)?; - let encoded_operand = self.mir_encoder.encode_operand_expr(operand) + let encoded_operand = self + .mir_encoder + .encode_operand_expr(operand) .with_span(span)?; - let len: usize = self.encoder.const_eval_intlike(mir::ConstantKind::Ty(times)).with_span(span)? - .to_u64().unwrap().try_into().unwrap(); - let lookup_ret_ty = self.encoder.encode_snapshot_type(sequence_types.elem_ty_rs) + let len: usize = self + .encoder + .const_eval_intlike(mir::ConstantKind::Ty(times)) + .with_span(span)? + .to_u64() + .unwrap() + .try_into() + .unwrap(); + let lookup_ret_ty = self + .encoder + .encode_snapshot_type(sequence_types.elem_ty_rs) .with_span(span)?; let inhaled_operand = if lookup_ret_ty.is_domain() || lookup_ret_ty.is_snapshot() { @@ -6322,20 +6619,19 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { encoded_operand }; - let mut stmts = self.encode_havoc_and_initialization(&encoded_lhs).with_span(span)?; + let mut stmts = self + .encode_havoc_and_initialization(&encoded_lhs) + .with_span(span)?; let idx: vir::Expr = vir_local! { i: Int }.into(); - let indices = vir_expr! { ([vir::Expr::from(0usize)] <= [idx]) && ([idx] < [vir::Expr::from(len)]) }; - let lookup_pure_call = sequence_types.encode_lookup_pure_call( - self.encoder, - encoded_lhs, - idx, - lookup_ret_ty, - ); - stmts.push(vir::Stmt::Inhale( vir::Inhale { + let indices = + vir_expr! { ([vir::Expr::from(0usize)] <= [idx]) && ([idx] < [vir::Expr::from(len)]) }; + let lookup_pure_call = + sequence_types.encode_lookup_pure_call(self.encoder, encoded_lhs, idx, lookup_ret_ty); + stmts.push(vir::Stmt::Inhale(vir::Inhale { expr: vir_expr! { - forall i: Int :: - { [lookup_pure_call] } :: - ([indices] ==> ([lookup_pure_call] == [inhaled_operand])) } + forall i: Int :: + { [lookup_pure_call] } :: + ([indices] ==> ([lookup_pure_call] == [inhaled_operand])) }, })); Ok(stmts) @@ -6358,8 +6654,12 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let havoc_ref_method_name = self .encoder .encode_builtin_method_use(BuiltinMethodKind::HavocRef)?; - if let vir::Expr::Local( vir::Local {variable: ref dst_local_var, ..} ) = dst { - Ok(vec![vir::Stmt::MethodCall( vir::MethodCall { + if let vir::Expr::Local(vir::Local { + variable: ref dst_local_var, + .. + }) = dst + { + Ok(vec![vir::Stmt::MethodCall(vir::MethodCall { method_name: havoc_ref_method_name, arguments: vec![], targets: vec![dst_local_var.clone()], @@ -6367,12 +6667,12 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } else { let tmp_var = self.get_auxiliary_local_var("havoc", dst.get_type().clone()); Ok(vec![ - vir::Stmt::MethodCall( vir::MethodCall { + vir::Stmt::MethodCall(vir::MethodCall { method_name: havoc_ref_method_name, arguments: vec![], targets: vec![tmp_var.clone()], }), - vir::Stmt::Assign( vir::Assign { + vir::Stmt::Assign(vir::Assign { target: dst.clone(), source: tmp_var.into(), kind: vir::AssignKind::Move, @@ -6383,15 +6683,19 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { /// Havoc and assume permission on fields #[tracing::instrument(level = "debug", skip(self))] - fn encode_havoc_and_initialization(&mut self, dst: &vir::Expr) -> EncodingResult> { + fn encode_havoc_and_initialization( + &mut self, + dst: &vir::Expr, + ) -> EncodingResult> { let mut stmts = vec![]; // Havoc `dst` stmts.extend(self.encode_havoc(dst)?); // Initialize `dst` - stmts.push(vir::Stmt::Inhale( vir::Inhale { - expr: self.mir_encoder - .encode_place_predicate_permission(dst.clone(), vir::PermAmount::Write) - .unwrap(), + stmts.push(vir::Stmt::Inhale(vir::Inhale { + expr: self + .mir_encoder + .encode_place_predicate_permission(dst.clone(), vir::PermAmount::Write) + .unwrap(), })); Ok(stmts) } @@ -6406,16 +6710,14 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { field: vir::Field, location: mir::Location, vir_assign_kind: vir::AssignKind, - can_copy_reference: bool // reference copies are allowed if the field is a constant + can_copy_reference: bool, // reference copies are allowed if the field is a constant ) -> SpannedEncodingResult> { let span = self.mir_encoder.get_span_of_location(location); if !self.init_info.is_vir_place_accessible(&dst, location) { let mut alloc_stmts = self.encode_havoc(&dst).with_span(span)?; let dst_field = dst.clone().field(field.clone()); let acc = vir::Expr::acc_permission(dst_field.clone(), vir::PermAmount::Write); - alloc_stmts.push(vir::Stmt::Inhale( vir::Inhale { - expr: acc, - })); + alloc_stmts.push(vir::Stmt::Inhale(vir::Inhale { expr: acc })); match vir_assign_kind { vir::AssignKind::Copy => { if field.typ.is_typed_ref_or_type_var() { @@ -6423,11 +6725,9 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let pred_acc = vir::Expr::predicate_access_predicate( field.typ, dst_field, - vir::PermAmount::Read + vir::PermAmount::Read, ); - alloc_stmts.push(vir::Stmt::Inhale( - vir::Inhale { expr: pred_acc } - )); + alloc_stmts.push(vir::Stmt::Inhale(vir::Inhale { expr: pred_acc })); } else { // TODO: Inhale the predicate rooted at dst_field return Err(SpannedEncodingError::unsupported( @@ -6463,9 +6763,9 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { field.clone(), location, vir::AssignKind::Copy, - false + false, )?; - stmts.push(vir::Stmt::Assign( vir::Assign { + stmts.push(vir::Stmt::Assign(vir::Assign { target: lhs.field(field), source: rhs, kind: vir::AssignKind::Copy, @@ -6484,12 +6784,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ) -> SpannedEncodingResult> { let span = self.mir_encoder.get_span_of_location(location); let field = self.encoder.encode_value_field(ty).with_span(span)?; - self.encode_copy_value_assign2( - dst, - src.field(field.clone()), - field, - location - ) + self.encode_copy_value_assign2(dst, src.field(field.clone()), field, location) } /// Copy a value by inhaling snapshot equality. @@ -6499,11 +6794,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { dst: vir::Expr, ) -> EncodingResult> { let mut stmts = self.encode_havoc_and_initialization(&dst)?; - stmts.push(vir::Stmt::Inhale( vir::Inhale { - expr: vir::Expr::eq_cmp( - vir::Expr::snap_app(src), - vir::Expr::snap_app(dst), - ), + stmts.push(vir::Stmt::Inhale(vir::Inhale { + expr: vir::Expr::eq_cmp(vir::Expr::snap_app(src), vir::Expr::snap_app(dst)), })); Ok(stmts) } @@ -6521,9 +6813,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { | ty::TyKind::Int(_) | ty::TyKind::Uint(_) | ty::TyKind::Float(_) - | ty::TyKind::Char => { - self.encode_copy_primitive_value(src, dst, self_ty, location)? - } + | ty::TyKind::Char => self.encode_copy_primitive_value(src, dst, self_ty, location)?, ty::TyKind::Adt(_, _) | ty::TyKind::Closure(_, _) @@ -6535,8 +6825,11 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { _ => { return Err(SpannedEncodingError::unsupported( - format!("copy operation for an unsupported type {:?}", self_ty.kind()), - span + format!( + "copy operation for an unsupported type {:?}", + self_ty.kind() + ), + span, )); } }; @@ -6586,7 +6879,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { if adt_def.is_union() { return Err(SpannedEncodingError::unsupported( "unions are not supported", - span + span, )); } let num_variants = adt_def.variants().len(); @@ -6597,10 +6890,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // Handle *signed* discriminats let discr_type = adt_def.repr().discr_type(); let discr_value: vir::Expr = if discr_type.is_signed() { - let bit_size = - Integer::from_attr(&tcx, discr_type) - .size() - .bits(); + let bit_size = Integer::from_attr(&tcx, discr_type).size().bits(); let shift = 128 - bit_size; let unsigned_discr = adt_def.discriminant_for_variant(tcx, variant_index).val; @@ -6617,20 +6907,24 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let discriminant = self .encoder .encode_discriminant_func_app(dst.clone(), adt_def)?; - stmts.push(vir::Stmt::Inhale( vir::Inhale { + stmts.push(vir::Stmt::Inhale(vir::Inhale { expr: vir::Expr::eq_cmp(discriminant, discr_value), })); let variant_name = variant_def.ident(tcx).to_string(); let new_dst_base = dst_base.variant(&variant_name); - let variant_field = if let vir::Expr::Variant( vir::Variant {variant_index: ref field, ..}) = new_dst_base { + let variant_field = if let vir::Expr::Variant(vir::Variant { + variant_index: ref field, + .. + }) = new_dst_base + { field.clone() } else { unreachable!() }; if !variant_def.fields.is_empty() { - stmts.push(vir::Stmt::Downcast( vir::Downcast { + stmts.push(vir::Stmt::Downcast(vir::Downcast { base: dst.clone(), field: variant_field, })); @@ -6642,7 +6936,8 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let operand = &operands[field_index]; let field_name = field.ident(tcx).to_string(); let field_ty = field.ty(tcx, subst); - let encoded_field = self.encoder + let encoded_field = self + .encoder .encode_struct_field(&field_name, field_ty) .with_span(span)?; stmts.extend(self.encode_assign_operand( @@ -6655,12 +6950,16 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { mir::AggregateKind::Closure(def_id, substs) => { // TODO: might need to assert history invariants? - assert!(!self.encoder.is_spec_closure(def_id), "spec closure: {def_id:?}"); + assert!( + !self.encoder.is_spec_closure(def_id), + "spec closure: {def_id:?}" + ); let cl_substs = substs.as_closure(); for (field_index, field_ty) in cl_substs.upvar_tys().enumerate() { let operand = &operands[field_index]; let field_name = format!("closure_{field_index}"); - let encoded_field = self.encoder + let encoded_field = self + .encoder .encode_raw_ref_field(field_name, field_ty) .with_span(span)?; stmts.extend(self.encode_assign_operand( @@ -6673,14 +6972,22 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { mir::AggregateKind::Array(..) => { let sequence_types = self.encoder.encode_sequence_types(ty).with_span(span)?; - let lookup_ret_ty = self.encoder.encode_snapshot_type(sequence_types.elem_ty_rs) + let lookup_ret_ty = self + .encoder + .encode_snapshot_type(sequence_types.elem_ty_rs) .with_span(span)?; for (idx, operand) in operands.iter().enumerate() { - let lookup_pure_call = sequence_types - .encode_lookup_pure_call(self.encoder, dst.clone(), idx.into(), lookup_ret_ty.clone()); + let lookup_pure_call = sequence_types.encode_lookup_pure_call( + self.encoder, + dst.clone(), + idx.into(), + lookup_ret_ty.clone(), + ); - let encoded_operand = self.mir_encoder.encode_operand_expr(operand) + let encoded_operand = self + .mir_encoder + .encode_operand_expr(operand) .with_span(span)?; stmts.push( @@ -6694,7 +7001,7 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { mir::AggregateKind::Generator(..) => { return Err(SpannedEncodingError::unsupported( "construction of generators is not supported", - span + span, )); } } @@ -6713,7 +7020,10 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } #[tracing::instrument(level = "debug", skip(self))] - fn get_label_after_location(&mut self, location: mir::Location) -> SpannedEncodingResult> { + fn get_label_after_location( + &mut self, + location: mir::Location, + ) -> SpannedEncodingResult> { let opt_label = self.label_after_location.get(&location); if opt_label.is_none() { if config::allow_unreachable_unsupported_code() { @@ -6752,9 +7062,12 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { location: mir::Location, ) -> SpannedEncodingResult<(vir::Expr, Vec, ty::Ty<'tcx>, Option)> { let span = self.mir_encoder.get_span_of_location(location); - let (encoded_place, ty, variant_idx) = self.mir_encoder.encode_place(place).with_span(span)?; + let (encoded_place, ty, variant_idx) = + self.mir_encoder.encode_place(place).with_span(span)?; trace!("encode_place(ty={:?})", ty); - let (encoded_expr, encoding_stmts) = self.postprocess_place_encoding(encoded_place, encode_kind).with_span(span)?; + let (encoded_expr, encoding_stmts) = self + .postprocess_place_encoding(encoded_place, encode_kind) + .with_span(span)?; Ok((encoded_expr, encoding_stmts, ty, variant_idx)) } @@ -6767,13 +7080,23 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ) -> EncodingResult<(vir::Expr, Vec)> { let sequence_types = self.encoder.encode_sequence_types(sequence_ty)?; - let lookup_res: vir::Expr = self.cfg_method.add_fresh_local_var(sequence_types.elem_pred_type.clone()).into(); - let lookup_res_val_field = self.encoder.encode_value_expr(lookup_res.clone(), sequence_types.elem_ty_rs)?; - let snap_lookup_res_val_field = self.encoder.patch_snapshots(vir::Expr::snap_app(lookup_res_val_field))?; + let lookup_res: vir::Expr = self + .cfg_method + .add_fresh_local_var(sequence_types.elem_pred_type.clone()) + .into(); + let lookup_res_val_field = self + .encoder + .encode_value_expr(lookup_res.clone(), sequence_types.elem_ty_rs)?; + let snap_lookup_res_val_field = self + .encoder + .patch_snapshots(vir::Expr::snap_app(lookup_res_val_field))?; - let lookup_ret_ty = self.encoder.encode_snapshot_type(sequence_types.elem_ty_rs)?; + let lookup_ret_ty = self + .encoder + .encode_snapshot_type(sequence_types.elem_ty_rs)?; - let (encoded_base_expr, mut stmts) = self.postprocess_place_encoding(base, ArrayAccessKind::Shared)?; + let (encoded_base_expr, mut stmts) = + self.postprocess_place_encoding(base, ArrayAccessKind::Shared)?; stmts.extend(self.encode_havoc_and_initialization(&lookup_res)?); let idx_val_int = self.encoder.patch_snapshots(vir::Expr::snap_app(index))?; @@ -6785,13 +7108,17 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { lookup_ret_ty, ); - stmts.push(vir::Stmt::Assert( vir::Assert { - expr: vir::Expr::predicate_access_predicate(sequence_types.sequence_pred_type, encoded_base_expr, vir::PermAmount::Read), + stmts.push(vir::Stmt::Assert(vir::Assert { + expr: vir::Expr::predicate_access_predicate( + sequence_types.sequence_pred_type, + encoded_base_expr, + vir::PermAmount::Read, + ), position: vir::Position::default(), })); - stmts.push(vir::Stmt::Inhale( vir::Inhale { - expr: vir_expr!{ [ lookup_pure_call ] == [ snap_lookup_res_val_field ] } + stmts.push(vir::Stmt::Inhale(vir::Inhale { + expr: vir_expr! { [ lookup_pure_call ] == [ snap_lookup_res_val_field ] }, })); Ok((lookup_res, stmts)) } @@ -6806,11 +7133,19 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { ) -> EncodingResult<(vir::Expr, Vec)> { let sequence_types = self.encoder.encode_sequence_types(sequence_ty)?; - let res: vir::Expr = self.cfg_method.add_fresh_local_var(sequence_types.elem_pred_type.clone()).into(); - let res_val_field = self.encoder.encode_value_expr(res.clone(), sequence_types.elem_ty_rs)?; - let snap_res_val_field = self.encoder.patch_snapshots(vir::Expr::snap_app(res_val_field.clone()))?; + let res: vir::Expr = self + .cfg_method + .add_fresh_local_var(sequence_types.elem_pred_type.clone()) + .into(); + let res_val_field = self + .encoder + .encode_value_expr(res.clone(), sequence_types.elem_ty_rs)?; + let snap_res_val_field = self + .encoder + .patch_snapshots(vir::Expr::snap_app(res_val_field.clone()))?; - let (encoded_base_expr, mut stmts) = self.postprocess_place_encoding(base, ArrayAccessKind::Mutable(None, location))?; + let (encoded_base_expr, mut stmts) = + self.postprocess_place_encoding(base, ArrayAccessKind::Mutable(None, location))?; let idx_val_int = self.encoder.patch_snapshots(vir::Expr::snap_app(index))?; @@ -6828,26 +7163,26 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { })); // exhale Array$4$i32(self) - let array_access_pred = vir::Expr::pred_permission( - encoded_base_expr.clone(), - vir::PermAmount::Write, - ).unwrap(); + let array_access_pred = + vir::Expr::pred_permission(encoded_base_expr.clone(), vir::PermAmount::Write).unwrap(); - stmts.push(vir_stmt!{ exhale [array_access_pred] }); + stmts.push(vir_stmt! { exhale [array_access_pred] }); - let old = |e| { vir::Expr::labelled_old(&before_label, e) }; - let old_lhs = |e| { vir::Expr::labelled_old("lhs", e) }; + let old = |e| vir::Expr::labelled_old(&before_label, e); + let old_lhs = |e| vir::Expr::labelled_old("lhs", e); // value of res - let lookup_ret_ty = self.encoder.encode_snapshot_type(sequence_types.elem_ty_rs)?; + let lookup_ret_ty = self + .encoder + .encode_snapshot_type(sequence_types.elem_ty_rs)?; let lookup_pure_call = sequence_types.encode_lookup_pure_call( self.encoder, encoded_base_expr.clone(), idx_val_int.clone(), lookup_ret_ty.clone(), ); - stmts.push(vir::Stmt::Inhale( vir::Inhale { - expr: vir_expr!{ [ old(lookup_pure_call) ] == [ snap_res_val_field ] } + stmts.push(vir::Stmt::Inhale(vir::Inhale { + expr: vir_expr! { [ old(lookup_pure_call) ] == [ snap_res_val_field ] }, })); // inhale magic wand @@ -6864,21 +7199,22 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { // NEEDSWORK: the vir! macro can't do everything we want here.. // TODO: de-duplicate with encode_array_direct_assign - let i_var: vir::Expr = vir_local!{ i: Int }.into(); + let i_var: vir::Expr = vir_local! { i: Int }.into(); - let zero_le_i = vir_expr!{ [ vir::Expr::from(0) ] <= [ i_var ] }; - let i_lt_len = vir_expr!{ [ i_var ] < [ sequence_types.len(self.encoder, encoded_base_expr.clone()) ] }; - let i_ne_idx = vir_expr!{ [ i_var ] != [ old(idx_val_int.clone()) ] }; - let idx_conditions = vir_expr!{ [zero_le_i] && ([i_lt_len] && [i_ne_idx]) }; + let zero_le_i = vir_expr! { [ vir::Expr::from(0) ] <= [ i_var ] }; + let i_lt_len = vir_expr! { [ i_var ] < [ sequence_types.len(self.encoder, encoded_base_expr.clone()) ] }; + let i_ne_idx = vir_expr! { [ i_var ] != [ old(idx_val_int.clone()) ] }; + let idx_conditions = vir_expr! { [zero_le_i] && ([i_lt_len] && [i_ne_idx]) }; let lookup_array_i = sequence_types.encode_lookup_pure_call( self.encoder, encoded_base_expr.clone(), i_var, lookup_ret_ty.clone(), ); - let lookup_same_as_old = vir_expr!{ [lookup_array_i] == [old(lookup_array_i.clone())] }; - let forall_body = vir_expr!{ [idx_conditions] ==> [lookup_same_as_old] }; - let all_others_unchanged = vir_expr!{ forall i: Int :: { [lookup_array_i] } :: [ forall_body ] }; + let lookup_same_as_old = vir_expr! { [lookup_array_i] == [old(lookup_array_i.clone())] }; + let forall_body = vir_expr! { [idx_conditions] ==> [lookup_same_as_old] }; + let all_others_unchanged = + vir_expr! { forall i: Int :: { [lookup_array_i] } :: [ forall_body ] }; let indexed_lookup_pure = sequence_types.encode_lookup_pure_call( self.encoder, encoded_base_expr.clone(), @@ -6886,14 +7222,15 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { lookup_ret_ty, ); // TODO: old inside snapshot or the other way around? - let snap_old_res_val_field = self.encoder.patch_snapshots(vir::Expr::snap_app(res_val_field.clone()))?; - let indexed_updated = vir_expr!{ [ indexed_lookup_pure ] == [ old_lhs(snap_old_res_val_field) ] }; + let snap_old_res_val_field = self + .encoder + .patch_snapshots(vir::Expr::snap_app(res_val_field.clone()))?; + let indexed_updated = + vir_expr! { [ indexed_lookup_pure ] == [ old_lhs(snap_old_res_val_field) ] }; - let magic_wand_rhs = vir_expr!{ [all_others_unchanged] && [indexed_updated] }; - self.array_magic_wand_at.insert( - location, - (res_val_field, encoded_base_expr, magic_wand_rhs) - ); + let magic_wand_rhs = vir_expr! { [all_others_unchanged] && [indexed_updated] }; + self.array_magic_wand_at + .insert(location, (res_val_field, encoded_base_expr, magic_wand_rhs)); Ok((res, stmts)) } @@ -6910,8 +7247,18 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { let (expr, stmts) = self.postprocess_place_encoding(base, array_encode_kind)?; (expr.field(field), stmts) } - PlaceEncoding::SliceAccess { base, index, rust_slice_ty: rust_ty, .. } | - PlaceEncoding::ArrayAccess { base, index, rust_array_ty: rust_ty, .. } => { + PlaceEncoding::SliceAccess { + base, + index, + rust_slice_ty: rust_ty, + .. + } + | PlaceEncoding::ArrayAccess { + base, + index, + rust_array_ty: rust_ty, + .. + } => { if let ArrayAccessKind::Mutable(loan, location) = array_encode_kind { self.encode_sequence_lookup_mut(*base, index, rust_ty, loan, location)? } else { @@ -6920,17 +7267,23 @@ impl<'p, 'v: 'p, 'tcx: 'v> ProcedureEncoder<'p, 'v, 'tcx> { } PlaceEncoding::Variant { box base, field } => { let (expr, stmts) = self.postprocess_place_encoding(base, array_encode_kind)?; - (vir::Expr::Variant( vir::Variant { - base: box expr, - variant_index: field, - position: vir::Position::default(), - }), - stmts) + ( + vir::Expr::Variant(vir::Variant { + base: box expr, + variant_index: field, + position: vir::Position::default(), + }), + stmts, + ) } }) } - fn register_error + Debug>(&self, span: T, error_ctxt: ErrorCtxt) -> vir::Position { + fn register_error + Debug>( + &self, + span: T, + error_ctxt: ErrorCtxt, + ) -> vir::Position { self.mir_encoder.register_error(span, error_ctxt) } } @@ -6950,17 +7303,19 @@ fn convert_loans_to_borrows(loans: &[facts::Loan]) -> Vec { /// len: Length of borrow_infos fn assert_one_magic_wand(len: usize) -> EncodingResult<()> { if len > 1 { - Err(EncodingError::internal( - format!("We can have at most one magic wand in the postcondition. But we have {len:?}") - )) - } else { Ok(()) } + Err(EncodingError::internal(format!( + "We can have at most one magic wand in the postcondition. But we have {len:?}" + ))) + } else { + Ok(()) + } } // Checks if a type is a reference to a string, or a reference to a reference to a string, etc. fn is_str(ty: ty::Ty<'_>) -> bool { match ty.kind() { ty::TyKind::Ref(_, inner, _) => inner.is_str() || is_str(*inner), - _ => false + _ => false, } } diff --git a/viper-sys/build.rs b/viper-sys/build.rs index 84622597571..c7adc079135 100644 --- a/viper-sys/build.rs +++ b/viper-sys/build.rs @@ -813,6 +813,56 @@ fn main() { method!("field"), method!("perm"), method!("entry"), + ]), + // SIF + java_class!("viper.silver.sif.SIFExtendedTransformer$", vec![ + object_getter!(), + method!("transform"), + ]), + java_class!("viper.silver.sif.SIFReturnStmt", vec![ + constructor!(), + ]), + java_class!("viper.silver.sif.SIFBreakStmt", vec![ + constructor!(), + ]), + java_class!("viper.silver.sif.SIFContinueStmt", vec![ + constructor!(), + ]), + java_class!("viper.silver.sif.SIFRaiseStmt", vec![ + constructor!(), + ]), + java_class!("viper.silver.sif.SIFExceptionHandler", vec![ + constructor!(), + ]), + java_class!("viper.silver.sif.SIFTryCatchStmt", vec![ + constructor!(), + ]), + java_class!("viper.silver.sif.SIFDeclassifyStmt", vec![ + constructor!(), + ]), + java_class!("viper.silver.sif.SIFInlinedCallStmt", vec![ + constructor!(), + ]), + java_class!("viper.silver.sif.SIFAssertNoException", vec![ + constructor!(), + ]), + java_class!("viper.silver.sif.SIFLowExp", vec![ + constructor!(), + ]), + java_class!("viper.silver.sif.SIFLowEventExp", vec![ + constructor!(), + ]), + java_class!("viper.silver.sif.SIFLowExitExp", vec![ + constructor!(), + ]), + java_class!("viper.silver.sif.SIFTerminatesExp", vec![ + constructor!(), + ]), + java_class!("viper.silver.sif.SIFInfo", vec![ + constructor!(), + ]), + java_class!("viper.silver.sif.SIFDynCheckInfo", vec![ + constructor!(), ]) ]) .generate(&generated_dir) diff --git a/viper/src/ast_factory/expression.rs b/viper/src/ast_factory/expression.rs index 1c34c15f89f..3b53b7d6225 100644 --- a/viper/src/ast_factory/expression.rs +++ b/viper/src/ast_factory/expression.rs @@ -7,7 +7,7 @@ use crate::ast_factory::{ AstFactory, }; use jni::objects::JObject; -use viper_sys::wrappers::viper::silver::ast; +use viper_sys::wrappers::viper::silver::{ast, sif}; // Floating-Point Operations #[derive(Debug, Clone, Copy)] @@ -1489,4 +1489,20 @@ impl<'a> AstFactory<'a> { ); Expr::new(obj) } + + pub fn low_with_pos(&self, expr: Expr, pos: Position) -> Expr<'a> { + build_ast_node_with_pos!( + self, + Expr, + sif::SIFLowExp, + expr.to_jobject(), + self.jni.new_option(None), + self.jni.new_map(&[]), + pos.to_jobject() + ) + } + + pub fn low_event(&self) -> Expr<'a> { + build_ast_node!(self, Expr, sif::SIFLowEventExp) + } } diff --git a/viper/src/lib.rs b/viper/src/lib.rs index 747f7fc1261..db6d955cd25 100644 --- a/viper/src/lib.rs +++ b/viper/src/lib.rs @@ -21,6 +21,7 @@ mod verification_backend; mod verification_context; mod verification_result; mod verifier; +pub mod sif_transformer; mod viper; pub use crate::{ diff --git a/viper/src/sif_transformer.rs b/viper/src/sif_transformer.rs new file mode 100644 index 00000000000..27c01ac1e45 --- /dev/null +++ b/viper/src/sif_transformer.rs @@ -0,0 +1,60 @@ +// © 2023, ETH Zurich +// +// This Source Code Form is subject to the terms of the Mozilla Public +// License, v. 2.0. If a copy of the MPL was not distributed with this +// file, You can obtain one at http://mozilla.org/MPL/2.0/. + +use crate::{ast_factory::Program, ast_utils::AstUtils, jni_utils::JniUtils}; +use jni::JNIEnv; +use log::debug; +use viper_sys::wrappers::viper::silver; + +pub struct SIFTransformer<'a> { + env: &'a JNIEnv<'a>, + jni: JniUtils<'a>, + ast_utils: AstUtils<'a>, +} + +impl<'a> SIFTransformer<'a> { + pub fn new(env: &'a JNIEnv) -> Self { + let jni = JniUtils::new(env); + let ast_utils = AstUtils::new(env); + SIFTransformer { + env, + jni, + ast_utils, + } + } + + #[must_use] + pub fn sif_transformation(self, program: Program<'a>) -> Program<'a> { + debug!( + "Program before transformation:\n{}", + self.ast_utils.pretty_print(program) + ); + + let sif_transformer = silver::sif::SIFExtendedTransformer_object::with(self.env); + + run_timed!("SIF transformation", debug, + let transformed_program = self.jni.unwrap_result( + sif_transformer + .call_transform(self.jni.unwrap_result(sif_transformer.singleton()), program.to_jobject(), false), + ); + ); + + debug_assert!(self + .env + .is_instance_of( + transformed_program, + self.env.find_class("viper/silver/ast/Program").unwrap(), + ) + .unwrap()); + + let transformed_program = Program::new(transformed_program); + debug!( + "Program after transformation:\n{}", + self.ast_utils.pretty_print(transformed_program) + ); + transformed_program + } +} diff --git a/viper/src/verification_context.rs b/viper/src/verification_context.rs index 2537431c837..deaeefc0063 100644 --- a/viper/src/verification_context.rs +++ b/viper/src/verification_context.rs @@ -5,7 +5,8 @@ // file, You can obtain one at http://mozilla.org/MPL/2.0/. use crate::{ - ast_factory::*, ast_utils::*, verification_backend::VerificationBackend, verifier::Verifier, + ast_factory::*, ast_utils::*, sif_transformer::SIFTransformer, + verification_backend::VerificationBackend, verifier::Verifier, }; use jni::AttachGuard; use log::{debug, info}; @@ -99,4 +100,8 @@ impl<'a> VerificationContext<'a> { .parse_command_line(&verifier_args) .start() } + + pub fn new_sif_transformer(&self) -> SIFTransformer { + SIFTransformer::new(&self.env) + } } diff --git a/vir/defs/polymorphic/ast/expr.rs b/vir/defs/polymorphic/ast/expr.rs index 12e3183dfaf..3fd79f4aae8 100644 --- a/vir/defs/polymorphic/ast/expr.rs +++ b/vir/defs/polymorphic/ast/expr.rs @@ -70,6 +70,10 @@ pub enum Expr { SnapApp(SnapApp), /// Cast from one type into another. Cast(Cast), + /// Set the given expression to the low security level + Low(Low), + /// The expression must always be reached by none or all execussion + LowEvent, } impl fmt::Display for Expr { @@ -102,6 +106,8 @@ impl fmt::Display for Expr { Expr::Map(map) => map.fmt(f), Expr::Downcast(downcast_expr) => downcast_expr.fmt(f), Expr::Cast(expr) => expr.fmt(f), + Expr::Low(low) => low.fmt(f), + Expr::LowEvent => write!(f, "low_event"), } } } @@ -184,8 +190,10 @@ impl Expr { | Expr::ContainerOp(ContainerOp { position, .. }) | Expr::Cast(Cast { position, .. }) | Expr::Map(Map { position, .. }) + | Expr::Low(Low { position, .. }) | Expr::Seq(Seq { position, .. }) => *position, Expr::Downcast(DowncastExpr { base, .. }) => base.pos(), + Expr::LowEvent => Default::default(), } } @@ -209,6 +217,7 @@ impl Expr { ..inner }), Expr::Downcast(..) => self, + Expr::LowEvent => self, } } } @@ -235,7 +244,8 @@ impl Expr { DomainFuncApp, InhaleExhale, SnapApp, - Cast + Cast, + Low ) } @@ -473,6 +483,18 @@ impl Expr { }) } + pub fn low(expr: Expr) -> Self { + let pos = expr.pos(); + Expr::Low(Low { + base: Box::new(expr), + position: pos, + }) + } + + pub fn low_event() -> Self { + Expr::LowEvent + } + pub fn find(&self, sub_target: &Expr) -> bool { pub struct ExprFinder<'a> { sub_target: &'a Expr, @@ -1047,7 +1069,7 @@ impl Expr { assert_eq!(typ1, typ2, "expr: {:?}", self); typ1 } - Expr::ForAll(..) | Expr::Exists(..) => &Type::Bool, + Expr::ForAll(..) | Expr::Exists(..) | Expr::Low(..) | Expr::LowEvent => &Type::Bool, Expr::MagicWand(..) | Expr::PredicateAccessPredicate(..) | Expr::FieldAccessPredicate(..) @@ -1530,6 +1552,8 @@ impl Expr { | Expr::Seq(..) | Expr::Map(..) | Expr::SnapApp(..) + | Expr::Low(..) + | Expr::LowEvent | Expr::Cast(..) => true.into(), } } @@ -1751,6 +1775,23 @@ impl Expr { let mut remover = ReadPermRemover {}; remover.fold(self) } + + pub fn is_relational(&self) -> bool { + struct RelationalFinder { + found: bool, + } + impl ExprWalker for RelationalFinder { + fn walk_low(&mut self, _expr: &Low) { + self.found = true; + } + + fn walk_low_event(&mut self) { + self.found = true; + } + } + let founder = RelationalFinder { found: false }; + founder.found + } } /// A component that can be used to represent a place as a vector. @@ -2608,6 +2649,30 @@ impl Hash for Cast { } } +#[derive(Debug, Clone, Eq, serde::Serialize, serde::Deserialize, PartialOrd, Ord)] +pub struct Low { + pub base: Box, + pub position: Position, +} + +impl fmt::Display for Low { + fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result { + write!(f, "low({})", self.base) + } +} + +impl PartialEq for Low { + fn eq(&self, other: &Self) -> bool { + self.base == other.base + } +} + +impl Hash for Low { + fn hash(&self, state: &mut H) { + self.base.hash(state); + } +} + #[derive(Debug, Clone, Eq, serde::Serialize, serde::Deserialize, PartialOrd, Ord)] pub struct SnapApp { pub base: Box, diff --git a/vir/defs/polymorphic/ast/expr_transformers.rs b/vir/defs/polymorphic/ast/expr_transformers.rs index c6fe744b85f..dbd425da188 100644 --- a/vir/defs/polymorphic/ast/expr_transformers.rs +++ b/vir/defs/polymorphic/ast/expr_transformers.rs @@ -380,6 +380,18 @@ pub trait ExprFolder: Sized { }) } + fn fold_low(&mut self, expr: Low) -> Expr { + let Low { base, position } = expr; + Expr::Low(Low { + base: self.fold_boxed(base), + position, + }) + } + + fn fold_low_event(&mut self) -> Expr { + Expr::LowEvent + } + fn fold_snap_app(&mut self, expr: SnapApp) -> Expr { let SnapApp { base, position } = expr; Expr::SnapApp(SnapApp { @@ -470,6 +482,8 @@ pub fn default_fold_expr(this: &mut T, e: Expr) -> Expr { Expr::Seq(seq) => this.fold_seq(seq), Expr::Map(map) => this.fold_map(map), Expr::Cast(cast) => this.fold_cast(cast), + Expr::Low(low) => this.fold_low(low), + Expr::LowEvent => this.fold_low_event(), } } @@ -663,6 +677,14 @@ pub trait ExprWalker: Sized { self.walk(base); self.walk(enum_place); } + + fn walk_low(&mut self, expr: &Low) { + let Low { base, .. } = expr; + self.walk(base); + } + + fn walk_low_event(&mut self) {} + fn walk_snap_app(&mut self, expr: &SnapApp) { let SnapApp { base, .. } = expr; self.walk(base); @@ -728,6 +750,8 @@ pub fn default_walk_expr(this: &mut T, e: &Expr) { Expr::Seq(seq) => this.walk_seq(seq), Expr::Map(map) => this.walk_map(map), Expr::Cast(cast) => this.walk_cast(cast), + Expr::Low(low) => this.walk_low(low), + Expr::LowEvent => this.walk_low_event(), } } @@ -1068,6 +1092,18 @@ pub trait FallibleExprFolder: Sized { })) } + fn fallible_fold_low(&mut self, expr: Low) -> Result { + let Low { base, position } = expr; + Ok(Expr::Low(Low { + base: self.fallible_fold_boxed(base)?, + position, + })) + } + + fn fallible_fold_low_event(&mut self) -> Result { + Ok(Expr::LowEvent) + } + fn fallible_fold_snap_app(&mut self, expr: SnapApp) -> Result { let SnapApp { base, position } = expr; Ok(Expr::SnapApp(SnapApp { @@ -1172,6 +1208,8 @@ pub fn default_fallible_fold_expr>( Expr::Seq(seq) => this.fallible_fold_seq(seq), Expr::Map(map) => this.fallible_fold_map(map), Expr::Cast(cast) => this.fallible_fold_cast(cast), + Expr::Low(low) => this.fallible_fold_low(low), + Expr::LowEvent => this.fallible_fold_low_event(), } } @@ -1384,6 +1422,15 @@ pub trait FallibleExprWalker: Sized { Ok(()) } + fn fallible_walk_low(&mut self, expr: &Low) -> Result<(), Self::Error> { + let Low { base, .. } = expr; + self.fallible_walk(base) + } + + fn fallible_walk_low_event(&mut self) -> Result<(), Self::Error> { + Ok(()) + } + fn fallible_walk_snap_app(&mut self, expr: &SnapApp) -> Result<(), Self::Error> { let SnapApp { base, .. } = expr; self.fallible_walk(base) @@ -1453,5 +1500,7 @@ pub fn default_fallible_walk_expr>( Expr::Seq(seq) => this.fallible_walk_seq(seq), Expr::Map(map) => this.fallible_walk_map(map), Expr::Cast(cast) => this.fallible_walk_cast(cast), + Expr::Low(low) => this.fallible_walk_low(low), + Expr::LowEvent => this.fallible_walk_low_event(), } } diff --git a/vir/defs/polymorphic/ast/stmt.rs b/vir/defs/polymorphic/ast/stmt.rs index e6d25ce2434..5c94c1b08bb 100644 --- a/vir/defs/polymorphic/ast/stmt.rs +++ b/vir/defs/polymorphic/ast/stmt.rs @@ -624,6 +624,7 @@ pub trait StmtFolder { kind, }) } + fn fold_fold(&mut self, statement: Fold) -> Stmt { let Fold { predicate, diff --git a/vir/src/converter/polymorphic_to_legacy.rs b/vir/src/converter/polymorphic_to_legacy.rs index 912aa5de3f1..f39be4dd3a5 100644 --- a/vir/src/converter/polymorphic_to_legacy.rs +++ b/vir/src/converter/polymorphic_to_legacy.rs @@ -406,6 +406,10 @@ impl From for legacy::Expr { Box::new((*cast.base).into()), cast.position.into(), ), + polymorphic::Expr::Low(low) => { + legacy::Expr::Low(Box::new((*low.base).into()), low.position.into()) + } + polymorphic::Expr::LowEvent => legacy::Expr::LowEvent, } } } diff --git a/vir/src/converter/type_substitution.rs b/vir/src/converter/type_substitution.rs index 24f2a36ef5e..c0e1de516ee 100644 --- a/vir/src/converter/type_substitution.rs +++ b/vir/src/converter/type_substitution.rs @@ -169,6 +169,8 @@ impl Generic for Expr { Expr::Downcast(down_cast) => Expr::Downcast(down_cast.substitute(map)), Expr::SnapApp(snap_app) => Expr::SnapApp(snap_app.substitute(map)), Expr::Cast(cast) => Expr::Cast(cast.substitute(map)), + Expr::Low(low) => Expr::Low(low.substitute(map)), + Expr::LowEvent => Expr::LowEvent, } } } @@ -437,6 +439,13 @@ impl Generic for Cast { } } +impl Generic for Low { + fn substitute(mut self, map: &FxHashMap) -> Self { + *self.base = self.base.substitute(map); + self + } +} + // function impl Generic for Function { fn substitute(self, map: &FxHashMap) -> Self { diff --git a/vir/src/legacy/ast/expr.rs b/vir/src/legacy/ast/expr.rs index 32e2824fc8f..899471bb41e 100644 --- a/vir/src/legacy/ast/expr.rs +++ b/vir/src/legacy/ast/expr.rs @@ -75,6 +75,10 @@ pub enum Expr { /// Snapshot call to convert from a Ref to a snapshot value SnapApp(Box, Position), Cast(CastKind, Box, Position), + /// check that the given expression is at the low security level + Low(Box, Position), + /// check that both execution reach this point or none + LowEvent, } /// A component that can be used to represent a place as a vector. @@ -306,6 +310,8 @@ impl fmt::Display for Expr { Expr::SnapApp(ref expr, _) => write!(f, "snap({expr})"), Expr::Cast(ref kind, ref base, _) => write!(f, "cast<{kind:?}>({base})"), + Expr::Low(ref base, _) => write!(f, "low({base})"), + Expr::LowEvent => write!(f, "low_event"), } } } @@ -388,9 +394,11 @@ impl Expr { | Expr::Seq(_, _, p) | Expr::Map(_, _, p) | Expr::Cast(_, _, p) + | Expr::Low(_, p) | Expr::SnapApp(_, p) => *p, // TODO Expr::DomainFuncApp(_, _, _, _, _, p) => p, Expr::Downcast(box ref base, ..) => base.pos(), + Expr::LowEvent => Position::default(), } } @@ -1243,7 +1251,7 @@ impl Expr { assert_eq!(typ1, typ2, "expr: {self:?}"); typ1? } - Expr::ForAll(..) | Expr::Exists(..) => &Type::Bool, + Expr::ForAll(..) | Expr::Exists(..) | Expr::Low(..) | Expr::LowEvent => &Type::Bool, Expr::MagicWand(..) | Expr::PredicateAccessPredicate(..) | Expr::FieldAccessPredicate(..) @@ -1716,6 +1724,8 @@ impl Expr { | Expr::Seq(..) | Expr::Map(..) | Expr::SnapApp(..) + | Expr::Low(..) + | Expr::LowEvent | Expr::Cast(..) => true.into(), } } diff --git a/vir/src/legacy/ast/expr_transformers.rs b/vir/src/legacy/ast/expr_transformers.rs index aa490dd0e8f..4f348cb1cd7 100644 --- a/vir/src/legacy/ast/expr_transformers.rs +++ b/vir/src/legacy/ast/expr_transformers.rs @@ -338,7 +338,15 @@ pub trait ExprFolder: Sized { } fn fold_cast(&mut self, kind: CastKind, base: Box, p: Position) -> Expr { - Expr::Cast(kind, base, self.fold_position(p)) + Expr::Cast(kind, self.fold_boxed(base), self.fold_position(p)) + } + + fn fold_low(&mut self, expr: Box, p: Position) -> Expr { + Expr::Low(self.fold_boxed(expr), self.fold_position(p)) + } + + fn fold_low_event(&mut self) -> Expr { + Expr::LowEvent } fn fold_position(&mut self, p: Position) -> Position { @@ -378,6 +386,8 @@ pub fn default_fold_expr(this: &mut T, e: Expr) -> Expr { Expr::Seq(x, y, p) => this.fold_seq(x, y, p), Expr::Map(x, y, p) => this.fold_map(x, y, p), Expr::Cast(kind, base, p) => this.fold_cast(kind, base, p), + Expr::Low(e, p) => this.fold_low(e, p), + Expr::LowEvent => this.fold_low_event(), } } @@ -606,6 +616,13 @@ pub trait ExprWalker: Sized { self.walk_position(pos); } + fn walk_low(&mut self, expr: &Expr, pos: &Position) { + self.walk(expr); + self.walk_position(pos); + } + + fn walk_low_event(&mut self) {} + fn walk_position(&mut self, _pos: &Position) {} } @@ -641,6 +658,8 @@ pub fn default_walk_expr(this: &mut T, e: &Expr) { Expr::Seq(ref ty, ref elems, ref p) => this.walk_seq(ty, elems, p), Expr::Map(ref ty, ref elems, ref p) => this.walk_map(ty, elems, p), Expr::Cast(ref kind, ref base, ref p) => this.walk_cast(kind, base, p), + Expr::Low(ref expr, ref p) => this.walk_low(expr, p), + Expr::LowEvent => this.walk_low_event(), } } @@ -964,6 +983,14 @@ pub trait FallibleExprFolder: Sized { ) -> Result { Ok(Expr::Cast(kind, self.fallible_fold_boxed(base)?, p)) } + + fn fallible_fold_low(&mut self, expr: Box, p: Position) -> Result { + Ok(Expr::Low(self.fallible_fold_boxed(expr)?, p)) + } + + fn fallible_fold_low_event(&mut self) -> Result { + Ok(Expr::LowEvent) + } } pub fn default_fallible_fold_expr>( @@ -1001,5 +1028,7 @@ pub fn default_fallible_fold_expr>( Expr::Seq(x, y, p) => this.fallible_fold_seq(x, y, p), Expr::Map(x, y, p) => this.fallible_fold_map(x, y, p), Expr::Cast(kind, base, p) => this.fallible_fold_cast(kind, base, p), + Expr::Low(expr, p) => this.fallible_fold_low(expr, p), + Expr::LowEvent => this.fallible_fold_low_event(), } }