1//===-- AArch64Arm64ECCallLowering.cpp - Lower Arm64EC calls ----*- C++ -*-===//
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/// \file
10/// This file contains the IR transform to lower external or indirect calls for
11/// the ARM64EC calling convention. Such calls must go through the runtime, so
12/// we can translate the calling convention for calls into the emulator.
13///
14/// This subsumes Control Flow Guard handling.
15///
16//===----------------------------------------------------------------------===//
17
18#include "AArch64.h"
19#include "AArch64Subtarget.h"
20#include "llvm/ADT/SetVector.h"
21#include "llvm/ADT/SmallString.h"
22#include "llvm/ADT/SmallVector.h"
23#include "llvm/ADT/Statistic.h"
24#include "llvm/IR/CallingConv.h"
25#include "llvm/IR/DiagnosticInfo.h"
26#include "llvm/IR/GlobalAlias.h"
27#include "llvm/IR/IRBuilder.h"
28#include "llvm/IR/Instruction.h"
29#include "llvm/IR/Mangler.h"
30#include "llvm/IR/Module.h"
31#include "llvm/Object/COFF.h"
32#include "llvm/Pass.h"
33#include "llvm/TargetParser/Triple.h"
34
35using namespace llvm;
36using namespace llvm::COFF;
37
38using OperandBundleDef = OperandBundleDefT<Value *>;
39
40#define DEBUG_TYPE "arm64eccalllowering"
41
42STATISTIC(Arm64ECCallsLowered, "Number of Arm64EC calls lowered");
43
44namespace {
45
46enum ThunkArgTranslation : uint8_t {
47 Direct,
48 Bitcast,
49 PointerIndirection,
50};
51
52struct ThunkArgInfo {
53 Type *Arm64Ty;
54 Type *X64Ty;
55 ThunkArgTranslation Translation;
56};
57
58class AArch64Arm64ECCallLowering : public ModulePass {
59public:
60 static char ID;
61 AArch64Arm64ECCallLowering() : ModulePass(ID) {}
62
63 Function *buildExitThunk(FunctionType *FnTy, AttributeList Attrs);
64 Function *buildEntryThunk(Function *F);
65 void lowerCall(CallBase *CB);
66 Function *buildGuestExitThunk(Function *F);
67 Function *buildPatchableThunk(GlobalAlias *UnmangledAlias,
68 GlobalAlias *MangledAlias);
69 bool processFunction(Function &F, SetVector<GlobalValue *> &DirectCalledFns,
70 DenseMap<GlobalAlias *, GlobalAlias *> &FnsMap);
71 bool runOnModule(Module &M) override;
72
73private:
74 ControlFlowGuardMode CFGuardModuleFlag = ControlFlowGuardMode::Disabled;
75 FunctionType *GuardFnType = nullptr;
76 FunctionType *DispatchFnType = nullptr;
77 Constant *GuardFnCFGlobal = nullptr;
78 Constant *GuardFnGlobal = nullptr;
79 Constant *DispatchFnGlobal = nullptr;
80 Module *M = nullptr;
81
82 Type *PtrTy;
83 Type *I64Ty;
84 Type *VoidTy;
85
86 void getThunkType(FunctionType *FT, AttributeList AttrList,
87 Arm64ECThunkType TT, raw_ostream &Out,
88 FunctionType *&Arm64Ty, FunctionType *&X64Ty,
89 SmallVector<ThunkArgTranslation> &ArgTranslations);
90 void getThunkRetType(FunctionType *FT, AttributeList AttrList,
91 raw_ostream &Out, Type *&Arm64RetTy, Type *&X64RetTy,
92 SmallVectorImpl<Type *> &Arm64ArgTypes,
93 SmallVectorImpl<Type *> &X64ArgTypes,
94 SmallVector<ThunkArgTranslation> &ArgTranslations,
95 bool &HasSretPtr);
96 void getThunkArgTypes(FunctionType *FT, AttributeList AttrList,
97 Arm64ECThunkType TT, raw_ostream &Out,
98 SmallVectorImpl<Type *> &Arm64ArgTypes,
99 SmallVectorImpl<Type *> &X64ArgTypes,
100 SmallVectorImpl<ThunkArgTranslation> &ArgTranslations,
101 bool HasSretPtr);
102 ThunkArgInfo canonicalizeThunkType(Type *T, Align Alignment, bool Ret,
103 uint64_t ArgSizeBytes, raw_ostream &Out);
104};
105
106} // end anonymous namespace
107
108void AArch64Arm64ECCallLowering::getThunkType(
109 FunctionType *FT, AttributeList AttrList, Arm64ECThunkType TT,
110 raw_ostream &Out, FunctionType *&Arm64Ty, FunctionType *&X64Ty,
111 SmallVector<ThunkArgTranslation> &ArgTranslations) {
112 Out << (TT == Arm64ECThunkType::Entry ? "$ientry_thunk$cdecl$"
113 : "$iexit_thunk$cdecl$");
114
115 Type *Arm64RetTy;
116 Type *X64RetTy;
117
118 SmallVector<Type *> Arm64ArgTypes;
119 SmallVector<Type *> X64ArgTypes;
120
121 // The first argument to a thunk is the called function, stored in x9.
122 // For exit thunks, we pass the called function down to the emulator;
123 // for entry/guest exit thunks, we just call the Arm64 function directly.
124 if (TT == Arm64ECThunkType::Exit)
125 Arm64ArgTypes.push_back(Elt: PtrTy);
126 X64ArgTypes.push_back(Elt: PtrTy);
127
128 bool HasSretPtr = false;
129 getThunkRetType(FT, AttrList, Out, Arm64RetTy, X64RetTy, Arm64ArgTypes,
130 X64ArgTypes, ArgTranslations, HasSretPtr);
131
132 getThunkArgTypes(FT, AttrList, TT, Out, Arm64ArgTypes, X64ArgTypes,
133 ArgTranslations, HasSretPtr);
134
135 Arm64Ty = FunctionType::get(Result: Arm64RetTy, Params: Arm64ArgTypes, isVarArg: false);
136
137 X64Ty = FunctionType::get(Result: X64RetTy, Params: X64ArgTypes, isVarArg: false);
138}
139
140void AArch64Arm64ECCallLowering::getThunkArgTypes(
141 FunctionType *FT, AttributeList AttrList, Arm64ECThunkType TT,
142 raw_ostream &Out, SmallVectorImpl<Type *> &Arm64ArgTypes,
143 SmallVectorImpl<Type *> &X64ArgTypes,
144 SmallVectorImpl<ThunkArgTranslation> &ArgTranslations, bool HasSretPtr) {
145
146 Out << "$";
147 if (FT->isVarArg()) {
148 // We treat the variadic function's thunk as a normal function
149 // with the following type on the ARM side:
150 // rettype exitthunk(
151 // ptr x9, ptr x0, i64 x1, i64 x2, i64 x3, ptr x4, i64 x5)
152 //
153 // that can coverage all types of variadic function.
154 // x9 is similar to normal exit thunk, store the called function.
155 // x0-x3 is the arguments be stored in registers.
156 // x4 is the address of the arguments on the stack.
157 // x5 is the size of the arguments on the stack.
158 //
159 // On the x64 side, it's the same except that x5 isn't set.
160 //
161 // If both the ARM and X64 sides are sret, there are only three
162 // arguments in registers.
163 //
164 // If the X64 side is sret, but the ARM side isn't, we pass an extra value
165 // to/from the X64 side, and let SelectionDAG transform it into a memory
166 // location.
167 Out << "varargs";
168
169 // x0-x3
170 for (int i = HasSretPtr ? 1 : 0; i < 4; i++) {
171 Arm64ArgTypes.push_back(Elt: I64Ty);
172 X64ArgTypes.push_back(Elt: I64Ty);
173 ArgTranslations.push_back(Elt: ThunkArgTranslation::Direct);
174 }
175
176 // x4
177 Arm64ArgTypes.push_back(Elt: PtrTy);
178 X64ArgTypes.push_back(Elt: PtrTy);
179 ArgTranslations.push_back(Elt: ThunkArgTranslation::Direct);
180 // x5
181 Arm64ArgTypes.push_back(Elt: I64Ty);
182 if (TT != Arm64ECThunkType::Entry) {
183 // FIXME: x5 isn't actually used by the x64 side; revisit once we
184 // have proper isel for varargs
185 X64ArgTypes.push_back(Elt: I64Ty);
186 ArgTranslations.push_back(Elt: ThunkArgTranslation::Direct);
187 }
188 return;
189 }
190
191 unsigned I = 0;
192 if (HasSretPtr)
193 I++;
194
195 if (I == FT->getNumParams()) {
196 Out << "v";
197 return;
198 }
199
200 for (unsigned E = FT->getNumParams(); I != E; ++I) {
201#if 0
202 // FIXME: Need more information about argument size; see
203 // https://reviews.llvm.org/D132926
204 uint64_t ArgSizeBytes = AttrList.getParamArm64ECArgSizeBytes(I);
205 Align ParamAlign = AttrList.getParamAlignment(I).valueOrOne();
206#else
207 uint64_t ArgSizeBytes = 0;
208 Align ParamAlign = Align();
209#endif
210 auto [Arm64Ty, X64Ty, ArgTranslation] =
211 canonicalizeThunkType(T: FT->getParamType(i: I), Alignment: ParamAlign,
212 /*Ret*/ false, ArgSizeBytes, Out);
213 Arm64ArgTypes.push_back(Elt: Arm64Ty);
214 X64ArgTypes.push_back(Elt: X64Ty);
215 ArgTranslations.push_back(Elt: ArgTranslation);
216 }
217}
218
219void AArch64Arm64ECCallLowering::getThunkRetType(
220 FunctionType *FT, AttributeList AttrList, raw_ostream &Out,
221 Type *&Arm64RetTy, Type *&X64RetTy, SmallVectorImpl<Type *> &Arm64ArgTypes,
222 SmallVectorImpl<Type *> &X64ArgTypes,
223 SmallVector<ThunkArgTranslation> &ArgTranslations, bool &HasSretPtr) {
224 Type *T = FT->getReturnType();
225#if 0
226 // FIXME: Need more information about argument size; see
227 // https://reviews.llvm.org/D132926
228 uint64_t ArgSizeBytes = AttrList.getRetArm64ECArgSizeBytes();
229#else
230 int64_t ArgSizeBytes = 0;
231#endif
232 if (T->isVoidTy()) {
233 if (FT->getNumParams()) {
234 Attribute SRetAttr0 = AttrList.getParamAttr(ArgNo: 0, Kind: Attribute::StructRet);
235 Attribute InRegAttr0 = AttrList.getParamAttr(ArgNo: 0, Kind: Attribute::InReg);
236 Attribute SRetAttr1, InRegAttr1;
237 if (FT->getNumParams() > 1) {
238 // Also check the second parameter (for class methods, the first
239 // parameter is "this", and the second parameter is the sret pointer.)
240 // It doesn't matter which one is sret.
241 SRetAttr1 = AttrList.getParamAttr(ArgNo: 1, Kind: Attribute::StructRet);
242 InRegAttr1 = AttrList.getParamAttr(ArgNo: 1, Kind: Attribute::InReg);
243 }
244 if ((SRetAttr0.isValid() && InRegAttr0.isValid()) ||
245 (SRetAttr1.isValid() && InRegAttr1.isValid())) {
246 // sret+inreg indicates a call that returns a C++ class value. This is
247 // actually equivalent to just passing and returning a void* pointer
248 // as the first or second argument. Translate it that way, instead of
249 // trying to model "inreg" in the thunk's calling convention; this
250 // simplfies the rest of the code, and matches MSVC mangling.
251 Out << "i8";
252 Arm64RetTy = I64Ty;
253 X64RetTy = I64Ty;
254 return;
255 }
256 if (SRetAttr0.isValid()) {
257 // FIXME: Sanity-check the sret type; if it's an integer or pointer,
258 // we'll get screwy mangling/codegen.
259 // FIXME: For large struct types, mangle as an integer argument and
260 // integer return, so we can reuse more thunks, instead of "m" syntax.
261 // (MSVC mangles this case as an integer return with no argument, but
262 // that's a miscompile.)
263 Type *SRetType = SRetAttr0.getValueAsType();
264 Align SRetAlign = AttrList.getParamAlignment(ArgNo: 0).valueOrOne();
265 canonicalizeThunkType(T: SRetType, Alignment: SRetAlign, /*Ret*/ true, ArgSizeBytes,
266 Out);
267 Arm64RetTy = VoidTy;
268 X64RetTy = VoidTy;
269 Arm64ArgTypes.push_back(Elt: FT->getParamType(i: 0));
270 X64ArgTypes.push_back(Elt: FT->getParamType(i: 0));
271 ArgTranslations.push_back(Elt: ThunkArgTranslation::Direct);
272 HasSretPtr = true;
273 return;
274 }
275 }
276
277 Out << "v";
278 Arm64RetTy = VoidTy;
279 X64RetTy = VoidTy;
280 return;
281 }
282
283 auto info =
284 canonicalizeThunkType(T, Alignment: Align(), /*Ret*/ true, ArgSizeBytes, Out);
285 Arm64RetTy = info.Arm64Ty;
286 X64RetTy = info.X64Ty;
287 if (X64RetTy->isPointerTy()) {
288 // If the X64 type is canonicalized to a pointer, that means it's
289 // passed/returned indirectly. For a return value, that means it's an
290 // sret pointer.
291 X64ArgTypes.push_back(Elt: X64RetTy);
292 X64RetTy = VoidTy;
293 }
294}
295
296ThunkArgInfo AArch64Arm64ECCallLowering::canonicalizeThunkType(
297 Type *T, Align Alignment, bool Ret, uint64_t ArgSizeBytes,
298 raw_ostream &Out) {
299
300 auto direct = [](Type *T) {
301 return ThunkArgInfo{.Arm64Ty: T, .X64Ty: T, .Translation: ThunkArgTranslation::Direct};
302 };
303
304 auto bitcast = [this](Type *Arm64Ty, uint64_t SizeInBytes) {
305 return ThunkArgInfo{.Arm64Ty: Arm64Ty,
306 .X64Ty: llvm::Type::getIntNTy(C&: M->getContext(), N: SizeInBytes * 8),
307 .Translation: ThunkArgTranslation::Bitcast};
308 };
309
310 auto pointerIndirection = [this](Type *Arm64Ty) {
311 return ThunkArgInfo{.Arm64Ty: Arm64Ty, .X64Ty: PtrTy,
312 .Translation: ThunkArgTranslation::PointerIndirection};
313 };
314
315 if (T->isHalfTy()) {
316 // Prefix with `llvm` since MSVC doesn't specify `_Float16`
317 Out << "__llvm_h__";
318 return direct(T);
319 }
320
321 if (T->isBFloatTy()) {
322 // Prefix with `llvm` since MSVC doesn't specify `__bf16`
323 Out << "__llvm_bf16__";
324 return direct(T);
325 }
326
327 if (T->isFloatTy()) {
328 Out << "f";
329 return direct(T);
330 }
331
332 if (T->isDoubleTy()) {
333 Out << "d";
334 return direct(T);
335 }
336
337 if (T->isFP128Ty()) {
338 // Prefix with `llvm` since MSVC doesn't specify `_Float128`
339 Out << "__llvm_q__";
340 // On windows f128 is passed indirectly, and Clang/LLVM
341 // returns using sret for compatibility with GCC.
342 return pointerIndirection(T);
343 }
344
345 if (T->isFloatingPointTy()) {
346 report_fatal_error(
347 reason: "Only half, bfloat16, float, double, and fp128 are supported "
348 "for ARM64EC thunks");
349 }
350
351 auto &DL = M->getDataLayout();
352
353 if (auto *StructTy = dyn_cast<StructType>(Val: T))
354 if (StructTy->getNumElements() == 1)
355 T = StructTy->getElementType(N: 0);
356
357 if (T->isArrayTy()) {
358 Type *ElementTy = T->getArrayElementType();
359 uint64_t ElementCnt = T->getArrayNumElements();
360 uint64_t ElementSizePerBytes = DL.getTypeSizeInBits(Ty: ElementTy) / 8;
361 uint64_t TotalSizeBytes = ElementCnt * ElementSizePerBytes;
362 if (ElementTy->isHalfTy() || ElementTy->isBFloatTy() ||
363 ElementTy->isFloatTy() || ElementTy->isDoubleTy() ||
364 ElementTy->isFP128Ty()) {
365 if (ElementTy->isHalfTy())
366 // Prefix with `llvm` since MSVC doesn't specify `_Float16`
367 Out << "__llvm_H__";
368 else if (ElementTy->isBFloatTy())
369 // Prefix with `llvm` since MSVC doesn't specify `__bf16`
370 Out << "__llvm_BF16__";
371 else if (ElementTy->isFloatTy())
372 Out << "F";
373 else if (ElementTy->isDoubleTy())
374 Out << "D";
375 else if (ElementTy->isFP128Ty())
376 // Prefix with `llvm` since MSVC doesn't specify `_Float128`
377 Out << "__llvm_Q__";
378 Out << TotalSizeBytes;
379 if (Alignment.value() >= 16 && !Ret)
380 Out << "a" << Alignment.value();
381 if (TotalSizeBytes <= 8) {
382 // Arm64 returns small structs of float/double in float registers;
383 // X64 uses RAX.
384 return bitcast(T, TotalSizeBytes);
385 } else {
386 // Struct is passed directly on Arm64, but indirectly on X64.
387 return pointerIndirection(T);
388 }
389 } else if (ElementTy->isFloatingPointTy()) {
390 report_fatal_error(
391 reason: "Only half, bfloat16, float, double, and fp128 are supported "
392 "for ARM64EC thunks");
393 }
394 }
395
396 if ((T->isIntegerTy() || T->isPointerTy()) && DL.getTypeSizeInBits(Ty: T) <= 64) {
397 Out << "i8";
398 return direct(I64Ty);
399 }
400
401 unsigned TypeSize = ArgSizeBytes;
402 if (TypeSize == 0)
403 TypeSize = DL.getTypeSizeInBits(Ty: T) / 8;
404 Out << "m";
405 if (TypeSize != 4)
406 Out << TypeSize;
407 if (Alignment.value() >= 16 && !Ret)
408 Out << "a" << Alignment.value();
409 // FIXME: Try to canonicalize Arm64Ty more thoroughly?
410 if (TypeSize == 1 || TypeSize == 2 || TypeSize == 4 || TypeSize == 8) {
411 // Pass directly in an integer register
412 return bitcast(T, TypeSize);
413 } else {
414 // Passed directly on Arm64, but indirectly on X64.
415 return pointerIndirection(T);
416 }
417}
418
419// This function builds the "exit thunk", a function which translates
420// arguments and return values when calling x64 code from AArch64 code.
421Function *AArch64Arm64ECCallLowering::buildExitThunk(FunctionType *FT,
422 AttributeList Attrs) {
423 SmallString<256> ExitThunkName;
424 llvm::raw_svector_ostream ExitThunkStream(ExitThunkName);
425 FunctionType *Arm64Ty, *X64Ty;
426 SmallVector<ThunkArgTranslation> ArgTranslations;
427 getThunkType(FT, AttrList: Attrs, TT: Arm64ECThunkType::Exit, Out&: ExitThunkStream, Arm64Ty,
428 X64Ty, ArgTranslations);
429 if (Function *F = M->getFunction(Name: ExitThunkName))
430 return F;
431
432 Function *F = Function::Create(Ty: Arm64Ty, Linkage: GlobalValue::LinkOnceODRLinkage, AddrSpace: 0,
433 N: ExitThunkName, M);
434 F->setCallingConv(CallingConv::ARM64EC_Thunk_Native);
435 F->setSection(".wowthk$aa");
436 F->setComdat(M->getOrInsertComdat(Name: ExitThunkName));
437 // Copy MSVC, and always set up a frame pointer. (Maybe this isn't necessary.)
438 F->addFnAttr(Kind: "frame-pointer", Val: "all");
439 // Only copy sret from the first argument. For C++ instance methods, clang can
440 // stick an sret marking on a later argument, but it doesn't actually affect
441 // the ABI, so we can omit it. This avoids triggering a verifier assertion.
442 if (FT->getNumParams()) {
443 auto SRet = Attrs.getParamAttr(ArgNo: 0, Kind: Attribute::StructRet);
444 auto InReg = Attrs.getParamAttr(ArgNo: 0, Kind: Attribute::InReg);
445 if (SRet.isValid() && !InReg.isValid())
446 F->addParamAttr(ArgNo: 1, Attr: SRet);
447 }
448 // FIXME: Copy anything other than sret? Shouldn't be necessary for normal
449 // C ABI, but might show up in other cases.
450 BasicBlock *BB = BasicBlock::Create(Context&: M->getContext(), Name: "", Parent: F);
451 IRBuilder<> IRB(BB);
452 Value *CalleePtr =
453 M->getOrInsertGlobal(Name: "__os_arm64x_dispatch_call_no_redirect", Ty: PtrTy);
454 Value *Callee = IRB.CreateLoad(Ty: PtrTy, Ptr: CalleePtr);
455 auto &DL = M->getDataLayout();
456 SmallVector<Value *> Args;
457 FunctionType *DispatcherCallTy = X64Ty;
458 // If we have a vararg function, the SelectionDAG lowering will need to
459 // recognize this so it can copy the arguments described by x4 (pointer) and
460 // x5 (length) to set up the x86-64 context correctly.
461 if (FT->isVarArg())
462 DispatcherCallTy =
463 FunctionType::get(Result: X64Ty->getReturnType(), Params: X64Ty->params(),
464 /*isVarArg=*/true);
465
466 // Pass the called function in x9.
467 auto X64TyOffset = 1;
468 Args.push_back(Elt: F->arg_begin());
469
470 Type *RetTy = Arm64Ty->getReturnType();
471 if (RetTy != X64Ty->getReturnType()) {
472 // If the return type is an array or struct, translate it. Values of size
473 // 8 or less go into RAX; bigger values go into memory, and we pass a
474 // pointer.
475 if (DL.getTypeStoreSize(Ty: RetTy) > 8) {
476 Args.push_back(Elt: IRB.CreateAlloca(Ty: RetTy));
477 X64TyOffset++;
478 }
479 }
480
481 for (auto [Arg, X64ArgType, ArgTranslation] : llvm::zip_equal(
482 t: make_range(x: F->arg_begin() + 1, y: F->arg_end()),
483 u: make_range(x: X64Ty->param_begin() + X64TyOffset, y: X64Ty->param_end()),
484 args&: ArgTranslations)) {
485 // Translate arguments from AArch64 calling convention to x86 calling
486 // convention.
487 //
488 // For simple types, we don't need to do any translation: they're
489 // represented the same way. (Implicit sign extension is not part of
490 // either convention.)
491 //
492 // The big thing we have to worry about is struct types... but
493 // fortunately AArch64 clang is pretty friendly here: the cases that need
494 // translation are always passed as a struct or array. (If we run into
495 // some cases where this doesn't work, we can teach clang to mark it up
496 // with an attribute.)
497 //
498 // The first argument is the called function, stored in x9.
499 if (ArgTranslation != ThunkArgTranslation::Direct) {
500 Value *Mem = IRB.CreateAlloca(Ty: Arg.getType());
501 IRB.CreateStore(Val: &Arg, Ptr: Mem);
502 if (ArgTranslation == ThunkArgTranslation::Bitcast) {
503 Type *IntTy = IRB.getIntNTy(N: DL.getTypeStoreSizeInBits(Ty: Arg.getType()));
504 Args.push_back(Elt: IRB.CreateLoad(Ty: IntTy, Ptr: Mem));
505 } else {
506 assert(ArgTranslation == ThunkArgTranslation::PointerIndirection);
507 Args.push_back(Elt: Mem);
508 }
509 } else {
510 Args.push_back(Elt: &Arg);
511 }
512 assert(Args.back()->getType() == X64ArgType);
513 }
514 // FIXME: Transfer necessary attributes? sret? anything else?
515
516 CallInst *Call = IRB.CreateCall(FTy: DispatcherCallTy, Callee, Args);
517 Call->setCallingConv(CallingConv::ARM64EC_Thunk_X64);
518
519 Value *RetVal = Call;
520 if (RetTy != X64Ty->getReturnType()) {
521 // If we rewrote the return type earlier, convert the return value to
522 // the proper type.
523 if (DL.getTypeStoreSize(Ty: RetTy) > 8) {
524 RetVal = IRB.CreateLoad(Ty: RetTy, Ptr: Args[1]);
525 } else {
526 Value *CastAlloca = IRB.CreateAlloca(Ty: RetTy);
527 IRB.CreateStore(Val: Call, Ptr: CastAlloca);
528 RetVal = IRB.CreateLoad(Ty: RetTy, Ptr: CastAlloca);
529 }
530 }
531
532 if (RetTy->isVoidTy())
533 IRB.CreateRetVoid();
534 else
535 IRB.CreateRet(V: RetVal);
536 return F;
537}
538
539// This function builds the "entry thunk", a function which translates
540// arguments and return values when calling AArch64 code from x64 code.
541Function *AArch64Arm64ECCallLowering::buildEntryThunk(Function *F) {
542 SmallString<256> EntryThunkName;
543 llvm::raw_svector_ostream EntryThunkStream(EntryThunkName);
544 FunctionType *Arm64Ty, *X64Ty;
545 SmallVector<ThunkArgTranslation> ArgTranslations;
546 getThunkType(FT: F->getFunctionType(), AttrList: F->getAttributes(),
547 TT: Arm64ECThunkType::Entry, Out&: EntryThunkStream, Arm64Ty, X64Ty,
548 ArgTranslations);
549 if (Function *F = M->getFunction(Name: EntryThunkName))
550 return F;
551
552 Function *Thunk = Function::Create(Ty: X64Ty, Linkage: GlobalValue::LinkOnceODRLinkage, AddrSpace: 0,
553 N: EntryThunkName, M);
554 Thunk->setCallingConv(CallingConv::ARM64EC_Thunk_X64);
555 Thunk->setSection(".wowthk$aa");
556 Thunk->setComdat(M->getOrInsertComdat(Name: EntryThunkName));
557 // Copy MSVC, and always set up a frame pointer. (Maybe this isn't necessary.)
558 Thunk->addFnAttr(Kind: "frame-pointer", Val: "all");
559
560 BasicBlock *BB = BasicBlock::Create(Context&: M->getContext(), Name: "", Parent: Thunk);
561 IRBuilder<> IRB(BB);
562
563 Type *RetTy = Arm64Ty->getReturnType();
564 Type *X64RetType = X64Ty->getReturnType();
565
566 bool TransformDirectToSRet = X64RetType->isVoidTy() && !RetTy->isVoidTy();
567 unsigned ThunkArgOffset = TransformDirectToSRet ? 2 : 1;
568 unsigned PassthroughArgSize =
569 (F->isVarArg() ? 5 : Thunk->arg_size()) - ThunkArgOffset;
570 assert(ArgTranslations.size() == (F->isVarArg() ? 5 : PassthroughArgSize));
571
572 // Translate arguments to call.
573 SmallVector<Value *> Args;
574 for (unsigned i = 0; i != PassthroughArgSize; ++i) {
575 Value *Arg = Thunk->getArg(i: i + ThunkArgOffset);
576 Type *ArgTy = Arm64Ty->getParamType(i);
577 ThunkArgTranslation ArgTranslation = ArgTranslations[i];
578 if (ArgTranslation != ThunkArgTranslation::Direct) {
579 // Translate array/struct arguments to the expected type.
580 if (ArgTranslation == ThunkArgTranslation::Bitcast) {
581 Value *CastAlloca = IRB.CreateAlloca(Ty: ArgTy);
582 IRB.CreateStore(Val: Arg, Ptr: CastAlloca);
583 Arg = IRB.CreateLoad(Ty: ArgTy, Ptr: CastAlloca);
584 } else {
585 assert(ArgTranslation == ThunkArgTranslation::PointerIndirection);
586 Arg = IRB.CreateLoad(Ty: ArgTy, Ptr: Arg);
587 }
588 }
589 assert(Arg->getType() == ArgTy);
590 Args.push_back(Elt: Arg);
591 }
592
593 if (F->isVarArg()) {
594 // The 5th argument to variadic entry thunks is used to model the x64 sp
595 // which is passed to the thunk in x4, this can be passed to the callee as
596 // the variadic argument start address after skipping over the 32 byte
597 // shadow store.
598
599 // The EC thunk CC will assign any argument marked as InReg to x4.
600 Thunk->addParamAttr(ArgNo: 5, Kind: Attribute::InReg);
601 Value *Arg = Thunk->getArg(i: 5);
602 Arg = IRB.CreatePtrAdd(Ptr: Arg, Offset: IRB.getInt64(C: 0x20));
603 Args.push_back(Elt: Arg);
604
605 // Pass in a zero variadic argument size (in x5).
606 Args.push_back(Elt: IRB.getInt64(C: 0));
607 }
608
609 // Call the function passed to the thunk.
610 Value *Callee = Thunk->getArg(i: 0);
611 CallInst *Call = IRB.CreateCall(FTy: Arm64Ty, Callee, Args);
612
613 auto SRetAttr = F->getAttributes().getParamAttr(ArgNo: 0, Kind: Attribute::StructRet);
614 auto InRegAttr = F->getAttributes().getParamAttr(ArgNo: 0, Kind: Attribute::InReg);
615 if (SRetAttr.isValid() && !InRegAttr.isValid()) {
616 Thunk->addParamAttr(ArgNo: 1, Attr: SRetAttr);
617 Call->addParamAttr(ArgNo: 0, Attr: SRetAttr);
618 }
619
620 Value *RetVal = Call;
621 if (TransformDirectToSRet) {
622 // The x64 side returns this value indirectly via a hidden pointer (sret).
623 // Mark the thunk's pointer arg with sret so that ISel saves it and copies
624 // it into x8 (RAX) on return, matching the x64 calling convention.
625 Thunk->addParamAttr(
626 ArgNo: 1, Attr: Attribute::getWithStructRetType(Context&: M->getContext(), Ty: RetTy));
627 IRB.CreateStore(Val: RetVal, Ptr: Thunk->getArg(i: 1));
628 } else if (X64RetType != RetTy) {
629 Value *CastAlloca = IRB.CreateAlloca(Ty: X64RetType);
630 IRB.CreateStore(Val: Call, Ptr: CastAlloca);
631 RetVal = IRB.CreateLoad(Ty: X64RetType, Ptr: CastAlloca);
632 }
633
634 // Return to the caller. Note that the isel has code to translate this
635 // "ret" to a tail call to __os_arm64x_dispatch_ret. (Alternatively, we
636 // could emit a tail call here, but that would require a dedicated calling
637 // convention, which seems more complicated overall.)
638 if (X64RetType->isVoidTy())
639 IRB.CreateRetVoid();
640 else
641 IRB.CreateRet(V: RetVal);
642
643 return Thunk;
644}
645
646std::optional<std::string> getArm64ECMangledFunctionName(GlobalValue &GV) {
647 if (!GV.hasName()) {
648 GV.setName("__unnamed");
649 }
650
651 return llvm::getArm64ECMangledFunctionName(Name: GV.getName());
652}
653
654// Builds the "guest exit thunk", a helper to call a function which may or may
655// not be an exit thunk. (We optimistically assume non-dllimport function
656// declarations refer to functions defined in AArch64 code; if the linker
657// can't prove that, we use this routine instead.)
658Function *AArch64Arm64ECCallLowering::buildGuestExitThunk(Function *F) {
659 llvm::raw_null_ostream NullThunkName;
660 FunctionType *Arm64Ty, *X64Ty;
661 SmallVector<ThunkArgTranslation> ArgTranslations;
662 getThunkType(FT: F->getFunctionType(), AttrList: F->getAttributes(),
663 TT: Arm64ECThunkType::GuestExit, Out&: NullThunkName, Arm64Ty, X64Ty,
664 ArgTranslations);
665 auto MangledName = getArm64ECMangledFunctionName(GV&: *F);
666 assert(MangledName && "Can't guest exit to function that's already native");
667 std::string ThunkName = *MangledName;
668 if (ThunkName[0] == '?' && ThunkName.find(s: "@") != std::string::npos) {
669 ThunkName.insert(pos: ThunkName.find(s: "@"), s: "$exit_thunk");
670 } else {
671 ThunkName.append(s: "$exit_thunk");
672 }
673 Function *GuestExit =
674 Function::Create(Ty: Arm64Ty, Linkage: GlobalValue::WeakODRLinkage, AddrSpace: 0, N: ThunkName, M);
675 GuestExit->setComdat(M->getOrInsertComdat(Name: ThunkName));
676 GuestExit->setSection(".wowthk$aa");
677 GuestExit->addMetadata(
678 Kind: "arm64ec_unmangled_name",
679 MD&: *MDNode::get(Context&: M->getContext(),
680 MDs: MDString::get(Context&: M->getContext(), Str: F->getName())));
681 GuestExit->setMetadata(
682 Kind: "arm64ec_ecmangled_name",
683 Node: MDNode::get(Context&: M->getContext(),
684 MDs: MDString::get(Context&: M->getContext(), Str: *MangledName)));
685 F->setMetadata(Kind: "arm64ec_hasguestexit", Node: MDNode::get(Context&: M->getContext(), MDs: {}));
686 BasicBlock *BB = BasicBlock::Create(Context&: M->getContext(), Name: "", Parent: GuestExit);
687 IRBuilder<> B(BB);
688
689 // Create new call instruction. The call check should always be a call,
690 // even if the original CallBase is an Invoke or CallBr instructio.
691 // This is treated as a direct call, so do not use GuardFnCFGlobal.
692 LoadInst *GuardCheckLoad = B.CreateLoad(Ty: PtrTy, Ptr: GuardFnGlobal);
693 Function *Thunk = buildExitThunk(FT: F->getFunctionType(), Attrs: F->getAttributes());
694 CallInst *GuardCheck = B.CreateCall(
695 FTy: GuardFnType, Callee: GuardCheckLoad, Args: {F, Thunk});
696 Value *GuardCheckDest = B.CreateExtractValue(Agg: GuardCheck, Idxs: 0);
697 Value *GuardFinalDest = B.CreateExtractValue(Agg: GuardCheck, Idxs: 1);
698
699 // Ensure that the first argument is passed in the correct register.
700 GuardCheck->setCallingConv(CallingConv::CFGuard_Check);
701
702 SmallVector<Value *> Args(llvm::make_pointer_range(Range: GuestExit->args()));
703 OperandBundleDef OB("cfguardtarget", GuardFinalDest);
704 CallInst *Call = B.CreateCall(FTy: Arm64Ty, Callee: GuardCheckDest, Args, OpBundles: OB);
705 Call->setTailCallKind(llvm::CallInst::TCK_MustTail);
706
707 if (Call->getType()->isVoidTy())
708 B.CreateRetVoid();
709 else
710 B.CreateRet(V: Call);
711
712 auto SRetAttr = F->getAttributes().getParamAttr(ArgNo: 0, Kind: Attribute::StructRet);
713 auto InRegAttr = F->getAttributes().getParamAttr(ArgNo: 0, Kind: Attribute::InReg);
714 if (SRetAttr.isValid() && !InRegAttr.isValid()) {
715 GuestExit->addParamAttr(ArgNo: 0, Attr: SRetAttr);
716 Call->addParamAttr(ArgNo: 0, Attr: SRetAttr);
717 }
718
719 return GuestExit;
720}
721
722Function *
723AArch64Arm64ECCallLowering::buildPatchableThunk(GlobalAlias *UnmangledAlias,
724 GlobalAlias *MangledAlias) {
725 llvm::raw_null_ostream NullThunkName;
726 FunctionType *Arm64Ty, *X64Ty;
727 Function *F = cast<Function>(Val: MangledAlias->getAliasee());
728 SmallVector<ThunkArgTranslation> ArgTranslations;
729 getThunkType(FT: F->getFunctionType(), AttrList: F->getAttributes(),
730 TT: Arm64ECThunkType::GuestExit, Out&: NullThunkName, Arm64Ty, X64Ty,
731 ArgTranslations);
732 std::string ThunkName(MangledAlias->getName());
733 if (ThunkName[0] == '?' && ThunkName.find(s: "@") != std::string::npos) {
734 ThunkName.insert(pos: ThunkName.find(s: "@"), s: "$hybpatch_thunk");
735 } else {
736 ThunkName.append(s: "$hybpatch_thunk");
737 }
738
739 Function *GuestExit =
740 Function::Create(Ty: Arm64Ty, Linkage: GlobalValue::WeakODRLinkage, AddrSpace: 0, N: ThunkName, M);
741 GuestExit->setComdat(M->getOrInsertComdat(Name: ThunkName));
742 GuestExit->setSection(".wowthk$aa");
743 BasicBlock *BB = BasicBlock::Create(Context&: M->getContext(), Name: "", Parent: GuestExit);
744 IRBuilder<> B(BB);
745
746 // Load the global symbol as a pointer to the check function.
747 LoadInst *DispatchLoad = B.CreateLoad(Ty: PtrTy, Ptr: DispatchFnGlobal);
748
749 // Create new dispatch call instruction.
750 Function *ExitThunk =
751 buildExitThunk(FT: F->getFunctionType(), Attrs: F->getAttributes());
752 CallInst *Dispatch =
753 B.CreateCall(FTy: DispatchFnType, Callee: DispatchLoad,
754 Args: {UnmangledAlias, ExitThunk, UnmangledAlias->getAliasee()});
755
756 // Ensure that the first arguments are passed in the correct registers.
757 Dispatch->setCallingConv(CallingConv::CFGuard_Check);
758
759 SmallVector<Value *> Args(llvm::make_pointer_range(Range: GuestExit->args()));
760 CallInst *Call = B.CreateCall(FTy: Arm64Ty, Callee: Dispatch, Args);
761 Call->setTailCallKind(llvm::CallInst::TCK_MustTail);
762
763 if (Call->getType()->isVoidTy())
764 B.CreateRetVoid();
765 else
766 B.CreateRet(V: Call);
767
768 auto SRetAttr = F->getAttributes().getParamAttr(ArgNo: 0, Kind: Attribute::StructRet);
769 auto InRegAttr = F->getAttributes().getParamAttr(ArgNo: 0, Kind: Attribute::InReg);
770 if (SRetAttr.isValid() && !InRegAttr.isValid()) {
771 GuestExit->addParamAttr(ArgNo: 0, Attr: SRetAttr);
772 Call->addParamAttr(ArgNo: 0, Attr: SRetAttr);
773 }
774
775 MangledAlias->setAliasee(GuestExit);
776 return GuestExit;
777}
778
779// Lower an indirect call with inline code.
780void AArch64Arm64ECCallLowering::lowerCall(CallBase *CB) {
781 IRBuilder<> B(CB);
782 Value *CalledOperand = CB->getCalledOperand();
783
784 // If the indirect call is called within catchpad or cleanuppad,
785 // we need to copy "funclet" bundle of the call.
786 SmallVector<llvm::OperandBundleDef, 1> Bundles;
787 if (auto Bundle = CB->getOperandBundle(ID: LLVMContext::OB_funclet))
788 Bundles.push_back(Elt: OperandBundleDef(*Bundle));
789
790 // Load the global symbol as a pointer to the check function.
791 Value *GuardFn;
792 if ((CFGuardModuleFlag == ControlFlowGuardMode::Enabled) &&
793 !CB->hasFnAttr(Kind: "guard_nocf"))
794 GuardFn = GuardFnCFGlobal;
795 else
796 GuardFn = GuardFnGlobal;
797 LoadInst *GuardCheckLoad = B.CreateLoad(Ty: PtrTy, Ptr: GuardFn);
798
799 // Create new call instruction. The CFGuard check should always be a call,
800 // even if the original CallBase is an Invoke or CallBr instruction.
801 Function *Thunk = buildExitThunk(FT: CB->getFunctionType(), Attrs: CB->getAttributes());
802 CallInst *GuardCheck =
803 B.CreateCall(FTy: GuardFnType, Callee: GuardCheckLoad, Args: {CalledOperand, Thunk},
804 OpBundles: Bundles);
805 Value *GuardCheckDest = B.CreateExtractValue(Agg: GuardCheck, Idxs: 0);
806 Value *GuardFinalDest = B.CreateExtractValue(Agg: GuardCheck, Idxs: 1);
807
808 // Ensure that the first argument is passed in the correct register.
809 GuardCheck->setCallingConv(CallingConv::CFGuard_Check);
810
811 // Update the call: set the callee, and add a bundle with the final
812 // destination,
813 CB->setCalledOperand(GuardCheckDest);
814 OperandBundleDef OB("cfguardtarget", GuardFinalDest);
815 auto *NewCall = CallBase::addOperandBundle(CB, ID: LLVMContext::OB_cfguardtarget,
816 OB, InsertPt: CB->getIterator());
817 NewCall->copyMetadata(SrcInst: *CB);
818 CB->replaceAllUsesWith(V: NewCall);
819 CB->eraseFromParent();
820}
821
822bool AArch64Arm64ECCallLowering::runOnModule(Module &Mod) {
823 if (!AArch64Options::Global.arm64ec_generate_thunks)
824 return false;
825
826 M = &Mod;
827
828 // Check if this module has the cfguard flag and read its value.
829 CFGuardModuleFlag = M->getControlFlowGuardMode();
830
831 // Warn if the module flag requests an unsupported CFGuard mechanism.
832 if (CFGuardModuleFlag == ControlFlowGuardMode::Enabled) {
833 if (auto *CI = mdconst::dyn_extract_or_null<ConstantInt>(
834 MD: Mod.getModuleFlag(Key: "cfguard-mechanism"))) {
835 auto MechanismOverride =
836 static_cast<ControlFlowGuardMechanism>(CI->getZExtValue());
837 if (MechanismOverride != ControlFlowGuardMechanism::Automatic &&
838 MechanismOverride != ControlFlowGuardMechanism::Check)
839 Mod.getContext().diagnose(
840 DI: DiagnosticInfoGeneric("only the Check Control Flow Guard mechanism "
841 "is supported for Arm64EC",
842 DS_Warning));
843 }
844 }
845
846 PtrTy = PointerType::getUnqual(C&: M->getContext());
847 I64Ty = Type::getInt64Ty(C&: M->getContext());
848 VoidTy = Type::getVoidTy(C&: M->getContext());
849
850 GuardFnType =
851 FunctionType::get(Result: StructType::get(elt1: PtrTy, elts: PtrTy), Params: {PtrTy, PtrTy}, isVarArg: false);
852 DispatchFnType = FunctionType::get(Result: PtrTy, Params: {PtrTy, PtrTy, PtrTy}, isVarArg: false);
853 GuardFnCFGlobal = M->getOrInsertGlobal(Name: "__os_arm64x_check_icall_cfg", Ty: PtrTy);
854 GuardFnGlobal = M->getOrInsertGlobal(Name: "__os_arm64x_check_icall", Ty: PtrTy);
855 DispatchFnGlobal = M->getOrInsertGlobal(Name: "__os_arm64x_dispatch_call", Ty: PtrTy);
856
857 // Mangle names of function aliases and add the alias name to
858 // arm64ec_unmangled_name metadata to ensure a weak anti-dependency symbol is
859 // emitted for the alias as well. Do this early, before handling
860 // hybrid_patchable functions, to avoid mangling their aliases.
861 for (GlobalAlias &A : Mod.aliases()) {
862 auto F = dyn_cast_or_null<Function>(Val: A.getAliaseeObject());
863 if (!F)
864 continue;
865 if (std::optional<std::string> MangledName =
866 getArm64ECMangledFunctionName(GV&: A)) {
867 F->addMetadata(Kind: "arm64ec_unmangled_name",
868 MD&: *MDNode::get(Context&: M->getContext(),
869 MDs: MDString::get(Context&: M->getContext(), Str: A.getName())));
870 A.setName(MangledName.value());
871 }
872 }
873
874 DenseMap<GlobalAlias *, GlobalAlias *> FnsMap;
875 SetVector<GlobalAlias *> PatchableFns;
876
877 for (Function &F : Mod) {
878 if (F.hasPersonalityFn()) {
879 GlobalValue *PersFn =
880 cast<GlobalValue>(Val: F.getPersonalityFn()->stripPointerCasts());
881 if (PersFn->getValueType() && PersFn->getValueType()->isFunctionTy()) {
882 if (std::optional<std::string> MangledName =
883 getArm64ECMangledFunctionName(GV&: *PersFn)) {
884 PersFn->setName(MangledName.value());
885 }
886 }
887 }
888
889 if (!F.hasFnAttribute(Kind: Attribute::HybridPatchable) ||
890 F.isDeclarationForLinker() || F.hasLocalLinkage() ||
891 F.getName().ends_with(Suffix: HybridPatchableTargetSuffix))
892 continue;
893
894 // Rename hybrid patchable functions and change callers to use a global
895 // alias instead.
896 if (std::optional<std::string> MangledName =
897 getArm64ECMangledFunctionName(GV&: F)) {
898 std::string OrigName(F.getName());
899 F.setName(MangledName.value() + HybridPatchableTargetSuffix);
900
901 // The unmangled symbol is a weak alias to an undefined symbol with the
902 // "EXP+" prefix. This undefined symbol is resolved by the linker by
903 // creating an x86 thunk that jumps back to the actual EC target. Since we
904 // can't represent that in IR, we create an alias to the target instead.
905 // The "EXP+" symbol is set as metadata, which is then used by
906 // emitGlobalAlias to emit the right alias.
907 auto *A =
908 GlobalAlias::create(Linkage: GlobalValue::LinkOnceODRLinkage, Name: OrigName, Aliasee: &F);
909 auto *AM = GlobalAlias::create(Linkage: GlobalValue::LinkOnceODRLinkage,
910 Name: MangledName.value(), Aliasee: &F);
911 F.replaceUsesWithIf(New: AM,
912 ShouldReplace: [](Use &U) { return isa<GlobalAlias>(Val: U.getUser()); });
913 F.replaceAllUsesWith(V: A);
914 F.setMetadata(Kind: "arm64ec_exp_name",
915 Node: MDNode::get(Context&: M->getContext(),
916 MDs: MDString::get(Context&: M->getContext(),
917 Str: "EXP+" + MangledName.value())));
918 A->setAliasee(&F);
919 AM->setAliasee(&F);
920
921 if (F.hasDLLExportStorageClass()) {
922 A->setDLLStorageClass(GlobalValue::DLLExportStorageClass);
923 F.setDLLStorageClass(GlobalValue::DefaultStorageClass);
924 }
925
926 FnsMap[A] = AM;
927 PatchableFns.insert(X: A);
928 }
929 }
930
931 SetVector<GlobalValue *> DirectCalledFns;
932 for (Function &F : Mod)
933 if (!F.isDeclarationForLinker() &&
934 F.getCallingConv() != CallingConv::ARM64EC_Thunk_Native &&
935 F.getCallingConv() != CallingConv::ARM64EC_Thunk_X64)
936 processFunction(F, DirectCalledFns, FnsMap);
937
938 struct ThunkInfo {
939 Constant *Src;
940 Constant *Dst;
941 Arm64ECThunkType Kind;
942 };
943 SmallVector<ThunkInfo> ThunkMapping;
944 for (Function &F : Mod) {
945 if (!F.isDeclarationForLinker() &&
946 (!F.hasLocalLinkage() || F.hasAddressTaken()) &&
947 F.getCallingConv() != CallingConv::ARM64EC_Thunk_Native &&
948 F.getCallingConv() != CallingConv::ARM64EC_Thunk_X64) {
949 if (!F.hasComdat())
950 F.setComdat(Mod.getOrInsertComdat(Name: F.getName()));
951 ThunkMapping.push_back(
952 Elt: {.Src: &F, .Dst: buildEntryThunk(F: &F), .Kind: Arm64ECThunkType::Entry});
953 }
954 }
955 for (GlobalValue *O : DirectCalledFns) {
956 auto GA = dyn_cast<GlobalAlias>(Val: O);
957 auto F = dyn_cast<Function>(Val: GA ? GA->getAliasee() : O);
958 ThunkMapping.push_back(
959 Elt: {.Src: O, .Dst: buildExitThunk(FT: F->getFunctionType(), Attrs: F->getAttributes()),
960 .Kind: Arm64ECThunkType::Exit});
961 if (!GA && !F->hasDLLImportStorageClass())
962 ThunkMapping.push_back(
963 Elt: {.Src: buildGuestExitThunk(F), .Dst: F, .Kind: Arm64ECThunkType::GuestExit});
964 }
965 for (GlobalAlias *A : PatchableFns) {
966 Function *Thunk = buildPatchableThunk(UnmangledAlias: A, MangledAlias: FnsMap[A]);
967 ThunkMapping.push_back(Elt: {.Src: Thunk, .Dst: A, .Kind: Arm64ECThunkType::GuestExit});
968 }
969
970 if (!ThunkMapping.empty()) {
971 SmallVector<Constant *> ThunkMappingArrayElems;
972 for (ThunkInfo &Thunk : ThunkMapping) {
973 ThunkMappingArrayElems.push_back(Elt: ConstantStruct::getAnon(
974 V: {Thunk.Src, Thunk.Dst,
975 ConstantInt::get(Context&: M->getContext(), V: APInt(32, uint8_t(Thunk.Kind)))}));
976 }
977 Constant *ThunkMappingArray = ConstantArray::get(
978 T: llvm::ArrayType::get(ElementType: ThunkMappingArrayElems[0]->getType(),
979 NumElements: ThunkMappingArrayElems.size()),
980 V: ThunkMappingArrayElems);
981 new GlobalVariable(Mod, ThunkMappingArray->getType(), /*isConstant*/ false,
982 GlobalValue::ExternalLinkage, ThunkMappingArray,
983 "llvm.arm64ec.symbolmap");
984 }
985
986 return true;
987}
988
989bool AArch64Arm64ECCallLowering::processFunction(
990 Function &F, SetVector<GlobalValue *> &DirectCalledFns,
991 DenseMap<GlobalAlias *, GlobalAlias *> &FnsMap) {
992 SmallVector<CallBase *, 8> IndirectCalls;
993
994 // For ARM64EC targets, a function definition's name is mangled differently
995 // from the normal symbol. We currently have no representation of this sort
996 // of symbol in IR, so we change the name to the mangled name, then store
997 // the unmangled name as metadata. Later passes that need the unmangled
998 // name (emitting the definition) can grab it from the metadata.
999 //
1000 // FIXME: Handle functions with weak linkage?
1001 if (!F.hasLocalLinkage() || F.hasAddressTaken()) {
1002 if (std::optional<std::string> MangledName =
1003 getArm64ECMangledFunctionName(GV&: F)) {
1004 F.addMetadata(Kind: "arm64ec_unmangled_name",
1005 MD&: *MDNode::get(Context&: M->getContext(),
1006 MDs: MDString::get(Context&: M->getContext(), Str: F.getName())));
1007 if (F.hasComdat() && F.getComdat()->getName() == F.getName()) {
1008 Comdat *MangledComdat = M->getOrInsertComdat(Name: MangledName.value());
1009 SmallVector<GlobalObject *> ComdatUsers =
1010 to_vector(Range: F.getComdat()->getUsers());
1011 for (GlobalObject *User : ComdatUsers)
1012 User->setComdat(MangledComdat);
1013 }
1014 F.setName(MangledName.value());
1015 }
1016 }
1017
1018 // Iterate over the instructions to find all indirect call/invoke/callbr
1019 // instructions. Make a separate list of pointers to indirect
1020 // call/invoke/callbr instructions because the original instructions will be
1021 // deleted as the checks are added.
1022 for (BasicBlock &BB : F) {
1023 for (Instruction &I : BB) {
1024 auto *CB = dyn_cast<CallBase>(Val: &I);
1025 if (!CB || CB->getCallingConv() == CallingConv::ARM64EC_Thunk_X64 ||
1026 CB->isInlineAsm())
1027 continue;
1028
1029 // We need to instrument any call that isn't directly calling an
1030 // ARM64 function.
1031 //
1032 // FIXME: getCalledFunction() fails if there's a bitcast (e.g.
1033 // unprototyped functions in C)
1034 if (Function *F = CB->getCalledFunction()) {
1035 if (!AArch64Options::Global.arm64ec_lower_direct_to_indirect ||
1036 F->hasLocalLinkage() || F->isIntrinsic() ||
1037 !F->isDeclarationForLinker())
1038 continue;
1039
1040 DirectCalledFns.insert(X: F);
1041 continue;
1042 }
1043
1044 // Use mangled global alias for direct calls to patchable functions.
1045 if (GlobalAlias *A = dyn_cast<GlobalAlias>(Val: CB->getCalledOperand())) {
1046 auto I = FnsMap.find(Val: A);
1047 if (I != FnsMap.end()) {
1048 CB->setCalledOperand(I->second);
1049 DirectCalledFns.insert(X: I->first);
1050 continue;
1051 }
1052 }
1053
1054 IndirectCalls.push_back(Elt: CB);
1055 ++Arm64ECCallsLowered;
1056 }
1057 }
1058
1059 if (IndirectCalls.empty())
1060 return false;
1061
1062 for (CallBase *CB : IndirectCalls)
1063 lowerCall(CB);
1064
1065 return true;
1066}
1067
1068char AArch64Arm64ECCallLowering::ID = 0;
1069INITIALIZE_PASS(AArch64Arm64ECCallLowering, "Arm64ECCallLowering",
1070 "AArch64Arm64ECCallLowering", false, false)
1071
1072ModulePass *llvm::createAArch64Arm64ECCallLoweringPass() {
1073 return new AArch64Arm64ECCallLowering;
1074}
1075