11#include "llvm/ADT/STLExtras.h"
12#include "llvm/ADT/ScopeExit.h"
13#include "llvm/Support/Threading.h"
16#include <shared_mutex>
21 mlir::MLIRContext &context)
23 addConversion([&](mlir::Type type) -> mlir::Type {
return type; });
26 addConversion([&](cir::PointerType type) -> mlir::Type {
27 mlir::Type loweredPointeeType = convertType(type.getPointee());
28 if (!loweredPointeeType)
30 return cir::PointerType::get(type.getContext(), loweredPointeeType,
33 addConversion([&](cir::ArrayType type) -> mlir::Type {
34 mlir::Type loweredElementType = convertType(type.getElementType());
35 if (!loweredElementType)
37 return cir::ArrayType::get(loweredElementType, type.getSize());
41 addConversion([&](cir::FuncType type) -> mlir::Type {
43 loweredInputTypes.reserve(type.getNumInputs());
44 if (mlir::failed(convertTypes(type.getInputs(), loweredInputTypes)))
47 mlir::Type loweredReturnType = convertType(type.getReturnType());
48 if (!loweredReturnType)
51 return cir::FuncType::get(loweredInputTypes, loweredReturnType,
54 addConversion([&](cir::StructType type) -> mlir::Type {
55 return convertRecordType(type);
57 addConversion([&](cir::UnionType type) -> mlir::Type {
58 return convertRecordType(type);
63 std::unique_lock<
decltype(recordTypeMutex)> lock(recordTypeMutex);
65 for (
auto rt : convertedRecordTypes)
66 rt.removeABIConversionNamePrefix();
74RecordRewritingTypeConverter::getCurrentThreadRecursiveStack() {
77 std::shared_lock<
decltype(callStackMutex)> lock(callStackMutex,
79 if (context.isMultithreadingEnabled())
81 auto recursiveStack = conversionCallStack.find(llvm::get_threadid());
82 if (recursiveStack != conversionCallStack.end())
83 return *recursiveStack->second;
88 std::unique_lock<
decltype(callStackMutex)> lock(callStackMutex);
89 auto recursiveStackInserted = conversionCallStack.insert(
90 std::make_pair(llvm::get_threadid(),
92 return *recursiveStackInserted.first->second;
95void RecordRewritingTypeConverter::addConvertedRecordType(cir::RecordType rt) {
96 std::unique_lock<
decltype(recordTypeMutex)> lock(recordTypeMutex);
97 convertedRecordTypes.push_back(rt);
100llvm::SmallVector<mlir::Type>
101RecordRewritingTypeConverter::convertRecordMemberTypes(cir::RecordType type) {
102 llvm::SmallVector<mlir::Type> loweredMemberTypes;
103 loweredMemberTypes.reserve(
type.getNumElements());
105 if (mlir::failed(convertTypes(
type.getMembers(), loweredMemberTypes)))
108 return loweredMemberTypes;
112RecordRewritingTypeConverter::convertRecordType(cir::RecordType type) {
116 if (!
type.getName()) {
117 llvm::SmallVector<mlir::Type> converted = convertRecordMemberTypes(type);
118 assert(converted.size() ==
type.getNumElements() &&
119 "member conversion must be one type in, one type out for the "
120 "kinds to carry over by index");
121 if (
auto u = mlir::dyn_cast<cir::UnionType>(type)) {
122 mlir::Type loweredPadding;
123 if (mlir::Type pad = u.getPadding())
124 loweredPadding = convertType(pad);
125 return cir::UnionType::get(
type.getContext(), converted,
type.getPacked(),
126 loweredPadding, u.getMemberKinds());
128 auto s = mlir::cast<cir::StructType>(type);
129 return cir::StructType::get(
type.getContext(), converted,
type.getPacked(),
130 s.getIsClass(), s.getMemberKinds());
133 assert(!
type.isIncomplete() ||
type.getMembers().empty());
138 if (
type.isIncomplete() ||
type.isABIConvertedRecord())
141 llvm::SmallVectorImpl<cir::RecordType> &recursiveStack =
142 getCurrentThreadRecursiveStack();
144 cir::RecordType convertedType;
145 if (mlir::isa<cir::UnionType>(type))
147 cir::UnionType::get(
type.getContext(),
type.getABIConvertedName());
150 cir::StructType::get(
type.getContext(),
type.getABIConvertedName(),
151 mlir::cast<cir::StructType>(type).getIsClass());
155 return convertedType;
161 if (llvm::is_contained(recursiveStack, type))
162 return convertedType;
164 recursiveStack.push_back(type);
165 llvm::scope_exit popConvertingType(
166 [&recursiveStack]() { recursiveStack.pop_back(); });
168 llvm::SmallVector<mlir::Type> convertedMembers =
169 convertRecordMemberTypes(type);
170 assert(convertedMembers.size() ==
type.getNumElements() &&
171 "member conversion must be one type in, one type out for the kinds "
172 "to carry over by index");
174 mlir::Type loweredPadding;
175 if (
auto u = mlir::dyn_cast<cir::UnionType>(type))
176 if (mlir::Type pad = u.getPadding())
177 loweredPadding = convertType(pad);
178 convertedType.
complete(convertedMembers,
type.getPacked(), loweredPadding,
179 type.getMemberKinds());
180 addConvertedRecordType(convertedType);
181 return convertedType;
void restoreRecordTypeNames()
Remove the temporary name of every record rebuilt by this converter.
RecordRewritingTypeConverter(mlir::MLIRContext &context)
void complete(llvm::ArrayRef< mlir::Type > members, bool packed, mlir::Type padding, llvm::ArrayRef< RecordMemberKind > memberKinds)
padding is union-only.
const internal::VariadicAllOfMatcher< Type > type
Matches Types in the clang AST.