1//==-- X86LoadValueInjectionLoadHardening.cpp - LVI load hardening for x86 --=//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8///
9/// Description: This pass finds Load Value Injection (LVI) gadgets consisting
10/// of a load from memory (i.e., SOURCE), and any operation that may transmit
11/// the value loaded from memory over a covert channel, or use the value loaded
12/// from memory to determine a branch/call target (i.e., SINK). After finding
13/// all such gadgets in a given function, the pass minimally inserts LFENCE
14/// instructions in such a manner that the following property is satisfied: for
15/// all SOURCE+SINK pairs, all paths in the CFG from SOURCE to SINK contain at
16/// least one LFENCE instruction. The algorithm that implements this minimal
17/// insertion is influenced by an academic paper that minimally inserts memory
18/// fences for high-performance concurrent programs:
19/// http://www.cs.ucr.edu/~lesani/companion/oopsla15/OOPSLA15.pdf
20/// The algorithm implemented in this pass is as follows:
21/// 1. Build a condensed CFG (i.e., a GadgetGraph) consisting only of the
22/// following components:
23/// - SOURCE instructions (also includes function arguments)
24/// - SINK instructions
25/// - Basic block entry points
26/// - Basic block terminators
27/// - LFENCE instructions
28/// 2. Analyze the GadgetGraph to determine which SOURCE+SINK pairs (i.e.,
29/// gadgets) are already mitigated by existing LFENCEs. If all gadgets have been
30/// mitigated, go to step 6.
31/// 3. Use a heuristic or plugin to approximate minimal LFENCE insertion.
32/// 4. Insert one LFENCE along each CFG edge that was cut in step 3.
33/// 5. Go to step 2.
34/// 6. If any LFENCEs were inserted, return `true` from runOnMachineFunction()
35/// to tell LLVM that the function was modified.
36///
37//===----------------------------------------------------------------------===//
38
39#include "ImmutableGraph.h"
40#include "X86.h"
41#include "X86Subtarget.h"
42#include "X86TargetMachine.h"
43#include "llvm/ADT/DenseMap.h"
44#include "llvm/ADT/STLExtras.h"
45#include "llvm/ADT/SmallSet.h"
46#include "llvm/ADT/Statistic.h"
47#include "llvm/ADT/StringRef.h"
48#include "llvm/CodeGen/MachineBasicBlock.h"
49#include "llvm/CodeGen/MachineDominanceFrontier.h"
50#include "llvm/CodeGen/MachineDominators.h"
51#include "llvm/CodeGen/MachineFunction.h"
52#include "llvm/CodeGen/MachineFunctionPass.h"
53#include "llvm/CodeGen/MachineInstr.h"
54#include "llvm/CodeGen/MachineInstrBuilder.h"
55#include "llvm/CodeGen/MachineLoopInfo.h"
56#include "llvm/CodeGen/RDFGraph.h"
57#include "llvm/CodeGen/RDFLiveness.h"
58#include "llvm/InitializePasses.h"
59#include "llvm/Support/DOTGraphTraits.h"
60#include "llvm/Support/Debug.h"
61#include "llvm/Support/DynamicLibrary.h"
62#include "llvm/Support/GraphWriter.h"
63#include "llvm/Support/raw_ostream.h"
64
65using namespace llvm;
66
67#define PASS_KEY "x86-lvi-load"
68#define DEBUG_TYPE PASS_KEY
69
70STATISTIC(NumFences, "Number of LFENCEs inserted for LVI mitigation");
71STATISTIC(NumFunctionsConsidered, "Number of functions analyzed");
72STATISTIC(NumFunctionsMitigated, "Number of functions for which mitigations "
73 "were deployed");
74STATISTIC(NumGadgets, "Number of LVI gadgets detected during analysis");
75
76static llvm::sys::DynamicLibrary OptimizeDL;
77typedef int (*OptimizeCutT)(unsigned int *Nodes, unsigned int NodesSize,
78 unsigned int *Edges, int *EdgeValues,
79 int *CutEdges /* out */, unsigned int EdgesSize);
80static OptimizeCutT OptimizeCut = nullptr;
81
82namespace {
83
84struct MachineGadgetGraph : ImmutableGraph<MachineInstr *, int> {
85 static constexpr int GadgetEdgeSentinel = -1;
86 static constexpr MachineInstr *const ArgNodeSentinel = nullptr;
87
88 using GraphT = ImmutableGraph<MachineInstr *, int>;
89 using Node = GraphT::Node;
90 using Edge = GraphT::Edge;
91 using size_type = GraphT::size_type;
92 MachineGadgetGraph(std::unique_ptr<Node[]> Nodes,
93 std::unique_ptr<Edge[]> Edges, size_type NodesSize,
94 size_type EdgesSize, int NumFences = 0, int NumGadgets = 0)
95 : GraphT(std::move(Nodes), std::move(Edges), NodesSize, EdgesSize),
96 NumFences(NumFences), NumGadgets(NumGadgets) {}
97 static inline bool isCFGEdge(const Edge &E) {
98 return E.getValue() != GadgetEdgeSentinel;
99 }
100 static inline bool isGadgetEdge(const Edge &E) {
101 return E.getValue() == GadgetEdgeSentinel;
102 }
103 int NumFences;
104 int NumGadgets;
105};
106
107constexpr StringRef X86LVILHPassName =
108 "X86 Load Value Injection (LVI) Load Hardening";
109
110class X86LoadValueInjectionLoadHardeningLegacy : public MachineFunctionPass {
111public:
112 X86LoadValueInjectionLoadHardeningLegacy() : MachineFunctionPass(ID) {}
113
114 StringRef getPassName() const override { return X86LVILHPassName; }
115 void getAnalysisUsage(AnalysisUsage &AU) const override;
116 bool runOnMachineFunction(MachineFunction &MF) override;
117
118 static char ID;
119};
120
121class X86LoadValueInjectionLoadHardeningImpl {
122public:
123 X86LoadValueInjectionLoadHardeningImpl() = default;
124
125 bool run(MachineFunction &MF, const MachineLoopInfo &MLI,
126 const MachineDominatorTree &MDT,
127 const MachineDominanceFrontier &MDF);
128
129private:
130 using GraphBuilder = ImmutableGraphBuilder<MachineGadgetGraph>;
131 using Edge = MachineGadgetGraph::Edge;
132 using Node = MachineGadgetGraph::Node;
133 using EdgeSet = MachineGadgetGraph::EdgeSet;
134 using NodeSet = MachineGadgetGraph::NodeSet;
135
136 const X86Subtarget *STI = nullptr;
137 const TargetInstrInfo *TII = nullptr;
138 const TargetRegisterInfo *TRI = nullptr;
139
140 std::unique_ptr<MachineGadgetGraph>
141 getGadgetGraph(MachineFunction &MF, const MachineLoopInfo &MLI,
142 const MachineDominatorTree &MDT,
143 const MachineDominanceFrontier &MDF) const;
144 int hardenLoadsWithPlugin(MachineFunction &MF,
145 std::unique_ptr<MachineGadgetGraph> Graph) const;
146 int hardenLoadsWithHeuristic(MachineFunction &MF,
147 std::unique_ptr<MachineGadgetGraph> Graph) const;
148 int elimMitigatedEdgesAndNodes(MachineGadgetGraph &G,
149 EdgeSet &ElimEdges /* in, out */,
150 NodeSet &ElimNodes /* in, out */) const;
151 std::unique_ptr<MachineGadgetGraph>
152 trimMitigatedEdges(std::unique_ptr<MachineGadgetGraph> Graph) const;
153 int insertFences(MachineFunction &MF, MachineGadgetGraph &G,
154 EdgeSet &CutEdges /* in, out */) const;
155 bool instrUsesRegToAccessMemory(const MachineInstr &I, Register Reg) const;
156 bool instrUsesRegToBranch(const MachineInstr &I, Register Reg) const;
157 inline bool isFence(const MachineInstr *MI) const {
158 return MI && (MI->getOpcode() == X86::LFENCE ||
159 (STI->useLVIControlFlowIntegrity() && MI->isCall()));
160 }
161};
162
163} // end anonymous namespace
164
165namespace llvm {
166
167template <>
168struct GraphTraits<MachineGadgetGraph *>
169 : GraphTraits<ImmutableGraph<MachineInstr *, int> *> {};
170
171template <>
172struct DOTGraphTraits<MachineGadgetGraph *> : DefaultDOTGraphTraits {
173 using GraphType = MachineGadgetGraph;
174 using Traits = llvm::GraphTraits<GraphType *>;
175 using NodeRef = Traits::NodeRef;
176 using EdgeRef = Traits::EdgeRef;
177 using ChildIteratorType = Traits::ChildIteratorType;
178 using ChildEdgeIteratorType = Traits::ChildEdgeIteratorType;
179
180 DOTGraphTraits(bool IsSimple = false) : DefaultDOTGraphTraits(IsSimple) {}
181
182 std::string getNodeLabel(NodeRef Node, GraphType *) {
183 if (Node->getValue() == MachineGadgetGraph::ArgNodeSentinel)
184 return "ARGS";
185
186 std::string Str;
187 raw_string_ostream OS(Str);
188 OS << *Node->getValue();
189 return OS.str();
190 }
191
192 static std::string getNodeAttributes(NodeRef Node, GraphType *) {
193 MachineInstr *MI = Node->getValue();
194 if (MI == MachineGadgetGraph::ArgNodeSentinel)
195 return "color = blue";
196 if (MI->getOpcode() == X86::LFENCE)
197 return "color = green";
198 return "";
199 }
200
201 static std::string getEdgeAttributes(NodeRef, ChildIteratorType E,
202 GraphType *) {
203 int EdgeVal = (*E.getCurrent()).getValue();
204 return EdgeVal >= 0 ? "label = " + std::to_string(val: EdgeVal)
205 : "color = red, style = \"dashed\"";
206 }
207};
208
209} // end namespace llvm
210
211char X86LoadValueInjectionLoadHardeningLegacy::ID = 0;
212
213void X86LoadValueInjectionLoadHardeningLegacy::getAnalysisUsage(
214 AnalysisUsage &AU) const {
215 MachineFunctionPass::getAnalysisUsage(AU);
216 AU.addRequired<MachineLoopInfoWrapperPass>();
217 AU.addRequired<MachineDominatorTreeWrapperPass>();
218 AU.addRequired<MachineDominanceFrontierWrapperPass>();
219 AU.setPreservesCFG();
220}
221
222static void writeGadgetGraph(raw_ostream &OS, MachineFunction &MF,
223 MachineGadgetGraph *G) {
224 WriteGraph(O&: OS, G, /*ShortNames*/ false,
225 Title: "Speculative gadgets for \"" + MF.getName() + "\" function");
226}
227
228bool X86LoadValueInjectionLoadHardeningImpl::run(
229 MachineFunction &MF, const MachineLoopInfo &MLI,
230 const MachineDominatorTree &MDT, const MachineDominanceFrontier &MDF) {
231 LLVM_DEBUG(dbgs() << "***** " << X86LVILHPassName << " : " << MF.getName()
232 << " *****\n");
233 STI = &MF.getSubtarget<X86Subtarget>();
234 const X86Options &CLOpts = STI->getCLOpts();
235
236 // FIXME: support 32-bit
237 if (!STI->is64Bit())
238 report_fatal_error(reason: "LVI load hardening is only supported on 64-bit", gen_crash_diag: false);
239
240 ++NumFunctionsConsidered;
241 TII = STI->getInstrInfo();
242 TRI = STI->getRegisterInfo();
243 LLVM_DEBUG(dbgs() << "Building gadget graph...\n");
244 std::unique_ptr<MachineGadgetGraph> Graph = getGadgetGraph(MF, MLI, MDT, MDF);
245 LLVM_DEBUG(dbgs() << "Building gadget graph... Done\n");
246 if (Graph == nullptr)
247 return false; // didn't find any gadgets
248
249 if (CLOpts.lvi_load_dot_verify) {
250 writeGadgetGraph(OS&: outs(), MF, G: Graph.get());
251 return false;
252 }
253
254 if (CLOpts.lvi_load_dot || CLOpts.lvi_load_dot_only) {
255 LLVM_DEBUG(dbgs() << "Emitting gadget graph...\n");
256 std::error_code FileError;
257 std::string FileName = "lvi.";
258 FileName += MF.getName();
259 FileName += ".dot";
260 raw_fd_ostream FileOut(FileName, FileError);
261 if (FileError)
262 errs() << FileError.message();
263 writeGadgetGraph(OS&: FileOut, MF, G: Graph.get());
264 FileOut.close();
265 LLVM_DEBUG(dbgs() << "Emitting gadget graph... Done\n");
266 if (CLOpts.lvi_load_dot_only)
267 return false;
268 }
269
270 int FencesInserted;
271 if (!CLOpts.lvi_load_opt_plugin.empty()) {
272 if (!OptimizeDL.isValid()) {
273 std::string ErrorMsg;
274 OptimizeDL = llvm::sys::DynamicLibrary::getPermanentLibrary(
275 filename: CLOpts.lvi_load_opt_plugin.str().c_str(), errMsg: &ErrorMsg);
276 if (!ErrorMsg.empty())
277 report_fatal_error(reason: Twine("Failed to load opt plugin: \"") + ErrorMsg +
278 "\"");
279 OptimizeCut = (OptimizeCutT)OptimizeDL.getAddressOfSymbol(symbolName: "optimize_cut");
280 if (!OptimizeCut)
281 report_fatal_error(reason: "Invalid optimization plugin");
282 }
283 FencesInserted = hardenLoadsWithPlugin(MF, Graph: std::move(Graph));
284 } else { // Use the default greedy heuristic
285 FencesInserted = hardenLoadsWithHeuristic(MF, Graph: std::move(Graph));
286 }
287
288 if (FencesInserted > 0)
289 ++NumFunctionsMitigated;
290 NumFences += FencesInserted;
291 return (FencesInserted > 0);
292}
293
294std::unique_ptr<MachineGadgetGraph>
295X86LoadValueInjectionLoadHardeningImpl::getGadgetGraph(
296 MachineFunction &MF, const MachineLoopInfo &MLI,
297 const MachineDominatorTree &MDT,
298 const MachineDominanceFrontier &MDF) const {
299 using namespace rdf;
300
301 // Build the Register Dataflow Graph using the RDF framework
302 DataFlowGraph DFG{MF, *TII, *TRI, MDT, MDF};
303 DFG.build();
304 Liveness L{MF.getRegInfo(), DFG};
305 L.computePhiInfo();
306
307 GraphBuilder Builder;
308 using GraphIter = GraphBuilder::BuilderNodeRef;
309 DenseMap<MachineInstr *, GraphIter> NodeMap;
310 int FenceCount = 0, GadgetCount = 0;
311 auto MaybeAddNode = [&NodeMap, &Builder](MachineInstr *MI) {
312 auto [Ref, Inserted] = NodeMap.try_emplace(Key: MI);
313 if (Inserted) {
314 auto I = Builder.addVertex(V: MI);
315 Ref->second = I;
316 return std::pair<GraphIter, bool>{I, true};
317 }
318 return std::pair<GraphIter, bool>{Ref->getSecond(), false};
319 };
320
321 // The `Transmitters` map memoizes transmitters found for each def. If a def
322 // has not yet been analyzed, then it will not appear in the map. If a def
323 // has been analyzed and was determined not to have any transmitters, then
324 // its list of transmitters will be empty.
325 DenseMap<NodeId, std::vector<NodeId>> Transmitters;
326
327 // Analyze all machine instructions to find gadgets and LFENCEs, adding
328 // each interesting value to `Nodes`
329 auto AnalyzeDef = [&](NodeAddr<DefNode *> SourceDef) {
330 SmallSet<NodeId, 8> UsesVisited, DefsVisited;
331 std::function<void(NodeAddr<DefNode *>)> AnalyzeDefUseChain =
332 [&](NodeAddr<DefNode *> Def) {
333 if (Transmitters.contains(Val: Def.Id))
334 return; // Already analyzed `Def`
335
336 // Use RDF to find all the uses of `Def`
337 rdf::NodeSet Uses;
338 RegisterRef DefReg = Def.Addr->getRegRef(G: DFG);
339 for (auto UseID : L.getAllReachedUses(RefRR: DefReg, DefA: Def)) {
340 auto Use = DFG.addr<UseNode *>(N: UseID);
341 if (Use.Addr->getFlags() & NodeAttrs::PhiRef) { // phi node
342 NodeAddr<PhiNode *> Phi = Use.Addr->getOwner(G: DFG);
343 for (const auto& I : L.getRealUses(P: Phi.Id)) {
344 if (DFG.getPRI().alias(RA: RegisterRef(I.first), RB: DefReg)) {
345 for (const auto &UA : I.second)
346 Uses.emplace(args: UA.first);
347 }
348 }
349 } else { // not a phi node
350 Uses.emplace(args&: UseID);
351 }
352 }
353
354 // For each use of `Def`, we want to know whether:
355 // (1) The use can leak the Def'ed value,
356 // (2) The use can further propagate the Def'ed value to more defs
357 for (auto UseID : Uses) {
358 if (!UsesVisited.insert(V: UseID).second)
359 continue; // Already visited this use of `Def`
360
361 auto Use = DFG.addr<UseNode *>(N: UseID);
362 assert(!(Use.Addr->getFlags() & NodeAttrs::PhiRef));
363 MachineOperand &UseMO = Use.Addr->getOp();
364 MachineInstr &UseMI = *UseMO.getParent();
365 assert(UseMO.isReg());
366
367 // We naively assume that an instruction propagates any loaded
368 // uses to all defs unless the instruction is a call, in which
369 // case all arguments will be treated as gadget sources during
370 // analysis of the callee function.
371 if (UseMI.isCall())
372 continue;
373
374 // Check whether this use can transmit (leak) its value.
375 if (instrUsesRegToAccessMemory(I: UseMI, Reg: UseMO.getReg()) ||
376 (!STI->getCLOpts().lvi_load_no_cbranch &&
377 instrUsesRegToBranch(I: UseMI, Reg: UseMO.getReg()))) {
378 Transmitters[Def.Id].push_back(x: Use.Addr->getOwner(G: DFG).Id);
379 if (UseMI.mayLoad())
380 continue; // Found a transmitting load -- no need to continue
381 // traversing its defs (i.e., this load will become
382 // a new gadget source anyways).
383 }
384
385 // Check whether the use propagates to more defs.
386 NodeAddr<InstrNode *> Owner{Use.Addr->getOwner(G: DFG)};
387 for (const auto &ChildDef :
388 Owner.Addr->members_if(P: DataFlowGraph::IsDef, G: DFG)) {
389 if (!DefsVisited.insert(V: ChildDef.Id).second)
390 continue; // Already visited this def
391 if (Def.Addr->getAttrs() & NodeAttrs::Dead)
392 continue;
393 if (Def.Id == ChildDef.Id)
394 continue; // `Def` uses itself (e.g., increment loop counter)
395
396 AnalyzeDefUseChain(ChildDef);
397
398 // `Def` inherits all of its child defs' transmitters.
399 for (auto TransmitterId : Transmitters[ChildDef.Id])
400 Transmitters[Def.Id].push_back(x: TransmitterId);
401 }
402 }
403
404 // Note that this statement adds `Def.Id` to the map if no
405 // transmitters were found for `Def`.
406 auto &DefTransmitters = Transmitters[Def.Id];
407
408 // Remove duplicate transmitters
409 llvm::sort(C&: DefTransmitters);
410 DefTransmitters.erase(first: llvm::unique(R&: DefTransmitters),
411 last: DefTransmitters.end());
412 };
413
414 // Find all of the transmitters
415 AnalyzeDefUseChain(SourceDef);
416 auto &SourceDefTransmitters = Transmitters[SourceDef.Id];
417 if (SourceDefTransmitters.empty())
418 return; // No transmitters for `SourceDef`
419
420 MachineInstr *Source = SourceDef.Addr->getFlags() & NodeAttrs::PhiRef
421 ? MachineGadgetGraph::ArgNodeSentinel
422 : SourceDef.Addr->getOp().getParent();
423 auto GadgetSource = MaybeAddNode(Source);
424 // Each transmitter is a sink for `SourceDef`.
425 for (auto TransmitterId : SourceDefTransmitters) {
426 MachineInstr *Sink = DFG.addr<StmtNode *>(N: TransmitterId).Addr->getCode();
427 auto GadgetSink = MaybeAddNode(Sink);
428 // Add the gadget edge to the graph.
429 Builder.addEdge(E: MachineGadgetGraph::GadgetEdgeSentinel,
430 From: GadgetSource.first, To: GadgetSink.first);
431 ++GadgetCount;
432 }
433 };
434
435 LLVM_DEBUG(dbgs() << "Analyzing def-use chains to find gadgets\n");
436 // Analyze function arguments
437 NodeAddr<BlockNode *> EntryBlock = DFG.getFunc().Addr->getEntryBlock(G: DFG);
438 for (NodeAddr<PhiNode *> ArgPhi :
439 EntryBlock.Addr->members_if(P: DataFlowGraph::IsPhi, G: DFG)) {
440 NodeList Defs = ArgPhi.Addr->members_if(P: DataFlowGraph::IsDef, G: DFG);
441 llvm::for_each(Range&: Defs, F: AnalyzeDef);
442 }
443 // Analyze every instruction in MF
444 for (NodeAddr<BlockNode *> BA : DFG.getFunc().Addr->members(G: DFG)) {
445 for (NodeAddr<StmtNode *> SA :
446 BA.Addr->members_if(P: DataFlowGraph::IsCode<NodeAttrs::Stmt>, G: DFG)) {
447 MachineInstr *MI = SA.Addr->getCode();
448 if (isFence(MI)) {
449 MaybeAddNode(MI);
450 ++FenceCount;
451 } else if (MI->mayLoad()) {
452 NodeList Defs = SA.Addr->members_if(P: DataFlowGraph::IsDef, G: DFG);
453 llvm::for_each(Range&: Defs, F: AnalyzeDef);
454 }
455 }
456 }
457 LLVM_DEBUG(dbgs() << "Found " << FenceCount << " fences\n");
458 LLVM_DEBUG(dbgs() << "Found " << GadgetCount << " gadgets\n");
459 if (GadgetCount == 0)
460 return nullptr;
461 NumGadgets += GadgetCount;
462
463 // Traverse CFG to build the rest of the graph
464 SmallPtrSet<MachineBasicBlock *, 8> BlocksVisited;
465 std::function<void(MachineBasicBlock *, GraphIter, unsigned)> TraverseCFG =
466 [&](MachineBasicBlock *MBB, GraphIter GI, unsigned ParentDepth) {
467 unsigned LoopDepth = MLI.getLoopDepth(BB: MBB);
468 auto NI = MBB->getFirstNonDebugInstr(/*SkipPseudoOp=*/false);
469 if (NI != MBB->end()) {
470 // Always add the first non-debug instruction in each block.
471 auto BeginBB = MaybeAddNode(&*NI);
472 Builder.addEdge(E: ParentDepth, From: GI, To: BeginBB.first);
473 if (!BlocksVisited.insert(Ptr: MBB).second)
474 return;
475
476 // Add any instructions within the block that are gadget components
477 GI = BeginBB.first;
478 while (++NI != MBB->end()) {
479 auto Ref = NodeMap.find(Val: &*NI);
480 if (Ref != NodeMap.end()) {
481 Builder.addEdge(E: LoopDepth, From: GI, To: Ref->getSecond());
482 GI = Ref->getSecond();
483 }
484 }
485
486 // Always add the terminator instruction, if one exists
487 auto T = MBB->getFirstTerminator();
488 if (T != MBB->end()) {
489 auto EndBB = MaybeAddNode(&*T);
490 if (EndBB.second)
491 Builder.addEdge(E: LoopDepth, From: GI, To: EndBB.first);
492 GI = EndBB.first;
493 }
494 }
495 for (MachineBasicBlock *Succ : MBB->successors())
496 TraverseCFG(Succ, GI, LoopDepth);
497 };
498 // ArgNodeSentinel is a pseudo-instruction that represents MF args in the
499 // GadgetGraph
500 GraphIter ArgNode = MaybeAddNode(MachineGadgetGraph::ArgNodeSentinel).first;
501 TraverseCFG(&MF.front(), ArgNode, 0);
502 std::unique_ptr<MachineGadgetGraph> G{Builder.get(Args&: FenceCount, Args&: GadgetCount)};
503 LLVM_DEBUG(dbgs() << "Found " << G->nodes_size() << " nodes\n");
504 return G;
505}
506
507// Returns the number of remaining gadget edges that could not be eliminated
508int X86LoadValueInjectionLoadHardeningImpl::elimMitigatedEdgesAndNodes(
509 MachineGadgetGraph &G, EdgeSet &ElimEdges /* in, out */,
510 NodeSet &ElimNodes /* in, out */) const {
511 if (G.NumFences > 0) {
512 // Eliminate fences and CFG edges that ingress and egress the fence, as
513 // they are trivially mitigated.
514 for (const Edge &E : G.edges()) {
515 const Node *Dest = E.getDest();
516 if (isFence(MI: Dest->getValue())) {
517 ElimNodes.insert(N: *Dest);
518 ElimEdges.insert(E);
519 for (const Edge &DE : Dest->edges())
520 ElimEdges.insert(E: DE);
521 }
522 }
523 }
524
525 // Find and eliminate gadget edges that have been mitigated.
526 int RemainingGadgets = 0;
527 NodeSet ReachableNodes{G};
528 for (const Node &RootN : G.nodes()) {
529 if (llvm::none_of(Range: RootN.edges(), P: MachineGadgetGraph::isGadgetEdge))
530 continue; // skip this node if it isn't a gadget source
531
532 // Find all of the nodes that are CFG-reachable from RootN using DFS
533 ReachableNodes.clear();
534 std::function<void(const Node *, bool)> FindReachableNodes =
535 [&](const Node *N, bool FirstNode) {
536 if (!FirstNode)
537 ReachableNodes.insert(N: *N);
538 for (const Edge &E : N->edges()) {
539 const Node *Dest = E.getDest();
540 if (MachineGadgetGraph::isCFGEdge(E) && !ElimEdges.contains(E) &&
541 !ReachableNodes.contains(N: *Dest))
542 FindReachableNodes(Dest, false);
543 }
544 };
545 FindReachableNodes(&RootN, true);
546
547 // Any gadget whose sink is unreachable has been mitigated
548 for (const Edge &E : RootN.edges()) {
549 if (MachineGadgetGraph::isGadgetEdge(E)) {
550 if (ReachableNodes.contains(N: *E.getDest())) {
551 // This gadget's sink is reachable
552 ++RemainingGadgets;
553 } else { // This gadget's sink is unreachable, and therefore mitigated
554 ElimEdges.insert(E);
555 }
556 }
557 }
558 }
559 return RemainingGadgets;
560}
561
562std::unique_ptr<MachineGadgetGraph>
563X86LoadValueInjectionLoadHardeningImpl::trimMitigatedEdges(
564 std::unique_ptr<MachineGadgetGraph> Graph) const {
565 NodeSet ElimNodes{*Graph};
566 EdgeSet ElimEdges{*Graph};
567 int RemainingGadgets =
568 elimMitigatedEdgesAndNodes(G&: *Graph, ElimEdges, ElimNodes);
569 if (ElimEdges.empty() && ElimNodes.empty()) {
570 Graph->NumFences = 0;
571 Graph->NumGadgets = RemainingGadgets;
572 } else {
573 Graph = GraphBuilder::trim(G: *Graph, TrimNodes: ElimNodes, TrimEdges: ElimEdges, Args: 0 /* NumFences */,
574 Args&: RemainingGadgets);
575 }
576 return Graph;
577}
578
579int X86LoadValueInjectionLoadHardeningImpl::hardenLoadsWithPlugin(
580 MachineFunction &MF, std::unique_ptr<MachineGadgetGraph> Graph) const {
581 int FencesInserted = 0;
582
583 do {
584 LLVM_DEBUG(dbgs() << "Eliminating mitigated paths...\n");
585 Graph = trimMitigatedEdges(Graph: std::move(Graph));
586 LLVM_DEBUG(dbgs() << "Eliminating mitigated paths... Done\n");
587 if (Graph->NumGadgets == 0)
588 break;
589
590 LLVM_DEBUG(dbgs() << "Cutting edges...\n");
591 EdgeSet CutEdges{*Graph};
592 auto Nodes = std::make_unique<unsigned int[]>(num: Graph->nodes_size() +
593 1 /* terminator node */);
594 auto Edges = std::make_unique<unsigned int[]>(num: Graph->edges_size());
595 auto EdgeCuts = std::make_unique<int[]>(num: Graph->edges_size());
596 auto EdgeValues = std::make_unique<int[]>(num: Graph->edges_size());
597 for (const Node &N : Graph->nodes()) {
598 Nodes[Graph->getNodeIndex(N)] = Graph->getEdgeIndex(E: *N.edges_begin());
599 }
600 Nodes[Graph->nodes_size()] = Graph->edges_size(); // terminator node
601 for (const Edge &E : Graph->edges()) {
602 Edges[Graph->getEdgeIndex(E)] = Graph->getNodeIndex(N: *E.getDest());
603 EdgeValues[Graph->getEdgeIndex(E)] = E.getValue();
604 }
605 OptimizeCut(Nodes.get(), Graph->nodes_size(), Edges.get(), EdgeValues.get(),
606 EdgeCuts.get(), Graph->edges_size());
607 for (int I = 0; I < Graph->edges_size(); ++I)
608 if (EdgeCuts[I])
609 CutEdges.set(I);
610 LLVM_DEBUG(dbgs() << "Cutting edges... Done\n");
611 LLVM_DEBUG(dbgs() << "Cut " << CutEdges.count() << " edges\n");
612
613 LLVM_DEBUG(dbgs() << "Inserting LFENCEs...\n");
614 FencesInserted += insertFences(MF, G&: *Graph, CutEdges);
615 LLVM_DEBUG(dbgs() << "Inserting LFENCEs... Done\n");
616 LLVM_DEBUG(dbgs() << "Inserted " << FencesInserted << " fences\n");
617
618 Graph = GraphBuilder::trim(G: *Graph, TrimNodes: NodeSet{*Graph}, TrimEdges: CutEdges);
619 } while (true);
620
621 return FencesInserted;
622}
623
624int X86LoadValueInjectionLoadHardeningImpl::hardenLoadsWithHeuristic(
625 MachineFunction &MF, std::unique_ptr<MachineGadgetGraph> Graph) const {
626 // If `MF` does not have any fences, then no gadgets would have been
627 // mitigated at this point.
628 if (Graph->NumFences > 0) {
629 LLVM_DEBUG(dbgs() << "Eliminating mitigated paths...\n");
630 Graph = trimMitigatedEdges(Graph: std::move(Graph));
631 LLVM_DEBUG(dbgs() << "Eliminating mitigated paths... Done\n");
632 }
633
634 if (Graph->NumGadgets == 0)
635 return 0;
636
637 LLVM_DEBUG(dbgs() << "Cutting edges...\n");
638 EdgeSet CutEdges{*Graph};
639
640 // Begin by collecting all ingress CFG edges for each node
641 DenseMap<const Node *, SmallVector<const Edge *, 2>> IngressEdgeMap;
642 for (const Edge &E : Graph->edges())
643 if (MachineGadgetGraph::isCFGEdge(E))
644 IngressEdgeMap[E.getDest()].push_back(Elt: &E);
645
646 // For each gadget edge, make cuts that guarantee the gadget will be
647 // mitigated. A computationally efficient way to achieve this is to either:
648 // (a) cut all egress CFG edges from the gadget source, or
649 // (b) cut all ingress CFG edges to the gadget sink.
650 //
651 // Moreover, the algorithm tries not to make a cut into a loop by preferring
652 // to make a (b)-type cut if the gadget source resides at a greater loop depth
653 // than the gadget sink, or an (a)-type cut otherwise.
654 for (const Node &N : Graph->nodes()) {
655 for (const Edge &E : N.edges()) {
656 if (!MachineGadgetGraph::isGadgetEdge(E))
657 continue;
658
659 SmallVector<const Edge *, 2> EgressEdges;
660 SmallVector<const Edge *, 2> &IngressEdges = IngressEdgeMap[E.getDest()];
661 for (const Edge &EgressEdge : N.edges())
662 if (MachineGadgetGraph::isCFGEdge(E: EgressEdge))
663 EgressEdges.push_back(Elt: &EgressEdge);
664
665 int EgressCutCost = 0, IngressCutCost = 0;
666 for (const Edge *EgressEdge : EgressEdges)
667 if (!CutEdges.contains(E: *EgressEdge))
668 EgressCutCost += EgressEdge->getValue();
669 for (const Edge *IngressEdge : IngressEdges)
670 if (!CutEdges.contains(E: *IngressEdge))
671 IngressCutCost += IngressEdge->getValue();
672
673 auto &EdgesToCut =
674 IngressCutCost < EgressCutCost ? IngressEdges : EgressEdges;
675 for (const Edge *E : EdgesToCut)
676 CutEdges.insert(E: *E);
677 }
678 }
679 LLVM_DEBUG(dbgs() << "Cutting edges... Done\n");
680 LLVM_DEBUG(dbgs() << "Cut " << CutEdges.count() << " edges\n");
681
682 LLVM_DEBUG(dbgs() << "Inserting LFENCEs...\n");
683 int FencesInserted = insertFences(MF, G&: *Graph, CutEdges);
684 LLVM_DEBUG(dbgs() << "Inserting LFENCEs... Done\n");
685 LLVM_DEBUG(dbgs() << "Inserted " << FencesInserted << " fences\n");
686
687 return FencesInserted;
688}
689
690int X86LoadValueInjectionLoadHardeningImpl::insertFences(
691 MachineFunction &MF, MachineGadgetGraph &G,
692 EdgeSet &CutEdges /* in, out */) const {
693 int FencesInserted = 0;
694 for (const Node &N : G.nodes()) {
695 for (const Edge &E : N.edges()) {
696 if (CutEdges.contains(E)) {
697 MachineInstr *MI = N.getValue(), *Prev;
698 MachineBasicBlock *MBB; // Insert an LFENCE in this MBB
699 MachineBasicBlock::iterator InsertionPt; // ...at this point
700 if (MI == MachineGadgetGraph::ArgNodeSentinel) {
701 // insert LFENCE at beginning of entry block
702 MBB = &MF.front();
703 InsertionPt = MBB->begin();
704 Prev = nullptr;
705 } else if (MI->isBranch()) { // insert the LFENCE before the branch
706 MBB = MI->getParent();
707 InsertionPt = MI;
708 Prev = MI->getPrevNode();
709 // Remove all egress CFG edges from this branch because the inserted
710 // LFENCE prevents gadgets from crossing the branch.
711 for (const Edge &E : N.edges()) {
712 if (MachineGadgetGraph::isCFGEdge(E))
713 CutEdges.insert(E);
714 }
715 } else { // insert the LFENCE after the instruction
716 MBB = MI->getParent();
717 InsertionPt = MI->getNextNode() ? MI->getNextNode() : MBB->end();
718 Prev = InsertionPt == MBB->end()
719 ? (MBB->empty() ? nullptr : &MBB->back())
720 : InsertionPt->getPrevNode();
721 }
722 // Ensure this insertion is not redundant (two LFENCEs in sequence).
723 if ((InsertionPt == MBB->end() || !isFence(MI: &*InsertionPt)) &&
724 (!Prev || !isFence(MI: Prev))) {
725 BuildMI(BB&: *MBB, I: InsertionPt, MIMD: DebugLoc(), MCID: TII->get(Opcode: X86::LFENCE));
726 ++FencesInserted;
727 }
728 }
729 }
730 }
731 return FencesInserted;
732}
733
734bool X86LoadValueInjectionLoadHardeningImpl::instrUsesRegToAccessMemory(
735 const MachineInstr &MI, Register Reg) const {
736 if (!MI.mayLoadOrStore() || MI.getOpcode() == X86::MFENCE ||
737 MI.getOpcode() == X86::SFENCE || MI.getOpcode() == X86::LFENCE)
738 return false;
739
740 const int MemRefBeginIdx = X86::getFirstAddrOperandIdx(MI);
741 if (MemRefBeginIdx < 0) {
742 LLVM_DEBUG(dbgs() << "Warning: unable to obtain memory operand for loading "
743 "instruction:\n";
744 MI.print(dbgs()); dbgs() << '\n';);
745 return false;
746 }
747
748 const MachineOperand &BaseMO =
749 MI.getOperand(i: MemRefBeginIdx + X86::AddrBaseReg);
750 const MachineOperand &IndexMO =
751 MI.getOperand(i: MemRefBeginIdx + X86::AddrIndexReg);
752 return (BaseMO.isReg() && BaseMO.getReg().isValid() &&
753 TRI->regsOverlap(RegA: BaseMO.getReg(), RegB: Reg)) ||
754 (IndexMO.isReg() && IndexMO.getReg().isValid() &&
755 TRI->regsOverlap(RegA: IndexMO.getReg(), RegB: Reg));
756}
757
758bool X86LoadValueInjectionLoadHardeningImpl::instrUsesRegToBranch(
759 const MachineInstr &MI, Register Reg) const {
760 if (!MI.isConditionalBranch())
761 return false;
762 for (const MachineOperand &Use : MI.uses())
763 if (Use.isReg() && Use.getReg() == Reg)
764 return true;
765 return false;
766}
767
768bool X86LoadValueInjectionLoadHardeningLegacy::runOnMachineFunction(
769 MachineFunction &MF) {
770 // Don't skip functions with the "optnone" attr but participate in opt-bisect.
771 // Note: Not needed for new PM impl, where it is handled at the PM level.
772 const Function &F = MF.getFunction();
773 if (!F.hasOptNone() && skipFunction(F))
774 return false;
775
776 // Bail early (without computing analyses) if LVI load hardening is disabled.
777 if (!MF.getSubtarget<X86Subtarget>().useLVILoadHardening()) {
778 return false;
779 }
780
781 const auto &MLI = getAnalysis<MachineLoopInfoWrapperPass>().getLI();
782 const auto &MDT = getAnalysis<MachineDominatorTreeWrapperPass>().getDomTree();
783 const auto &MDF = getAnalysis<MachineDominanceFrontierWrapperPass>().getMDF();
784
785 X86LoadValueInjectionLoadHardeningImpl Impl;
786 return Impl.run(MF, MLI, MDT, MDF);
787}
788
789PreservedAnalyses X86LoadValueInjectionLoadHardeningPass::run(
790 MachineFunction &MF, MachineFunctionAnalysisManager &MFAM) {
791 // Bail early (without computing analyses) if LVI load hardening is disabled.
792 if (!MF.getSubtarget<X86Subtarget>().useLVILoadHardening()) {
793 return PreservedAnalyses::all();
794 }
795
796 const auto &MLI = MFAM.getResult<MachineLoopAnalysis>(IR&: MF);
797 const auto &MDT = MFAM.getResult<MachineDominatorTreeAnalysis>(IR&: MF);
798 const auto &MDF = MFAM.getResult<MachineDominanceFrontierAnalysis>(IR&: MF);
799
800 X86LoadValueInjectionLoadHardeningImpl Impl;
801 const bool Modified = Impl.run(MF, MLI, MDT, MDF);
802 return Modified ? getMachineFunctionPassPreservedAnalyses()
803 .preserveSet<CFGAnalyses>()
804 : PreservedAnalyses::all();
805}
806
807INITIALIZE_PASS_BEGIN(X86LoadValueInjectionLoadHardeningLegacy, PASS_KEY,
808 "X86 LVI load hardening", false, false)
809INITIALIZE_PASS_DEPENDENCY(MachineLoopInfoWrapperPass)
810INITIALIZE_PASS_DEPENDENCY(MachineDominatorTreeWrapperPass)
811INITIALIZE_PASS_DEPENDENCY(MachineDominanceFrontierWrapperPass)
812INITIALIZE_PASS_END(X86LoadValueInjectionLoadHardeningLegacy, PASS_KEY,
813 "X86 LVI load hardening", false, false)
814
815FunctionPass *llvm::createX86LoadValueInjectionLoadHardeningLegacyPass() {
816 return new X86LoadValueInjectionLoadHardeningLegacy();
817}
818