/// \file osra.cpp /// \brief // // This file is distributed under the MIT License. See LICENSE.md for details. // // Standard includes #include #include // LLVM includes #include "llvm/Analysis/ConstantFolding.h" #include "llvm/IR/AssemblyAnnotationWriter.h" #include "llvm/IR/Constants.h" #include "llvm/IR/DataLayout.h" #include "llvm/IR/Module.h" #include "llvm/Support/FormattedStream.h" #include "llvm/Support/raw_os_ostream.h" #include "llvm/Pass.h" // Local includes #include "datastructures.h" #include "debug.h" #include "memoryaccess.h" #include "revamb.h" #include "ir-helpers.h" #include "osra.h" using namespace llvm; using Predicate = CmpInst::Predicate; using OSR = OSRAPass::OSR; using BoundedValue = OSRAPass::BoundedValue; using CE = ConstantExpr; using CI = ConstantInt; using std::pair; using std::make_pair; using std::numeric_limits; const BoundedValue::MergeType AndMerge = BoundedValue::And; const BoundedValue::MergeType OrMerge = BoundedValue::Or; template static auto skip(unsigned ToSkip, C &Container) -> iterator_range { auto Begin = std::begin(Container); while (ToSkip --> 0) Begin++; return make_range(Begin, std::end(Container)); } char OSRAPass::ID = 0; static RegisterPass X("osra", "OSRA Pass", true, true); Constant *OSR::evaluate(Constant *Value, Type *Int64) const { Constant *BaseC = CI::get(Int64, Base, BV->isSigned()); Constant *FactorC = CI::get(Int64, Factor, BV->isSigned()); return CE::getAdd(BaseC, CE::getMul(FactorC, Value)); } static bool isPositive(Constant *C, const DataLayout &DL) { auto *Zero = CI::get(C->getType(), 0, true); auto *Compare = CE::getCompare(CmpInst::ICMP_SGE, C, Zero); return getConstValue(Compare, DL)->getLimitedValue(); } pair OSR::boundaries(Type *Int64, const DataLayout &DL) const { Constant *Min = nullptr; Constant *Max = nullptr; std::tie(Min, Max) = BV->actualBoundaries(Int64); Min = evaluate(Min, Int64); Max = evaluate(Max, Int64); return { Min, Max }; } /// \brief Combine two constants using \p Opcode operation /// /// \param Opcode the opcode of the binary operator. /// \param Signed whether the operands are signed or not. /// \param Op1 the first operand. /// \param Op2 the second operand. /// \param T the type of the operands the result. /// \param DL the DataLayout to compute the result. /// \return the result of the operation. static uint64_t combineImpl(unsigned Opcode, bool Signed, Constant *Op1, Constant *Op2, IntegerType *T, const DataLayout &DL) { auto *R = ConstantFoldInstOperands(Opcode, T, { Op1, Op2 }, DL); return getExtValue(R, Signed, DL); } static uint64_t combineImpl(unsigned Opcode, bool Signed, uint64_t Op1, Constant *Op2, IntegerType *T, const DataLayout &DL) { return combineImpl(Opcode, Signed, CI::get(T, Op1, Signed), Op2, T, DL); } static uint64_t combineImpl(unsigned Opcode, bool Signed, Constant *Op1, uint64_t Op2, IntegerType *T, const DataLayout &DL) { return combineImpl(Opcode, Signed, Op1, CI::get(T, Op2, Signed), T, DL); } uint64_t BoundedValue::performOp(uint64_t Op1, unsigned Opcode, uint64_t Op2, const DataLayout &DL) const { assert(Value != nullptr); // Obtain the type IntegerType *Ty = dyn_cast(Value->getType()); // If it's not an integer type it must be a Store instruction if (Ty == nullptr) { auto *Store = cast(Value); Ty = cast(Store->getValueOperand()->getType()); } // Build operands bool IsSigned = isSigned(); auto *COp1 = CI::get(Ty, Op1, IsSigned); auto *COp2 = CI::get(Ty, Op2, IsSigned); // Compute the result auto *Result = ConstantFoldInstOperands(Opcode, Ty, { COp1, COp2 }, DL); return getExtValue(Result, IsSigned, DL); } BoundedValue BoundedValue::moveTo(llvm::Value *V, const DataLayout &DL, uint64_t Offset, uint64_t Multiplier) const { BoundedValue Result = *this; Result.Value = V; using I = Instruction; if (Result.LowerBound != Result.lowerExtreme()) { Result.LowerBound = performOp(Result.LowerBound, I::Mul, Multiplier, DL); Result.LowerBound = performOp(Result.LowerBound, I::Add, Offset, DL); } if (Result.UpperBound != Result.upperExtreme()) { Result.UpperBound = performOp(Result.UpperBound, I::Mul, Multiplier, DL); Result.UpperBound = performOp(Result.UpperBound, I::Add, Offset, DL); } return Result; } bool OSR::combine(unsigned Opcode, Constant *Operand, unsigned FreeOpIndex, const DataLayout &DL) { using I = Instruction; auto *TheType = cast(Operand->getType()); bool Multiplicative = !(Opcode == I::Add || Opcode == I::Sub); bool Signed = (Opcode == I::SDiv || Opcode == I::AShr); Operand = getConstValue(Operand, DL); uint64_t OldValue = Base; uint64_t OldFactor = Factor; bool Changed = false; // Handle the only case of non-commutative operation with first operand // constant that we handle: subtraction if (!I::isCommutative(Opcode) && FreeOpIndex != 0) { assert(Opcode == I::Sub); // c - x // x = a + b * y // (c - a) + (-b) * y Base = combineImpl(Opcode, Signed, Operand, Base, TheType, DL); Changed |= Base != OldValue; auto *MinusOne = Constant::getAllOnesValue(TheType); Factor = combineImpl(I::Mul, Signed, MinusOne, Factor, TheType, DL); Changed |= OldFactor != Factor; } else { // Commutative/second operand constant case Base = combineImpl(Opcode, Signed, Base, Operand, TheType, DL); Changed |= Base != OldValue; if (Multiplicative) { Factor = combineImpl(Opcode, Signed, Factor, Operand, TheType, DL); Changed |= OldFactor != Factor; } } return Changed; } class OSRAnnotationWriter : public AssemblyAnnotationWriter { public: OSRAnnotationWriter(OSRAPass &JTFC) : JTFC(JTFC) { } virtual void emitInstructionAnnot(const Instruction *I, formatted_raw_ostream &Output) { JTFC.describe(Output, I); } virtual void emitBasicBlockStartAnnot(const BasicBlock *BB, formatted_raw_ostream &Output) { JTFC.describe(Output, BB); } private: OSRAPass &JTFC; }; void OSR::describe(formatted_raw_ostream &O) const { O << "[" << static_cast(Base) << " + " << static_cast(Factor) << " * x, with x = "; if (BV == nullptr) O << "null"; else BV->describe(O); O << "]"; } void BoundedValue::describe(formatted_raw_ostream &O) const { if (Negated) O << "NOT "; O << "("; O << getName(Value); O << ", "; switch (Sign) { case AnySignedness: O << "*"; break; case UnknownSignedness: O << "?"; break; case Signed: O << "s"; break; case Unsigned: O << "u"; break; case InconsistentSignedness: O << "x"; break; } if (Bottom) { O << ", bottom"; } else if (!isUninitialized()) { O << ", "; if (!isConstant() && LowerBound == lowerExtreme()) { O << "min"; } else { O << LowerBound; } O << ", "; if (!isConstant() && UpperBound == upperExtreme()) { O << "max"; } else { O << UpperBound; } } O << ")"; } void OSRAPass::describe(formatted_raw_ostream &O, const BasicBlock *BB) const { BVs.describe(O, BB); } void OSRAPass::describe(formatted_raw_ostream &O, const Instruction *I) const { auto OSRIt = OSRs.find(I); auto ConstraintsIt = Constraints.find(I); if (OSRIt == OSRs.end() && ConstraintsIt == Constraints.end()) return; if (OSRIt != OSRs.end()) { O << " ; "; OSRIt->second.describe(O); O << "\n"; } if (ConstraintsIt != Constraints.end()) { O << " ;"; for (auto Constraint : ConstraintsIt->second) { O << " "; Constraint.describe(O); } O << "\n"; } if (auto *Load = dyn_cast(I)) { auto LoadReachersIt = LoadReachers.find(Load); if (LoadReachersIt != LoadReachers.end()) { O << " ; "; for (auto P : LoadReachersIt->second) { O << "{" << getName(P.first) << ", "; P.second.describe(O); O << "} "; } O << "\n"; } } } Constant *OSR::solveEquation(Constant *KnownTerm, bool CeilingRounding, const DataLayout &DL) { // (KnownTerm - Base) udiv Factor bool IsSigned = BV->isSigned(); auto *BaseConst = CI::get(KnownTerm->getType(), Base, IsSigned); auto *Numerator = CE::getSub(KnownTerm, BaseConst); auto *Denominator = CI::get(KnownTerm->getType(), Factor, IsSigned); Constant *Remainder = nullptr; Constant *Division = nullptr; if (IsSigned) { Remainder = CE::getSRem(Numerator, Denominator); Division = CE::getSDiv(Numerator, Denominator); } else { Remainder = CE::getURem(Numerator, Denominator); Division = CE::getUDiv(Numerator, Denominator); } if (isa(Division)) return Division; bool HasRemainder = getConstValue(Remainder, DL)->getLimitedValue() != 0; if (CeilingRounding && HasRemainder) Division = CE::getAdd(Division, CI::get(Division->getType(), 1)); return Division; } OSR OSRAPass::createOSR(Value *V, BasicBlock *BB) { auto OtherOSRIt = OSRs.find(V); if (OtherOSRIt != OSRs.end()) return switchBlock(OtherOSRIt->second, BB); else return OSR(&BVs.get(BB, V)); } /// Helper function to check if two BV vectors are identical static bool differ(SmallVector &Old, SmallVector &New) { if (Old.size() != New.size()) return true; for (auto &OldConstraint : Old) { bool Found = false; for (auto &NewConstraint : New) { if (OldConstraint.value() == NewConstraint.value()) { Found = true; if (!(OldConstraint == NewConstraint)) return true; } } if (!Found) return true; } return false; } template static bool mergeBVVectors(OSRAPass::BVVector &Base, OSRAPass::BVVector &New, const DataLayout &DL, Type *Int64) { bool Result = false; // Merge the two BV vectors for (auto &NewConstraint : New) { bool Found = false; for (auto &BaseConstraint : Base) { if (NewConstraint.value() == BaseConstraint.value()) { Result |= BaseConstraint.merge(NewConstraint, DL, Int64); Found = true; break; } } if (!Found) { Result = true; Base.push_back(NewConstraint); } } return Result; } /// Given an instruction, identifies, if possible, the constant operand. If /// both operands are constant, it returns a Constant with the folded operation /// and nullptr. If only one is constant, it return the constant and a reference /// to the free operand. If none of the operands are constant returns { nullptr, /// nullptr }. It also returns { nullptr, nullptr } if I is not commutative and /// only the first operand is constant. std::pair OSRAPass::identifyOperands(const Instruction *I, const DataLayout &DL) { assert(I->getNumOperands() == 2); Value *FirstOp = I->getOperand(0); Value *SecondOp = I->getOperand(1); Constant *Constants[2] = { dyn_cast(FirstOp), dyn_cast(SecondOp) }; // Is the first operand constant? if (auto *Operand = dyn_cast(FirstOp)) { auto OSRIt = OSRs.find(Operand); if (OSRIt != OSRs.end() && OSRIt->second.isConstant()) Constants[0] = CI::get(Operand->getType(), OSRIt->second.constant()); } // Is the second operand constant? if (auto *Operand = dyn_cast(SecondOp)) { auto OSRIt = OSRs.find(Operand); if (OSRIt != OSRs.end() && OSRIt->second.isConstant()) Constants[1] = CI::get(Operand->getType(), OSRIt->second.constant()); } // No constant operands if (Constants[0] == nullptr && Constants[1] == nullptr) return { nullptr, nullptr }; // Both operands are constant, constant fold them if (Constants[0] != nullptr && Constants[1] != nullptr) { Instruction *Clone = I->clone(); Clone->setOperand(0, Constants[0]); Clone->setOperand(1, Constants[1]); Constant *Result = ConstantFoldInstruction(Clone, DL); if (isa(Result)) return { nullptr, nullptr }; else return { Result, nullptr }; } // Only one operand is constant if (Constants[0] != nullptr) return { Constants[0], SecondOp }; else return { Constants[1], FirstOp }; } // TODO: check also undefined behaviors due to shifts static bool isSupportedOperation(unsigned Opcode, Constant *ConstantOp, unsigned FreeOpIndex, const DataLayout &DL) { // Division by zero if ((Opcode == Instruction::SDiv || Opcode == Instruction::UDiv) && getZExtValue(ConstantOp, DL) == 0) return false; // Shift too much auto *OperandTy = dyn_cast(ConstantOp->getType()); if ((Opcode == Instruction::Shl || Opcode == Instruction::LShr || Opcode == Instruction::AShr) && getZExtValue(ConstantOp, DL) >= OperandTy->getBitWidth()) return false; // 128-bit operand auto *ConstantOpTy = dyn_cast(ConstantOp->getType()); if (ConstantOpTy != nullptr && ConstantOpTy->getBitWidth() > 64) return false; if (!Instruction::isCommutative(Opcode) && FreeOpIndex != 0 && Opcode != Instruction::Sub) return false; return true; } bool OSRAPass::updateLoadReacher(LoadInst *Load, Instruction *I, OSR NewOSR) { // Check if the instruction propagating the OSR is already a // component of this load or not auto ReachersIt = LoadReachers.find(Load); if (ReachersIt != LoadReachers.end()) { auto &Reachers = ReachersIt->second; auto Pred = [I] (const std::pair &P) { return P.first == I; }; auto ReacherIt = std::find_if(Reachers.begin(), Reachers.end(), Pred); if (ReacherIt != Reachers.end()) { // We've already propagated I to Load in the past, check if we have new // information if (ReacherIt->second == NewOSR || ReacherIt->second.boundedValue()->value() == Load) { return false; } else { const Value *ReacherValue = ReacherIt->second.boundedValue()->value(); assert(!(Reachers.size() > 1 && ReacherValue == Load && ReacherValue != NewOSR.boundedValue()->value())); *ReacherIt = make_pair(I, NewOSR); return true; } } } LoadReachers[Load].push_back({ I, NewOSR }); return true; } bool OSRAPass::isDead(Instruction *I) const { while (I != nullptr) { if (!I->hasOneUse()) return false; auto *U = dyn_cast(*I->user_begin()); if (U == nullptr) return false; switch (U->getOpcode()) { case Instruction::ZExt: case Instruction::SExt: case Instruction::IntToPtr: case Instruction::PtrToInt: I = dyn_cast(U); break; case Instruction::Store: { auto *Store = cast(U); if (Store->getValueOperand() != I) return false; bool Used = false; auto *State = dyn_cast(Store->getPointerOperand()); if (State == nullptr) return false; auto Visitor = [State, &Used] (BasicBlockRange R) { for (Instruction &I : R) { if (auto *Load = dyn_cast(&I)) { if (Load->getPointerOperand() == State) { Used = true; return StopNow; } } else if (auto *Store = dyn_cast(&I)) { if (Store->getPointerOperand() == State) { return NoSuccessors; } } } return Continue; }; visitSuccessors(Store, make_blacklist(BlockBlackList), Visitor); return !Used; } default: return false; } } return false; } void OSRAPass::mergeLoadReacher(LoadInst *Load) { auto &Reachers = LoadReachers[Load]; assert(Reachers.size() > 0); OSRs.erase(Load); // TODO: implement a real merge strategy, considering input boundaries OSR Result = Reachers[0].second; for (auto P : skip(1, Reachers)) { OSR ReachingOSR = P.second; if (ReachingOSR != Result) { OSR FreeOSR = createOSR(Load, Load->getParent()); if (Reachers.size() == RDP->getReachingDefinitionsCount(Load)) BVs.forceBV(Load, pathSensitiveMerge(Load)); OSRs.insert({ Load, FreeOSR }); return; } } OSRs.insert({ Load, Result }); return; } /// \brief State of a definition reaching a load while being processed by /// OSRAPass::pathSensitiveMerge class Reacher { public: Reacher(LoadInst *Reached, Instruction *Reacher, OSR &ReachingOSR) : Summary(BoundedValue(ReachingOSR.boundedValue()->value())), LastMergeHeight(0), ReachingOSR(ReachingOSR), LTR(std::set { Reacher->getParent() }), LastActiveHeight(Active) { } /// \brief Notify that the stack has grown void newHeight(unsigned NewHeight) { LastActiveHeight = std::min(LastActiveHeight, NewHeight); LastMergeHeight = std::min(LastMergeHeight, NewHeight); } /// \brief Check if the reacher is active at the current stack height bool isActive(unsigned CurrentHeight) const { return CurrentHeight <= LastActiveHeight; } /// \brief Check if \p BB leads to the definition represented by this object bool isLTR(BasicBlock *BB) const { return LTR.count(BB) != 0; } /// \brief Register \p BB as a basic block leading to this definition bool registerLTR(BasicBlock *BB) { return LTR.insert(BB).second; } /// \brief Mark this Reacher as active at the current height void setActive() { LastActiveHeight = Active; } /// \brief Mark this Reacher as inactive at height \p Height void setInactive(unsigned Height) { LastActiveHeight = Height; } /// \brief Set the last height of the stack when a merge was performed void setLastMerge(unsigned Height) { LastMergeHeight = Height; } /// \brief Retrieve the last height of the stack when a merge was performed unsigned lastMerge() const { return LastMergeHeight; } /// Compute a BV relative to \p V by applying the OSR associated to this /// definition and the constraints accumulated in Summary BoundedValue computeBV(Value *V, const DataLayout &DL, Type *Int64) const { auto Result = ReachingOSR.apply(Summary, V, DL); if (!Result.hasSignedness()) Result.setBottom(); if (!Result.isUninitialized() && !Result.isBottom()) { using Cmp = CmpInst; auto Predicate = Result.isSigned() ? Cmp::ICMP_SLE : Cmp::ICMP_ULE; Constant *Compare = CE::getCompare(Predicate, Result.lower(Int64), Result.upper(Int64)); if (getZExtValue(Compare, DL) == 0) Result.setBottom(); } return Result; } /// \brief Rreturn the OSR associated to this definition const OSR &osr() const { return ReachingOSR; } public: BoundedValue Summary; ///< BV representing the known constraints on the /// reaching definition's value private: unsigned LastMergeHeight; OSR &ReachingOSR; const unsigned Active = std::numeric_limits::max(); std::set LTR; unsigned LastActiveHeight; }; BoundedValue OSRAPass::pathSensitiveMerge(LoadInst *Reached) { // Initialization steps const unsigned MaxDepth = 10; Module *M = Reached->getParent()->getParent()->getParent(); const DataLayout &DL = M->getDataLayout(); Type *Int64 = IntegerType::get(M->getContext(), 64); MemoryAccess ReachedMA(Reached, DL); // Debug support raw_os_ostream OsOstream(dbg); formatted_raw_ostream FormattedStream(OsOstream); FormattedStream.SetUnbuffered(); DBG("psm", dbg << "Performing PSM for " << getName(Reached) << "\n";); std::vector Reachers; Reachers.reserve(LoadReachers[Reached].size()); unsigned ReacherIndex = 0; for (auto &P : LoadReachers[Reached]) { ReacherIndex++; // TODO: isConstant? if (P.second.factor() == 0) return BoundedValue(Reached); Reachers.emplace_back(Reached, P.first, P.second); DBG("psm", dbg << " Reacher " << std::dec << ReacherIndex << " is " << getName(P.first) << " (relative to " << getName(P.second.boundedValue()->value()) << ")\n";); } assert(Reachers.size() > 0); struct State { BasicBlock *BB; pred_iterator PredecessorIt; }; std::set InStack; std::vector Stack; State Initial = { Reached->getParent(), getValidPred(Reached->getParent()), }; if (Initial.PredecessorIt == pred_end(Initial.BB)) return BoundedValue(Reached); Stack.push_back(Initial); InStack.insert(Reached->getParent()); while (!Stack.empty()) { State &S = Stack.back(); unsigned Height = Stack.size(); BasicBlock *Pred = *S.PredecessorIt; std::string Indent(Height * 2, ' '); DBG("psm", dbg << Indent << "Exploring " << getName(Pred) << "\n";); // Check if any store in Pred can alias ReachedMA bool MayAlias = MemoryAccess::mayAlias(Pred, ReachedMA, DL); // Hold whether we should proceed to the predecessors or not // Initialize to false, the code handling the various reacher will enable // this flag if at least one of the reachers is active bool Proceed = false; // Reacher-specific handling ReacherIndex = 0; for (Reacher &R : Reachers) { ReacherIndex++; // Check if this reacher has been deactivated if (!R.isActive(Height)) continue; // Is this a BB leading to the reacher? if (R.isLTR(Pred)) { DBG("psm", dbg << Indent << " Merging reacher " << ReacherIndex << " (relative to " << getName(R.Summary.value()) << ")\n";); // Insert everything is on the stack, but stop if we meet one that's // already there for (State &NewLTRState : Stack) if (!R.registerLTR(NewLTRState.BB)) break; // Perform merge from the top to last merge height BoundedValue Result = R.Summary; auto Range = make_range(Stack.begin() + R.lastMerge(), Stack.end()); for (State &ToMerge : Range) { // Obtain the constraint from the appropriate edge BoundedValue *EdgeBV = BVs.getEdge(ToMerge.BB, *ToMerge.PredecessorIt, Result.value()); // if (EdgeBV == nullptr) // return BoundedValue(Reached); if (EdgeBV != nullptr) { // And-merge Result.merge(*EdgeBV, DL, Int64); DBG("psm", { dbg << Indent << " Got "; EdgeBV->describe(FormattedStream); dbg << " from the " << getName(*ToMerge.PredecessorIt) << " -> " << getName(ToMerge.BB) << " edge: "; Result.describe(FormattedStream); dbg << "\n"; }); if (Result.isBottom()) break; } else { DBG("psm", { dbg << Indent << " Got no info" << " from the " << getName(*ToMerge.PredecessorIt) << " -> " << getName(ToMerge.BB) << " edge\n"; }); } } // If result is bottom, we went through a contradictory branch, ignore // it and deactivate if (!Result.isBottom()) { R.Summary = Result; // Register the current height as the last merge R.setLastMerge(Height); } else { DBG("psm", dbg << Indent << " We got an incoherent situation, ignore it\n";); } // Deactivate R.setInactive(Height); } else if (MayAlias) { DBG("psm", dbg << Indent << " Deactivating reacher " << ReacherIndex << "\n";); // We don't know if it's an LTR, check if it may alias, and if so, // deactivate this reacher R.setInactive(Height); } else { // Activate R.setActive(); // At least one of the reacher is active, we have to proceed to the // predecessor Proceed = true; } } // Check it's not already in stack Proceed &= InStack.count(Pred) == 0; DBG("psm", if (!(InStack.count(Pred) == 0)) { dbg << Indent << " It's already on the stack\n"; }); // Check we're not exceeding the maximum allowed depth Proceed &= Height < MaxDepth; DBG("psm", if (!(Height < MaxDepth)) { dbg << Indent << " We exceeded the maximum depth\n"; }); // Check we have at least a non-dispatcher predecessor pred_iterator NewPredIt = getValidPred(Pred); Proceed &= NewPredIt != pred_end(Pred); DBG("psm", if (!(NewPredIt != pred_end(Pred))) { dbg << Indent << " No predecessors\n"; }); if (Proceed) { // We have to go deeper State NewState = { Pred, NewPredIt }; Stack.push_back(NewState); InStack.insert(Pred); } else { // Pop until the stack is empty or we still have unexplored predecessors unsigned OldHeight = Stack.size(); while (Stack.size() != 0) { State &Top = Stack.back(); auto End = pred_end(Top.BB); if (nextValidPred(++Top.PredecessorIt, End) != End) break; InStack.erase(Top.BB); Stack.pop_back(); } // If we popped something make sure we update all the heights unsigned NewHeight = Stack.size(); if (NewHeight < OldHeight) for (Reacher &R : Reachers) R.newHeight(NewHeight); } } // Or-merge all the collected BVs // TODO: adding the OSR offset is safe, but the multiplier? BoundedValue FinalBV = Reachers[0].computeBV(Reached, DL, Int64); DBG("psm", { unsigned I = 0; for (const Reacher &R : Reachers) { BoundedValue ReacherBV = R.computeBV(Reached, DL, Int64); dbg << "Reacher " << ++I << ": "; ReacherBV.describe(FormattedStream); dbg << " (from "; R.osr().describe(FormattedStream); dbg << ")\n"; } }); for (Reacher &R : skip(1, Reachers)) { BoundedValue ReacherBV = R.computeBV(Reached, DL, Int64); DBG("psm", { dbg << ""; FinalBV.describe(FormattedStream); dbg << " += "; ReacherBV.describe(FormattedStream); dbg << " (from "; R.osr().describe(FormattedStream); dbg << ")\n"; }); if (FinalBV.isBottom()) return BoundedValue(Reached); FinalBV.merge(ReacherBV, DL, Int64); } if (FinalBV.isUninitialized() || FinalBV.isTop() || FinalBV.isBottom()) return BoundedValue(Reached); DBG("psm", { dbg << "FinalBV: "; FinalBV.describe(FormattedStream); dbg << "\n"; }); assert(!FinalBV.isUninitialized()); return FinalBV; } // Terminology: // * OSR: Offseted Shifted Range, our main data flow value which represents the // result of an instruction as another value, which lies withing a // certain range of values, multiplied by a factor and with an // offset, e.g. 100 + 4 * x, with 0 < x < 4. // * free value: a value we can't represent as an OSR of another value // * bounded variable (or BV): a free value and the range within which it lies. bool OSRAPass::runOnFunction(Function &F) { DBG("passes", { dbg << "Starting OSRAPass\n"; }); const DataLayout DL = F.getParent()->getDataLayout(); RDP = &getAnalysis(); auto &SCP = getAnalysis(); // The Overtaken map keeps track of which load/store instructions have been // overtaken by another load/store, meaning that they are not "free" but can // be expressed in terms of another stored/loaded value std::map Overtaken; auto *Int64 = Type::getInt64Ty(F.getParent()->getContext()); using UpdateFunc = std::function; for (auto &BB : F) { if (!BB.empty()) { if (auto *Call = dyn_cast(&*BB.begin())) { Function *Callee = Call->getCalledFunction(); // TODO: comparing with "newpc" string is sad if (Callee != nullptr && Callee->getName() == "newpc") break; } } BlockBlackList.insert(&BB); } // Cleanup all the data freeContainer(OSRs); BVs = BVMap(&BlockBlackList, &DL, Int64); freeContainer(Constraints); // Initialize the WorkList with all the instructions in the function UniquedQueue WorkList; auto &BBList = F.getBasicBlockList(); for (auto &BB : make_range(BBList.begin(), BBList.end())) if (BlockBlackList.find(&BB) == BlockBlackList.end()) for (auto &I : make_range(BB.begin(), BB.end())) WorkList.insert(&I); // TODO: make these member functions auto InBlackList = [this] (BasicBlock *BB) { return BlockBlackList.find(BB) != BlockBlackList.end(); }; auto EnqueueUsers = [this, &WorkList] (Instruction *I) { for (User *U : I->users()) if (auto *UI = dyn_cast(U)) if (BlockBlackList.find(UI->getParent()) == BlockBlackList.end()) { WorkList.insert(UI); } }; auto PropagateConstraints = [this, &EnqueueUsers] (Instruction *I, Value *Operand, UpdateFunc Updater) { // We want to propagate contraints through zero-extensions if (auto *OperandInst = dyn_cast(Operand)) { auto OperandConstraintIt = Constraints.find(OperandInst); auto InstrConstraintIt = Constraints.find(I); // Does the operand have constraints? if (OperandConstraintIt != Constraints.end()) { auto New = Updater(OperandConstraintIt->second); // Does the instruction already had a constraint? if (InstrConstraintIt != Constraints.end()) { // Did the constraint changed? if (!differ(New, InstrConstraintIt->second)) return; Constraints.erase(InstrConstraintIt); } Constraints.insert({ I, New }); EnqueueUsers(I); } } }; while (!WorkList.empty()) { Instruction *I = WorkList.pop(); // TODO: create a member function for each group of opcodes unsigned Opcode = I->getOpcode(); switch (Opcode) { case Instruction::Add: case Instruction::Sub: case Instruction::Mul: case Instruction::Shl: case Instruction::SDiv: case Instruction::UDiv: case Instruction::LShr: case Instruction::AShr: { // Check if it's a free value auto OldOSRIt = OSRs.find(I); bool IsFree = OldOSRIt == OSRs.end(); bool Changed = false; Constant *ConstantOp = nullptr; Value *OtherOp = nullptr; std::tie(ConstantOp, OtherOp) = identifyOperands(I, DL); if (OtherOp == nullptr) { if (ConstantOp != nullptr) { // If OtherOp is nullptr but ConstantOp is not it means we were able // to fold the operation in a constant if (!IsFree) OSRs.erase(I); uint64_t Constant = getZExtValue(ConstantOp, DL); BoundedValue ConstantBV = BoundedValue::createConstant(I, Constant); auto &BV = BVs.forceBV(I, ConstantBV); OSR ConstantOSR(&BV); OSRs.emplace(make_pair(I, ConstantOSR)); EnqueueUsers(I); } // In any case, break break; } // Get or create an OSR for the non-constant operator, this // will be our starting point OSR NewOSR = createOSR(OtherOp, I->getParent()); if (!IsFree) { if (NewOSR.isRelativeTo(OldOSRIt->second.boundedValue()->value())) { break; } else { Changed = true; } } // Check we're not depending on ourselves, if we are leave us as a free // value if (NewOSR.isRelativeTo(I)) { assert(IsFree); break; } // TODO: this is probably a bad idea if (NewOSR.boundedValue()->isBottom()) { if (!IsFree) OSRs.erase(OldOSRIt); break; } // TODO: skip this if isDead(I) // Update signedness information if the given operation is // sign-aware if (Opcode == Instruction::SDiv || Opcode == Instruction::UDiv || Opcode == Instruction::LShr || Opcode == Instruction::AShr) { BVs.setSignedness(I->getParent(), NewOSR.boundedValue()->value(), Opcode == Instruction::SDiv || Opcode == Instruction::AShr); } // Check for undefined behaviors unsigned FreeOpIndex = OtherOp == I->getOperand(0) ? 0 : 1; if (!isSupportedOperation(Opcode, ConstantOp, FreeOpIndex, DL)) { NewOSR = OSR(&BVs.get(I->getParent(), I)); Changed = true; } else { // Combine the base OSR with the new operation Changed |= NewOSR.combine(Opcode, ConstantOp, FreeOpIndex, DL); } // Check if the OSR has changed if (IsFree || Changed) { // Update the OSR and enqueue all I's uses if (!IsFree) OSRs.erase(I); OSRs.emplace(make_pair(I, NewOSR)); EnqueueUsers(I); } break; } case Instruction::ICmp: { // TODO: this part is quite ugly, try to improve it auto SimplifiedComparison = SCP.getComparison(cast(I)); ICmpInst *Comparison = new ICmpInst(SimplifiedComparison.Predicate, SimplifiedComparison.LHS, SimplifiedComparison.RHS); std::unique_ptr SimplifiedCmpInst(Comparison); Predicate P = Comparison->getPredicate(); Value *LHS = Comparison->getOperand(0); Value *RHS = Comparison->getOperand(1); Constant *ConstOp = nullptr; Value *FreeOpValue = nullptr; Instruction *FreeOp = nullptr; std::tie(ConstOp, FreeOpValue) = identifyOperands(Comparison, DL); if (FreeOpValue != nullptr) { FreeOp = dyn_cast(FreeOpValue); if (FreeOp == nullptr) break; } if (isDead(I)) break; // Comparison for equality and inequality are handled to propagate // constraints in case of test of the result of a comparison (e.g., (x < // 3) == 0). if (ConstOp != nullptr && FreeOp != nullptr && Constraints.find(FreeOp) != Constraints.end() && (P == CmpInst::ICMP_EQ || P == CmpInst::ICMP_NE)) { // If we're comparing with 0 for equality or inequality and the // non-constant operand has constraints, propagate them flipping them // (if necessary). if (getZExtValue(ConstOp, DL) == 0) { if (P == CmpInst::ICMP_EQ) { PropagateConstraints(I, FreeOp, [] (BVVector &Constraints) { BVVector Result = Constraints; // TODO: This is wrong! !(a & b) == !a || !b, // not !a && !b for (auto &Constraint : Result) Constraint.flip(); return Result; }); } else { PropagateConstraints(I, FreeOp, [] (BVVector &Constraints) { return Constraints; }); } // Do not proceed break; } } // Compute a new constraint // Check the comparison operator is a supported one if (P != CmpInst::ICMP_UGT && P != CmpInst::ICMP_UGE && P != CmpInst::ICMP_SGT && P != CmpInst::ICMP_SGE && P != CmpInst::ICMP_ULT && P != CmpInst::ICMP_ULE && P != CmpInst::ICMP_SLT && P != CmpInst::ICMP_SLE && P != CmpInst::ICMP_EQ && P != CmpInst::ICMP_NE) break; auto OldBVsIt = Constraints.find(I); bool HasConstraints = OldBVsIt != Constraints.end(); BVVector NewConstraints; if (FreeOp == nullptr) { if (ConstOp == nullptr) { // Both operands are free, give up // TODO: are we sure this is what we want? if (HasConstraints) Constraints.erase(OldBVsIt); HasConstraints = false; break; } else { // FreeOpValue is nullptr but ConstOp is not: we were able to fold // the operation into a constant if (getZExtValue(ConstOp, DL) != 0) { // The comparison holds, we're saying nothing useful (e.g. 2 < 3), // remove any constraint if (HasConstraints) Constraints.erase(OldBVsIt); HasConstraints = false; } else { // The comparison does not hold, move to bottom all the involved // BVs auto *FirstOp = dyn_cast(LHS); if (FirstOp != nullptr) { auto FirstOSRIt = OSRs.find(FirstOp); if (FirstOSRIt != OSRs.end()) { auto FirstOSR = FirstOSRIt->second; NewConstraints.push_back(*FirstOSR.boundedValue()); } } if (auto *SecondOp = dyn_cast(RHS)) { auto SecondOSRIt = OSRs.find(SecondOp); if (SecondOSRIt != OSRs.end()) { auto SecondOSR = SecondOSRIt->second; NewConstraints.push_back(*SecondOSR.boundedValue()); } } for (auto &Constraint : NewConstraints) Constraint.setBottom(); } } } else { // We have a constant operand and a free one BasicBlock *BB = I->getParent(); auto Handle = [&] (OSR &BaseOp) { if (BaseOp.boundedValue()->isBottom() || BaseOp.isRelativeTo(I) || BaseOp.factor() == 0) return; // Notify the BV about the sign we're going to use, unless it's a // comparison of (in)equality bool IsSigned; if (P != CmpInst::ICMP_EQ && P != CmpInst::ICMP_NE) { IsSigned = Comparison->isSigned(); BVs.setSignedness(BB, BaseOp.boundedValue()->value(), IsSigned); } else { // TODO: we don't know what sign to use here, so we ignore it, // should we switch to AnySignedness? if (!BaseOp.boundedValue()->hasSignedness()) return; IsSigned = BaseOp.boundedValue()->isSigned(); } // Setting the sign might lead to bottom if (BaseOp.boundedValue()->isBottom()) return; // Create a copy of the current value of the BV BoundedValue NewBV = *(BaseOp.boundedValue()); auto Merge = [&] (Predicate P, Constant *ConstOp) { // Solve the equation to obtain the new boundary value // x < 1.5 == x < 2 (Ceiling) // x <= 1.5 == x <= 1 (Floor) // x > 1.5 == x > 1 (Floor) // x >= 1.5 == x >= 2 (Ceiling) bool RoundUp = (P == CmpInst::ICMP_UGE || P == CmpInst::ICMP_SGE || P == CmpInst::ICMP_ULT || P == CmpInst::ICMP_SLT); Constant *NewBoundC = BaseOp.solveEquation(ConstOp, RoundUp, DL); if (isa(NewBoundC)) return false; uint64_t NewBound = getExtValue(NewBoundC, IsSigned, DL); // TODO: this is an hack if (NewBound == 0 && (P == CmpInst::ICMP_ULT || P == CmpInst::ICMP_UGE)) return true; using BV = BoundedValue; switch (P) { case CmpInst::ICMP_UGT: case CmpInst::ICMP_UGE: case CmpInst::ICMP_SGT: case CmpInst::ICMP_SGE: if (CmpInst::isFalseWhenEqual(P)) NewBound++; NewBV.merge(BV::createGE(NewBV.value(), NewBound, IsSigned), DL, Int64); break; case CmpInst::ICMP_ULT: case CmpInst::ICMP_ULE: case CmpInst::ICMP_SLT: case CmpInst::ICMP_SLE: if (CmpInst::isFalseWhenEqual(P)) NewBound--; NewBV.merge(BV::createLE(NewBV.value(), NewBound, IsSigned), DL, Int64); break; case CmpInst::ICMP_EQ: NewBV.merge(BV::createEQ(NewBV.value(), NewBound, NewBV.isSigned()), DL, Int64); break; case CmpInst::ICMP_NE: NewBV.merge(BV::createNE(NewBV.value(), NewBound, NewBV.isSigned()), DL, Int64); break; default: assert(false); break; } return true; }; bool Result = Merge(P, ConstOp); if (!Result) return; // Unsigned inequations implictly say that both operands are greater // than or equal to zero. This means that if we have `x - 5 < 10`, // we don't just know that `x < 15` but also that `x - 5 >= 0`, // i.e., `x >= 5`. if (P == CmpInst::ICMP_ULT || P == CmpInst::ICMP_ULE) { auto *Zero = ConstantInt::get(ConstOp->getType(), 0); Result = Merge(CmpInst::ICMP_UGE, Zero); } if (!Result) return; NewConstraints.push_back(NewBV); }; OSR TheOSR = createOSR(FreeOp, BB); NewConstraints.clear(); // Handle the base case Handle(TheOSR); // Handle all the reaching definitions, if it's referred to a load const Value *BaseValue = nullptr; if (TheOSR.boundedValue() != nullptr) BaseValue = TheOSR.boundedValue()->value(); if (BaseValue != nullptr) { if (auto *Load = dyn_cast(BaseValue)) { const OSR *LoadOSR = getOSR(Load); auto &Reachers = LoadReachers[Load]; // Register this instruction to be visited again when Load changes Subscriptions[Load].insert(I); if (Reachers.size() > 1 && LoadOSR != nullptr && LoadOSR->boundedValue()->value() == Load) { for (auto &P : Reachers) { if (!P.second.isConstant() && P.second.boundedValue() != nullptr && P.second.boundedValue()->value() != nullptr && P.second.boundedValue()->value() != Load) { OSR TheOSR = switchBlock(P.second, BB); Handle(TheOSR); } } } } } } bool Changed = true; // Check against the old constraints associated with this comparison if (HasConstraints) { BVVector &OldBVsVector = OldBVsIt->second; if (NewConstraints.size() == OldBVsVector.size()) { bool Different = false; auto OldIt = OldBVsVector.begin(); auto NewIt = NewConstraints.begin(); // Loop over all the elements until a different one is found or we // reached the end while (!Different && OldIt != OldBVsVector.end()) { Different |= *OldIt != *NewIt; OldIt++; NewIt++; } Changed = Different; } } // If something changed replace the BV vector and re-enqueue all the // users if (Changed) { Constraints[I] = NewConstraints; EnqueueUsers(I); } break; } case Instruction::ZExt: case Instruction::Trunc: { // Associate OSR only if the operand has an OSR and always enqueue the // users auto *Operand = I->getOperand(0); auto OpOSRIt = OSRs.find(Operand); if (OpOSRIt != OSRs.end()) { OSR NewOSR = createOSR(Operand, I->getParent()); if (NewOSR.isRelativeTo(I)) break; OSRs.emplace(make_pair(I, NewOSR)); EnqueueUsers(I); } PropagateConstraints(I, Operand, [] (BVVector &BV) { return BV; }); break; } case Instruction::And: case Instruction::Or: { Instruction *FirstOperand = dyn_cast(I->getOperand(0)); Instruction *SecondOperand = dyn_cast(I->getOperand(1)); if (FirstOperand == nullptr || SecondOperand == nullptr) break; auto FirstConstraintIt = Constraints.find(FirstOperand); auto SecondConstraintIt = Constraints.find(SecondOperand); // We can merge the BVs only if both operands have one if (FirstConstraintIt == Constraints.end() || SecondConstraintIt == Constraints.end()) break; // Initialize the new boundaries with the first operand auto NewConstraints = FirstConstraintIt->second; auto &OtherConstraints = SecondConstraintIt->second; if (Opcode == Instruction::And) mergeBVVectors(NewConstraints, OtherConstraints, DL, Int64); else mergeBVVectors(NewConstraints, OtherConstraints, DL, Int64); bool Changed = true; // If this instruction already had constraints, compare them with the // new ones auto OldConstraintsIt = Constraints.find(I); if (OldConstraintsIt != Constraints.end()) Changed = differ(OldConstraintsIt->second, NewConstraints); // If something changed, register the new constraints and re-enqueue all // the users of the instruction if (Changed) { Constraints[I] = NewConstraints; EnqueueUsers(I); } break; } case Instruction::Br: { auto *Branch = cast(I); // Unconditional branches bring no useful information if (Branch->isUnconditional()) break; auto *Condition = dyn_cast(Branch->getCondition()); if (Condition == nullptr) break; // Were we able to handle the condition? auto BranchConstraintsIt = Constraints.find(Condition); if (BranchConstraintsIt == Constraints.end()) break; // Take a reference to the constraints, and produce a complementary // version auto &BranchConstraints = BranchConstraintsIt->second; BVVector FlippedBranchConstraints = BranchConstraintsIt->second; // TODO: This is wrong! !(a & b) == !a || !b, not !a && !b for (auto &BranchConstraint : FlippedBranchConstraints) BranchConstraint.flip(); // Create and initialize the worklist with the positive constraints for // the true branch, and the negated constraints for the false branch struct WLEntry { WLEntry(BasicBlock *Target, BasicBlock *Origin, BVVector Constraints) : Target(Target), Origin(Origin), Constraints(Constraints) { } BasicBlock *Target; BasicBlock *Origin; BVVector Constraints; }; std::vector ConstraintsWL; if (!InBlackList(Branch->getSuccessor(0))) { ConstraintsWL.push_back(WLEntry(Branch->getSuccessor(0), Branch->getParent(), BranchConstraints)); } if (!InBlackList(Branch->getSuccessor(1))) { ConstraintsWL.push_back(WLEntry(Branch->getSuccessor(1), Branch->getParent(), FlippedBranchConstraints)); } // TODO: can we do this in a DFA way? // Process the worklist while (!ConstraintsWL.empty()) { auto Entry = ConstraintsWL.back(); ConstraintsWL.pop_back(); assert(BlockBlackList.find(Entry.Target) == BlockBlackList.end()); // Merge each changed bound with the existing one for (auto ConstraintIt = Entry.Constraints.begin(); ConstraintIt != Entry.Constraints.end();) { auto Result = BVs.update(Entry.Target, Entry.Origin, *ConstraintIt); bool Changed = Result.first; BoundedValue &NewBV = Result.second; if (Changed) { // From now we propagate the updated constraint *ConstraintIt = NewBV; ConstraintIt++; } else { ConstraintIt = Entry.Constraints.erase(ConstraintIt); } } // Compute the set of affected values llvm::SmallSet Affected; for (BoundedValue &Constraint : Entry.Constraints) Affected.insert(Constraint.value()); // Look for instructions using constraints that have changed for (Instruction &ConstraintUser : *Entry.Target) { // Avoid looking up instructions that simply cannot be there switch (ConstraintUser.getOpcode()) { case Instruction::ICmp: case Instruction::And: case Instruction::Or: { // Ignore instructions without an associated constraint auto ConstraintIt = Constraints.find(&ConstraintUser); if (ConstraintIt == Constraints.end()) continue; // If it's using one of the changed variables, insert it in the // worklist BVVector &InstructionConstraints = ConstraintIt->second; for (BoundedValue &Constraint : InstructionConstraints) { if (Affected.count(Constraint.value()) != 0) { WorkList.insert(&ConstraintUser); break; } } break; } case Instruction::Load: { // Check if any of the reaching definitions of this load is // affected by the constraints being propagated LoadInst *Load = cast(&ConstraintUser); auto ReachersIt = LoadReachers.find(Load); if (ReachersIt == LoadReachers.end()) break; auto &Reachers = ReachersIt->second; for (auto &P : Reachers) { const Value *ReacherValue = nullptr; if (P.second.boundedValue() != nullptr) ReacherValue = P.second.boundedValue()->value(); if (Affected.count(ReacherValue) != 0) { // We're affected, update mergeLoadReacher(Load); WorkList.insert(Load); EnqueueUsers(Load); Affected.insert(Load); break; } } break; } default: break; } } // Propagate the new constraints to the successors (except for the // dispatcher) if (Entry.Constraints.size() != 0) for (BasicBlock *Successor : successors(Entry.Target)) if (BlockBlackList.find(Successor) == BlockBlackList.end()) ConstraintsWL.push_back(WLEntry(Successor, Entry.Target, Entry.Constraints)); } break; } case Instruction::Store: case Instruction::Load: { // Create the OSR to propagate MemoryAccess MA; // TODO: rename SelfOSR (it's not always self) OSR SelfOSR; BVVector TheConstraints; bool HasConstraints = false; if (auto *TheLoad = dyn_cast(I)) { // It's a load MA = MemoryAccess(TheLoad, DL); auto OSRIt = OSRs.find(I); if (OSRIt != OSRs.end()) SelfOSR = OSRIt->second; else SelfOSR = OSR(&BVs.get(I->getParent(), I)); } else if (auto *TheStore = dyn_cast(I)) { // It's a store MA = MemoryAccess(TheStore, DL); Value *ValueOp = TheStore->getValueOperand(); if (auto *ConstantOp = dyn_cast(ValueOp)) { // We're storing a constant, create a constant OSR uint64_t Constant = getZExtValue(ConstantOp, DL); BoundedValue ConstantBV = BoundedValue::createConstant(ConstantOp, Constant); auto &BV = BVs.forceBV(I->getParent(), ConstantOp, ConstantBV); SelfOSR = OSR(&BV); } else if (auto *ToStore = dyn_cast(ValueOp)) { // Compute the OSR to propagate: either the one of the value to // store, or an OSR relative to the value being stored auto OSRIt = OSRs.find(ToStore); if (OSRIt != OSRs.end()) SelfOSR = OSRIt->second; else SelfOSR = OSR(&BVs.get(I->getParent(), ToStore)); // Check if the value we're storing has constraints auto ConstraintIt = Constraints.find(ToStore); HasConstraints = ConstraintIt != Constraints.end(); if (HasConstraints) TheConstraints = ConstraintIt->second; } } auto ReachedLoads = RDP->getReachedLoads(I); for (LoadInst *ReachedLoad : ReachedLoads) { assert(ReachedLoad != I); // OSR propagation first // Take the reference OSR (SelfOSR) and "contextualize" it in // the reached load's basic block OSR NewOSR = switchBlock(SelfOSR, ReachedLoad->getParent()); bool Changed = updateLoadReacher(ReachedLoad, I, NewOSR); if (Changed) mergeLoadReacher(ReachedLoad); // Constraints propagation if (HasConstraints) { // Does the reached load carries any constraints already? auto ReachedLoadConstraintIt = Constraints.find(ReachedLoad); if (ReachedLoadConstraintIt != Constraints.end()) { // Merge the constraints (using the `or` logic) directly in-place // in the reached load's BVVector using BV = BoundedValue; Changed |= mergeBVVectors(ReachedLoadConstraintIt->second, TheConstraints, DL, Int64); } else { // The reached load has no constraints, simply propagate the input // ones Constraints.insert({ ReachedLoad, TheConstraints }); Changed = true; } } // If OSR or constraints have changed, mark the reached load and its // uses to be visited again if (Changed) { WorkList.insert(ReachedLoad); EnqueueUsers(ReachedLoad); for (Instruction *Subscriber : Subscriptions[ReachedLoad]) WorkList.insert(Subscriber); } } break; } default: break; } } DBG("osr", { BVs.prepareDescribe(); raw_os_ostream OutputStream(dbg); F.getParent()->print(OutputStream, new OSRAnnotationWriter(*this)); }); // Free up memory not part of the analysis result freeContainer(Constraints); freeContainer(LoadReachers); freeContainer(BlockBlackList); freeContainer(Subscriptions); DBG("passes", { dbg << "Ending OSRAPass\n"; }); return false; } void OSRAPass::BVMap::describe(formatted_raw_ostream &O, const BasicBlock *BB) const { if (BBMap.find(BB) != BBMap.end()) for (MapValue &MV : BBMap[BB]) { O << " ; "; { auto &BVO = MV.Summary; O << "<"; BVO.describe(O); O << ">"; } if (MV.Components.size() > 0) O << " = "; for (auto &BVO : MV.Components) { O << "<"; O << getName(BVO.first); O << ", "; BVO.second.describe(O); O << "> || "; } O << "\n"; } O << "\n"; } std::pair OSRAPass::BVMap::update(BasicBlock *Target, BasicBlock *Origin, BoundedValue NewBV) { auto Index = make_pair(Target, NewBV.value()); auto MapIt = TheMap.find(Index); MapValue *BVOVector = nullptr; // Have we ever seen this value for this basic block? if (MapIt == TheMap.end()) { // No, just insert it MapValue NewBVOVector; NewBVOVector.Components.push_back({ make_pair(Origin, NewBV) }); BVOVector = &TheMap.insert({ Index, NewBVOVector }).first->second; return { true, summarize(Target, BVOVector) }; } else if (isForced(MapIt)) { return { false, MapIt->second.Summary }; } else { bool Changed = true; BVOVector = &MapIt->second; // Look for an entry with the given origin BoundedValue *Base = nullptr; for (BVWithOrigin &BVO : BVOVector->Components) if (BVO.first == Origin) Base = &BVO.second; // Did we ever see this Origin? if (Base == nullptr) BVOVector->Components.push_back({ Origin, NewBV }); else Changed = Base->merge(NewBV, *DL, Int64); // Re-merge all the entries auto &Result = summarize(Target, BVOVector); return { Changed, Result }; } // TODO: should Changed be false if isForced? } BoundedValue &OSRAPass::BVMap::summarize(BasicBlock *Target, MapValue *BVOVector) { if (BVOVector->Components.size() == 0) return BVOVector->Summary; // Initialize the summary BV with the first BV BVOVector->Summary = BVOVector->Components[0].second; unsigned PredecessorsCount = 0; for (auto *Predecessor : predecessors(Target)) if (BlockBlackList->find(Predecessor) == BlockBlackList->end() && !pred_empty(Predecessor)) PredecessorsCount++; // Do we have a constraint for each predecessor? if (BVOVector->Components.size() == PredecessorsCount) { // Yes, we can populate the summary by merging all the components for (auto &BVO : skip(1, BVOVector->Components)) BVOVector->Summary.merge(BVO.second, *DL, Int64); } else { // No, keep the summary at top BVOVector->Summary.setTop(); } return BVOVector->Summary; } bool OSR::compare(unsigned short P, Constant *C, const DataLayout &DL, Type *Int64) { Constant *BaseConstant = CI::get(Int64, Base); Constant *Compare = CE::getCompare(P, BaseConstant, C); return getConstValue(Compare, DL)->getLimitedValue() != 0; } void BoundedValue::setSignedness(bool IsSigned) { // TODO: assert? if (Bottom) return; // If we're already inconsistent just return if (Sign == InconsistentSignedness) return; Signedness NewSign = IsSigned ? Signed : Unsigned; if (Sign == UnknownSignedness) { assert(LowerBound == 0 && UpperBound == 0); Sign = NewSign; if (IsSigned) { LowerBound = numeric_limits::min(); UpperBound = numeric_limits::max(); } else { LowerBound = numeric_limits::min(); UpperBound = numeric_limits::max(); } } else if (Sign == AnySignedness) { Sign = NewSign; } else if (Sign != NewSign) { Sign = InconsistentSignedness; // TODO: handle top case if (LowerBound > numeric_limits::max() || UpperBound > numeric_limits::max()) { setBottom(); } } } template bool BoundedValue::merge(const BoundedValue &Other, const DataLayout &DL, Type *Int64) { if (Bottom) return false; if (Other.Bottom) { setBottom(); return true; } if (isTop() && Other.isTop()) { return false; } else if (MT == And && isTop()) { LowerBound = Other.LowerBound; UpperBound = Other.UpperBound; Sign = Other.Sign; Negated = Other.Negated; return true; } else if (MT == And && Other.isTop()) { return false; } else if (MT == Or && isTop()) { return false; } else if (MT == Or && Other.isTop()) { setTop(); return true; } if (Sign == AnySignedness && Other.Sign == AnySignedness) { setBottom(); return true; } if (Sign == AnySignedness || Other.Sign == AnySignedness) { if (Sign == AnySignedness) Sign = Other.Sign; } else { setSignedness(Other.isSigned()); } if (Bottom) return true; // We don't handle this case for now if (Sign == InconsistentSignedness) { setBottom(); return true; } // TODO: reimplement all of this using a simple and sane range merging // approach Predicate LE = isSigned() ? CmpInst::ICMP_SLE : CmpInst::ICMP_ULE; Predicate LT = isSigned() ? CmpInst::ICMP_SLT : CmpInst::ICMP_ULT; Predicate GE = isSigned() ? CmpInst::ICMP_SGE : CmpInst::ICMP_UGE; Predicate GT = isSigned() ? CmpInst::ICMP_SGT : CmpInst::ICMP_UGT; auto Compare = [&Int64, &DL] (uint64_t A, Predicate P, int64_t B) { Constant *Compare = CE::getCompare(P, CI::get(Int64, A), CI::get(Int64, B)); return getZExtValue(Compare, DL) != 0; }; const BoundedValue *LeftmostOp = this; const BoundedValue *RightmostOp = &Other; // Check that the LB of the lefmost is <= of the rightmost LB if (Compare(LeftmostOp->LowerBound, GT, RightmostOp->LowerBound)) std::swap(LeftmostOp, RightmostOp); // If they both start at the same point, LeftmostOp is the largest if (Compare(LeftmostOp->LowerBound, CmpInst::ICMP_EQ, RightmostOp->LowerBound) && Compare(RightmostOp->UpperBound, GT, LeftmostOp->UpperBound)) std::swap(LeftmostOp, RightmostOp); enum { Disjoint, Overlapping } Overlap; bool LowerLT = Compare(LeftmostOp->LowerBound, LT, RightmostOp->LowerBound); bool LowerLE = Compare(LeftmostOp->LowerBound, LE, RightmostOp->LowerBound); bool UpperGT = Compare(LeftmostOp->UpperBound, GT, RightmostOp->UpperBound); bool UpperGE = Compare(LeftmostOp->UpperBound, GE, RightmostOp->UpperBound); bool StrictlyIncluded = LowerLT && UpperGT; bool Included = LowerLE && UpperGE; if (Compare(LeftmostOp->UpperBound, LT, RightmostOp->LowerBound)) Overlap = Disjoint; else Overlap = Overlapping; const BoundedValue *NegatedOp = nullptr; const BoundedValue *NonNegatedOp = nullptr; enum { NoNegated, OneNegated, BothNegated } Operands; if (!Negated && !Other.Negated) { Operands = NoNegated; } else if (Negated && Other.Negated) { Operands = BothNegated; } else { Operands = OneNegated; if (Negated) { NegatedOp = this; NonNegatedOp = &Other; } else { NegatedOp = &Other; NonNegatedOp = this; } } uint64_t OldLowerBound = LowerBound; uint64_t OldUpperBound = UpperBound; bool OldNegated = Negated; // In the following table we report all the possible situations and the // relative result we produce: // // type overlap op1 op2 result // ====================================== // and disjoint + + bottom // and disjoint + - op1 // and disjoint - - bottom // and overlapping + + intersection // and overlapping + - op1-op2 // and overlapping - - !union // or disjoint + + bottom // or disjoint + - op2 // or disjoint - - top // or overlapping + + union // or overlapping + - !(op2-op1) // or overlapping - - !intersection // bool Changed = false; if (MT == And) { switch(Overlap) { case Disjoint: switch (Operands) { case NoNegated: setBottom(); Changed = true; break; case BothNegated: if (LeftmostOp->LowerBound == LeftmostOp->lowerExtreme() && RightmostOp->UpperBound == RightmostOp->upperExtreme()) { std::tie(LowerBound, UpperBound) = make_pair(LeftmostOp->UpperBound, RightmostOp->LowerBound); LowerBound++; UpperBound--; Negated = false; if (!Compare(LowerBound, LE, UpperBound)) { LowerBound = 0; UpperBound = 0; setBottom(); } break; } else if (LeftmostOp->UpperBound + 1 == RightmostOp->LowerBound) { setBound(CI::get(Int64, Other.LowerBound), DL); if (!Bottom) setBound(CI::get(Int64, Other.UpperBound), DL); Negated = true; } else { setBottom(); Changed = true; } break; case OneNegated: // Assign to NotNegated if (this != NonNegatedOp) { LowerBound = Other.LowerBound; UpperBound = Other.UpperBound; Negated = Other.Negated; } break; } break; case Overlapping: switch (Operands) { case NoNegated: // Intersection setBound(CI::get(Int64, Other.LowerBound), DL); if (!Bottom) setBound(CI::get(Int64, Other.UpperBound), DL); Negated = false; break; case OneNegated: // TODO: If one of the two is strictly included go to bottom if (StrictlyIncluded || (LowerBound == Other.LowerBound && UpperBound == Other.UpperBound) || (Included && LeftmostOp == NegatedOp)) { setBottom(); Changed = true; break; } // NonNegated - Negated // [5,10] - ![8,12] => NonNegated.Up = Negated.Down - 1 // [5,10] - ![1,12] == [0,10] - ([_,0] | [13,_]) // [5,10] - ![0,7] => NonNegated.Down = Negated.Up + 1 // [5,10] - ![4,7] // [5,10] - ![5,7] // [5,10] - ![6,12] // Check if NonNegated is after Negated uint64_t NewLowerBound, NewUpperBound; if (Compare(NonNegatedOp->LowerBound, GE, NegatedOp->LowerBound)) { NewLowerBound = NegatedOp->UpperBound + 1; NewUpperBound = NonNegatedOp->UpperBound; } else { NewLowerBound = NonNegatedOp->LowerBound; NewUpperBound = NegatedOp->LowerBound - 1; } LowerBound = NewLowerBound; UpperBound = NewUpperBound; Negated = false; break; case BothNegated: // Negated union setBound(CI::get(Int64, Other.LowerBound), DL); if (!Bottom) setBound(CI::get(Int64, Other.UpperBound), DL); Negated = true; break; } break; } } else if (MT == Or) { switch(Overlap) { case Disjoint: switch (Operands) { case NoNegated: if (LeftmostOp->UpperBound + 1 == RightmostOp->LowerBound) { setBound(CI::get(Int64, Other.LowerBound), DL); if (!Bottom) setBound(CI::get(Int64, Other.UpperBound), DL); } else { setBottom(); Changed = true; } break; case OneNegated: // Assign to Negated if (this != NegatedOp) { LowerBound = Other.LowerBound; UpperBound = Other.UpperBound; Negated = Other.Negated; } break; case BothNegated: setTop(); Changed = true; break; } break; case Overlapping: switch (Operands) { case NoNegated: setBound(CI::get(Int64, Other.LowerBound), DL); if (!Bottom) setBound(CI::get(Int64, Other.UpperBound), DL); break; case OneNegated: // TODO: comment this if (StrictlyIncluded) { if (LeftmostOp == NonNegatedOp) setTop(); else setBottom(); Changed = true; break; } if ((LowerBound == Other.LowerBound && UpperBound == Other.UpperBound) || (Included && LeftmostOp == NonNegatedOp)) { setTop(); Changed = true; break; } // ![5,25] || [6,30] // ![5,25] || [5,10] // Check if NonNegated is before Negated uint64_t NewLowerBound, NewUpperBound; if (Compare(NonNegatedOp->LowerBound, LE, NegatedOp->LowerBound)) { NewLowerBound = NonNegatedOp->UpperBound + 1; NewUpperBound = NegatedOp->UpperBound; } else { NewLowerBound = NegatedOp->LowerBound; NewUpperBound = NonNegatedOp->LowerBound - 1; } LowerBound = NewLowerBound; UpperBound = NewUpperBound; Negated = true; break; case BothNegated: setBound(CI::get(Int64, Other.LowerBound), DL); if (!Bottom) setBound(CI::get(Int64, Other.UpperBound), DL); Negated = true; break; } break; } } Changed |= (OldLowerBound != LowerBound || OldUpperBound != UpperBound || OldNegated != Negated); assert(Compare(LowerBound, LE, UpperBound)); return Changed; } // Note: this function is implemented with lower bound restriction in mind, with // additional changes to support bound enlargement (logical `or`) or work on the // upper bound just set the template arguments appopriately template bool BoundedValue::setBound(Constant *NewValue, const DataLayout &DL) { assert(Sign != UnknownSignedness && Sign != AnySignedness && !Bottom); uint64_t &Bound = B == Lower ? LowerBound : UpperBound; // Create a Constant for the current bound Constant *OldValue = CI::get(NewValue->getType(), Bound, isSigned()); // If the signedness is inconsistent, check that the new value lies in the // signed positive area, otherwise go to bottom // Note: OldValue should already be in this range, thanks to `setSignedness`. if (Sign == InconsistentSignedness && !isPositive(NewValue, DL)) { setBottom(); return true; } // Update the lower bound only if NewValue > OldValue Predicate CompOp = (isSigned() ? CmpInst::ICMP_SGT : CmpInst::ICMP_UGT); // If we want a logical or, flip the direction of the comparison if (Type == Or) CompOp = CmpInst::getSwappedPredicate(CompOp); if (B == Upper) CompOp = CmpInst::getSwappedPredicate(CompOp); // Perform the comparison and, in case, update the LowerBound auto *Compare = CE::getCompare(CompOp, NewValue, OldValue); if (getConstValue(Compare, DL)->getLimitedValue()) { if (isSigned()) Bound = getSExtValue(NewValue, DL); else Bound = getZExtValue(NewValue, DL); return true; } return false; }