Merge pull request #3067 from neobrain/refactor_thunks

Thunks: Minor restructuring and small cleanups
This commit is contained in:
Ryan Houdek authored and GitHub committed 2023-09-07 20:16:17 -07:00
commit be07254935
9 files changed
+716 -609

No files matched your search

+1 -1
View File
@@ -1,7 +1,7 @@
find_package(Clang REQUIRED CONFIG)
find_package(OpenSSL REQUIRED COMPONENTS Crypto)
add_library(thunkgenlib gen.cpp)
add_library(thunkgenlib analysis.cpp gen.cpp)
target_include_directories(thunkgenlib INTERFACE ${CMAKE_CURRENT_SOURCE_DIR})
target_include_directories(thunkgenlib SYSTEM PUBLIC ${CLANG_INCLUDE_DIRS})
target_link_libraries(thunkgenlib PUBLIC clang-cpp LLVM)
+355
View File
@@ -0,0 +1,355 @@
#include "analysis.h"
#include "diagnostics.h"
#include <clang/AST/RecursiveASTVisitor.h>
#include <clang/Frontend/CompilerInstance.h>
#include <fmt/format.h>
struct NamespaceAnnotations {
std::optional<unsigned> version;
std::optional<std::string> load_host_endpoint_via;
bool generate_guest_symtable = false;
bool indirect_guest_calls = false;
};
static NamespaceAnnotations GetNamespaceAnnotations(clang::ASTContext& context, clang::CXXRecordDecl* decl) {
if (!decl->hasDefinition()) {
return {};
}
ErrorReporter report_error { context };
NamespaceAnnotations ret;
for (const clang::CXXBaseSpecifier& base : decl->bases()) {
auto annotation = base.getType().getAsString();
if (annotation == "fexgen::generate_guest_symtable") {
ret.generate_guest_symtable = true;
} else if (annotation == "fexgen::indirect_guest_calls") {
ret.indirect_guest_calls = true;
} else {
throw report_error(base.getSourceRange().getBegin(), "Unknown namespace annotation");
}
}
for (const clang::FieldDecl* field : decl->fields()) {
auto name = field->getNameAsString();
if (name == "load_host_endpoint_via") {
auto loader_function_expr = field->getInClassInitializer()->IgnoreCasts();
auto loader_function_str = llvm::dyn_cast_or_null<clang::StringLiteral>(loader_function_expr);
if (loader_function_expr && !loader_function_str) {
throw report_error(loader_function_expr->getBeginLoc(),
"Must initialize load_host_endpoint_via with a string");
}
if (loader_function_str) {
ret.load_host_endpoint_via = loader_function_str->getString();
}
} else if (name == "version") {
auto initializer = field->getInClassInitializer()->IgnoreCasts();
auto version_literal = llvm::dyn_cast_or_null<clang::IntegerLiteral>(initializer);
if (!initializer || !version_literal) {
throw report_error(field->getBeginLoc(), "No version given (expected integral typed member, e.g. \"int version = 5;\")");
}
ret.version = version_literal->getValue().getZExtValue();
} else {
throw report_error(field->getBeginLoc(), "Unknown namespace annotation");
}
}
return ret;
}
enum class CallbackStrategy {
Default,
Stub,
Guest,
};
struct Annotations {
bool custom_host_impl = false;
bool custom_guest_entrypoint = false;
bool returns_guest_pointer = false;
std::optional<clang::QualType> uniform_va_type;
CallbackStrategy callback_strategy = CallbackStrategy::Default;
};
static Annotations GetAnnotations(clang::ASTContext& context, clang::CXXRecordDecl* decl) {
ErrorReporter report_error { context };
Annotations ret;
for (const auto& base : decl->bases()) {
auto annotation = base.getType().getAsString();
if (annotation == "fexgen::returns_guest_pointer") {
ret.returns_guest_pointer = true;
} else if (annotation == "fexgen::custom_host_impl") {
ret.custom_host_impl = true;
} else if (annotation == "fexgen::callback_stub") {
ret.callback_strategy = CallbackStrategy::Stub;
} else if (annotation == "fexgen::callback_guest") {
ret.callback_strategy = CallbackStrategy::Guest;
} else if (annotation == "fexgen::custom_guest_entrypoint") {
ret.custom_guest_entrypoint = true;
} else {
throw report_error(base.getSourceRange().getBegin(), "Unknown annotation");
}
}
for (const auto& child_decl : decl->getPrimaryContext()->decls()) {
if (auto field = llvm::dyn_cast_or_null<clang::FieldDecl>(child_decl)) {
throw report_error(field->getBeginLoc(), "Unknown field annotation");
} else if (auto type_alias = llvm::dyn_cast_or_null<clang::TypedefNameDecl>(child_decl)) {
auto name = type_alias->getNameAsString();
if (name == "uniform_va_type") {
ret.uniform_va_type = type_alias->getUnderlyingType();
} else {
throw report_error(type_alias->getBeginLoc(), "Unknown type alias annotation");
}
}
}
return ret;
}
void AnalysisAction::ExecuteAction() {
clang::ASTFrontendAction::ExecuteAction();
// Post-processing happens here rather than in an overridden EndSourceFileAction implementation.
// We can't move the logic to the latter since this code might still raise errors, but
// clang's diagnostics engine is already shut down by the time EndSourceFileAction is called.
auto& context = getCompilerInstance().getASTContext();
if (context.getDiagnostics().hasErrorOccurred()) {
return;
}
decl_contexts.front() = context.getTranslationUnitDecl();
try {
ParseInterface(context);
EmitOutput(context);
} catch (ClangDiagnosticAsException& exception) {
exception.Report(context.getDiagnostics());
}
}
static clang::ClassTemplateDecl*
FindClassTemplateDeclByName(clang::DeclContext& decl_context, std::string_view symbol_name) {
auto& ast_context = decl_context.getParentASTContext();
auto* ident = &ast_context.Idents.get(symbol_name);
auto declname = ast_context.DeclarationNames.getIdentifier(ident);
auto result = decl_context.noload_lookup(declname);
if (result.empty()) {
return nullptr;
} else if (std::next(result.begin()) == result.end()) {
return llvm::dyn_cast<clang::ClassTemplateDecl>(*result.begin());
} else {
throw std::runtime_error("Found multiple matches to symbol " + std::string { symbol_name });
}
}
void AnalysisAction::ParseInterface(clang::ASTContext& context) {
ErrorReporter report_error { context };
if (auto template_decl = FindClassTemplateDeclByName(*context.getTranslationUnitDecl(), "fex_gen_type")) {
for (auto* decl : template_decl->specializations()) {
const auto& template_args = decl->getTemplateArgs();
assert(template_args.size() == 1);
// NOTE: Function types that are equivalent but use differently
// named types (e.g. GLuint/GLenum) are represented by
// different Type instances. The canonical type they refer
// to is unique, however.
auto type = context.getCanonicalType(template_args[0].getAsType()).getTypePtr();
funcptr_types.insert(type);
}
}
// Process declarations and specializations of fex_gen_config,
// i.e. the function descriptions of the thunked API
for (auto& decl_context : decl_contexts) {
if (const auto template_decl = FindClassTemplateDeclByName(*decl_context, "fex_gen_config")) {
// Gather general information about symbols in this namespace
const auto annotations = GetNamespaceAnnotations(context, template_decl->getTemplatedDecl());
auto namespace_decl = llvm::dyn_cast<clang::NamespaceDecl>(decl_context);
namespaces.push_back({ namespace_decl,
namespace_decl ? namespace_decl->getNameAsString() : "",
annotations.load_host_endpoint_via.value_or(""),
annotations.generate_guest_symtable,
annotations.indirect_guest_calls });
const auto namespace_idx = namespaces.size() - 1;
const NamespaceInfo& namespace_info = namespaces.back();
if (annotations.version) {
if (namespace_decl) {
throw report_error(template_decl->getBeginLoc(), "Library version must be defined in the global namespace");
}
lib_version = annotations.version;
}
// Process specializations of template fex_gen_config
// First, perform some validation and process member annotations
// In a second iteration, process the actual function API
for (auto* decl : template_decl->specializations()) {
if (decl->getSpecializationKind() == clang::TSK_ExplicitInstantiationDefinition) {
throw report_error(decl->getBeginLoc(), "fex_gen_config may not be partially specialized\n");
}
const auto& template_args = decl->getTemplateArgs();
assert(template_args.size() == 1);
const auto template_arg_loc = decl->getTypeAsWritten()->getTypeLoc().castAs<clang::TemplateSpecializationTypeLoc>().getArgLoc(0).getLocation();
if (auto emitted_function = llvm::dyn_cast<clang::FunctionDecl>(template_args[0].getAsDecl())) {
// Process later
} else {
throw report_error(template_arg_loc, "Cannot annotate this kind of symbol");
}
}
// Process API functions
for (auto* decl : template_decl->specializations()) {
if (decl->getSpecializationKind() == clang::TSK_ExplicitInstantiationDefinition) {
throw report_error(decl->getBeginLoc(), "fex_gen_config may not be partially specialized\n");
}
const auto& template_args = decl->getTemplateArgs();
assert(template_args.size() == 1);
const auto template_arg_loc = decl->getTypeAsWritten()->getTypeLoc().castAs<clang::TemplateSpecializationTypeLoc>().getArgLoc(0).getLocation();
auto emitted_function = llvm::dyn_cast<clang::FunctionDecl>(template_args[0].getAsDecl());
assert(emitted_function && "Argument is not a function");
auto return_type = emitted_function->getReturnType();
const auto annotations = GetAnnotations(context, decl);
if (return_type->isFunctionPointerType() && !annotations.returns_guest_pointer) {
throw report_error( template_arg_loc,
"Function pointer return types require explicit annotation\n");
}
// TODO: Use the types as written in the signature instead?
ThunkedFunction data;
data.function_name = emitted_function->getName().str();
data.return_type = return_type;
data.is_variadic = emitted_function->isVariadic();
data.decl = emitted_function;
data.custom_host_impl = annotations.custom_host_impl;
for (std::size_t param_idx = 0; param_idx < emitted_function->param_size(); ++param_idx) {
auto* param = emitted_function->getParamDecl(param_idx);
data.param_types.push_back(param->getType());
if (param->getType()->isFunctionPointerType()) {
auto funcptr = param->getFunctionType()->getAs<clang::FunctionProtoType>();
ThunkedCallback callback;
callback.return_type = funcptr->getReturnType();
for (auto& cb_param : funcptr->getParamTypes()) {
callback.param_types.push_back(cb_param);
}
callback.is_stub = annotations.callback_strategy == CallbackStrategy::Stub;
callback.is_guest = annotations.callback_strategy == CallbackStrategy::Guest;
callback.is_variadic = funcptr->isVariadic();
if (callback.is_guest && !data.custom_host_impl) {
throw report_error(template_arg_loc, "callback_guest can only be used with custom_host_impl");
}
data.callbacks.emplace(param_idx, callback);
if (!callback.is_stub && !callback.is_guest) {
funcptr_types.insert(context.getCanonicalType(funcptr));
}
if (data.callbacks.size() != 1) {
throw report_error(template_arg_loc, "Support for more than one callback is untested");
}
if (funcptr->isVariadic() && !callback.is_stub) {
throw report_error(template_arg_loc, "Variadic callbacks are not supported");
}
}
}
thunked_api.push_back(ThunkedAPIFunction { (const FunctionParams&)data, data.function_name, data.return_type,
namespace_info.host_loader.empty() ? "dlsym_default" : namespace_info.host_loader,
data.is_variadic || annotations.custom_guest_entrypoint,
data.is_variadic,
std::nullopt });
if (namespace_info.generate_guest_symtable) {
thunked_api.back().symtable_namespace = namespace_idx;
}
if (data.is_variadic) {
if (!annotations.uniform_va_type) {
throw report_error(decl->getBeginLoc(), "Variadic functions must be annotated with parameter type using uniform_va_type");
}
// Convert variadic argument list into a count + pointer pair
data.param_types.push_back(context.getSizeType());
data.param_types.push_back(context.getPointerType(*annotations.uniform_va_type));
}
if (data.is_variadic) {
// This function is thunked through an "_internal" symbol since its signature
// is different from the one in the native host/guest libraries.
data.function_name = data.function_name + "_internal";
if (data.custom_host_impl) {
throw report_error(decl->getBeginLoc(), "Custom host impl requested but this is implied by the function signature already");
}
data.custom_host_impl = true;
}
// For indirect calls, register the function signature as a function pointer type
if (namespace_info.indirect_guest_calls) {
funcptr_types.insert(context.getCanonicalType(emitted_function->getFunctionType()));
}
thunks.push_back(std::move(data));
}
}
}
}
class ASTVisitor : public clang::RecursiveASTVisitor<ASTVisitor> {
std::vector<clang::DeclContext*>& decl_contexts;
public:
ASTVisitor(std::vector<clang::DeclContext*>& decl_contexts_)
: decl_contexts(decl_contexts_) {
}
/**
* Matches "template<auto> struct fex_gen_config { ... }"
*/
bool VisitClassTemplateDecl(clang::ClassTemplateDecl* decl) {
if (decl->getName() != "fex_gen_config") {
return true;
}
if (llvm::dyn_cast<clang::NamespaceDecl>(decl->getDeclContext())) {
decl_contexts.push_back(decl->getDeclContext());
}
return true;
}
};
class ASTConsumer : public clang::ASTConsumer {
std::vector<clang::DeclContext*>& decl_contexts;
public:
ASTConsumer(std::vector<clang::DeclContext*>& decl_contexts_)
: decl_contexts(decl_contexts_) {
}
void HandleTranslationUnit(clang::ASTContext& context) override {
ASTVisitor { decl_contexts }.TraverseDecl(context.getTranslationUnitDecl());
}
};
std::unique_ptr<clang::ASTConsumer> AnalysisAction::CreateASTConsumer(clang::CompilerInstance&, clang::StringRef) {
return std::make_unique<ASTConsumer>(decl_contexts);
}
+123
View File
@@ -0,0 +1,123 @@
#pragma once
#include <clang/Frontend/FrontendAction.h>
#include <memory>
#include <optional>
#include <string>
#include <unordered_map>
#include <unordered_set>
#include <vector>
struct FunctionParams {
std::vector<clang::QualType> param_types;
};
struct ThunkedCallback : FunctionParams {
clang::QualType return_type;
bool is_stub = false; // Callback will be replaced by a stub that calls std::abort
bool is_guest = false; // Callback will never be called on the host
bool is_variadic = false;
};
/**
* Guest<->Host transition point.
*
* These are normally used to translate the public API of the guest to host
* function calls (ThunkedAPIFunction), but a thunk library may also define
* internal thunks that don't correspond to any function in the implemented
* API.
*/
struct ThunkedFunction : FunctionParams {
std::string function_name;
clang::QualType return_type;
// If true, param_types contains an extra size_t and the valist for marshalling through an internal function
bool is_variadic = false;
// If true, the unpacking function will call a custom fexfn_impl function
// to be provided manually instead of calling the host library function
// directly.
// This is implied e.g. for thunks generated for variadic functions
bool custom_host_impl = false;
std::string GetOriginalFunctionName() const {
const std::string suffix = "_internal";
assert(function_name.length() > suffix.size());
assert((std::string_view { &*function_name.end() - suffix.size(), suffix.size() } == suffix));
return function_name.substr(0, function_name.size() - suffix.size());
}
// Maps parameter index to ThunkedCallback
std::unordered_map<unsigned, ThunkedCallback> callbacks;
clang::FunctionDecl* decl;
};
/**
* Function that is part of the API of the thunked library.
*
* For each of these, there is:
* - A publicly visible guest entrypoint (usually auto-generated but may be manually defined)
* - A pointer to the native host library function loaded through dlsym (or a user-provided function specified via host_loader)
* - A ThunkedFunction with the same function_name (possibly suffixed with _internal)
*/
struct ThunkedAPIFunction : FunctionParams {
std::string function_name;
clang::QualType return_type;
// name of the function to load the native host symbol with
std::string host_loader;
// If true, no guest-side implementation of this function will be autogenerated
bool custom_guest_impl;
bool is_variadic;
// Index of the symbol table to store this export in (see guest_symtables).
// If empty, a library export is created, otherwise the function is entered into a function pointer array
std::optional<std::size_t> symtable_namespace;
};
struct NamespaceInfo {
clang::DeclContext* context;
std::string name;
// Function to load native host library functions with.
// This function must be defined manually with the signature "void* func(void*, const char*)"
std::string host_loader;
bool generate_guest_symtable;
bool indirect_guest_calls;
};
class AnalysisAction : public clang::ASTFrontendAction {
public:
AnalysisAction() {
decl_contexts.push_back(nullptr); // global namespace (replaced by getTranslationUnitDecl later)
}
void ExecuteAction() override;
std::unique_ptr<clang::ASTConsumer> CreateASTConsumer(clang::CompilerInstance&, clang::StringRef /*file*/) override;
protected:
// Build the internal API representation by processing fex_gen_config and other annotated entities
void ParseInterface(clang::ASTContext&);
// Called from ExecuteAction() after parsing is complete
virtual void EmitOutput(clang::ASTContext&) {};
std::vector<clang::DeclContext*> decl_contexts;
std::vector<ThunkedFunction> thunks;
std::vector<ThunkedAPIFunction> thunked_api;
std::unordered_set<const clang::Type*> funcptr_types;
std::optional<unsigned> lib_version;
std::vector<NamespaceInfo> namespaces;
};
+60
View File
@@ -0,0 +1,60 @@
#pragma once
#include <clang/AST/ASTContext.h>
#include <clang/Basic/Diagnostic.h>
#include <clang/Basic/SourceLocation.h>
#include <utility>
#include <vector>
struct ClangDiagnosticAsException {
std::pair<clang::SourceLocation, unsigned> diagnostic;
std::vector<ClangDiagnosticAsException> notes;
// List of callbacks that add an argument to a clang::DiagnosticBuilder
std::vector<std::function<void(clang::DiagnosticBuilder&)>> args;
ClangDiagnosticAsException& AddString(std::string str) {
args.push_back([arg=std::move(str)](clang::DiagnosticBuilder& db) {
db.AddString(arg);
});
return *this;
}
ClangDiagnosticAsException& AddTaggedVal(clang::QualType type) {
args.push_back([val=type](clang::DiagnosticBuilder& db) {
db.AddTaggedVal(reinterpret_cast<uintptr_t>(val.getAsOpaquePtr()), clang::DiagnosticsEngine::ak_qualtype);
});
return *this;
}
ClangDiagnosticAsException& addNote(ClangDiagnosticAsException diagnostic) {
notes.push_back(std::move(diagnostic));
return *this;
}
void Report(clang::DiagnosticsEngine& diagnostics) const {
{
auto builder = diagnostics.Report(diagnostic.first, diagnostic.second);
for (auto& arg_appender : args) {
arg_appender(builder);
}
}
for (auto& note : notes) {
note.Report(diagnostics);
}
}
};
// Helper class to build a custom DiagID from the given message and store it in a throwable object
struct ErrorReporter {
clang::ASTContext& context;
template<std::size_t N>
[[nodiscard]] ClangDiagnosticAsException operator()(clang::SourceLocation loc, const char (&message)[N],
clang::DiagnosticsEngine::Level level = clang::DiagnosticsEngine::Error) {
auto id = context.getDiagnostics().getCustomDiagID(level, message);
return { std::pair(loc, id) };
}
};
+26 -446
View File
@@ -1,283 +1,30 @@
#include "clang/AST/RecursiveASTVisitor.h"
#include "clang/Frontend/CompilerInstance.h"
#include "analysis.h"
#include "interface.h"
#include <clang/Frontend/CompilerInstance.h>
#include <fstream>
#include <numeric>
#include <iostream>
#include <string_view>
#include <unordered_map>
#include <unordered_set>
#include <fmt/format.h>
#include <fmt/ostream.h>
#include <openssl/sha.h>
#include "interface.h"
struct FunctionParams {
std::vector<clang::QualType> param_types;
};
struct ThunkedCallback : FunctionParams {
clang::QualType return_type;
bool is_stub = false; // Callback will be replaced by a stub that calls std::abort
bool is_guest = false; // Callback will never be called on the host
bool is_variadic = false;
};
/**
* Guest<->Host transition point.
*
* These are normally used to translate the public API of the guest to host
* function calls (ThunkedAPIFunction), but a thunk library may also define
* internal thunks that don't correspond to any function in the implemented
* API.
*/
struct ThunkedFunction : FunctionParams {
std::string function_name;
clang::QualType return_type;
// If true, param_types contains an extra size_t and the valist for marshalling through an internal function
bool is_variadic = false;
// If true, the unpacking function will call a custom fexfn_impl function
// to be provided manually instead of calling the host library function
// directly.
// This is implied e.g. for thunks generated for variadic functions
bool custom_host_impl = false;
std::string GetOriginalFunctionName() const {
const std::string suffix = "_internal";
assert(function_name.length() > suffix.size());
assert((std::string_view { &*function_name.end() - suffix.size(), suffix.size() } == suffix));
return function_name.substr(0, function_name.size() - suffix.size());
}
// Maps parameter index to ThunkedCallback
std::unordered_map<unsigned, ThunkedCallback> callbacks;
clang::FunctionDecl* decl;
};
/**
* Function that is part of the API of the thunked library.
*
* For each of these, there is:
* - A publicly visible guest entrypoint (usually auto-generated but may be manually defined)
* - A pointer to the native host library function loaded through dlsym (or a user-provided function specified via host_loader)
* - A ThunkedFunction with the same function_name (possibly suffixed with _internal)
*/
struct ThunkedAPIFunction : FunctionParams {
std::string function_name;
clang::QualType return_type;
// name of the function to load the native host symbol with
std::string host_loader;
// If true, no guest-side implementation of this function will be autogenerated
bool custom_guest_impl;
bool is_variadic;
// Index of the symbol table to store this export in (see guest_symtables).
// If empty, a library export is created, otherwise the function is entered into a function pointer array
std::optional<std::size_t> symtable_namespace;
};
struct NamespaceInfo {
clang::DeclContext* context;
std::string name;
// Function to load native host library functions with.
// This function must be defined manually with the signature "void* func(void*, const char*)"
std::string host_loader;
bool generate_guest_symtable;
bool indirect_guest_calls;
};
static std::vector<clang::DeclContext*> decl_contexts;
struct ClangDiagnosticAsException {
std::pair<clang::SourceLocation, unsigned> diagnostic;
void Report(clang::DiagnosticsEngine& diagnostics) const {
diagnostics.Report(diagnostic.first, diagnostic.second);
}
};
// Helper class to build a custom DiagID from the given message and store it in a throwable object
struct ErrorReporter {
clang::ASTContext& context;
template<std::size_t N>
[[nodiscard]] ClangDiagnosticAsException operator()(clang::SourceLocation loc, const char (&message)[N]) {
auto id = context.getDiagnostics().getCustomDiagID(clang::DiagnosticsEngine::Error, message);
return { std::pair(loc, id) };
}
};
struct NamespaceAnnotations {
std::optional<unsigned> version;
std::optional<std::string> load_host_endpoint_via;
bool generate_guest_symtable = false;
bool indirect_guest_calls = false;
};
static NamespaceAnnotations GetNamespaceAnnotations(clang::ASTContext& context, clang::CXXRecordDecl* decl) {
if (!decl->hasDefinition()) {
return {};
}
ErrorReporter report_error { context };
NamespaceAnnotations ret;
for (const clang::CXXBaseSpecifier& base : decl->bases()) {
auto annotation = base.getType().getAsString();
if (annotation == "fexgen::generate_guest_symtable") {
ret.generate_guest_symtable = true;
} else if (annotation == "fexgen::indirect_guest_calls") {
ret.indirect_guest_calls = true;
} else {
throw report_error(base.getSourceRange().getBegin(), "Unknown namespace annotation");
}
}
for (const clang::FieldDecl* field : decl->fields()) {
auto name = field->getNameAsString();
if (name == "load_host_endpoint_via") {
auto loader_function_expr = field->getInClassInitializer()->IgnoreCasts();
auto loader_function_str = llvm::dyn_cast_or_null<clang::StringLiteral>(loader_function_expr);
if (loader_function_expr && !loader_function_str) {
throw report_error(loader_function_expr->getBeginLoc(),
"Must initialize load_host_endpoint_via with a string");
}
if (loader_function_str) {
ret.load_host_endpoint_via = loader_function_str->getString();
}
} else if (name == "version") {
auto initializer = field->getInClassInitializer()->IgnoreCasts();
auto version_literal = llvm::dyn_cast_or_null<clang::IntegerLiteral>(initializer);
if (!initializer || !version_literal) {
throw report_error(field->getBeginLoc(), "No version given (expected integral typed member, e.g. \"int version = 5;\")");
}
ret.version = version_literal->getValue().getZExtValue();
} else {
throw report_error(field->getBeginLoc(), "Unknown namespace annotation");
}
}
return ret;
}
enum class CallbackStrategy {
Default,
Stub,
Guest,
};
struct Annotations {
bool custom_host_impl = false;
bool custom_guest_entrypoint = false;
bool returns_guest_pointer = false;
std::optional<clang::QualType> uniform_va_type;
CallbackStrategy callback_strategy = CallbackStrategy::Default;
};
static Annotations GetAnnotations(clang::ASTContext& context, clang::CXXRecordDecl* decl) {
ErrorReporter report_error { context };
Annotations ret;
for (const auto& base : decl->bases()) {
auto annotation = base.getType().getAsString();
if (annotation == "fexgen::returns_guest_pointer") {
ret.returns_guest_pointer = true;
} else if (annotation == "fexgen::custom_host_impl") {
ret.custom_host_impl = true;
} else if (annotation == "fexgen::callback_stub") {
ret.callback_strategy = CallbackStrategy::Stub;
} else if (annotation == "fexgen::callback_guest") {
ret.callback_strategy = CallbackStrategy::Guest;
} else if (annotation == "fexgen::custom_guest_entrypoint") {
ret.custom_guest_entrypoint = true;
} else {
throw report_error(base.getSourceRange().getBegin(), "Unknown annotation");
}
}
for (const auto& child_decl : decl->getPrimaryContext()->decls()) {
if (auto field = llvm::dyn_cast_or_null<clang::FieldDecl>(child_decl)) {
throw report_error(field->getBeginLoc(), "Unknown field annotation");
} else if (auto type_alias = llvm::dyn_cast_or_null<clang::TypedefNameDecl>(child_decl)) {
auto name = type_alias->getNameAsString();
if (name == "uniform_va_type") {
ret.uniform_va_type = type_alias->getUnderlyingType();
} else {
throw report_error(type_alias->getBeginLoc(), "Unknown type alias annotation");
}
}
}
return ret;
}
class ASTVisitor : public clang::RecursiveASTVisitor<ASTVisitor> {
public:
/**
* Matches "template<auto> struct fex_gen_config { ... }"
*/
bool VisitClassTemplateDecl(clang::ClassTemplateDecl* decl) {
if (decl->getName() != "fex_gen_config") {
return true;
}
if (llvm::dyn_cast<clang::NamespaceDecl>(decl->getDeclContext())) {
decl_contexts.push_back(decl->getDeclContext());
}
return true;
}
};
class ASTConsumer : public clang::ASTConsumer {
public:
void HandleTranslationUnit(clang::ASTContext& context) override {
ASTVisitor{}.TraverseDecl(context.getTranslationUnitDecl());
}
};
class GenerateThunkLibsAction : public clang::ASTFrontendAction {
class GenerateThunkLibsAction : public AnalysisAction {
public:
GenerateThunkLibsAction(const std::string& libname, const OutputFilenames&);
void ExecuteAction() override;
std::unique_ptr<clang::ASTConsumer> CreateASTConsumer(clang::CompilerInstance&, clang::StringRef /*file*/) override;
private:
// Build the internal API representation by processing fex_gen_config and other annotated entities
void ParseInterface(clang::ASTContext&);
// Generate helper code for thunk libraries and write them to the output file
void EmitOutput();
void EmitOutput(clang::ASTContext&) override;
const std::string& libfilename;
std::string libname; // sanitized filename, usable as part of emitted function names
const OutputFilenames& output_filenames;
std::vector<ThunkedFunction> thunks;
std::vector<ThunkedAPIFunction> thunked_api;
std::unordered_set<const clang::Type*> funcptr_types;
std::optional<unsigned> lib_version;
std::vector<NamespaceInfo> namespaces;
};
GenerateThunkLibsAction::GenerateThunkLibsAction(const std::string& libname_, const OutputFilenames& output_filenames_)
@@ -287,9 +34,6 @@ GenerateThunkLibsAction::GenerateThunkLibsAction(const std::string& libname_, co
c = '_';
}
}
decl_contexts.clear();
decl_contexts.push_back(nullptr); // global namespace (replaced by getTranslationUnitDecl later)
}
template<typename Fn>
@@ -303,185 +47,7 @@ static std::string format_function_args(const FunctionParams& params, Fn&& forma
return ret;
};
static clang::ClassTemplateDecl*
FindClassTemplateDeclByName(clang::DeclContext& decl_context, std::string_view symbol_name) {
auto& ast_context = decl_context.getParentASTContext();
auto* ident = &ast_context.Idents.get(symbol_name);
auto declname = ast_context.DeclarationNames.getIdentifier(ident);
auto result = decl_context.noload_lookup(declname);
if (result.empty()) {
return nullptr;
} else if (std::next(result.begin()) == result.end()) {
return llvm::dyn_cast<clang::ClassTemplateDecl>(*result.begin());
} else {
throw std::runtime_error("Found multiple matches to symbol " + std::string { symbol_name });
}
}
void GenerateThunkLibsAction::ExecuteAction() {
clang::ASTFrontendAction::ExecuteAction();
// Post-processing happens here rather than in an overridden EndSourceFileAction implementation.
// We can't move the logic to the latter since this code might still raise errors, but
// clang's diagnostics engine is already shut down by the time EndSourceFileAction is called.
auto& context = getCompilerInstance().getASTContext();
if (context.getDiagnostics().hasErrorOccurred()) {
return;
}
decl_contexts.front() = context.getTranslationUnitDecl();
try {
ParseInterface(context);
EmitOutput();
} catch (ClangDiagnosticAsException& exception) {
exception.Report(context.getDiagnostics());
}
}
void GenerateThunkLibsAction::ParseInterface(clang::ASTContext& context) {
ErrorReporter report_error { context };
if (auto template_decl = FindClassTemplateDeclByName(*context.getTranslationUnitDecl(), "fex_gen_type")) {
for (auto* decl : template_decl->specializations()) {
const auto& template_args = decl->getTemplateArgs();
assert(template_args.size() == 1);
// NOTE: Function types that are equivalent but use differently
// named types (e.g. GLuint/GLenum) are represented by
// different Type instances. The canonical type they refer
// to is unique, however.
auto type = context.getCanonicalType(template_args[0].getAsType()).getTypePtr();
funcptr_types.insert(type);
}
}
// Process declarations and specializations of fex_gen_config,
// i.e. the function descriptions of the thunked API
for (auto& decl_context : decl_contexts) {
if (const auto template_decl = FindClassTemplateDeclByName(*decl_context, "fex_gen_config")) {
// Gather general information about symbols in this namespace
const auto annotations = GetNamespaceAnnotations(context, template_decl->getTemplatedDecl());
auto namespace_decl = llvm::dyn_cast<clang::NamespaceDecl>(decl_context);
namespaces.push_back({ namespace_decl,
namespace_decl ? namespace_decl->getNameAsString() : "",
annotations.load_host_endpoint_via.value_or(""),
annotations.generate_guest_symtable,
annotations.indirect_guest_calls });
const auto namespace_idx = namespaces.size() - 1;
const NamespaceInfo& namespace_info = namespaces.back();
if (annotations.version) {
if (namespace_decl) {
throw report_error(template_decl->getBeginLoc(), "Library version must be defined in the global namespace");
}
lib_version = annotations.version;
}
// Process specializations of template fex_gen_config
for (auto* decl : template_decl->specializations()) {
if (decl->getSpecializationKind() == clang::TSK_ExplicitInstantiationDefinition) {
throw report_error(decl->getBeginLoc(), "fex_gen_config may not be partially specialized\n");
}
const auto& template_args = decl->getTemplateArgs();
assert(template_args.size() == 1);
auto emitted_function = llvm::dyn_cast<clang::FunctionDecl>(template_args[0].getAsDecl());
assert(emitted_function && "Argument is not a function");
auto return_type = emitted_function->getReturnType();
const auto annotations = GetAnnotations(context, decl);
if (return_type->isFunctionPointerType() && !annotations.returns_guest_pointer) {
throw report_error( decl->getBeginLoc(),
"Function pointer return types require explicit annotation\n");
}
// TODO: Use the types as written in the signature instead?
ThunkedFunction data;
data.function_name = emitted_function->getName().str();
data.return_type = return_type;
data.is_variadic = emitted_function->isVariadic();
data.decl = emitted_function;
data.custom_host_impl = annotations.custom_host_impl;
for (std::size_t param_idx = 0; param_idx < emitted_function->param_size(); ++param_idx) {
auto* param = emitted_function->getParamDecl(param_idx);
data.param_types.push_back(param->getType());
if (param->getType()->isFunctionPointerType()) {
auto funcptr = param->getFunctionType()->getAs<clang::FunctionProtoType>();
ThunkedCallback callback;
callback.return_type = funcptr->getReturnType();
for (auto& cb_param : funcptr->getParamTypes()) {
callback.param_types.push_back(cb_param);
}
callback.is_stub = annotations.callback_strategy == CallbackStrategy::Stub;
callback.is_guest = annotations.callback_strategy == CallbackStrategy::Guest;
callback.is_variadic = funcptr->isVariadic();
if (callback.is_guest && !data.custom_host_impl) {
throw report_error(decl->getBeginLoc(), "callback_guest can only be used with custom_host_impl");
}
data.callbacks.emplace(param_idx, callback);
if (!callback.is_stub && !callback.is_guest) {
funcptr_types.insert(context.getCanonicalType(funcptr));
}
if (data.callbacks.size() != 1) {
throw report_error(decl->getBeginLoc(), "Support for more than one callback is untested");
}
if (funcptr->isVariadic() && !callback.is_stub) {
throw report_error(decl->getBeginLoc(), "Variadic callbacks are not supported");
}
}
}
thunked_api.push_back(ThunkedAPIFunction { (const FunctionParams&)data, data.function_name, data.return_type,
namespace_info.host_loader.empty() ? "dlsym_default" : namespace_info.host_loader,
data.is_variadic || annotations.custom_guest_entrypoint,
data.is_variadic,
std::nullopt });
if (namespace_info.generate_guest_symtable) {
thunked_api.back().symtable_namespace = namespace_idx;
}
if (data.is_variadic) {
if (!annotations.uniform_va_type) {
throw report_error(decl->getBeginLoc(), "Variadic functions must be annotated with parameter type using uniform_va_type");
}
// Convert variadic argument list into a count + pointer pair
data.param_types.push_back(context.getSizeType());
data.param_types.push_back(context.getPointerType(*annotations.uniform_va_type));
}
if (data.is_variadic) {
// This function is thunked through an "_internal" symbol since its signature
// is different from the one in the native host/guest libraries.
data.function_name = data.function_name + "_internal";
if (data.custom_host_impl) {
throw report_error(decl->getBeginLoc(), "Custom host impl requested but this is implied by the function signature already");
}
data.custom_host_impl = true;
}
// For indirect calls, register the function signature as a function pointer type
if (namespace_info.indirect_guest_calls) {
funcptr_types.insert(context.getCanonicalType(emitted_function->getFunctionType()));
}
thunks.push_back(std::move(data));
}
}
}
}
void GenerateThunkLibsAction::EmitOutput() {
void GenerateThunkLibsAction::EmitOutput(clang::ASTContext& context) {
static auto format_decl = [](clang::QualType type, const std::string_view& name) {
clang::QualType innermostPointee = type;
while (innermostPointee->isPointerType()) {
@@ -797,10 +363,24 @@ void GenerateThunkLibsAction::EmitOutput() {
}
}
std::unique_ptr<clang::ASTConsumer> GenerateThunkLibsAction::CreateASTConsumer(clang::CompilerInstance&, clang::StringRef) {
return std::make_unique<ASTConsumer>();
}
bool GenerateThunkLibsActionFactory::runInvocation(
std::shared_ptr<clang::CompilerInvocation> Invocation, clang::FileManager *Files,
std::shared_ptr<clang::PCHContainerOperations> PCHContainerOps,
clang::DiagnosticConsumer *DiagConsumer) {
clang::CompilerInstance Compiler(std::move(PCHContainerOps));
Compiler.setInvocation(std::move(Invocation));
Compiler.setFileManager(Files);
std::unique_ptr<clang::FrontendAction> GenerateThunkLibsActionFactory::create() {
return std::make_unique<GenerateThunkLibsAction>(libname, output_filenames);
GenerateThunkLibsAction Action(libname, output_filenames);
Compiler.createDiagnostics(DiagConsumer, false);
if (!Compiler.hasDiagnostics())
return false;
Compiler.createSourceManager(*Files);
const bool Success = Compiler.ExecuteAction(Action);
Files->clearStatCache();
return Success;
}
+5 -2
View File
@@ -8,13 +8,16 @@ struct OutputFilenames {
std::string guest;
};
class GenerateThunkLibsActionFactory : public clang::tooling::FrontendActionFactory {
class GenerateThunkLibsActionFactory : public clang::tooling::ToolAction {
public:
GenerateThunkLibsActionFactory(std::string_view libname_, OutputFilenames output_filenames_)
: libname(std::move(libname_)), output_filenames(std::move(output_filenames_)) {
}
std::unique_ptr<clang::FrontendAction> create() override;
bool runInvocation(
std::shared_ptr<clang::CompilerInvocation> Invocation, clang::FileManager *Files,
std::shared_ptr<clang::PCHContainerOperations> PCHContainerOps,
clang::DiagnosticConsumer *DiagConsumer) override;
private:
std::string libname;
+59 -88
View File
@@ -2,6 +2,7 @@
#include <cstdint>
#include <type_traits>
#include <utility>
template<typename Result, typename... Args>
struct PackedArguments;
@@ -102,97 +103,67 @@ struct PackedArguments<void, A0, A1, A2, A3, A4, A5, A6, A7, A8, A9, A10, A11, A
template<typename A0, typename A1, typename A2, typename A3, typename A4, typename A5, typename A6, typename A7, typename A8, typename A9, typename A10, typename A11, typename A12, typename A13, typename A14, typename A15, typename A16, typename A17, typename A18>
struct PackedArguments<void, A0, A1, A2, A3, A4, A5, A6, A7, A8, A9, A10, A11, A12, A13, A14, A15, A16, A17, A18> { A0 a0; A1 a1; A2 a2; A3 a3; A4 a4; A5 a5; A6 a6; A7 a7; A8 a8; A9 a9; A10 a10; A11 a11; A12 a12; A13 a13; A14 a14; A15 a15; A16 a16; A17 a17; A18 a18; };
// Helper struct that allows assigning the result of a function to a variable, even if that result is a void type.
//
// For non-void result types, the overloaded the comma operator will always returns its left argument.
// For void types, the overloaded comma operator is *not* used. Instead, a dummy object is returned.
struct Regularize {};
template<typename T>
T&& operator,(T&& t, Regularize) {
return std::forward<T>(t);
}
template<typename Result, typename... Args>
void Invoke(Result(*func)(Args...), PackedArguments<Result, Args...>& args) {
constexpr auto NumArgs = sizeof...(Args);
static_assert(NumArgs <= 19 || NumArgs == 24);
if constexpr (std::is_void_v<Result>) {
if constexpr (NumArgs == 0) {
func();
} else if constexpr (NumArgs == 1) {
func(args.a0);
} else if constexpr (NumArgs == 2) {
func(args.a0, args.a1);
} else if constexpr (NumArgs == 3) {
func(args.a0, args.a1, args.a2);
} else if constexpr (NumArgs == 4) {
func(args.a0, args.a1, args.a2, args.a3);
} else if constexpr (NumArgs == 5) {
func(args.a0, args.a1, args.a2, args.a3, args.a4);
} else if constexpr (NumArgs == 6) {
func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5);
} else if constexpr (NumArgs == 7) {
func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6);
} else if constexpr (NumArgs == 8) {
func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7);
} else if constexpr (NumArgs == 9) {
func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8);
} else if constexpr (NumArgs == 10) {
func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9);
} else if constexpr (NumArgs == 11) {
func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10);
} else if constexpr (NumArgs == 12) {
func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11);
} else if constexpr (NumArgs == 13) {
func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12);
} else if constexpr (NumArgs == 14) {
func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13);
} else if constexpr (NumArgs == 15) {
func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13, args.a14);
} else if constexpr (NumArgs == 16) {
func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13, args.a14, args.a15);
} else if constexpr (NumArgs == 17) {
func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13, args.a14, args.a15, args.a16);
} else if constexpr (NumArgs == 18) {
func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13, args.a14, args.a15, args.a16, args.a17);
} else if constexpr (NumArgs == 19) {
func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13, args.a14, args.a15, args.a16, args.a17, args.a18);
} else if constexpr (NumArgs == 24) {
func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13, args.a14, args.a15, args.a16, args.a17, args.a18, args.a19, args.a20, args.a21, args.a22, args.a23);
}
} else {
if constexpr (NumArgs == 0) {
args.rv = func();
} else if constexpr (NumArgs == 1) {
args.rv = func(args.a0);
} else if constexpr (NumArgs == 2) {
args.rv = func(args.a0, args.a1);
} else if constexpr (NumArgs == 3) {
args.rv = func(args.a0, args.a1, args.a2);
} else if constexpr (NumArgs == 4) {
args.rv = func(args.a0, args.a1, args.a2, args.a3);
} else if constexpr (NumArgs == 5) {
args.rv = func(args.a0, args.a1, args.a2, args.a3, args.a4);
} else if constexpr (NumArgs == 6) {
args.rv = func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5);
} else if constexpr (NumArgs == 7) {
args.rv = func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6);
} else if constexpr (NumArgs == 8) {
args.rv = func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7);
} else if constexpr (NumArgs == 9) {
args.rv = func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8);
} else if constexpr (NumArgs == 10) {
args.rv = func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9);
} else if constexpr (NumArgs == 11) {
args.rv = func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10);
} else if constexpr (NumArgs == 12) {
args.rv = func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11);
} else if constexpr (NumArgs == 13) {
args.rv = func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12);
} else if constexpr (NumArgs == 14) {
args.rv = func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13);
} else if constexpr (NumArgs == 15) {
args.rv = func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13, args.a14);
} else if constexpr (NumArgs == 16) {
args.rv = func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13, args.a14, args.a15);
} else if constexpr (NumArgs == 17) {
args.rv = func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13, args.a14, args.a15, args.a16);
} else if constexpr (NumArgs == 18) {
args.rv = func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13, args.a14, args.a15, args.a16, args.a17);
} else if constexpr (NumArgs == 19) {
args.rv = func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13, args.a14, args.a15, args.a16, args.a17, args.a18);
} else if constexpr (NumArgs == 24) {
args.rv = func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13, args.a14, args.a15, args.a16, args.a17, args.a18, args.a19, args.a20, args.a21, args.a22, args.a23);
}
std::conditional_t<std::is_void_v<Result>, Regularize, Result> rv;
if constexpr (NumArgs == 0) {
rv = (func(), Regularize{});
} else if constexpr (NumArgs == 1) {
rv = (func(args.a0), Regularize {});
} else if constexpr (NumArgs == 2) {
rv = (func(args.a0, args.a1), Regularize{});
} else if constexpr (NumArgs == 3) {
rv = (func(args.a0, args.a1, args.a2), Regularize{});
} else if constexpr (NumArgs == 4) {
rv = (func(args.a0, args.a1, args.a2, args.a3), Regularize{});
} else if constexpr (NumArgs == 5) {
rv = (func(args.a0, args.a1, args.a2, args.a3, args.a4), Regularize{});
} else if constexpr (NumArgs == 6) {
rv = (func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5), Regularize{});
} else if constexpr (NumArgs == 7) {
rv = (func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6), Regularize{});
} else if constexpr (NumArgs == 8) {
rv = (func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7), Regularize{});
} else if constexpr (NumArgs == 9) {
rv = (func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8), Regularize{});
} else if constexpr (NumArgs == 10) {
rv = (func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9), Regularize{});
} else if constexpr (NumArgs == 11) {
rv = (func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10), Regularize{});
} else if constexpr (NumArgs == 12) {
rv = (func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11), Regularize{});
} else if constexpr (NumArgs == 13) {
rv = (func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12), Regularize{});
} else if constexpr (NumArgs == 14) {
rv = (func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13), Regularize{});
} else if constexpr (NumArgs == 15) {
rv = (func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13, args.a14), Regularize{});
} else if constexpr (NumArgs == 16) {
rv = (func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13, args.a14, args.a15), Regularize{});
} else if constexpr (NumArgs == 17) {
rv = (func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13, args.a14, args.a15, args.a16), Regularize{});
} else if constexpr (NumArgs == 18) {
rv = (func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13, args.a14, args.a15, args.a16, args.a17), Regularize{});
} else if constexpr (NumArgs == 19) {
rv = (func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13, args.a14, args.a15, args.a16, args.a17, args.a18), Regularize{});
} else if constexpr (NumArgs == 24) {
rv = (func(args.a0, args.a1, args.a2, args.a3, args.a4, args.a5, args.a6, args.a7, args.a8, args.a9, args.a10, args.a11, args.a12, args.a13, args.a14, args.a15, args.a16, args.a17, args.a18, args.a19, args.a20, args.a21, args.a22, args.a23), Regularize{});
}
if constexpr (!std::is_void_v<Result>) {
args.rv = rv;
}
}
+85
View File
@@ -0,0 +1,85 @@
#pragma once
#include <clang/Frontend/TextDiagnosticPrinter.h>
#include <clang/Tooling/Tooling.h>
#include <llvm/Support/raw_os_ostream.h>
#include <optional>
/**
* Prints diagnostics to console like clang::TextDiagnosticPrinter.
* A copy of the first error message is stored so that it can be queried
* after compiling.
*/
class TestDiagnosticConsumer : public clang::TextDiagnosticPrinter {
bool silent;
std::optional<std::string> first_error;
public:
TestDiagnosticConsumer(bool silent_) : clang::TextDiagnosticPrinter(llvm::errs(), new clang::DiagnosticOptions), silent(silent_) {
}
void HandleDiagnostic(clang::DiagnosticsEngine::Level level,
const clang::Diagnostic& diag) override {
if (level >= clang::DiagnosticsEngine::Error && !first_error) {
llvm::SmallVector<char, 64> message;
diag.FormatDiagnostic(message);
first_error = std::string(message.begin(), message.end());
}
if (silent && level != clang::DiagnosticsEngine::Fatal) {
return;
}
clang::TextDiagnosticPrinter::HandleDiagnostic(level, diag);
}
std::optional<std::string> GetFirstError() const {
return first_error;
}
};
/**
* Run the given ToolAction on the input code.
*
* The "silent" parameter is used to suppress non-fatal diagnostics in tests that expect failure
*/
inline void run_tool(clang::tooling::ToolAction& action, std::string_view code, bool silent = false) {
const char* memory_filename = "gen_input.cpp";
auto adjuster = clang::tooling::getClangStripDependencyFileAdjuster();
std::vector<std::string> args = { "clang-tool", "-fsyntax-only", "-std=c++17", "-Werror", "-I.", memory_filename };
// Corresponds to the content of GeneratorInterface.h
const char* common_header_code = R"(namespace fexgen {
struct returns_guest_pointer {};
struct custom_host_impl {};
struct callback_annotation_base { bool prevent_multiple; };
struct callback_stub : callback_annotation_base {};
struct callback_guest : callback_annotation_base {};
} // namespace fexgen
)";
llvm::IntrusiveRefCntPtr<llvm::vfs::OverlayFileSystem> overlay_fs(new llvm::vfs::OverlayFileSystem(llvm::vfs::getRealFileSystem()));
llvm::IntrusiveRefCntPtr<llvm::vfs::InMemoryFileSystem> memory_fs(new llvm::vfs::InMemoryFileSystem);
overlay_fs->pushOverlay(memory_fs);
memory_fs->addFile(memory_filename, 0, llvm::MemoryBuffer::getMemBufferCopy(code));
memory_fs->addFile("thunks_common.h", 0, llvm::MemoryBuffer::getMemBufferCopy(common_header_code));
llvm::IntrusiveRefCntPtr<clang::FileManager> files(new clang::FileManager(clang::FileSystemOptions(), overlay_fs));
auto invocation = clang::tooling::ToolInvocation(args, &action, files.get(), std::make_shared<clang::PCHContainerOperations>());
TestDiagnosticConsumer consumer(silent);
invocation.setDiagnosticConsumer(&consumer);
invocation.run();
if (auto error = consumer.GetFirstError()) {
throw std::runtime_error(*error);
}
}
inline void run_tool(std::unique_ptr<clang::tooling::ToolAction> action, std::string_view code, bool silent = false) {
return run_tool(*action, code, silent);
}
+2 -72
View File
@@ -3,17 +3,16 @@
#include <clang/ASTMatchers/ASTMatchers.h>
#include <clang/ASTMatchers/ASTMatchFinder.h>
#include <clang/Frontend/CompilerInstance.h>
#include <clang/Frontend/TextDiagnosticPrinter.h>
#include <clang/Tooling/Tooling.h>
#include <llvm/Support/raw_os_ostream.h>
#include <interface.h>
#include <filesystem>
#include <fstream>
#include <string_view>
#include "common.h"
/**
* This class parses its input code and stores it alongside its AST representation.
*
@@ -177,75 +176,6 @@ public:
}
};
class TestDiagnosticConsumer : public clang::TextDiagnosticPrinter {
bool silent;
std::optional<std::string> first_error;
public:
TestDiagnosticConsumer(bool silent_) : clang::TextDiagnosticPrinter(llvm::errs(), new clang::DiagnosticOptions), silent(silent_) {
}
void HandleDiagnostic(clang::DiagnosticsEngine::Level level,
const clang::Diagnostic& diag) override {
if (level >= clang::DiagnosticsEngine::Error && !first_error) {
llvm::SmallVector<char, 64> message;
diag.FormatDiagnostic(message);
first_error = std::string(message.begin(), message.end());
}
if (silent && level != clang::DiagnosticsEngine::Fatal) {
return;
}
clang::TextDiagnosticPrinter::HandleDiagnostic(level, diag);
}
std::optional<std::string> GetFirstError() const {
return first_error;
}
};
/**
* The "silent" parameter is used to suppress non-fatal diagnostics in tests that expect failure
*/
static void run_tool(clang::tooling::ToolAction& action, std::string_view code, bool silent = false) {
const char* memory_filename = "gen_input.cpp";
auto adjuster = clang::tooling::getClangStripDependencyFileAdjuster();
std::vector<std::string> args = { "clang-tool", "-fsyntax-only", "-std=c++17", "-Werror", "-I.", memory_filename };
const char* common_header_code = R"(namespace fexgen {
struct returns_guest_pointer {};
struct custom_host_impl {};
struct callback_annotation_base { bool prevent_multiple; };
struct callback_stub : callback_annotation_base {};
struct callback_guest : callback_annotation_base {};
} // namespace fexgen
)";
llvm::IntrusiveRefCntPtr<llvm::vfs::OverlayFileSystem> overlay_fs(new llvm::vfs::OverlayFileSystem(llvm::vfs::getRealFileSystem()));
llvm::IntrusiveRefCntPtr<llvm::vfs::InMemoryFileSystem> memory_fs(new llvm::vfs::InMemoryFileSystem);
overlay_fs->pushOverlay(memory_fs);
memory_fs->addFile(memory_filename, 0, llvm::MemoryBuffer::getMemBufferCopy(code));
memory_fs->addFile("thunks_common.h", 0, llvm::MemoryBuffer::getMemBufferCopy(common_header_code));
llvm::IntrusiveRefCntPtr<clang::FileManager> files(new clang::FileManager(clang::FileSystemOptions(), overlay_fs));
auto invocation = clang::tooling::ToolInvocation(args, &action, files.get(), std::make_shared<clang::PCHContainerOperations>());
TestDiagnosticConsumer consumer(silent);
invocation.setDiagnosticConsumer(&consumer);
invocation.run();
if (auto error = consumer.GetFirstError()) {
throw std::runtime_error(*error);
}
}
static void run_tool(std::unique_ptr<clang::tooling::ToolAction> action, std::string_view code, bool silent = false) {
return run_tool(*action, code, silent);
}
SourceWithAST::SourceWithAST(std::string_view input) : code(input) {
// Call run_tool with a ToolAction that assigns this->ast