diff --git a/resolve-cveassert/src/ArithmeticSanitizer.cpp b/resolve-cveassert/src/ArithmeticSanitizer.cpp index eb861fa7..befcf279 100644 --- a/resolve-cveassert/src/ArithmeticSanitizer.cpp +++ b/resolve-cveassert/src/ArithmeticSanitizer.cpp @@ -15,7 +15,7 @@ #include "CVEAssert.hpp" #include "IRUtils.hpp" -#include "Vulnerability.hpp" +#include "Remediation.hpp" #include #include @@ -90,8 +90,7 @@ static void widenIntOverflow(Function *F) { } } -void sanitizeDivideByZero(Function *F, - Vulnerability::RemediationStrategies strategy) { +void sanitizeDivideByZero(Function *F, RemediationStrategies strategy) { Module *M = F->getParent(); auto &Ctx = M->getContext(); auto usize_ty = Type::getInt64Ty(Ctx); @@ -100,15 +99,15 @@ void sanitizeDivideByZero(Function *F, std::vector worklist; switch (strategy) { - case Vulnerability::RemediationStrategies::CONTINUE: - case Vulnerability::RemediationStrategies::EXIT: - case Vulnerability::RemediationStrategies::RECOVER: + case RemediationStrategies::CONTINUE: + case RemediationStrategies::EXIT: + case RemediationStrategies::RECOVER: break; default: llvm::errs() << "[CVEAssert] Error: sanitizeDivideByZero does not support " << " remediation strategy defaulting to continue strategy!\n"; - strategy = Vulnerability::RemediationStrategies::CONTINUE; + strategy = RemediationStrategies::CONTINUE; break; } @@ -255,8 +254,7 @@ void sanitizeDivideByZero(Function *F, } } -void sanitizeIntOverflow(Function *F, - Vulnerability::RemediationStrategies strategy) { +void sanitizeIntOverflow(Function *F, RemediationStrategies strategy) { std::vector worklist; Module *M = F->getParent(); auto &Ctx = M->getContext(); @@ -265,20 +263,20 @@ void sanitizeIntOverflow(Function *F, IRBuilder<> builder(Ctx); switch (strategy) { - case Vulnerability::RemediationStrategies::WIDEN: + case RemediationStrategies::WIDEN: return widenIntOverflow(F); - case Vulnerability::RemediationStrategies::RECOVER: - case Vulnerability::RemediationStrategies::EXIT: - case Vulnerability::RemediationStrategies::WRAP: - case Vulnerability::RemediationStrategies::SAT: + case RemediationStrategies::RECOVER: + case RemediationStrategies::EXIT: + case RemediationStrategies::WRAP: + case RemediationStrategies::SAT: break; default: llvm::errs() << "[CVEAssert] Error: sanitizeIntOverflow does not support " "remediation strategy specified defaulting to wrap strategy!\n"; - strategy = Vulnerability::RemediationStrategies::WRAP; + strategy = RemediationStrategies::WRAP; break; } @@ -415,7 +413,7 @@ void sanitizeIntOverflow(Function *F, builder.CreateBr(joinResultBB); builder.SetInsertPoint(&*joinResultBB->begin()); - if (strategy == Vulnerability::RemediationStrategies::SAT) { + if (strategy == RemediationStrategies::SAT) { binary_inst->replaceAllUsesWith(satResult); } else { binary_inst->replaceAllUsesWith(safeResult); @@ -425,8 +423,7 @@ void sanitizeIntOverflow(Function *F, } } -void sanitizeBitShift(Function *F, - Vulnerability::RemediationStrategies strategy) { +void sanitizeBitShift(Function *F, RemediationStrategies strategy) { Module *M = F->getParent(); auto &Ctx = M->getContext(); auto usize_ty = Type::getInt64Ty(Ctx); @@ -435,14 +432,14 @@ void sanitizeBitShift(Function *F, std::vector worklist; switch (strategy) { - case Vulnerability::RemediationStrategies::EXIT: - case Vulnerability::RemediationStrategies::RECOVER: + case RemediationStrategies::EXIT: + case RemediationStrategies::RECOVER: break; default: llvm::errs() << "[CVEAssert] Error: sanitizeBitShift does not support " << " remediation strategy defaulting to EXIT strategy!\n"; - strategy = Vulnerability::RemediationStrategies::EXIT; + strategy = RemediationStrategies::EXIT; break; } diff --git a/resolve-cveassert/src/ArithmeticSanitizer.hpp b/resolve-cveassert/src/ArithmeticSanitizer.hpp index 0f44f3f6..135634c5 100644 --- a/resolve-cveassert/src/ArithmeticSanitizer.hpp +++ b/resolve-cveassert/src/ArithmeticSanitizer.hpp @@ -5,11 +5,8 @@ #pragma once -#include "Vulnerability.hpp" +#include "Remediation.hpp" #include "llvm/IR/Function.h" -void sanitizeDivideByZero(llvm::Function *F, - Vulnerability::RemediationStrategies strategy); -void sanitizeIntOverflow(llvm::Function *F, - Vulnerability::RemediationStrategies strategy); -void sanitizeBitShift(llvm::Function *F, - Vulnerability::RemediationStrategies strategy); +void sanitizeDivideByZero(llvm::Function *F, RemediationStrategies strategy); +void sanitizeIntOverflow(llvm::Function *F, RemediationStrategies strategy); +void sanitizeBitShift(llvm::Function *F, RemediationStrategies strategy); diff --git a/resolve-cveassert/src/BoundsCheck.cpp b/resolve-cveassert/src/BoundsCheck.cpp index c494dbd4..cb541087 100644 --- a/resolve-cveassert/src/BoundsCheck.cpp +++ b/resolve-cveassert/src/BoundsCheck.cpp @@ -18,7 +18,7 @@ #include "CVEAssert.hpp" #include "IRUtils.hpp" -#include "Vulnerability.hpp" +#include "Remediation.hpp" #include #include @@ -162,8 +162,7 @@ static Function *getOrCreateAccessOk(Module *M, BoundsClass cls) { } static Function *getOrCreateBoundsCheckLoadSanitizer( - Function *F, Type *ty, Vulnerability::RemediationStrategies strategy, - BoundsClass cls) { + Function *F, Type *ty, RemediationStrategies strategy, BoundsClass cls) { std::string handlerName = "__cve_bound_ld_" + getLLVMType(ty) + "_" + classTag(cls); Module *M = F->getParent(); @@ -223,8 +222,7 @@ static Function *getOrCreateBoundsCheckLoadSanitizer( } static Function *getOrCreateBoundsCheckStoreSanitizer( - Function *F, Type *ty, Vulnerability::RemediationStrategies strategy, - BoundsClass cls) { + Function *F, Type *ty, RemediationStrategies strategy, BoundsClass cls) { std::string handlerName = "__cve_bound_st_" + getLLVMType(ty) + "_" + classTag(cls); Module *M = F->getParent(); @@ -287,9 +285,10 @@ static Function *getOrCreateBoundsCheckStoreSanitizer( return resolveStoreFn; } -static Function *getOrCreateBoundsCheckMemcpySanitizer( - Function *F, Vulnerability::RemediationStrategies strategy, - BoundsClass srcCls, BoundsClass dstCls) { +static Function * +getOrCreateBoundsCheckMemcpySanitizer(Function *F, + RemediationStrategies strategy, + BoundsClass srcCls, BoundsClass dstCls) { std::string handlerName = std::string("__cve_memcpy_") + classTag(srcCls) + "_" + classTag(dstCls); Module *M = F->getParent(); @@ -357,9 +356,10 @@ static Function *getOrCreateBoundsCheckMemcpySanitizer( return resolveMemmoveFn; } -static Function *getOrCreateBoundsCheckMemmoveSanitizer( - Function *F, Vulnerability::RemediationStrategies strategy, - BoundsClass srcCls, BoundsClass dstCls) { +static Function * +getOrCreateBoundsCheckMemmoveSanitizer(Function *F, + RemediationStrategies strategy, + BoundsClass srcCls, BoundsClass dstCls) { std::string handlerName = std::string("__cve_memmove_") + classTag(srcCls) + "_" + classTag(dstCls); Module *M = F->getParent(); @@ -429,8 +429,7 @@ static Function *getOrCreateBoundsCheckMemmoveSanitizer( } static Function *getOrCreateBoundsCheckMemsetSanitizer( - Function *F, Vulnerability::RemediationStrategies strategy, - BoundsClass cls) { + Function *F, RemediationStrategies strategy, BoundsClass cls) { std::string handlerName = std::string("__cve_memset_") + classTag(cls); Module *M = F->getParent(); LLVMContext &Ctx = M->getContext(); @@ -633,8 +632,7 @@ void instrumentGep(Function *F) { } } -void instrumentMemcpy(Function *F, - Vulnerability::RemediationStrategies strategy) { +void instrumentMemcpy(Function *F, RemediationStrategies strategy) { LLVMContext &Ctx = F->getContext(); IRBuilder<> builder(Ctx); std::vector memcpyList; @@ -691,8 +689,7 @@ void instrumentMemcpy(Function *F, } } -void instrumentMemset(Function *F, - Vulnerability::RemediationStrategies strategy) { +void instrumentMemset(Function *F, RemediationStrategies strategy) { LLVMContext &Ctx = F->getContext(); IRBuilder<> builder(Ctx); std::vector memsetList; @@ -760,8 +757,7 @@ void instrumentMemset(Function *F, } } -void instrumentMemmove(Function *F, - Vulnerability::RemediationStrategies strategy) { +void instrumentMemmove(Function *F, RemediationStrategies strategy) { LLVMContext &Ctx = F->getContext(); IRBuilder<> builder(Ctx); std::vector memmoveList; @@ -818,8 +814,7 @@ void instrumentMemmove(Function *F, } } -void instrumentLoadStore(Function *F, - Vulnerability::RemediationStrategies strategy) { +void instrumentLoadStore(Function *F, RemediationStrategies strategy) { LLVMContext &Ctx = F->getContext(); IRBuilder<> builder(Ctx); @@ -827,16 +822,16 @@ void instrumentLoadStore(Function *F, std::vector storeList; switch (strategy) { - case Vulnerability::RemediationStrategies::CONTINUE: - case Vulnerability::RemediationStrategies::EXIT: - case Vulnerability::RemediationStrategies::RECOVER: + case RemediationStrategies::CONTINUE: + case RemediationStrategies::EXIT: + case RemediationStrategies::RECOVER: break; default: llvm::errs() << "[CVEAssert] Error: instrumentLoadStore does not support " "remediation strategy " << "defaulting to continue strategy!\n"; - strategy = Vulnerability::RemediationStrategies::CONTINUE; + strategy = RemediationStrategies::CONTINUE; break; } @@ -903,8 +898,7 @@ void instrumentLoadStore(Function *F, } } -void sanitizeMemInstBounds(Function *F, - Vulnerability::RemediationStrategies strategy) { +void sanitizeMemInstBounds(Function *F, RemediationStrategies strategy) { instrumentGep(F); instrumentMemcpy(F, strategy); instrumentMemmove(F, strategy); diff --git a/resolve-cveassert/src/BoundsCheck.hpp b/resolve-cveassert/src/BoundsCheck.hpp index 99e18a50..fa680c12 100644 --- a/resolve-cveassert/src/BoundsCheck.hpp +++ b/resolve-cveassert/src/BoundsCheck.hpp @@ -5,15 +5,10 @@ #pragma once -#include "Vulnerability.hpp" +#include "Remediation.hpp" #include "llvm/IR/Function.h" -void instrumentLoadStore(llvm::Function *F, - Vulnerability::RemediationStrategies strategy); -void instrumentMemcpy(llvm::Function *F, - Vulnerability::RemediationStrategies strategy); -void instrumentMemmove(llvm::Function *F, - Vulnerability::RemediationStrategies strategy); -void instrumentMemset(llvm::Function *F, - Vulnerability::RemediationStrategies strategy); -void sanitizeMemInstBounds(llvm::Function *F, - Vulnerability::RemediationStrategies strategy); +void instrumentLoadStore(llvm::Function *F, RemediationStrategies strategy); +void instrumentMemcpy(llvm::Function *F, RemediationStrategies strategy); +void instrumentMemmove(llvm::Function *F, RemediationStrategies strategy); +void instrumentMemset(llvm::Function *F, RemediationStrategies strategy); +void sanitizeMemInstBounds(llvm::Function *F, RemediationStrategies strategy); diff --git a/resolve-cveassert/src/CVEAssert.cpp b/resolve-cveassert/src/CVEAssert.cpp index 3fa215b6..34254ad2 100644 --- a/resolve-cveassert/src/CVEAssert.cpp +++ b/resolve-cveassert/src/CVEAssert.cpp @@ -35,6 +35,7 @@ #include "InstrumentAllocators.hpp" #include "NullPointerSanitizer.hpp" #include "OperationMasking.hpp" +#include "Remediation.hpp" #include "Vulnerability.hpp" using namespace llvm; @@ -114,8 +115,7 @@ struct LabelCVEPass : public PassInfoMixin { vulnerabilities = Vulnerability::parseVulnerabilityFile(); } - void applyAutomaticSanitizers(Function &F, - Vulnerability::RemediationStrategies strategy) { + void applyAutomaticSanitizers(Function &F, RemediationStrategies strategy) { /// applies all automatic sanitizers (operation masking excluded) sanitizeFreeOfNonHeap(&F, strategy); sanitizeMemInstBounds(&F, strategy); @@ -189,17 +189,17 @@ struct LabelCVEPass : public PassInfoMixin { out << F; out << "[CVEAssert] === Inserted Sanitizer Helpers === \n"; - if (vuln.UndesirableFunction.has_value()) { + if (vuln.Operation.has_value()) { /* NOTE: We are using '0' as a temporary this will be updated future PRs */ - sanitizeUndesirableOperationInFunction(&F, *vuln.UndesirableFunction, 0); + sanitizeContract(&F, *vuln.Operation, 0); result = PreservedAnalyses::none(); - out << "[CVEAssert] === Post Sanitization of Undesirable Operation IR " + out << "[CVEAssert] === Post Sanitization of Masked Operation IR " "=== \n"; out << F; } - if (vuln.Strategy == Vulnerability::RemediationStrategies::NONE) { + if (vuln.Strategy == RemediationStrategies::NONE) { out << "[CVEAssert] NONE strategy selected for " << vuln.TargetFileName << ":" << vuln.TargetFunctionName << "...\n"; out << "[CVEAssert] Skipping remediation\n"; @@ -345,7 +345,7 @@ struct LabelCVEPass : public PassInfoMixin { for (auto &vuln : vulns) { // Also skip instrumentation for skipped vulnerabilities - if (vuln.Strategy == Vulnerability::RemediationStrategies::NONE) { + if (vuln.Strategy == RemediationStrategies::NONE) { continue; } @@ -416,7 +416,7 @@ struct LabelCVEPass : public PassInfoMixin { std::vector moduleVulns; for (auto &vuln : vulnerabilities) { - if (vuln.Output == Vulnerability::RemediationOutput::PATCH) { + if (vuln.Output == RemediationOutput::PATCH) { patchVulns.push_back(vuln); } else { moduleVulns.push_back(vuln); diff --git a/resolve-cveassert/src/Contract.hpp b/resolve-cveassert/src/Contract.hpp new file mode 100644 index 00000000..1aaeb209 --- /dev/null +++ b/resolve-cveassert/src/Contract.hpp @@ -0,0 +1,31 @@ +/* + * Copyright (c) 2025 Riverside Research. + * LGPL-3; See LICENSE.txt in the repo root for details. + */ + +#pragma once + +#include "Remediation.hpp" +#include "llvm/Support/JSON.h" + +#include + +enum class PredicateKind { + InBounds, + NotEqual, + NotNull, + NonZero, +}; + +// Predicates tell the compiler what +// must be true before executing the operation +struct Predicate { + PredicateKind kind; + unsigned arg0; + unsigned arg1; +}; + +struct Contract { + std::vector preconditions; + RemediationStrategies strategy; +}; diff --git a/resolve-cveassert/src/FreeNonHeapMem.cpp b/resolve-cveassert/src/FreeNonHeapMem.cpp index 617f8727..2773a24d 100644 --- a/resolve-cveassert/src/FreeNonHeapMem.cpp +++ b/resolve-cveassert/src/FreeNonHeapMem.cpp @@ -7,8 +7,9 @@ #include "llvm/IR/IRBuilder.h" #include "llvm/IR/InlineAsm.h" +#include "CVEAssert.hpp" #include "IRUtils.hpp" -#include "Vulnerability.hpp" +#include "Remediation.hpp" using namespace llvm; @@ -57,8 +58,8 @@ Function *getOrCreateIsHeap(Function *F) { return cveIsHeapFn; } -Function *getOrCreateFreeOfNonHeapSanitizer( - Function *F, Vulnerability::RemediationStrategies strategy) { +Function *getOrCreateFreeOfNonHeapSanitizer(Function *F, + RemediationStrategies strategy) { std::string handlerName = "__cve_nonheap_free"; Module *M = F->getParent(); LLVMContext &Ctx = M->getContext(); @@ -119,8 +120,7 @@ Function *getOrCreateFreeOfNonHeapSanitizer( return cveFreeNonHeapFn; } -void sanitizeFreeOfNonHeap(Function *F, - Vulnerability::RemediationStrategies strategy) { +void sanitizeFreeOfNonHeap(Function *F, RemediationStrategies strategy) { LLVMContext &Ctx = F->getContext(); IRBuilder<> builder(Ctx); std::vector workList; diff --git a/resolve-cveassert/src/FreeNonHeapMem.hpp b/resolve-cveassert/src/FreeNonHeapMem.hpp index 657b728f..e85081b0 100644 --- a/resolve-cveassert/src/FreeNonHeapMem.hpp +++ b/resolve-cveassert/src/FreeNonHeapMem.hpp @@ -5,11 +5,11 @@ #pragma once -#include "Vulnerability.hpp" +#include "Remediation.hpp" #include "llvm/IR/Function.h" llvm::Function *getOrCreateIsHeap(llvm::Function *F); -llvm::Function *getOrCreateFreeOfNonHeapSanitizer( - llvm::Function *F, Vulnerability::RemediationStrategies strategy); -void sanitizeFreeOfNonHeap(llvm::Function *F, - Vulnerability::RemediationStrategies strategy); +llvm::Function * +getOrCreateFreeOfNonHeapSanitizer(llvm::Function *F, + RemediationStrategies strategy); +void sanitizeFreeOfNonHeap(llvm::Function *F, RemediationStrategies strategy); diff --git a/resolve-cveassert/src/IRUtils.cpp b/resolve-cveassert/src/IRUtils.cpp index e26c91d9..6720d7de 100644 --- a/resolve-cveassert/src/IRUtils.cpp +++ b/resolve-cveassert/src/IRUtils.cpp @@ -19,7 +19,7 @@ #include "CVEAssert.hpp" #include "IRUtils.hpp" -#include "Vulnerability.hpp" +#include "Remediation.hpp" #include #include @@ -524,9 +524,8 @@ Function *getOrCreateRecoverBufferFunction(Module *M) { return resolveRecoverFn; } -Function * -getOrCreateRemediationBehavior(Module *M, - Vulnerability::RemediationStrategies strategy) { +Function *getOrCreateRemediationBehavior(Module *M, + RemediationStrategies strategy) { auto &Ctx = M->getContext(); auto ptr_ty = PointerType::get(Ctx, 0); auto void_ty = Type::getVoidTy(Ctx); @@ -536,10 +535,10 @@ getOrCreateRemediationBehavior(Module *M, std::string fnName; switch (strategy) { - case Vulnerability::RemediationStrategies::EXIT: + case RemediationStrategies::EXIT: fnName = "__cve_exit"; break; - case Vulnerability::RemediationStrategies::RECOVER: + case RemediationStrategies::RECOVER: fnName = "__cve_recover"; break; default: @@ -556,7 +555,7 @@ getOrCreateRemediationBehavior(Module *M, IRBuilder<> builder(BB); switch (strategy) { - case Vulnerability::RemediationStrategies::EXIT: { + case RemediationStrategies::EXIT: { FunctionType *exitTy = FunctionType::get(void_ty, {i32_ty}, false); FunctionCallee exitFn = M->getOrInsertFunction("_exit", exitTy); builder.CreateCall(exitFn, {builder.getInt32(3)}); @@ -564,7 +563,7 @@ getOrCreateRemediationBehavior(Module *M, break; } - case Vulnerability::RemediationStrategies::RECOVER: { + case RemediationStrategies::RECOVER: { FunctionCallee longjmpFn = M->getOrInsertFunction( "longjmp", FunctionType::get(void_ty, {ptr_ty, i32_ty}, false)); diff --git a/resolve-cveassert/src/IRUtils.hpp b/resolve-cveassert/src/IRUtils.hpp index ed5e04e2..28b8e686 100644 --- a/resolve-cveassert/src/IRUtils.hpp +++ b/resolve-cveassert/src/IRUtils.hpp @@ -4,7 +4,8 @@ */ #pragma once -#include "Vulnerability.hpp" +#include "Remediation.hpp" + #include "llvm/ADT/StringRef.h" #include "llvm/IR/BasicBlock.h" #include "llvm/IR/Function.h" @@ -17,9 +18,8 @@ std::string getLLVMType(llvm::Type *ty); llvm::Function *getOrCreateResolveReportSanitizerTriggered(llvm::Module *M); -llvm::Function * -getOrCreateRemediationBehavior(llvm::Module *M, - Vulnerability::RemediationStrategies strategy); +llvm::Function *getOrCreateRemediationBehavior(llvm::Module *M, + RemediationStrategies strategy); llvm::Function * getOrCreateResolveHelper(llvm::Module *M, std::string fn_name, llvm::FunctionType *fn_type, diff --git a/resolve-cveassert/src/NullPointerSanitizer.cpp b/resolve-cveassert/src/NullPointerSanitizer.cpp index 185cca39..5ab5b8d9 100644 --- a/resolve-cveassert/src/NullPointerSanitizer.cpp +++ b/resolve-cveassert/src/NullPointerSanitizer.cpp @@ -13,13 +13,13 @@ #include "CVEAssert.hpp" #include "IRUtils.hpp" -#include "Vulnerability.hpp" +#include "Remediation.hpp" using namespace llvm; static Function * getOrCreateNullPtrLoadSanitizer(Function *F, Type *ty, - Vulnerability::RemediationStrategies strategy) { + RemediationStrategies strategy) { std::string handlerName = "__cve_null_check_ld_" + getLLVMType(ty); Module *M = F->getParent(); LLVMContext &Ctx = M->getContext(); @@ -62,12 +62,12 @@ getOrCreateNullPtrLoadSanitizer(Function *F, Type *ty, builder.SetInsertPoint(SanitizeNullPtrBB); switch (strategy) { - case Vulnerability::RemediationStrategies::CONTINUE: + case RemediationStrategies::CONTINUE: builder.CreateRet(Constant::getNullValue(ty)); break; - case Vulnerability::RemediationStrategies::EXIT: - case Vulnerability::RemediationStrategies::RECOVER: + case RemediationStrategies::EXIT: + case RemediationStrategies::RECOVER: builder.CreateCall(getOrCreateResolveReportSanitizerTriggered(M)); builder.CreateCall(getOrCreateRemediationBehavior(M, strategy)); builder.CreateUnreachable(); @@ -86,8 +86,9 @@ getOrCreateNullPtrLoadSanitizer(Function *F, Type *ty, return resolveNullPtrLdFn; } -static Function *getOrCreateNullPtrStoreSanitizer( - Function *F, Type *ty, Vulnerability::RemediationStrategies strategy) { +static Function * +getOrCreateNullPtrStoreSanitizer(Function *F, Type *ty, + RemediationStrategies strategy) { std::string handlerName = "__cve_null_check_st_" + getLLVMType(ty); Module *M = F->getParent(); LLVMContext &Ctx = M->getContext(); @@ -135,12 +136,12 @@ static Function *getOrCreateNullPtrStoreSanitizer( builder.SetInsertPoint(SanitizeNullPtrBB); switch (strategy) { - case Vulnerability::RemediationStrategies::CONTINUE: + case RemediationStrategies::CONTINUE: builder.CreateRetVoid(); break; - case Vulnerability::RemediationStrategies::EXIT: - case Vulnerability::RemediationStrategies::RECOVER: + case RemediationStrategies::EXIT: + case RemediationStrategies::RECOVER: builder.CreateCall(getOrCreateResolveReportSanitizerTriggered(M)); builder.CreateCall(getOrCreateRemediationBehavior(M, strategy)); builder.CreateUnreachable(); @@ -160,8 +161,7 @@ static Function *getOrCreateNullPtrStoreSanitizer( return resolveNullPtrStFn; } -void sanitizeNullPointers(Function *F, - Vulnerability::RemediationStrategies strategy) { +void sanitizeNullPointers(Function *F, RemediationStrategies strategy) { LLVMContext &Ctx = F->getContext(); IRBuilder<> builder(Ctx); @@ -169,16 +169,16 @@ void sanitizeNullPointers(Function *F, std::vector storeList; switch (strategy) { - case Vulnerability::RemediationStrategies::EXIT: - case Vulnerability::RemediationStrategies::RECOVER: - case Vulnerability::RemediationStrategies::CONTINUE: + case RemediationStrategies::EXIT: + case RemediationStrategies::RECOVER: + case RemediationStrategies::CONTINUE: break; default: llvm::errs() << "[CVEAssert] Error: sanitizeNullPointers does not support " "remediation strategy " << "defaulting to continue strategy!\n"; - strategy = Vulnerability::RemediationStrategies::CONTINUE; + strategy = RemediationStrategies::CONTINUE; break; } diff --git a/resolve-cveassert/src/NullPointerSanitizer.hpp b/resolve-cveassert/src/NullPointerSanitizer.hpp index fecdc01d..f8219ecd 100644 --- a/resolve-cveassert/src/NullPointerSanitizer.hpp +++ b/resolve-cveassert/src/NullPointerSanitizer.hpp @@ -5,7 +5,6 @@ #pragma once -#include "Vulnerability.hpp" +#include "Remediation.hpp" #include "llvm/IR/Function.h" -void sanitizeNullPointers(llvm::Function *F, - Vulnerability::RemediationStrategies strategy); +void sanitizeNullPointers(llvm::Function *F, RemediationStrategies strategy); diff --git a/resolve-cveassert/src/OperationMasking.cpp b/resolve-cveassert/src/OperationMasking.cpp index d13532c9..cd85c0ff 100644 --- a/resolve-cveassert/src/OperationMasking.cpp +++ b/resolve-cveassert/src/OperationMasking.cpp @@ -11,6 +11,7 @@ #include "llvm/Support/raw_ostream.h" #include "IRUtils.hpp" +#include "Remediation.hpp" #include #include @@ -18,14 +19,6 @@ using namespace llvm; -enum Cond { // Maybe adding an enum for all the possible conditions - EQ = 1, - GT = 2, - GT_EQ = 3, - LT = 4, - LT_EQ = 5 -}; - // Parameters // 1. Which arguments to return (or zero) // 2. Which arguments to test (if any) @@ -33,47 +26,58 @@ enum Cond { // Maybe adding an enum for all the possible conditions // We will continue generalizing this following eval-2 // Change this function name to be "replaceUndesirableOperation" more // generalized name -static Function *replaceUndesirableFunction(Module *M, CallInst *call, +static Function *getOrCreateContractWrapper(Module *M, CallInst *call, unsigned int argNum) { LLVMContext &Ctx = M->getContext(); IRBuilder<> builder(Ctx); - std::string handlerName = - "resolve_sanitized_" + call->getCalledFunction()->getName().str(); + Function *originalFn = call->getCalledFunction(); + + std::string handlerName = "__cve_contract_" + originalFn->getName().str(); - FunctionType *resolveSanitizedFnTy = - call->getCalledFunction()->getFunctionType(); + FunctionType *wrapperTy = originalFn->getFunctionType(); - Function *resolveSanitizedFn = - getOrCreateResolveHelper(M, handlerName, resolveSanitizedFnTy); + Function *resolveWrapperFn = + getOrCreateResolveHelper(M, handlerName, wrapperTy); + + SmallVector Args; + for (Argument &arg : resolveWrapperFn->args()) { + Args.push_back(&arg); + } - if (!resolveSanitizedFn->empty()) { - recordPatchFunction(resolveSanitizedFn); - return resolveSanitizedFn; + if (!resolveWrapperFn->empty()) { + recordPatchFunction(resolveWrapperFn); + return resolveWrapperFn; } - BasicBlock *EntryBB = BasicBlock::Create(Ctx, "", resolveSanitizedFn); + // TODO: Create 3 basic blocks + // 1. Preconditions + // 2. Valid path + // 3. Recovery path + BasicBlock *EntryBB = BasicBlock::Create(Ctx, "entry", resolveWrapperFn); // Insert a return instruction here. builder.SetInsertPoint(EntryBB); - builder.CreateRet(resolveSanitizedFn->getArg(argNum)); - validateIR(resolveSanitizedFn); - recordPatchFunction(resolveSanitizedFn); - return resolveSanitizedFn; + // TODO: Create helper to generate the llvm-ir for preconditions + // TODO: Create helper to generate valid path (call original operation + // contract) + // TODO: Create helper to generate recovery path + builder.CreateCall(originalFn, Args); + builder.CreateRet(resolveWrapperFn->getArg(argNum)); + + validateIR(resolveWrapperFn); + recordPatchFunction(resolveWrapperFn); + return resolveWrapperFn; } -void sanitizeUndesirableOperationInFunction(Function *F, std::string fnName, - unsigned int argNum) { +void sanitizeContract(Function *F, std::string fnName, unsigned int argNum) { Module *M = F->getParent(); LLVMContext &Ctx = M->getContext(); IRBuilder<> builder(Ctx); - // Container to store call insts std::vector callsToReplace; - // loop over each basic block in the vulnerable function for (auto &BB : *F) { - // loop over each instruction for (auto &inst : BB) { if (auto *call = dyn_cast(&inst)) { Function *calledFn = call->getCalledFunction(); @@ -93,24 +97,18 @@ void sanitizeUndesirableOperationInFunction(Function *F, std::string fnName, return; } - // Construct the resolve_sanitize_func function - Function *resolveSanitizedFn = - replaceUndesirableFunction(M, callsToReplace.front(), argNum); + Function *resolveWrapperFn = + getOrCreateContractWrapper(M, callsToReplace.front(), argNum); - // Replace calls at all callsites in the module for (auto call : callsToReplace) { builder.SetInsertPoint(call); - - // Get the arguments for the vulnerable function SmallVector fnArgs; for (unsigned int i = 0; i < call->arg_size(); ++i) { fnArgs.push_back(call->getOperand(i)); } - auto sanitizedCall = builder.CreateCall(resolveSanitizedFn, fnArgs); - - // replace all callsites - call->replaceAllUsesWith(sanitizedCall); + auto resolveWrapperCall = builder.CreateCall(resolveWrapperFn, fnArgs); + call->replaceAllUsesWith(resolveWrapperCall); call->eraseFromParent(); } } diff --git a/resolve-cveassert/src/OperationMasking.hpp b/resolve-cveassert/src/OperationMasking.hpp index b637d894..91276c5d 100644 --- a/resolve-cveassert/src/OperationMasking.hpp +++ b/resolve-cveassert/src/OperationMasking.hpp @@ -7,6 +7,5 @@ #include "llvm/IR/Function.h" #include -void sanitizeUndesirableOperationInFunction(llvm::Function *F, - std::string fnName, - unsigned int argNum); +void sanitizeContract(llvm::Function *F, std::string fnName, + unsigned int argNum); diff --git a/resolve-cveassert/src/Remediation.hpp b/resolve-cveassert/src/Remediation.hpp new file mode 100644 index 00000000..1a7ce9e0 --- /dev/null +++ b/resolve-cveassert/src/Remediation.hpp @@ -0,0 +1,19 @@ +/* + * Copyright (c) 2025 Riverside Research. + * LGPL-3; See LICENSE.txt in the repo root for details. + */ + +#pragma once + +enum class RemediationStrategies { + NONE, /* Skip remediation for this vulnerability */ + RECOVER, /* Applies setjmp and longjmp in vulnerable function */ + SAT, /* Uses saturating arithmetic */ + EXIT, /* Inserts exit function call with exit code */ + WRAP, /* Uses 2's complement arithmetic */ + CONTINUE, /* Invalid operations are ignored and return 0 */ + WIDEN /* Widen potentially overflowing intermediate operations */ +}; + +enum class RemediationOutput { INLINE, PATCH }; + diff --git a/resolve-cveassert/src/Vulnerability.hpp b/resolve-cveassert/src/Vulnerability.hpp index d169d4ef..003dd08e 100644 --- a/resolve-cveassert/src/Vulnerability.hpp +++ b/resolve-cveassert/src/Vulnerability.hpp @@ -6,6 +6,8 @@ #pragma once #include "CVEAssert.hpp" +#include "Contract.hpp" +#include "Remediation.hpp" #include "llvm/Support/JSON.h" #include "llvm/Support/raw_ostream.h" @@ -20,26 +22,47 @@ struct Vulnerability { - enum RemediationStrategies { - NONE, /* Skip remediation for this vulnerability */ - RECOVER, /* Applies setjmp and longjmp in vulnerable function */ - SAT, /* Uses saturating arithmetic */ - EXIT, /* Inserts exit function call with exit code */ - WRAP, /* Uses 2's complement arithmetic */ - CONTINUE, /* Invalid operations are ignored and return 0 */ - WIDEN /* Widen potentially overflowing intermediate operations */ - }; - - enum RemediationOutput { INLINE, PATCH }; - std::string TargetFileName; std::string TargetFunctionName; uint32_t WeaknessID; - std::optional UndesirableFunction; + std::optional Operation; RemediationStrategies Strategy; RemediationOutput Output; bool Gated; + static RemediationStrategies + parseRemediationStrategy(std::optional remediation) { + if (!remediation) { + llvm::errs() + << "[CVEAssert] Warning: remediation-strategy not specified. " + << "Defaulting to continue strategy.\n"; + return RemediationStrategies::CONTINUE; + } + + llvm::StringRef strategy = *remediation; + + if (remediation->str() == "none") { + return RemediationStrategies::NONE; + } else if (remediation->str() == "widen") { + return RemediationStrategies::WIDEN; + } else if (remediation->str() == "sat") { + return RemediationStrategies::SAT; + } else if (remediation->str() == "exit") { + return RemediationStrategies::EXIT; + } else if (remediation->str() == "continue") { + return RemediationStrategies::CONTINUE; + } else if (remediation->str() == "recover") { + return RemediationStrategies::RECOVER; + } else if (remediation->str() == "wrap") { + return RemediationStrategies::WRAP; + } else { + llvm::errs() + << "[CVEAssert] Warning: remediation-strategy is not recognized " + << "defaulting to continue strategy!\n"; + strategy = RemediationStrategies::CONTINUE; + } + } + static std::optional fromJson(llvm::json::Object *jsonObj) { auto getKey = [&](std::string key) -> std::optional { // Try the key as-is @@ -62,6 +85,14 @@ struct Vulnerability { return std::nullopt; }; + auto getArray[&](std::string)->std::optional { + if (auto arr = jsonObj->getArray(key)) { + return *arr; + } + std::replace(key.begin(), key.end(), '-', '_'); + return jsonObj->getArray(key); + }; + // Retrieve the target file name. auto targetFile = getKey("affected-file"); if (!targetFile) { @@ -84,41 +115,42 @@ struct Vulnerability { return std::nullopt; } - std::optional undesirableFunction = std::nullopt; - if (auto uf = getKey("undesirable-function")) { - undesirableFunction = uf->str(); + std::optional operation = std::nullopt; + if (auto op = getKey("operation")) { + operation = op->str(); } auto remediation = getKey("remediation-strategy"); - RemediationStrategies strategy; - - if (!remediation) { - llvm::errs() - << "[CVEAssert] Warning: remediation-strategy is not specified " - << "defaulting to continue strategy!\n"; - strategy = RemediationStrategies::CONTINUE; - - } else if (remediation->str() == "widen") { - strategy = RemediationStrategies::WIDEN; - } else if (remediation->str() == "sat") { - strategy = RemediationStrategies::SAT; - } else if (remediation->str() == "exit") { - strategy = RemediationStrategies::EXIT; - } else if (remediation->str() == "continue") { - strategy = RemediationStrategies::CONTINUE; - } else if (remediation->str() == "recover") { - strategy = RemediationStrategies::RECOVER; - } else if (remediation->str() == "none") { - strategy = RemediationStrategies::NONE; - } else if (remediation->str() == "wrap") { - strategy = RemediationStrategies::WRAP; - } else { - llvm::errs() - << "[CVEAssert] Warning: remediation-strategy is not recognized " - << "defaulting to continue strategy!\n"; - strategy = RemediationStrategies::CONTINUE; - } + auto strategy = parseRemediationStrategy(remediation); + // if (!remediation) { + // llvm::errs() + // << "[CVEAssert] Warning: remediation-strategy is not specified " + // << "defaulting to continue strategy!\n"; + // strategy = RemediationStrategies::CONTINUE; + // + // } else if (remediation->str() == "widen") { + // strategy = RemediationStrategies::WIDEN; + // } else if (remediation->str() == "sat") { + // strategy = RemediationStrategies::SAT; + // } else if (remediation->str() == "exit") { + // strategy = RemediationStrategies::EXIT; + // } else if (remediation->str() == "continue") { + // strategy = RemediationStrategies::CONTINUE; + // } else if (remediation->str() == "recover") { + // strategy = RemediationStrategies::RECOVER; + // } else if (remediation->str() == "none") { + // strategy = RemediationStrategies::NONE; + // } else if (remediation->str() == "wrap") { + // strategy = RemediationStrategies::WRAP; + // } else { + // llvm::errs() + // << "[CVEAssert] Warning: remediation-strategy is not recognized + // " + // << "defaulting to continue strategy!\n"; + // strategy = RemediationStrategies::CONTINUE; + // } + // auto outstr = getKey("output"); RemediationOutput output; @@ -144,11 +176,15 @@ struct Vulnerability { output = RemediationOutput::INLINE; } + // TODO: Move contract parsing logic into its own function. + // auto contract = parsePredicate(jsonObj); + Vulnerability vuln{ targetFile->str(), targetFunction->str(), static_cast(std::stoi(vulnID->str())), - undesirableFunction, + operation, + contract, strategy, output, gated,