================ @@ -0,0 +1,185 @@ +//===- RecordTypeConverter.cpp - Record-rebuilding type converter ---------===// +// +// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. +// See https://llvm.org/LICENSE.txt for license information. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +//===----------------------------------------------------------------------===// + +#include "RecordTypeConverter.h" + +#include "llvm/ADT/STLExtras.h" +#include "llvm/ADT/ScopeExit.h" +#include "llvm/Support/Threading.h" + +#include <mutex> +#include <shared_mutex> + +using namespace cir; + +RecordRewritingTypeConverter::RecordRewritingTypeConverter( + mlir::MLIRContext &context) + : context(context) { + addConversion([&](mlir::Type type) -> mlir::Type { return type; }); + // This is necessary in order to convert CIR pointer types that are pointing + // to CIR types that are being converted. + addConversion([&](cir::PointerType type) -> mlir::Type { + mlir::Type loweredPointeeType = convertType(type.getPointee()); + if (!loweredPointeeType) + return {}; + return cir::PointerType::get(type.getContext(), loweredPointeeType, + type.getAddrSpace()); + }); + addConversion([&](cir::ArrayType type) -> mlir::Type { + mlir::Type loweredElementType = convertType(type.getElementType()); + if (!loweredElementType) + return {}; + return cir::ArrayType::get(loweredElementType, type.getSize()); + }); + // This is necessary in order to convert CIR function types that have + // argument or return types that use CIR types that are being converted. + addConversion([&](cir::FuncType type) -> mlir::Type { + llvm::SmallVector<mlir::Type> loweredInputTypes; + loweredInputTypes.reserve(type.getNumInputs()); + if (mlir::failed(convertTypes(type.getInputs(), loweredInputTypes))) + return {}; + + mlir::Type loweredReturnType = convertType(type.getReturnType()); + if (!loweredReturnType) + return {}; + + return cir::FuncType::get(loweredInputTypes, loweredReturnType, + /*isVarArg=*/type.getVarArg()); + }); + addConversion([&](cir::StructType type) -> mlir::Type { + return convertRecordType(type); + }); + addConversion([&](cir::UnionType type) -> mlir::Type { + return convertRecordType(type); + }); +} + +void RecordRewritingTypeConverter::restoreRecordTypeNames() { + std::unique_lock<decltype(recordTypeMutex)> lock(recordTypeMutex); + + for (auto rt : convertedRecordTypes) + rt.removeABIConversionNamePrefix(); +} + +// This provides a stack for the RecordTypes being processed on the current +// thread, which lets us solve recursive conversions. This implementation is +// cribbed from the LLVMTypeConverter which solves a similar but not identical +// problem. +llvm::SmallVector<cir::RecordType> & +RecordRewritingTypeConverter::getCurrentThreadRecursiveStack() { + { + // Most of the time, the entry already exists in the map. + std::shared_lock<decltype(callStackMutex)> lock(callStackMutex, + std::defer_lock); + if (context.isMultithreadingEnabled()) + lock.lock(); + auto recursiveStack = conversionCallStack.find(llvm::get_threadid()); + if (recursiveStack != conversionCallStack.end()) + return *recursiveStack->second; + } + + // First time this thread gets here, we have to get an exclusive access to + // insert in the map + std::unique_lock<decltype(callStackMutex)> lock(callStackMutex); + auto recursiveStackInserted = conversionCallStack.insert( + std::make_pair(llvm::get_threadid(), + std::make_unique<llvm::SmallVector<cir::RecordType>>())); + return *recursiveStackInserted.first->second; +} + +void RecordRewritingTypeConverter::addConvertedRecordType(cir::RecordType rt) { + std::unique_lock<decltype(recordTypeMutex)> lock(recordTypeMutex); + convertedRecordTypes.push_back(rt); +} + +llvm::SmallVector<mlir::Type> +RecordRewritingTypeConverter::convertRecordMemberTypes(cir::RecordType type) { + llvm::SmallVector<mlir::Type> loweredMemberTypes; + loweredMemberTypes.reserve(type.getNumElements()); + + if (mlir::failed(convertTypes(type.getMembers(), loweredMemberTypes))) + return {}; + + return loweredMemberTypes; +} + +cir::RecordType +RecordRewritingTypeConverter::convertRecordType(cir::RecordType type) { + if (!shouldConvertRecord(type)) ---------------- koparasy wrote:
3f6587f https://github.com/llvm/llvm-project/pull/228599 _______________________________________________ cfe-commits mailing list [email protected] https://lists.llvm.org/cgi-bin/mailman/listinfo/cfe-commits
