From 289844eeddc802bc3a2b8b9daf6dca033ac09e81 Mon Sep 17 00:00:00 2001 From: Lucian Popescu Date: Tue, 29 Sep 2026 14:24:11 +0100 Subject: [PATCH 1/7] Move Mapper::ToString into printer.cpp --- cpp2rust/converter/converter.cpp | 22 +- cpp2rust/converter/converter_lib.cpp | 37 +- cpp2rust/converter/converter_lib.h | 5 + cpp2rust/converter/mapper.cpp | 404 +----------------- cpp2rust/converter/mapper.h | 14 - .../converter/models/converter_refcount.cpp | 5 +- cpp2rust/converter/printer.cpp | 380 ++++++++++++++++ cpp2rust/converter/printer.h | 23 + cpp2rust/cpp_rule_preprocessor.cpp | 38 +- 9 files changed, 487 insertions(+), 441 deletions(-) create mode 100644 cpp2rust/converter/printer.cpp create mode 100644 cpp2rust/converter/printer.h diff --git a/cpp2rust/converter/converter.cpp b/cpp2rust/converter/converter.cpp index d5c6ed73b..8a2ec46eb 100644 --- a/cpp2rust/converter/converter.cpp +++ b/cpp2rust/converter/converter.cpp @@ -22,6 +22,7 @@ #include "converter/converter_lib.h" #include "converter/lex.h" #include "converter/mapper.h" +#include "converter/printer.h" namespace cpp2rust { std::unordered_map Converter::inner_structs_; @@ -649,7 +650,7 @@ bool Converter::RecordDerivesDefault(const clang::RecordDecl *decl) { } // Records that contain std::array do not derive Default - if (Mapper::ToString(f->getType()).contains("std::array")) { + if (Printer::ToString(ctx_, f->getType()).contains("std::array")) { return false; } @@ -1590,7 +1591,7 @@ const clang::Expr *Converter::GetParentExpr(const clang::Expr *expr) { bool Converter::GetFmtArg(clang::Expr *arg, std::string &fmt, std::string &fmt_args, const char *&fmt_trait, std::string &fmt_width) { - std::string arg_str = Mapper::ToString(arg); + std::string arg_str = Printer::ToString(ctx_, arg); if (auto *str_lit = clang::dyn_cast(arg->IgnoreImplicit())) { if (!IsAsciiStringLiteral(str_lit)) { @@ -1637,7 +1638,7 @@ bool Converter::GetRawArg(clang::Expr *arg, std::string &raw_args) { std::string str = ToString(arg); raw_args += "(&(" + str + ").iter().take((" + str + ").len() - 1).map(|&c| c as u8).collect::>()[..]"; - } else if (Mapper::ToString(arg).contains("std::endl")) { + } else if (Printer::ToString(ctx_, arg).contains("std::endl")) { raw_args += "(&[b'\\n']"; } else if (clang::isa(arg->IgnoreImplicit())) { raw_args += "(b" + GetEscapedStringLiteral(arg); @@ -1724,7 +1725,7 @@ void Converter::ConvertCallToOstream(clang::CallExpr *expr) { void Converter::ConvertPrintf(clang::CallExpr *expr) { bool is_fprintf = - Mapper::ToString(expr->getCallee()).starts_with("int fprintf"); + Printer::ToString(ctx_, expr->getCallee()).starts_with("int fprintf"); StrCat("printf("); for (unsigned i = is_fprintf; i < expr->getNumArgs(); ++i) { @@ -2168,7 +2169,7 @@ std::optional Converter::ConvertCallExpr(clang::CallExpr *expr) { auto *callee = expr->getCallee(); - if (auto fn = Mapper::ToString(callee); + if (auto fn = Printer::ToString(ctx_, callee); fn.starts_with("int printf") || fn.starts_with("int fprintf")) { ConvertPrintf(expr); } else if (IsTransparentStdCall(expr)) { @@ -3957,7 +3958,7 @@ std::string Converter::GetArrayDefaultAsString(clang::QualType qual_type) { clang::dyn_cast(qual_type)) { return GetDefaultAsString(array_type->getElementType()); } - if (Mapper::ToString(qual_type).contains("std::array")) { + if (Printer::ToString(ctx_, qual_type).contains("std::array")) { assert(GetTemplateArgs(qual_type).has_value()); auto template_args = *GetTemplateArgs(qual_type); assert(template_args.size() == 2); @@ -4080,11 +4081,11 @@ Converter::GetOverloadedFunctionName(const clang::FunctionDecl *decl) { name += '_'; switch (arg.getKind()) { case clang::TemplateArgument::Type: - name += Mapper::ToRustName( + name += Printer::ToRustName( arg.getAsType().getCanonicalType().getAsString()); break; case clang::TemplateArgument::Integral: - name += Mapper::ToRustName( + name += Printer::ToRustName( std::string(GetNumAsString(arg.getAsIntegral()))); break; default: @@ -4128,7 +4129,8 @@ std::string Converter::GetRecordName(const clang::NamedDecl *decl) const { if (auto it = inner_structs_.find(ID); it != inner_structs_.end()) { return it->second; } - return Mapper::ToRustName(Mapper::ToString(Mapper::GetTypeForDecl(decl))); + return Printer::ToRustName( + Printer::ToString(ctx_, GetTypeForDecl(ctx_, decl))); } std::vector @@ -4201,7 +4203,7 @@ void Converter::ConvertVarInit(clang::QualType qual_type, clang::Expr *expr) { !Mapper::Contains( clang::cast(ctor->getArg(0)->IgnoreCasts()) ->getCallee()) && - Mapper::ToString(ctor->getConstructor()->getThisType()) == + Printer::ToString(ctx_, ctor->getConstructor()->getThisType()) == "std::string") { { PushParen paren(*this); diff --git a/cpp2rust/converter/converter_lib.cpp b/cpp2rust/converter/converter_lib.cpp index 4c01a1556..53d3c6f20 100644 --- a/cpp2rust/converter/converter_lib.cpp +++ b/cpp2rust/converter/converter_lib.cpp @@ -26,6 +26,7 @@ #include "converter/lex.h" #include "converter/mapper.h" +#include "converter/printer.h" // https://doc.rust-lang.org/reference/keywords.html static const char rust_keywords[][12] = { @@ -699,6 +700,7 @@ static std::string GetParamSignature(const clang::Decl *decl) { } static std::string GetLexicalSpecializationID(const clang::Decl *decl) { + auto &ctx = decl->getASTContext(); std::string id; if (const auto *var = clang::dyn_cast(decl)) { @@ -706,18 +708,18 @@ static std::string GetLexicalSpecializationID(const clang::Decl *decl) { } if (const auto *self = clang::dyn_cast(decl)) { - id += Mapper::ToString(Mapper::GetTypeForDecl(self)); + id += Printer::ToString(ctx, GetTypeForDecl(ctx, self)); } if (const auto *spec = clang::dyn_cast( decl->getLexicalDeclContext()); spec && decl->getLexicalDeclContext() != decl->getDeclContext()) { - id += Mapper::ToString(Mapper::GetTypeForDecl(spec)); + id += Printer::ToString(ctx, GetTypeForDecl(ctx, spec)); } for (const auto *dc = decl->getDeclContext(); dc; dc = dc->getParent()) { if (const auto *spec = clang::dyn_cast(dc)) { - id += Mapper::ToString(Mapper::GetTypeForDecl(spec)); + id += Printer::ToString(ctx, GetTypeForDecl(ctx, spec)); } if (const auto *fn = clang::dyn_cast(dc); fn && fn->getTemplateSpecializationArgs()) { @@ -737,6 +739,35 @@ std::string GetMethodID(const clang::CXXMethodDecl *decl) { return decl->getQualifiedNameAsString() + GetID(decl); } +clang::QualType GetTypeForDecl(clang::ASTContext &ctx, + const clang::NamedDecl *decl) { + if (const auto *spec = + llvm::dyn_cast(decl)) { + llvm::ArrayRef args = + spec->getTemplateArgs().asArray(); + llvm::SmallVector canon(args.begin(), + args.end()); + ctx.canonicalizeTemplateArguments(canon); + + return ctx.getTemplateSpecializationType( + clang::ElaboratedTypeKeyword::None, + clang::TemplateName(spec->getSpecializedTemplate()), args, canon); + } + + const auto *rdecl = llvm::dyn_cast(decl); + assert(rdecl && "Unsupported decl type"); + + return ctx.getTagType(clang::ElaboratedTypeKeyword::None, + rdecl->getQualifier(), rdecl, /*OwnsTag*/ false); +} + +bool HasFunctionParameterPack(const clang::FunctionDecl *decl) { + if (auto *primary = decl->getPrimaryTemplate()) { + decl = primary->getTemplatedDecl(); + } + return decl->getNumParams() && decl->parameters().back()->isParameterPack(); +} + std::string DisambiguateAnonymousTag(const clang::TagDecl *tag) { if (!tag) { return ""; diff --git a/cpp2rust/converter/converter_lib.h b/cpp2rust/converter/converter_lib.h index 6e0a8cc98..1cf00d4fd 100644 --- a/cpp2rust/converter/converter_lib.h +++ b/cpp2rust/converter/converter_lib.h @@ -167,6 +167,11 @@ std::string GetNamedDeclAsString(const clang::NamedDecl *decl); std::string DisambiguateAnonymousTag(const clang::TagDecl *tag); +clang::QualType GetTypeForDecl(clang::ASTContext &ctx, + const clang::NamedDecl *decl); + +bool HasFunctionParameterPack(const clang::FunctionDecl *decl); + const char *AccessSpecifierAsString(clang::AccessSpecifier spec); template llvm::SmallString<16> GetNumAsString(const T &num) { diff --git a/cpp2rust/converter/mapper.cpp b/cpp2rust/converter/mapper.cpp index 0464976ce..f2f1a0f82 100644 --- a/cpp2rust/converter/mapper.cpp +++ b/cpp2rust/converter/mapper.cpp @@ -4,21 +4,19 @@ #include "converter/mapper.h" #include -#include #include -#include #include #include #include #include #include -#include #include #include #include #include "converter/converter_lib.h" +#include "converter/printer.h" #include "converter/translation_rule.h" namespace cpp2rust::Mapper { @@ -34,18 +32,6 @@ std::unordered_multimap std::unordered_multimap types_; // src -> TypeRule -clang::PrintingPolicy getPrintPolicy() { - assert(ctx_); - clang::PrintingPolicy policy(ctx_->getLangOpts()); - policy.Bool = true; - policy.SuppressTagKeyword = true; - policy.SuppressScope = false; - policy.FullyQualifiedName = true; - policy.SuppressUnwrittenScope = true; - policy.UsePreferredNames = true; - return policy; -} - std::string GetExprMapKey(const std::string &str) { // Extract the function name from something like // const T1 & std::foo::fn_name(args) @@ -74,8 +60,6 @@ std::string GetExprMapKey(const std::string &str) { return result; } -constexpr const char kPackMarker[] = "&&..."; - std::string GetTypeMapKey(const std::string &str) { auto n = str.find_first_of("<["); if (n == std::string::npos || str[n] == '<') { @@ -394,7 +378,7 @@ TranslationRule::ExprRule *search(const clang::Expr *expr) { if (RefersToUserDefinedDecl(expr)) { return nullptr; } - auto qualified_name = ToString(expr); + auto qualified_name = Printer::ToString(*ctx_, expr); auto [rule, subs] = search(exprs_, qualified_name, GetExprMapKey(qualified_name)); log() << "search expr " << qualified_name << ", result:\n"; @@ -408,13 +392,14 @@ TranslationRule::ExprRule *search(const clang::Expr *expr) { std::pair>> search(clang::QualType qual_type) { - auto sugared = ToString(qual_type, ScalarSugar::kPreserve); + auto sugared = + Printer::ToString(*ctx_, qual_type, Printer::ScalarSugar::kPreserve); if (auto res = search(types_, sugared, GetTypeMapKey(sugared)); res.first) { log() << "search type " << sugared << ", result: " << res.first->type_info.type << '\n'; return res; } - auto type = ToString(qual_type); + auto type = Printer::ToString(*ctx_, qual_type); if (type == sugared) { log() << "search type " << type << ", result: None\n"; return {}; @@ -458,49 +443,6 @@ void addRulesFromDirectory(const std::filesystem::path &dir, Model model) { } } -clang::QualType normalizeQualType(clang::QualType qual_type) { - assert(ctx_); - - bool isLRef = qual_type->isLValueReferenceType(); - bool isRRef = qual_type->isRValueReferenceType(); - qual_type = qual_type.getNonReferenceType(); - - clang::Qualifiers qualifiers = qual_type.getQualifiers(); - - while (true) { - if (const auto *attributed = - llvm::dyn_cast(qual_type)) { - qual_type = attributed->getModifiedType(); - continue; - } - if (const auto *dcltype = llvm::dyn_cast(qual_type)) { - qual_type = dcltype->getUnderlyingType(); - continue; - } - break; - } - - if (llvm::isa(qual_type)) { - qual_type = qual_type.getCanonicalType(); - } - - qual_type = qual_type.withFastQualifiers(qualifiers.getFastQualifiers()); - if (qualifiers.hasNonFastQualifiers()) { - qual_type = ctx_->getQualifiedType(qual_type, qualifiers); - } - - if (isLRef) { - qual_type = ctx_->getLValueReferenceType(qual_type); - } - - if (isRRef) { - qual_type = ctx_->getRValueReferenceType(qual_type); - } - - return qual_type.getCanonicalType().getUnqualifiedType().getDesugaredType( - *ctx_); -} - std::string mapTypeStringRecursive(const std::string &cpp_type) { auto [rule, subs] = search(types_, cpp_type, GetTypeMapKey(cpp_type)); if (!rule) { @@ -515,24 +457,6 @@ std::string mapTypeStringRecursive(const std::string &cpp_type) { return instantiateTgt(subs, rule->type_info.type); } -std::string normalizeTranslationRule(std::string rule) { - // Detach pointer from double reference. Useful for matching translation - // rules. - ReplaceAll(rule, "*&&", "* &&"); - - static const std::array, 1> - normalization_rules{{ - // Ignore constant template parameters, i.e. replace them with _. - {std::regex(R"(\b\d+\b)"), "_"}, - }}; - - for (const auto &r : normalization_rules) { - rule = std::regex_replace(rule, r.first, r.second); - } - - return rule; -} - } // namespace PushASTContext::PushASTContext(clang::ASTContext &ctx) : prev_(ctx_) { @@ -566,7 +490,7 @@ bool IsLibcPassthrough(const clang::Expr *expr) { std::string MapFunctionName(const clang::FunctionDecl *decl) { assert(decl); if (!IsUserDefinedDecl(decl) && - exprs_.contains(GetExprMapKey(ToString(decl)))) { + exprs_.contains(GetExprMapKey(Printer::ToString(*ctx_, decl)))) { return std::format("libcc2rs::{}_{}", decl->getNameAsString(), model_ == Model::kRefCount ? "refcount" : "unsafe"); } @@ -574,7 +498,7 @@ std::string MapFunctionName(const clang::FunctionDecl *decl) { } std::string InstantiateTemplate(const clang::Expr *expr, unsigned n) { - auto expr_str = ToString(expr); + auto expr_str = Printer::ToString(*ctx_, expr); auto [rule, subs] = search(exprs_, expr_str, GetExprMapKey(expr_str)); auto text = std::format("T{}", n); if (!rule) { @@ -646,7 +570,7 @@ const TranslationRule::TypeInfo &GetParamInfo(const clang::Expr *expr, } std::string GetParamType(const clang::Expr *expr, unsigned index) { - auto expr_str = ToString(expr); + auto expr_str = Printer::ToString(*ctx_, expr); auto [rule, subs] = search(exprs_, expr_str, GetExprMapKey(expr_str)); for (auto &ty : subs) { if (ty) { @@ -660,30 +584,9 @@ bool ParamIsPointer(const clang::Expr *expr, unsigned index) { return GetParamInfo(expr, index).is_pointer(); } -clang::QualType GetTypeForDecl(const clang::NamedDecl *decl) { - if (const auto *spec = - llvm::dyn_cast(decl)) { - llvm::ArrayRef args = - spec->getTemplateArgs().asArray(); - llvm::SmallVector canon(args.begin(), - args.end()); - ctx_->canonicalizeTemplateArguments(canon); - - return ctx_->getTemplateSpecializationType( - clang::ElaboratedTypeKeyword::None, - clang::TemplateName(spec->getSpecializedTemplate()), args, canon); - } - - const auto *rdecl = llvm::dyn_cast(decl); - assert(rdecl && "Unsupported decl type"); - - return ctx_->getTagType(clang::ElaboratedTypeKeyword::None, - rdecl->getQualifier(), rdecl, /*OwnsTag*/ false); -} - void AddRuleForUserDefinedType(clang::NamedDecl *decl) { - auto cpp_name = ToString(GetTypeForDecl(decl)); - auto rs_name = ToRustName(cpp_name); + auto cpp_name = Printer::ToString(*ctx_, GetTypeForDecl(*ctx_, decl)); + auto rs_name = Printer::ToRustName(cpp_name); AddTypeRule(cpp_name, TranslationRule::TypeRule::Plain(rs_name)); @@ -725,293 +628,6 @@ void AddRuleForUserDefinedType(clang::NamedDecl *decl) { } } -std::string ToRustName(std::string name) { - ReplaceAll(name, "::", "_"); - ReplaceAll(name, "*", "ptr"); - ReplaceAll(name, "&", "ref"); - ReplaceAll(name, "[", "arr"); - ReplaceAll(name, "]", "arr"); - ReplaceAll(name, "-", "neg"); - for (auto &c : name) { - if (!std::isalnum(c) && c != '_') { - c = '_'; - } - } - return name; -} - -std::string ToString(clang::QualType qual_type, ScalarSugar sugar) { - assert(ctx_); - - if (sugar == ScalarSugar::kPreserve) { - clang::QualType t = qual_type; - if (const auto *decltype_type = - clang::dyn_cast(t.getTypePtr())) { - t = decltype_type->getUnderlyingType(); - } - if (const auto *typedef_type = t->getAs()) { - if (t.getCanonicalType()->isBuiltinType()) { - return typedef_type->getDecl()->getNameAsString(); - } - } else if (const auto *predef = t->getAs()) { - return predef->getIdentifier()->getName().str(); - } else if (const auto *ptr = t->getAs()) { - auto pointee = ptr->getPointeeType(); - auto canonical = pointee.getCanonicalType().getDesugaredType(*ctx_); - bool builtin_alias = canonical->isBuiltinType() && - (pointee->getAs() || - pointee->getAs()); - if (!builtin_alias && Map(pointee) == Map(canonical)) { - pointee = canonical; - } - std::string out; - llvm::raw_string_ostream os(out); - ctx_->getPointerType(pointee).print(os, getPrintPolicy()); - return normalizeTranslationRule(std::move(out)); - } - } - - if (auto cxx_record_decl = qual_type->getAsCXXRecordDecl()) { - if (cxx_record_decl->isLambda()) { - return ToString(cxx_record_decl->getLambdaCallOperator()); - } - } - - if (auto *tag = qual_type->getAsTagDecl(); - tag && !tag->getIdentifier() && !tag->getTypedefNameForAnonDecl()) { - return ToString(clang::cast(tag)); - } - - if (auto *tag = qual_type->getAsTagDecl(); - tag && tag->getIdentifier() && - tag->getDeclContext()->isFunctionOrMethod()) { - return GetNamedDeclAsString(tag); - } - - if (auto renamed = DisambiguateAnonymousTag(qual_type->getAsTagDecl()); - !renamed.empty()) { - return renamed; - } - - std::string type; - llvm::raw_string_ostream os(type); - normalizeQualType(qual_type).print(os, getPrintPolicy()); - return normalizeTranslationRule(std::move(type)); -} - -bool HasFunctionParameterPack(const clang::FunctionDecl *decl) { - if (auto *primary = decl->getPrimaryTemplate()) { - decl = primary->getTemplatedDecl(); - } - return decl->getNumParams() && decl->parameters().back()->isParameterPack(); -} - -std::string ToString(const clang::NamedDecl *decl) { - if (auto *record = clang::dyn_cast(decl); - record && !record->getIdentifier()) { - if (auto renamed = DisambiguateAnonymousTag(record); !renamed.empty()) { - return renamed; - } - if (auto *typedef_decl = record->getTypedefNameForAnonDecl()) { - return ToString(clang::cast(typedef_decl)); - } - return GetNamedDeclAsString(record); - } - - if (auto *enum_decl = clang::dyn_cast(decl)) { - if (auto renamed = DisambiguateAnonymousTag(enum_decl); !renamed.empty()) { - return renamed; - } - if (!enum_decl->getIdentifier() && - !enum_decl->getTypedefNameForAnonDecl()) { - return GetNamedDeclAsString(enum_decl); - } - } - - std::string out; - llvm::raw_string_ostream os(out); - - const clang::FunctionDecl *func_decl = nullptr; - if (auto *template_decl = llvm::dyn_cast(decl)) { - func_decl = template_decl->getTemplatedDecl(); - } else { - func_decl = llvm::dyn_cast_or_null(decl); - } - - if (!func_decl) { - decl->printQualifiedName(os, getPrintPolicy()); - return normalizeTranslationRule(std::move(out)); - } - - os << ToString(func_decl->getReturnType()) << ' '; - if (const auto op = func_decl->getOverloadedOperator(); - op >= clang::OverloadedOperatorKind::OO_LessLess && - op <= clang::OverloadedOperatorKind::OO_GreaterGreaterEqual) { - // ensure matchTemplate does not consider these operator names when matching - func_decl->getQualifier().print(os, getPrintPolicy()); - os << "operator "; - switch (op) { - case clang::OverloadedOperatorKind::OO_LessLess: - os << "shl"; - break; - case clang::OverloadedOperatorKind::OO_GreaterGreater: - os << "shr"; - break; - case clang::OverloadedOperatorKind::OO_LessLessEqual: - os << "shleq"; - break; - case clang::OverloadedOperatorKind::OO_GreaterGreaterEqual: - os << "shreq"; - break; - default: - assert(0 && "Unexpected overloaded operator kind"); - } - } else if (const auto *method_decl = - llvm::dyn_cast(func_decl)) { - if (method_decl->getParent()->isLambda() && - method_decl->getOverloadedOperator() == clang::OO_Call) { - func_decl->printName(os, getPrintPolicy()); - } else { - func_decl->printQualifiedName(os, getPrintPolicy()); - } - } else { - func_decl->printQualifiedName(os, getPrintPolicy()); - } - - bool has_pack = HasFunctionParameterPack(func_decl); - unsigned num_params = func_decl->getNumParams(); - if (has_pack) { - const auto *primary = func_decl->getPrimaryTemplate(); - num_params = - (primary ? primary->getTemplatedDecl() : func_decl)->getNumParams() - 1; - } - - os << '('; - for (unsigned i = 0; i < num_params; ++i) { - if (i) { - os << ", "; - } - os << ToString(func_decl->getParamDecl(i)->getType()); - } - if (has_pack) { - if (num_params) { - os << ", "; - } - os << kPackMarker; - } - if (func_decl->isVariadic()) { - if (func_decl->getNumParams()) { - os << ", "; - } - os << "..."; - } - os << ')'; - - if (const auto *method_decl = - llvm::dyn_cast(func_decl)) { - if (method_decl->isConst()) { - os << " const"; - } - if (method_decl->isVolatile()) { - os << " volatile"; - } - switch (method_decl->getRefQualifier()) { - case clang::RQ_LValue: - os << " &"; - break; - case clang::RQ_RValue: - os << " &&"; - break; - default: - break; - } - } - - return normalizeTranslationRule(std::move(out)); -} - -std::string ToString(const clang::Expr *expr) { - if (!expr) { - assert(0 && "!expr"); - } - - expr = expr->IgnoreParenImpCasts(); - - if (llvm::isa(expr) && - expr->getBeginLoc().isMacroID()) { - auto &sm = ctx_->getSourceManager(); - auto name = clang::Lexer::getImmediateMacroName(expr->getBeginLoc(), sm, - ctx_->getLangOpts()); - if (!name.empty()) { - return name.str(); - } - } - - if (const auto *CE = llvm::dyn_cast(expr)) { - if (const auto *decl = CE->getDirectCallee()) { - return ToString(decl); - } - } - - if (const auto *ctor = llvm::dyn_cast(expr)) { - if (const auto *ctor_decl = ctor->getConstructor()) { - return ToString(ctor_decl); - } - assert(0 && "expr is a CXXConstructExpr but could not get constructor"); - } - - if (const auto *ME = llvm::dyn_cast(expr)) { - if (const auto *member_decl = - llvm::dyn_cast(ME->getMemberDecl())) { - if (const auto *method_decl = - llvm::dyn_cast(member_decl)) { - return ToString(method_decl); - } - if (ME->isArrow()) { - auto *base = ME->getBase()->IgnoreParenImpCasts(); - if (auto *op = llvm::dyn_cast(base)) { - if (op->getOperator() == clang::OO_Arrow) { - return ToString(op->getArg(0)->getType()) + "->" + - ToString(member_decl); - } - } - } else if (auto for_range = GetParentForRange(*ctx_, ME)) { - if (ToString(for_range->getRangeInit()->getType()) - .starts_with("std::map<")) { - auto iter_type = GetForRangeIteratorType(for_range); - if (!iter_type.isNull()) { - return ToString(iter_type) + "->" + ToString(member_decl); - } - } - } - return ToString(member_decl); - } - assert(0 && "expr is a MemberExpr but could not get named decl"); - } - - if (const auto *decl_ref = llvm::dyn_cast(expr)) { - if (const auto *named_decl = - llvm::dyn_cast(decl_ref->getDecl())) { - if (const auto *tmpl_decl = - llvm::dyn_cast(named_decl)) { - return ToString(tmpl_decl->getTemplatedDecl()); - } - return ToString(named_decl); - } - return ""; - } - - if (const auto *uop = llvm::dyn_cast(expr)) { - auto sub = ToString(uop->getSubExpr()); - std::string_view opcode = - clang::UnaryOperator::getOpcodeStr(uop->getOpcode()); - return uop->isPostfix() ? std::format("{}{}", sub, opcode) - : std::format("{}{}", opcode, sub); - } - - return "Unhandled case in ToString"; -} - void LoadTranslationRules(Model model, clang::ASTContext &ctx, const std::string &rules_dir) { ctx_ = &ctx; diff --git a/cpp2rust/converter/mapper.h b/cpp2rust/converter/mapper.h index dbd6380f8..4864d299f 100644 --- a/cpp2rust/converter/mapper.h +++ b/cpp2rust/converter/mapper.h @@ -41,20 +41,6 @@ bool MapsToRefcountPointer(clang::QualType qual_type); const std::vector *MappedDerives(clang::QualType qual_type); void SetDerives(clang::QualType qual_type, std::vector derives); -enum class ScalarSugar { - kDesugar, - kPreserve, -}; - -bool HasFunctionParameterPack(const clang::FunctionDecl *decl); - -clang::QualType GetTypeForDecl(const clang::NamedDecl *decl); -std::string ToString(clang::QualType qual_type, - ScalarSugar sugar = ScalarSugar::kDesugar); -std::string ToString(const clang::Expr *expr); -std::string ToString(const clang::NamedDecl *decl); -std::string ToRustName(std::string name); - void LoadTranslationRules(Model model, clang::ASTContext &ctx, const std::string &rules_dir); void AddRuleForUserDefinedType(clang::NamedDecl *decl); diff --git a/cpp2rust/converter/models/converter_refcount.cpp b/cpp2rust/converter/models/converter_refcount.cpp index 6ce2ca605..c7a33aab1 100644 --- a/cpp2rust/converter/models/converter_refcount.cpp +++ b/cpp2rust/converter/models/converter_refcount.cpp @@ -15,6 +15,7 @@ #include "converter/converter_lib.h" #include "converter/lex.h" #include "converter/mapper.h" +#include "converter/printer.h" namespace cpp2rust { std::map @@ -1063,7 +1064,7 @@ static std::vector printf2fmt(std::string &format) { void ConverterRefCount::ConvertPrintf(clang::CallExpr *expr) { bool is_fprintf = - Mapper::ToString(expr->getCallee()).starts_with("int fprintf"); + Printer::ToString(ctx_, expr->getCallee()).starts_with("int fprintf"); std::string format; if (auto *str = clang::dyn_cast( expr->getArg(is_fprintf)->IgnoreImplicit())) { @@ -1076,7 +1077,7 @@ void ConverterRefCount::ConvertPrintf(clang::CallExpr *expr) { } bool ends_newline = format.ends_with("\\n\""); - auto fd = is_fprintf ? Mapper::ToString(expr->getArg(0)) : "stdout"; + auto fd = is_fprintf ? Printer::ToString(ctx_, expr->getArg(0)) : "stdout"; if (fd == "stdout" || fd == "__stdoutp") { StrCat(ends_newline ? "println!(" : "print!("); } else if (fd == "stderr" || fd == "__stderrp") { diff --git a/cpp2rust/converter/printer.cpp b/cpp2rust/converter/printer.cpp new file mode 100644 index 000000000..914ec5a9e --- /dev/null +++ b/cpp2rust/converter/printer.cpp @@ -0,0 +1,380 @@ +// Copyright (c) 2022-present INESC-ID. +// Distributed under the MIT license that can be found in the LICENSE file. + +#include "converter/printer.h" + +#include +#include +#include +#include + +#include +#include +#include +#include +#include + +#include "converter/converter_lib.h" +#include "converter/mapper.h" + +namespace cpp2rust::Printer { + +namespace { + +constexpr const char kPackMarker[] = "&&..."; + +clang::PrintingPolicy getPrintPolicy(clang::ASTContext &ctx) { + clang::PrintingPolicy policy(ctx.getLangOpts()); + policy.Bool = true; + policy.SuppressTagKeyword = true; + policy.SuppressScope = false; + policy.FullyQualifiedName = true; + policy.SuppressUnwrittenScope = true; + policy.UsePreferredNames = true; + return policy; +} + +clang::QualType normalizeQualType(clang::ASTContext &ctx, + clang::QualType qual_type) { + + bool isLRef = qual_type->isLValueReferenceType(); + bool isRRef = qual_type->isRValueReferenceType(); + qual_type = qual_type.getNonReferenceType(); + + clang::Qualifiers qualifiers = qual_type.getQualifiers(); + + while (true) { + if (const auto *attributed = + llvm::dyn_cast(qual_type)) { + qual_type = attributed->getModifiedType(); + continue; + } + if (const auto *dcltype = llvm::dyn_cast(qual_type)) { + qual_type = dcltype->getUnderlyingType(); + continue; + } + break; + } + + if (llvm::isa(qual_type)) { + qual_type = qual_type.getCanonicalType(); + } + + qual_type = qual_type.withFastQualifiers(qualifiers.getFastQualifiers()); + if (qualifiers.hasNonFastQualifiers()) { + qual_type = ctx.getQualifiedType(qual_type, qualifiers); + } + + if (isLRef) { + qual_type = ctx.getLValueReferenceType(qual_type); + } + + if (isRRef) { + qual_type = ctx.getRValueReferenceType(qual_type); + } + + return qual_type.getCanonicalType().getUnqualifiedType().getDesugaredType( + ctx); +} + +std::string normalizeTranslationRule(std::string rule) { + // Detach pointer from double reference. Useful for matching translation + // rules. + ReplaceAll(rule, "*&&", "* &&"); + + static const std::array, 1> + normalization_rules{{ + // Ignore constant template parameters, i.e. replace them with _. + {std::regex(R"(\b\d+\b)"), "_"}, + }}; + + for (const auto &r : normalization_rules) { + rule = std::regex_replace(rule, r.first, r.second); + } + + return rule; +} + +} // namespace + +std::string ToRustName(std::string name) { + ReplaceAll(name, "::", "_"); + ReplaceAll(name, "*", "ptr"); + ReplaceAll(name, "&", "ref"); + ReplaceAll(name, "[", "arr"); + ReplaceAll(name, "]", "arr"); + ReplaceAll(name, "-", "neg"); + for (auto &c : name) { + if (!std::isalnum(c) && c != '_') { + c = '_'; + } + } + return name; +} + +std::string ToString(clang::ASTContext &ctx, clang::QualType qual_type, + ScalarSugar sugar) { + + if (sugar == ScalarSugar::kPreserve) { + clang::QualType t = qual_type; + if (const auto *decltype_type = + clang::dyn_cast(t.getTypePtr())) { + t = decltype_type->getUnderlyingType(); + } + if (const auto *typedef_type = t->getAs()) { + if (t.getCanonicalType()->isBuiltinType()) { + return typedef_type->getDecl()->getNameAsString(); + } + } else if (const auto *predef = t->getAs()) { + return predef->getIdentifier()->getName().str(); + } else if (const auto *ptr = t->getAs()) { + auto pointee = ptr->getPointeeType(); + auto canonical = pointee.getCanonicalType().getDesugaredType(ctx); + bool builtin_alias = canonical->isBuiltinType() && + (pointee->getAs() || + pointee->getAs()); + if (!builtin_alias && Mapper::Map(pointee) == Mapper::Map(canonical)) { + pointee = canonical; + } + std::string out; + llvm::raw_string_ostream os(out); + ctx.getPointerType(pointee).print(os, getPrintPolicy(ctx)); + return normalizeTranslationRule(std::move(out)); + } + } + + if (auto cxx_record_decl = qual_type->getAsCXXRecordDecl()) { + if (cxx_record_decl->isLambda()) { + return ToString(ctx, cxx_record_decl->getLambdaCallOperator()); + } + } + + if (auto *tag = qual_type->getAsTagDecl(); + tag && !tag->getIdentifier() && !tag->getTypedefNameForAnonDecl()) { + return ToString(ctx, clang::cast(tag)); + } + + if (auto *tag = qual_type->getAsTagDecl(); + tag && tag->getIdentifier() && + tag->getDeclContext()->isFunctionOrMethod()) { + return GetNamedDeclAsString(tag); + } + + if (auto renamed = DisambiguateAnonymousTag(qual_type->getAsTagDecl()); + !renamed.empty()) { + return renamed; + } + + std::string type; + llvm::raw_string_ostream os(type); + normalizeQualType(ctx, qual_type).print(os, getPrintPolicy(ctx)); + return normalizeTranslationRule(std::move(type)); +} + +std::string ToString(clang::ASTContext &ctx, const clang::NamedDecl *decl) { + if (auto *record = clang::dyn_cast(decl); + record && !record->getIdentifier()) { + if (auto renamed = DisambiguateAnonymousTag(record); !renamed.empty()) { + return renamed; + } + if (auto *typedef_decl = record->getTypedefNameForAnonDecl()) { + return ToString(ctx, clang::cast(typedef_decl)); + } + return GetNamedDeclAsString(record); + } + + if (auto *enum_decl = clang::dyn_cast(decl)) { + if (auto renamed = DisambiguateAnonymousTag(enum_decl); !renamed.empty()) { + return renamed; + } + if (!enum_decl->getIdentifier() && + !enum_decl->getTypedefNameForAnonDecl()) { + return GetNamedDeclAsString(enum_decl); + } + } + + std::string out; + llvm::raw_string_ostream os(out); + + const clang::FunctionDecl *func_decl = nullptr; + if (auto *template_decl = llvm::dyn_cast(decl)) { + func_decl = template_decl->getTemplatedDecl(); + } else { + func_decl = llvm::dyn_cast_or_null(decl); + } + + if (!func_decl) { + decl->printQualifiedName(os, getPrintPolicy(ctx)); + return normalizeTranslationRule(std::move(out)); + } + + os << ToString(ctx, func_decl->getReturnType()) << ' '; + if (const auto op = func_decl->getOverloadedOperator(); + op >= clang::OverloadedOperatorKind::OO_LessLess && + op <= clang::OverloadedOperatorKind::OO_GreaterGreaterEqual) { + // ensure matchTemplate does not consider these operator names when matching + func_decl->getQualifier().print(os, getPrintPolicy(ctx)); + os << "operator "; + switch (op) { + case clang::OverloadedOperatorKind::OO_LessLess: + os << "shl"; + break; + case clang::OverloadedOperatorKind::OO_GreaterGreater: + os << "shr"; + break; + case clang::OverloadedOperatorKind::OO_LessLessEqual: + os << "shleq"; + break; + case clang::OverloadedOperatorKind::OO_GreaterGreaterEqual: + os << "shreq"; + break; + default: + assert(0 && "Unexpected overloaded operator kind"); + } + } else if (const auto *method_decl = + llvm::dyn_cast(func_decl)) { + if (method_decl->getParent()->isLambda() && + method_decl->getOverloadedOperator() == clang::OO_Call) { + func_decl->printName(os, getPrintPolicy(ctx)); + } else { + func_decl->printQualifiedName(os, getPrintPolicy(ctx)); + } + } else { + func_decl->printQualifiedName(os, getPrintPolicy(ctx)); + } + + bool has_pack = HasFunctionParameterPack(func_decl); + unsigned num_params = func_decl->getNumParams(); + if (has_pack) { + const auto *primary = func_decl->getPrimaryTemplate(); + num_params = + (primary ? primary->getTemplatedDecl() : func_decl)->getNumParams() - 1; + } + + os << '('; + for (unsigned i = 0; i < num_params; ++i) { + if (i) { + os << ", "; + } + os << ToString(ctx, func_decl->getParamDecl(i)->getType()); + } + if (has_pack) { + if (num_params) { + os << ", "; + } + os << kPackMarker; + } + if (func_decl->isVariadic()) { + if (func_decl->getNumParams()) { + os << ", "; + } + os << "..."; + } + os << ')'; + + if (const auto *method_decl = + llvm::dyn_cast(func_decl)) { + if (method_decl->isConst()) { + os << " const"; + } + if (method_decl->isVolatile()) { + os << " volatile"; + } + switch (method_decl->getRefQualifier()) { + case clang::RQ_LValue: + os << " &"; + break; + case clang::RQ_RValue: + os << " &&"; + break; + default: + break; + } + } + + return normalizeTranslationRule(std::move(out)); +} + +std::string ToString(clang::ASTContext &ctx, const clang::Expr *expr) { + if (!expr) { + assert(0 && "!expr"); + } + + expr = expr->IgnoreParenImpCasts(); + + if (llvm::isa(expr) && + expr->getBeginLoc().isMacroID()) { + auto &sm = ctx.getSourceManager(); + auto name = clang::Lexer::getImmediateMacroName(expr->getBeginLoc(), sm, + ctx.getLangOpts()); + if (!name.empty()) { + return name.str(); + } + } + + if (const auto *CE = llvm::dyn_cast(expr)) { + if (const auto *decl = CE->getDirectCallee()) { + return ToString(ctx, decl); + } + } + + if (const auto *ctor = llvm::dyn_cast(expr)) { + if (const auto *ctor_decl = ctor->getConstructor()) { + return ToString(ctx, ctor_decl); + } + assert(0 && "expr is a CXXConstructExpr but could not get constructor"); + } + + if (const auto *ME = llvm::dyn_cast(expr)) { + if (const auto *member_decl = + llvm::dyn_cast(ME->getMemberDecl())) { + if (const auto *method_decl = + llvm::dyn_cast(member_decl)) { + return ToString(ctx, method_decl); + } + if (ME->isArrow()) { + auto *base = ME->getBase()->IgnoreParenImpCasts(); + if (auto *op = llvm::dyn_cast(base)) { + if (op->getOperator() == clang::OO_Arrow) { + return ToString(ctx, op->getArg(0)->getType()) + "->" + + ToString(ctx, member_decl); + } + } + } else if (auto for_range = GetParentForRange(ctx, ME)) { + if (ToString(ctx, for_range->getRangeInit()->getType()) + .starts_with("std::map<")) { + auto iter_type = GetForRangeIteratorType(for_range); + if (!iter_type.isNull()) { + return ToString(ctx, iter_type) + "->" + ToString(ctx, member_decl); + } + } + } + return ToString(ctx, member_decl); + } + assert(0 && "expr is a MemberExpr but could not get named decl"); + } + + if (const auto *decl_ref = llvm::dyn_cast(expr)) { + if (const auto *named_decl = + llvm::dyn_cast(decl_ref->getDecl())) { + if (const auto *tmpl_decl = + llvm::dyn_cast(named_decl)) { + return ToString(ctx, tmpl_decl->getTemplatedDecl()); + } + return ToString(ctx, named_decl); + } + return ""; + } + + if (const auto *uop = llvm::dyn_cast(expr)) { + auto sub = ToString(ctx, uop->getSubExpr()); + std::string_view opcode = + clang::UnaryOperator::getOpcodeStr(uop->getOpcode()); + return uop->isPostfix() ? std::format("{}{}", sub, opcode) + : std::format("{}{}", opcode, sub); + } + + return "Unhandled case in ToString"; +} + +} // namespace cpp2rust::Printer diff --git a/cpp2rust/converter/printer.h b/cpp2rust/converter/printer.h new file mode 100644 index 000000000..a19a77dea --- /dev/null +++ b/cpp2rust/converter/printer.h @@ -0,0 +1,23 @@ +#pragma once + +// Copyright (c) 2022-present INESC-ID. +// Distributed under the MIT license that can be found in the LICENSE file. + +#include +#include +#include + +#include + +namespace cpp2rust::Printer { +enum class ScalarSugar { + kDesugar, + kPreserve, +}; + +std::string ToString(clang::ASTContext &ctx, clang::QualType qual_type, + ScalarSugar sugar = ScalarSugar::kDesugar); +std::string ToString(clang::ASTContext &ctx, const clang::Expr *expr); +std::string ToString(clang::ASTContext &ctx, const clang::NamedDecl *decl); +std::string ToRustName(std::string name); +} // namespace cpp2rust::Printer diff --git a/cpp2rust/cpp_rule_preprocessor.cpp b/cpp2rust/cpp_rule_preprocessor.cpp index 43a38ab6a..f24dc6b6c 100644 --- a/cpp2rust/cpp_rule_preprocessor.cpp +++ b/cpp2rust/cpp_rule_preprocessor.cpp @@ -34,6 +34,7 @@ #include "compat/platform_flags.h" #include "converter/converter_lib.h" #include "converter/mapper.h" +#include "converter/printer.h" namespace fs = std::filesystem; @@ -125,7 +126,8 @@ class Callback : public clang::ast_matchers::MatchFinder::MatchCallback { type = lookupType(tdecl); } } - auto src = Mapper::ToString(type, Mapper::ScalarSugar::kPreserve); + auto src = + Printer::ToString(*R.Context, type, Printer::ScalarSugar::kPreserve); out_.try_emplace(var->getQualifiedNameAsString(), std::move(src)); return; } @@ -137,7 +139,7 @@ class Callback : public clang::ast_matchers::MatchFinder::MatchCallback { if (const auto *fcall = R.Nodes.getNodeAs("fcall")) { if (fcall->getDirectCallee()) { - add(Mapper::ToString(fcall)); + add(Printer::ToString(*R.Context, fcall)); return; } @@ -145,46 +147,45 @@ class Callback : public clang::ast_matchers::MatchFinder::MatchCallback { clang::FunctionDecl *rule = nullptr; clang::FunctionDecl *decl = lookupCalledDecl( func->getDescribedFunctionTemplate(), lookup, &rule); - if (Mapper::HasFunctionParameterPack(func) && - Mapper::HasFunctionParameterPack(decl)) { + if (HasFunctionParameterPack(func) && HasFunctionParameterPack(decl)) { addPackRule(func, rule, decl); return; } - add(Mapper::ToString(decl)); + add(Printer::ToString(*R.Context, decl)); return; } if (const auto *ctor = R.Nodes.getNodeAs("ctor")) { if (ctor->getConstructor()) { - add(Mapper::ToString(ctor)); + add(Printer::ToString(*R.Context, ctor)); return; } } if (const auto *muse = R.Nodes.getNodeAs("muse")) { if (llvm::isa(muse->getMemberDecl())) { - add(Mapper::ToString(muse)); + add(Printer::ToString(*R.Context, muse)); return; } } if (const auto *um = R.Nodes.getNodeAs("umuse")) { - add(Mapper::ToString(um)); + add(Printer::ToString(*R.Context, um)); return; } if (R.Nodes.getNodeAs("declref")) { if (const auto *enum_val = R.Nodes.getNodeAs("enum_val")) { - add(Mapper::ToString(enum_val)); + add(Printer::ToString(*R.Context, enum_val)); return; } else if (const auto *decl = R.Nodes.getNodeAs("decl")) { - add(Mapper::ToString(decl)); + add(Printer::ToString(*R.Context, decl)); return; } } if (const auto *uop = R.Nodes.getNodeAs("udeclref")) { - add(Mapper::ToString(uop)); + add(Printer::ToString(*R.Context, uop)); return; } if (const auto *dsme = @@ -193,12 +194,12 @@ class Callback : public clang::ast_matchers::MatchFinder::MatchCallback { clang::MemberExpr *expr = lookupArrowAccess( func->getDescribedFunctionTemplate(), dsme->getMemberNameInfo(), dsme->getQualifierLoc()); - add(Mapper::ToString(expr)); + add(Printer::ToString(*R.Context, expr)); return; } clang::NamedDecl *decl = lookupMemberAccess( func->getDescribedFunctionTemplate(), dsme->getMember()); - add(Mapper::ToString(decl)); + add(Printer::ToString(*R.Context, decl)); return; } if (const auto *uctor = @@ -206,13 +207,13 @@ class Callback : public clang::ast_matchers::MatchFinder::MatchCallback { LookupInfo lookup(uctor); clang::NamedDecl *decl = lookupCalledDecl( func->getDescribedFunctionTemplate(), lookup, nullptr); - add(Mapper::ToString(decl)); + add(Printer::ToString(*R.Context, decl)); return; } if (const auto *lit = R.Nodes.getNodeAs("macro_int")) { if (lit->getBeginLoc().isMacroID()) { - add(Mapper::ToString(lit)); + add(Printer::ToString(*R.Context, lit)); } return; } @@ -226,7 +227,7 @@ class Callback : public clang::ast_matchers::MatchFinder::MatchCallback { void addPackRule(const clang::FunctionDecl *func, clang::FunctionDecl *rule, clang::FunctionDecl *callee) { - auto key = Mapper::ToString(callee); + auto key = Printer::ToString(sema_->Context, callee); auto init_type = getInitType(func, rule); if (init_type.isNull()) { out_.try_emplace(func->getQualifiedNameAsString(), std::move(key)); @@ -276,9 +277,10 @@ class Callback : public clang::ast_matchers::MatchFinder::MatchCallback { } } } - llvm::errs() << "ERROR: Init type " << Mapper::ToString(type) + llvm::errs() << "ERROR: Init type " + << Printer::ToString(sema_->Context, type) << " is not a template argument of " - << Mapper::ToString(callee) << '\n'; + << Printer::ToString(sema_->Context, callee) << '\n'; std::exit(EXIT_FAILURE); } From f8dc5d9764f92ea95233026a1aafe7ef02872929 Mon Sep 17 00:00:00 2001 From: Lucian Popescu Date: Tue, 29 Sep 2026 14:33:21 +0100 Subject: [PATCH 2/7] Move exprs_ and types_ in rule registry --- cpp2rust/converter/converter.cpp | 107 ++-- cpp2rust/converter/converter_lib.cpp | 35 +- cpp2rust/converter/converter_lib.h | 17 +- cpp2rust/converter/factory.cpp | 4 +- cpp2rust/converter/mapper.cpp | 552 ++---------------- cpp2rust/converter/mapper.h | 55 +- .../converter/models/converter_refcount.cpp | 85 +-- .../converter/models/converter_refcount.h | 8 +- cpp2rust/converter/printer.cpp | 5 +- cpp2rust/converter/rules/matching.cpp | 261 +++++++++ cpp2rust/converter/rules/matching.h | 16 + cpp2rust/converter/rules/registry.cpp | 246 ++++++++ cpp2rust/converter/rules/registry.h | 33 ++ cpp2rust/cpp_rule_preprocessor.cpp | 2 - 14 files changed, 765 insertions(+), 661 deletions(-) create mode 100644 cpp2rust/converter/rules/matching.cpp create mode 100644 cpp2rust/converter/rules/matching.h create mode 100644 cpp2rust/converter/rules/registry.cpp create mode 100644 cpp2rust/converter/rules/registry.h diff --git a/cpp2rust/converter/converter.cpp b/cpp2rust/converter/converter.cpp index 8a2ec46eb..3e6f7221e 100644 --- a/cpp2rust/converter/converter.cpp +++ b/cpp2rust/converter/converter.cpp @@ -23,6 +23,7 @@ #include "converter/lex.h" #include "converter/mapper.h" #include "converter/printer.h" +#include "converter/rules/registry.h" namespace cpp2rust { std::unordered_map Converter::inner_structs_; @@ -119,7 +120,7 @@ bool Converter::Convert(clang::QualType qual_type) { record_decls_.MarkReferenced(GetRecordName(decl)); } - auto mapped = Mapper::Map(qual_type); + auto mapped = Mapper::Map(ctx_, qual_type); if (!mapped.empty() && mapped != token::kIgnoreRule) { StrCat(mapped); return false; @@ -130,7 +131,7 @@ bool Converter::Convert(clang::QualType qual_type) { } bool Converter::ConvertMappedType(clang::QualType qual_type) { - std::string type_as_string = Mapper::Map(qual_type); + std::string type_as_string = Mapper::Map(ctx_, qual_type); if (type_as_string == token::kIgnoreRule) { return false; } @@ -152,7 +153,7 @@ std::string Converter::ConvertPointeeType(clang::QualType ptr_type) { } bool Converter::VisitBuiltinType(clang::BuiltinType *type) { - auto mapped = Mapper::Map(clang::QualType(type, 0)); + auto mapped = Mapper::Map(ctx_, clang::QualType(type, 0)); if (mapped.empty()) { llvm::report_fatal_error(llvm::Twine("no type rule for builtin type: ") + type->getName(ctx_.getPrintingPolicy())); @@ -179,7 +180,7 @@ bool Converter::VisitRecordType(clang::RecordType *type) { } StrCat(GetRecordName(decl)); - Mapper::AddRuleForUserDefinedType(decl); + RuleRegistry::AddRuleForUserDefinedType(ctx_, decl); return false; } @@ -672,14 +673,14 @@ bool Converter::RecordDerivesDefault(const clang::RecordDecl *decl) { } bool Converter::IsPassThroughRule(clang::Expr *expr) const { - const auto *rule = Mapper::GetExprRule(GetCalleeOrExpr(expr)); + const auto *rule = Mapper::GetExprRule(ctx_, GetCalleeOrExpr(expr)); return rule && rule->body.size() == 1 && std::holds_alternative( rule->body[0]); } bool Converter::RecordDerivesCopy(const clang::RecordDecl *decl) const { - auto *derives = Mapper::MappedDerives(ctx_.getCanonicalTagType(decl)); + auto *derives = Mapper::MappedDerives(ctx_, ctx_.getCanonicalTagType(decl)); return derives && std::find(derives->begin(), derives->end(), "Copy") != derives->end(); } @@ -692,7 +693,7 @@ bool Converter::RecordHasCopyableFields(const clang::RecordDecl *decl) { for (auto f : decl->fields()) { // Records that contain std::vector, std::array, std::string or anything // that is translated to Vec<>, do not derive Copy - auto mapped = Mapper::Map(f->getType()); + auto mapped = Mapper::Map(ctx_, f->getType()); if (mapped.starts_with("Vec<")) { return false; } @@ -740,7 +741,7 @@ bool Converter::VisitRecordDecl(clang::RecordDecl *decl) { return false; } - Mapper::AddRuleForUserDefinedType(decl); + RuleRegistry::AddRuleForUserDefinedType(ctx_, decl); EmitRustStructOrUnion(decl); return false; @@ -802,7 +803,7 @@ void Converter::EmitRustStructOrUnion(clang::RecordDecl *decl) { EmitReprC(decl); } auto attrs = GetStructAttributes(decl); - Mapper::SetDerives(ctx_.getCanonicalTagType(decl), + Mapper::SetDerives(ctx_, ctx_.getCanonicalTagType(decl), std::vector(attrs.begin(), attrs.end())); StrCat("#[derive("); for (auto *attr : attrs) { @@ -901,7 +902,7 @@ void Converter::EmitReprC(clang::RecordDecl *decl) { void Converter::EmitRustUnion(clang::RecordDecl *decl) { EmitReprC(decl); auto attrs = GetStructAttributes(decl); - Mapper::SetDerives(ctx_.getCanonicalTagType(decl), + Mapper::SetDerives(ctx_, ctx_.getCanonicalTagType(decl), std::vector(attrs.begin(), attrs.end())); StrCat("#[derive("); for (auto *attr : attrs) { @@ -927,7 +928,7 @@ void Converter::EmitRustUnion(clang::RecordDecl *decl) { bool Converter::VisitCXXRecordDecl(clang::CXXRecordDecl *decl) { decl->dump(log()); - Mapper::AddRuleForUserDefinedType(decl); + RuleRegistry::AddRuleForUserDefinedType(ctx_, decl); if (!IsConvertibleCXXRecordDecl(decl)) { return false; } @@ -1461,7 +1462,7 @@ void Converter::ConvertForRangeBody(clang::CXXForRangeStmt *stmt, bool Converter::VisitCXXForRangeStmt(clang::CXXForRangeStmt *stmt) { auto range_init_type = stmt->getRangeInit()->getType(); - if (!Mapper::Contains(range_init_type.getUnqualifiedType())) { + if (!Mapper::Contains(ctx_, range_init_type.getUnqualifiedType())) { // FIXME: improve error handling log() << "for range stmts only for types in std namespace\n"; } @@ -1482,7 +1483,7 @@ bool Converter::VisitCXXForRangeStmtMap(clang::CXXForRangeStmt *stmt) { auto loop_var_name = GetNamedDeclAsString(loop_var); StrCat("'loop_:"); - auto map_type = Mapper::Map(stmt->getRangeInit()->getType()); + auto map_type = Mapper::Map(ctx_, stmt->getRangeInit()->getType()); StrCat(keyword::kFor, loop_var_name, keyword::kIn, "UnsafeMapIterator::begin(&"); Convert(stmt->getRangeInit()); @@ -1558,7 +1559,7 @@ bool Converter::Convert(clang::Expr *expr, std::optional implicit_convert_to) { bool needs_conversion = expr && implicit_convert_to && - NeedsImplicitScalarCast(expr->IgnoreImplicit()->getType(), + NeedsImplicitScalarCast(ctx_, expr->IgnoreImplicit()->getType(), *implicit_convert_to); PushParen paren(*this, needs_conversion); computed_expr_type_ = ComputedExprType::Unknown; @@ -1614,7 +1615,7 @@ bool Converter::GetFmtArg(clang::Expr *arg, std::string &fmt, } else if (arg_str.contains("Setw")) { fmt_width = Trim(ToString(arg)); } else if (!arg->getType()->isCharType() && - Mapper::Map(arg->getType()) != + Mapper::Map(ctx_, arg->getType()) != std::format("Vec<{}>", CharRustType())) { fmt += ("{:" + fmt_width + fmt_trait + "}"); fmt_width.clear(); // Reset setw after first usage @@ -1632,7 +1633,7 @@ bool Converter::GetFmtArg(clang::Expr *arg, std::string &fmt, bool Converter::GetRawArg(clang::Expr *arg, std::string &raw_args) { if (arg->getType()->isCharType()) { raw_args += "(&[" + ToString(arg) + " as u8]"; - } else if (Mapper::Map(arg->getType()) == + } else if (Mapper::Map(ctx_, arg->getType()) == std::format("Vec<{}>", CharRustType())) { PushExprKind push(*this, ExprKind::RValue); std::string str = ToString(arg); @@ -1788,14 +1789,15 @@ bool Converter::VisitCallExpr(clang::CallExpr *expr) { return false; } - if (IsImplicitAssignmentCall(expr) && !Mapper::Contains(expr->getCallee())) { + if (IsImplicitAssignmentCall(expr) && + !Mapper::Contains(ctx_, expr->getCallee())) { auto *call = clang::cast(expr); ConvertAssignment(call->getImplicitObjectArgument(), call->getArg(0), "="); return false; } - if (Mapper::Contains(expr->getCallee())) { - if (Mapper::IsLibcPassthrough(GetCalleeOrExpr(expr))) { + if (Mapper::Contains(ctx_, expr->getCallee())) { + if (Mapper::IsLibcPassthrough(ctx_, GetCalleeOrExpr(expr))) { ConvertGenericCallExpr(expr); return false; } @@ -1836,7 +1838,7 @@ bool Converter::VisitCallExpr(clang::CallExpr *expr) { if (auto *opcall = clang::dyn_cast(expr); opcall && !IsUserOperatorCall(opcall) && - !Mapper::Contains(expr->getCallee())) { + !Mapper::Contains(ctx_, expr->getCallee())) { return ConvertCXXOperatorCallExpr(opcall); } @@ -1878,7 +1880,7 @@ std::string Converter::GetFunctionRefName(const clang::FunctionDecl *fn_decl) { return std::format("{}::{}", GetRecordName(method->getParent()), GetMethodName(method)); } - return Mapper::MapFunctionName(fn_decl); + return Mapper::MapFunctionName(ctx_, fn_decl); } void Converter::ConvertFunctionToFunctionPointer( @@ -1934,7 +1936,8 @@ Converter::CallInfo Converter::CollectCallInfo(clang::CallExpr *expr) { function ? function->getNumParams() : proto->getNumParams(); info.is_variadic = function ? function->isVariadic() : proto->isVariadic(); info.is_fn_ptr_call = !function; - info.is_libc_passthrough = Mapper::IsLibcPassthrough(GetCalleeOrExpr(expr)); + info.is_libc_passthrough = + Mapper::IsLibcPassthrough(ctx_, GetCalleeOrExpr(expr)); for (unsigned i = 0; i < num_named_params && i < num_args; ++i) { auto *arg = expr->getArg(i + arg_begin); @@ -2092,7 +2095,8 @@ void Converter::EmitArgList(const CallInfo &info) { ConvertParamTy(ca.param_type, ca.expr); if (info.is_libc_passthrough) { StrCat(std::format( - "as {}", Mapper::GetParamType(GetCalleeOrExpr(info.expr), i))); + "as {}", + Mapper::GetParamType(ctx_, GetCalleeOrExpr(info.expr), i))); } break; } @@ -2177,7 +2181,7 @@ Converter::ConvertCallExpr(clang::CallExpr *expr) { } else if (IsBuiltinConstantP(callee)) { StrCat(expr->getArg(0)->isCXX11ConstantExpr(ctx_) ? token::kOne : token::kZero); - } else if (Mapper::Contains(callee)) { + } else if (Mapper::Contains(ctx_, callee)) { auto **args = expr->getArgs(); auto num_args = expr->getNumArgs(); auto ctx = CollectRefBindingTempArgs(expr); @@ -2219,7 +2223,7 @@ std::string Converter::getIntegerLiteral(clang::IntegerLiteral *expr, if (ty->isFloatingType() || incl_type) { if (expr->getValue().isZero()) { - if (auto init = Mapper::MapInitializer(ty); !init.empty()) { + if (auto init = Mapper::MapInitializer(ctx_, ty); !init.empty()) { return init; } } @@ -2235,7 +2239,7 @@ bool Converter::VisitIntegerLiteral(clang::IntegerLiteral *expr) { computed_expr_type_ = ComputedExprType::FreshValue; return false; } - StrCat(getIntegerLiteral(expr, Mapper::Map(expr->getType()) != "i32")); + StrCat(getIntegerLiteral(expr, Mapper::Map(ctx_, expr->getType()) != "i32")); computed_expr_type_ = ComputedExprType::FreshValue; return false; } @@ -2417,7 +2421,7 @@ void Converter::ConvertIntegralToBooleanCast(clang::ImplicitCastExpr *expr) { bool Converter::IsCastRedundantInRust(clang::Expr *expr, clang::QualType target_type) { auto target = GetUnsafeTypeAsString(target_type); - if (const auto *rule = Mapper::GetExprRule(expr)) { + if (const auto *rule = Mapper::GetExprRule(ctx_, expr)) { return rule->return_type.type == target; } return GetUnsafeTypeAsString(expr->getType()) == target; @@ -2785,14 +2789,14 @@ void Converter::ConvertGenericBinaryOperator(clang::BinaryOperator *expr) { PushParen outer(*this); { PushParen lhs_paren(*this); - Convert(lhs, GetOperandImplicitConversionTarget(expr, lhs, rhs)); + Convert(lhs, GetOperandImplicitConversionTarget(ctx_, expr, lhs, rhs)); } StrCat(expr->getOpcodeStr()); { PushParen rhs_paren(*this); - Convert(rhs, GetOperandImplicitConversionTarget(expr, rhs, lhs)); + Convert(rhs, GetOperandImplicitConversionTarget(ctx_, expr, rhs, lhs)); } computed_expr_type_ = ComputedExprType::FreshValue; } @@ -3090,7 +3094,7 @@ bool Converter::ConvertCXXOperatorCallExpr(clang::CXXOperatorCallExpr *expr) { case clang::OverloadedOperatorKind::OO_Arrow: if (IsUniquePtr(expr->getArg(0)->getType())) { ConvertUniquePtrDeref(expr); - } else if (GetStrongestIteratorCategory(expr->getArg(0)->getType()) == + } else if (GetStrongestIteratorCategory(ctx_, expr->getArg(0)->getType()) == IteratorCategory::Bidirectional) { Convert(expr->getArg(0)); } else if (expr->getOperator() == clang::OverloadedOperatorKind::OO_Star) { @@ -3152,7 +3156,7 @@ bool Converter::VisitMemberExpr(clang::MemberExpr *expr) { return false; } if (auto *method = clang::dyn_cast(member); - method && IsMethodOnPtr(method) && !Mapper::Contains(expr)) { + method && IsMethodOnPtr(method) && !Mapper::Contains(ctx_, expr)) { SetUFCSReceiver(expr->getBase(), expr->isArrow(), method); StrCat(GetRecordName(method->getParent()), token::kDoubleColon, GetMethodName(method)); @@ -3274,7 +3278,7 @@ replaceNonUniformLibcField(clang::MemberExpr *expr) { void Converter::ConvertMemberExpr(clang::MemberExpr *expr) { if (auto mapped = GetMappedAsString(expr); !mapped.empty()) { - if (Mapper::ReturnsPointer(expr)) { + if (Mapper::ReturnsPointer(ctx_, expr)) { StrCat(token::kStar, mapped); } else { StrCat(mapped); @@ -3732,10 +3736,10 @@ bool Converter::VisitOffsetOfExpr(clang::OffsetOfExpr *expr) { bool Converter::VisitEnumDecl(clang::EnumDecl *decl) { ENSURE(decl_ids_.insert(GetID(decl)).second); - if (Mapper::Contains(ctx_.getCanonicalTagType(decl))) { + if (Mapper::Contains(ctx_, ctx_.getCanonicalTagType(decl))) { return false; } - Mapper::AddRuleForUserDefinedType(decl); + RuleRegistry::AddRuleForUserDefinedType(ctx_, decl); auto name = GetRecordName(decl); StrCat(std::format("pub type {} = {};", name, GetUnsafeTypeAsString(decl->getIntegerType()))); @@ -3999,7 +4003,7 @@ std::string Converter::GetDefaultAsString(clang::QualType qual_type) { return arr; } - if (auto init = Mapper::MapInitializer(qual_type); !init.empty()) { + if (auto init = Mapper::MapInitializer(ctx_, qual_type); !init.empty()) { computed_expr_type_ = ComputedExprType::FreshValue; return init; } @@ -4201,8 +4205,8 @@ void Converter::ConvertVarInit(clang::QualType qual_type, clang::Expr *expr) { ctor && ctor->getNumArgs() != 0 && IsReferenceType(ctor->getArg(0)) && clang::isa(ctor->getArg(0)->IgnoreCasts()) && !Mapper::Contains( - clang::cast(ctor->getArg(0)->IgnoreCasts()) - ->getCallee()) && + ctx_, clang::cast(ctor->getArg(0)->IgnoreCasts()) + ->getCallee()) && Printer::ToString(ctx_, ctor->getConstructor()->getThisType()) == "std::string") { { @@ -4225,9 +4229,10 @@ void Converter::ConvertVarInit(clang::QualType qual_type, clang::Expr *expr) { void Converter::ConvertUnsignedArithOperand(clang::Expr *expr, clang::QualType type) { - bool needs_cast = (expr->isIntegerConstantExpr(ctx_) && - !clang::isa(expr)) || - Mapper::Map(expr->getType()) != Mapper::Map(type); + bool needs_cast = + (expr->isIntegerConstantExpr(ctx_) && + !clang::isa(expr)) || + Mapper::Map(ctx_, expr->getType()) != Mapper::Map(ctx_, type); PushParen paren(*this, needs_cast); Convert(expr); if (needs_cast) { @@ -4327,7 +4332,7 @@ void Converter::ConvertArraySubscript(clang::Expr *base, clang::Expr *idx, Convert(idx); } - if (Mapper::Map(idx->getType()) != "usize") { + if (Mapper::Map(ctx_, idx->getType()) != "usize") { StrCat(keyword::kAs, "usize"); } } @@ -4813,7 +4818,7 @@ Converter::CollectRefBindingTempArgs(clang::CallExpr *expr) { for (unsigned i = 0; i < expr->getNumArgs() && i < fn->getNumParams(); ++i) { auto param_type = fn->getParamDecl(i)->getType(); - if (NeedsRefBindingTemp(expr->getArg(i), param_type)) { + if (NeedsRefBindingTemp(ctx_, expr->getArg(i), param_type)) { ctx.materialized_args[i] = param_type; } } @@ -4863,7 +4868,7 @@ std::string Converter::ConvertPlaceholder(clang::Expr *expr, clang::Expr *arg, if (ph_ctx.declared_in_rule_as_rust_ptr && arg->getType()->isArrayType()) { return std::format( "({} as {})", ConvertFreshPointer(arg), - Mapper::GetParamType(GetCalleeOrExpr(expr), ph_ctx.arg_idx)); + Mapper::GetParamType(ctx_, GetCalleeOrExpr(expr), ph_ctx.arg_idx)); } if (ph_ctx.needs_materialization()) { @@ -4879,7 +4884,7 @@ std::string Converter::ConvertPlaceholder(clang::Expr *expr, clang::Expr *arg, if (ph_ctx.needs_pointer_receiver()) { auto param_type = - Mapper::GetParamType(GetCalleeOrExpr(expr), ph_ctx.arg_idx); + Mapper::GetParamType(ctx_, GetCalleeOrExpr(expr), ph_ctx.arg_idx); return std::format("({} as {})", ConvertFreshObject(arg, param_type), param_type); } @@ -4945,7 +4950,7 @@ std::string Converter::ConvertMappedMethodCall( std::string Converter::GetMappedAsString(clang::Expr *expr, clang::Expr **args, unsigned num_args, TempMaterializationCtx *ctx) { - auto *tgt_ir = Mapper::GetExprRule(GetCalleeOrExpr(expr)); + auto *tgt_ir = Mapper::GetExprRule(ctx_, GetCalleeOrExpr(expr)); if (!tgt_ir) return {}; @@ -4969,7 +4974,7 @@ std::string Converter::ConvertIRFragment( if (auto *t = std::get_if(&frag)) { result += t->text; } else if (auto *g = std::get_if(&frag)) { - result += Mapper::InstantiateTemplate(GetCalleeOrExpr(expr), g->n); + result += Mapper::InstantiateTemplate(ctx_, GetCalleeOrExpr(expr), g->n); } else if (auto *ph = std::get_if(&frag)) { auto arg_idx = ph->n; assert(arg_idx < all_args.size()); @@ -4985,9 +4990,9 @@ std::string Converter::ConvertIRFragment( .access = ph->access, .is_receiver = is_receiver, .is_cpp_ptr = arg->getType()->isPointerType(), - .maps_to_rust_ptr = Mapper::MapsToPointer(arg->getType()), + .maps_to_rust_ptr = Mapper::MapsToPointer(ctx_, arg->getType()), .declared_in_rule_as_rust_ptr = - Mapper::ParamIsPointer(GetCalleeOrExpr(expr), arg_idx), + Mapper::ParamIsPointer(ctx_, GetCalleeOrExpr(expr), arg_idx), .is_index_base = ph->is_index_base, }; result += ConvertPlaceholder(expr, arg, ph_ctx); @@ -5007,7 +5012,7 @@ std::string Converter::ConvertIRFragment( std::string Converter::ConvertVariadicTail(clang::Expr *expr, const std::vector &all_args) { - const auto *tgt_ir = Mapper::GetExprRule(GetCalleeOrExpr(expr)); + const auto *tgt_ir = Mapper::GetExprRule(ctx_, GetCalleeOrExpr(expr)); unsigned fixed = tgt_ir ? tgt_ir->params.size() : 0; Buffer buf(*this); @@ -5026,7 +5031,7 @@ Converter::ConvertVariadicTail(clang::Expr *expr, std::string Converter::ConvertInitFragment(clang::Expr *expr, const std::vector &all_args) { - const auto *tgt_ir = Mapper::GetExprRule(GetCalleeOrExpr(expr)); + const auto *tgt_ir = Mapper::GetExprRule(ctx_, GetCalleeOrExpr(expr)); assert(tgt_ir && tgt_ir->init_type.valid()); auto *callee = clang::cast(expr)->getDirectCallee(); assert(callee); @@ -5066,7 +5071,7 @@ std::string Converter::AccessLValueObject(clang::MemberExpr *member) { if (member->isArrow()) { auto *op = clang::dyn_cast(object->IgnoreImplicit()); - if (op && GetStrongestIteratorCategory(op->getArg(0)->getType()) == + if (op && GetStrongestIteratorCategory(ctx_, op->getArg(0)->getType()) == IteratorCategory::Bidirectional) { return ToString(object); } diff --git a/cpp2rust/converter/converter_lib.cpp b/cpp2rust/converter/converter_lib.cpp index 53d3c6f20..c413a0cd9 100644 --- a/cpp2rust/converter/converter_lib.cpp +++ b/cpp2rust/converter/converter_lib.cpp @@ -1475,15 +1475,15 @@ std::optional GetParamImplicitConvertTarget(clang::Expr *expr, } std::optional -GetStrongestIteratorCategory(clang::QualType type) { +GetStrongestIteratorCategory(clang::ASTContext &ctx, clang::QualType type) { type = type.getNonReferenceType().getUnqualifiedType(); - if (!Mapper::Contains(type)) { + if (!Mapper::Contains(ctx, type)) { return std::nullopt; } - if (Mapper::MapsToRefcountPointer(type)) { + if (Mapper::MapsToRefcountPointer(ctx, type)) { return IteratorCategory::Contiguous; } - auto mapped = Mapper::Map(type); + auto mapped = Mapper::Map(ctx, type); if (mapped.empty()) { return std::nullopt; } @@ -1560,15 +1560,17 @@ bool IsBuiltinVaStart(const clang::CallExpr *expr) { return false; } -bool NeedsImplicitScalarCast(clang::QualType from, clang::QualType to) { +bool NeedsImplicitScalarCast(clang::ASTContext &ctx, clang::QualType from, + clang::QualType to) { return !from.isNull() && !to.isNull() && from->isIntegerType() && to->isIntegerType() && from.getCanonicalType().getUnqualifiedType() == to.getCanonicalType().getUnqualifiedType() && - Mapper::Map(from) != Mapper::Map(to); + Mapper::Map(ctx, from) != Mapper::Map(ctx, to); } -bool NeedsRefBindingTemp(const clang::Expr *arg, clang::QualType param_type) { +bool NeedsRefBindingTemp(clang::ASTContext &ctx, const clang::Expr *arg, + clang::QualType param_type) { if (!param_type->isReferenceType()) { return false; } @@ -1584,28 +1586,27 @@ bool NeedsRefBindingTemp(const clang::Expr *arg, clang::QualType param_type) { // void foo(const size_t &) {} <-- size_t -> usize // unsigned long x = 1; foo(x); <-- unsigned long -> u64 return param_type->getPointeeType().isConstQualified() && - NeedsImplicitScalarCast(arg->IgnoreImplicit()->getType(), + NeedsImplicitScalarCast(ctx, arg->IgnoreImplicit()->getType(), param_type.getNonReferenceType()); } -bool IsSizeType(clang::QualType type) { - auto rust_type = Mapper::Map(type); +bool IsSizeType(clang::ASTContext &ctx, clang::QualType type) { + auto rust_type = Mapper::Map(ctx, type); return rust_type == "usize" || rust_type == "isize"; } -std::optional -GetOperandImplicitConversionTarget(const clang::BinaryOperator *op, - const clang::Expr *operand, - const clang::Expr *sibling) { +std::optional GetOperandImplicitConversionTarget( + clang::ASTContext &ctx, const clang::BinaryOperator *op, + const clang::Expr *operand, const clang::Expr *sibling) { if (op->isComparisonOp()) { - if (NeedsImplicitScalarCast(operand->getType(), sibling->getType()) && - IsSizeType(sibling->getType())) { + if (NeedsImplicitScalarCast(ctx, operand->getType(), sibling->getType()) && + IsSizeType(ctx, sibling->getType())) { return sibling->getType(); } return std::nullopt; } if ((op->isAdditiveOp() || op->isMultiplicativeOp() || op->isBitwiseOp()) && - NeedsImplicitScalarCast(operand->getType(), op->getType())) { + NeedsImplicitScalarCast(ctx, operand->getType(), op->getType())) { return op->getType(); } return std::nullopt; diff --git a/cpp2rust/converter/converter_lib.h b/cpp2rust/converter/converter_lib.h index 1cf00d4fd..2cf01c92d 100644 --- a/cpp2rust/converter/converter_lib.h +++ b/cpp2rust/converter/converter_lib.h @@ -35,7 +35,7 @@ enum class IteratorCategory { }; std::optional -GetStrongestIteratorCategory(clang::QualType type); +GetStrongestIteratorCategory(clang::ASTContext &ctx, clang::QualType type); bool IsBuiltinConstantP(const clang::Expr *expr); bool IsGlobalVar(const clang::VarDecl *decl); @@ -275,16 +275,17 @@ std::string GetClassName(clang::QualType type); bool IsVaListType(clang::QualType type); -bool NeedsImplicitScalarCast(clang::QualType from, clang::QualType to); +bool NeedsImplicitScalarCast(clang::ASTContext &ctx, clang::QualType from, + clang::QualType to); -bool NeedsRefBindingTemp(const clang::Expr *arg, clang::QualType param_type); +bool NeedsRefBindingTemp(clang::ASTContext &ctx, const clang::Expr *arg, + clang::QualType param_type); -bool IsSizeType(clang::QualType type); +bool IsSizeType(clang::ASTContext &ctx, clang::QualType type); -std::optional -GetOperandImplicitConversionTarget(const clang::BinaryOperator *op, - const clang::Expr *operand, - const clang::Expr *sibling); +std::optional GetOperandImplicitConversionTarget( + clang::ASTContext &ctx, const clang::BinaryOperator *op, + const clang::Expr *operand, const clang::Expr *sibling); bool IsBuiltinVaStart(const clang::CallExpr *expr); diff --git a/cpp2rust/converter/factory.cpp b/cpp2rust/converter/factory.cpp index c27963154..4a34e2c9b 100644 --- a/cpp2rust/converter/factory.cpp +++ b/cpp2rust/converter/factory.cpp @@ -3,15 +3,15 @@ #include "converter/factory.h" -#include "converter/mapper.h" #include "converter/models/converter_refcount.h" +#include "converter/rules/registry.h" namespace cpp2rust { std::unique_ptr CreateConverter(std::string &rs_code, clang::ASTContext &ctx, Model model, const std::string &rules_dir) { - Mapper::LoadTranslationRules(model, ctx, rules_dir); + RuleRegistry::Load(model, rules_dir); switch (model) { case Model::kUnsafe: return std::make_unique(rs_code, ctx); diff --git a/cpp2rust/converter/mapper.cpp b/cpp2rust/converter/mapper.cpp index f2f1a0f82..98067eddb 100644 --- a/cpp2rust/converter/mapper.cpp +++ b/cpp2rust/converter/mapper.cpp @@ -5,323 +5,23 @@ #include #include -#include #include #include #include #include -#include #include #include #include "converter/converter_lib.h" #include "converter/printer.h" +#include "converter/rules/registry.h" #include "converter/translation_rule.h" namespace cpp2rust::Mapper { namespace { -clang::ASTContext *ctx_ = nullptr; -Model model_ = Model::kUnsafe; -bool translation_rules_loaded_ = false; - -std::unordered_multimap - exprs_; // src -> ExprRule -std::unordered_multimap - types_; // src -> TypeRule - -std::string GetExprMapKey(const std::string &str) { - // Extract the function name from something like - // const T1 & std::foo::fn_name(args) - auto n = str.find_first_of('('); - if (n == std::string::npos) { - n = str.size(); - } - - // Walk backwards from '(' tracking <> depth: - // - skip characters inside template arguments (depth > 0) - // - stop at the first space outside all angle brackets - std::string result; - int depth = 0; - for (int i = (int)n - 1; i >= 0; --i) { - char c = str[i]; - if (c == '>') - ++depth; - else if (c == '<') - --depth; - else if (c == ' ' && depth == 0) - break; - else if (depth == 0) - result += c; - } - std::reverse(result.begin(), result.end()); - return result; -} - -std::string GetTypeMapKey(const std::string &str) { - auto n = str.find_first_of("<["); - if (n == std::string::npos || str[n] == '<') { - return str.substr(0, n); - } - // something like int[][] or T1[] -> [] - return str.substr(n + 1); -} - -void AddTypeRule(std::string src, TranslationRule::TypeRule &&rule) { - auto key = GetTypeMapKey(src); - rule.src = std::move(src); - types_.emplace(std::move(key), std::move(rule)); -} - -// Attempts to unify an instantiated C++ type or function signature with a -// corresponding template pattern. If the two match structurally, it returns -// a mapping from template parameter names (e.g., "T1") to their concrete -// instantiated types (e.g., "int"). If no match is possible, returns nullopt. -// -// Example: -// template_str = "std::vector::vector()" -// instantiated = "std::vector::vector()" -// result = { "int" } -std::optional>> -matchTemplate(const std::string &template_str, - const std::string &instantiated) { - auto matchLiteralAt = [&](const std::string &input_str, size_t pos, - std::string_view literal, size_t &end_pos) -> bool { - size_t i = pos; - size_t j = 0; - - while (true) { - while (i < input_str.size() && std::isspace(input_str[i])) { - i++; - } - - while (j < literal.size() && std::isspace(literal[j])) { - j++; - } - - if (j == literal.size()) { - end_pos = i; - return true; - } - - if (i >= input_str.size()) { - return false; - } - - if (input_str[i] != literal[j]) { - return false; - } - - i++; - j++; - } - }; - - auto findNextLiteralSameDepth = [&](const std::string &s, size_t start, - std::string_view lit) -> size_t { - int ang = 0; - int par = 0; - int sq = 0; - - for (size_t i = 0; i < s.size() && i < start; i++) { - switch (s[i]) { - case '<': { - ang++; - break; - } - case '>': { - ang--; - break; - } - case '(': { - par++; - break; - } - case ')': { - par--; - break; - } - case '[': { - sq++; - break; - } - case ']': { - sq--; - break; - } - default: - break; - } - assert(ang >= 0 && par >= 0 && sq >= 0 && "Unbalanced ang, par or sq"); - } - - int base_ang = ang; - int base_par = par; - int base_sq = sq; - - for (size_t i = start; i <= s.size(); i++) { - if (ang == base_ang && par == base_par && sq == base_sq) { - size_t end_i = 0; - if (matchLiteralAt(s, i, lit, end_i)) { - return i; - } - } - - if (i == s.size()) { - break; - } - - char c = s[i]; - switch (c) { - case '<': { - ang++; - break; - } - case '>': { - ang--; - break; - } - case '(': { - par++; - break; - } - case ')': { - par--; - break; - } - case '[': { - sq++; - break; - } - case ']': { - sq--; - break; - } - default: - break; - } - - if (ang < 0 || par < 0 || sq < 0) { - return std::string::npos; - } - } - - return std::string::npos; - }; - - std::vector> captured; - - size_t ti = 0; - size_t si = 0; - - while (ti < template_str.size()) { - if (template_str[ti] == 'T' && ti + 1 < template_str.size() && - std::isdigit(template_str[ti + 1])) { - size_t tj = ti + 2; - while (tj < template_str.size() && std::isdigit(template_str[tj])) { - tj++; - } - - size_t type_idx = std::stoi(&template_str[ti + 1]) - 1; - assert(type_idx < TranslationRule::kMaxGenerics && - "template placeholder exceeds kMaxGenerics"); - ti = tj; - - std::string_view nextLit; - size_t scan = ti; - while (scan < template_str.size()) { - if (template_str[scan] == 'T' && scan + 1 < template_str.size() && - std::isdigit(template_str[scan + 1])) { - break; - } - scan++; - } - nextLit = std::string_view(template_str).substr(ti, scan - ti); - - captured.resize(std::max(captured.size(), type_idx + 1)); - auto &repl = captured[type_idx]; - if (repl.has_value()) { - size_t end_pos = 0; - if (!matchLiteralAt(instantiated, si, *repl, end_pos)) { - return std::nullopt; - } - si = end_pos; - } else { - if (!nextLit.empty()) { - size_t k = findNextLiteralSameDepth(instantiated, si, nextLit); - if (k == std::string::npos) { - return std::nullopt; - } - - size_t a = si; - size_t b = k; - - while (a < b && std::isspace(instantiated[a])) { - a++; - } - while (b > a && std::isspace(instantiated[b - 1])) { - b--; - } - - repl = instantiated.substr(a, b - a); - si = k; - } else { - size_t a = si; - size_t b = instantiated.size(); - - while (a < b && std::isspace(instantiated[a])) { - a++; - } - while (b > a && std::isspace(instantiated[b - 1])) { - b--; - } - - repl = instantiated.substr(a, b - a); - si = instantiated.size(); - } - } - - if (!nextLit.empty()) { - size_t end_pos = 0; - if (!matchLiteralAt(instantiated, si, nextLit, end_pos)) { - return std::nullopt; - } - si = end_pos; - ti += nextLit.size(); - } - } else { - size_t tj = ti; - while (tj < template_str.size()) { - if (template_str[tj] == 'T' && tj + 1 < template_str.size() && - std::isdigit(template_str[tj + 1])) { - break; - } - ++tj; - } - - auto lit = std::string_view(template_str).substr(ti, tj - ti); - size_t end_pos = 0; - if (!matchLiteralAt(instantiated, si, lit, end_pos)) { - return std::nullopt; - } - si = end_pos; - ti = tj; - } - } - - while (si < instantiated.size() && std::isspace(instantiated[si])) { - si++; - } - - if (si != instantiated.size()) { - return std::nullopt; - } - - return captured; -} - // Substitutes concrete types into a target template string using the provided // type mapping. Each template parameter in `tgt_template` is replaced with its // corresponding instantiated type from `types`. @@ -330,7 +30,7 @@ matchTemplate(const std::string &template_str, // types = { {"i32"} } // tgt_template = "Vec" // result = "Vec" -std::string instantiateTgt(const std::vector> &types, +std::string instantiateTgt(const Matching::Bindings &types, const std::string &tgt_template) { assert(types.size() <= TranslationRule::kMaxGenerics && "template placeholder exceeds kMaxGenerics"); @@ -351,103 +51,11 @@ std::string instantiateTgt(const std::vector> &types, return instantiated_template; } -template -std::pair>> -search(std::unordered_multimap &map, const std::string &txt, - const std::string &key) { - auto [it, end] = map.equal_range(key); - T *rule = nullptr; - std::vector> subs; - - for (; it != end; ++it) { - auto &this_rule = it->second; - auto this_subs = matchTemplate(this_rule.src, txt); - if (!this_subs) { - continue; - } - // tie breaker: prefer more specific rules (usually the longer ones) - if (!rule || this_rule.src.size() > rule->src.size()) { - rule = &this_rule; - subs = *std::move(this_subs); - } - } - return {rule, std::move(subs)}; -} - -TranslationRule::ExprRule *search(const clang::Expr *expr) { - if (RefersToUserDefinedDecl(expr)) { - return nullptr; - } - auto qualified_name = Printer::ToString(*ctx_, expr); - auto [rule, subs] = - search(exprs_, qualified_name, GetExprMapKey(qualified_name)); - log() << "search expr " << qualified_name << ", result:\n"; - if (rule) { - rule->dump(); - } else { - log() << "None\n"; - } - return rule; -} - -std::pair>> -search(clang::QualType qual_type) { - auto sugared = - Printer::ToString(*ctx_, qual_type, Printer::ScalarSugar::kPreserve); - if (auto res = search(types_, sugared, GetTypeMapKey(sugared)); res.first) { - log() << "search type " << sugared - << ", result: " << res.first->type_info.type << '\n'; - return res; - } - auto type = Printer::ToString(*ctx_, qual_type); - if (type == sugared) { - log() << "search type " << type << ", result: None\n"; - return {}; - } - auto res = search(types_, type, GetTypeMapKey(type)); - log() << "search type " << type - << ", result: " << (res.first ? res.first->type_info.type : "None") - << '\n'; - return res; -} - -void addRulesFromDirectory(const std::filesystem::path &dir, Model model) { - namespace fs = std::filesystem; - for (const auto &entry : fs::directory_iterator(dir)) { - const auto &path = entry.path(); - assert(fs::exists(path / "ir_src.json") && - (fs::exists(path / "ir_unsafe.json") || - fs::exists(path / "ir_refcount.json"))); - auto [expr_rules, type_rules] = TranslationRule::Load(path, model); - if (expr_rules.empty() && type_rules.empty()) { - log() << "No rules found in " << path << '\n'; - continue; - } - for (auto &[_, rule] : expr_rules) { - exprs_.emplace(GetExprMapKey(rule.src), std::move(rule)); - } - for (auto &[_, rule] : type_rules) { - auto key = GetTypeMapKey(rule.src); - auto [begin, end] = types_.equal_range(key); - for (auto it = begin; it != end; ++it) { - if (it->second.src == rule.src) { - llvm::errs() << "ERROR: duplicate type rule for C++ type '" - << rule.src << "': maps to both '" - << it->second.type_info.type << "' and '" - << rule.type_info.type << "'\n"; - std::exit(EXIT_FAILURE); - } - } - types_.emplace(std::move(key), std::move(rule)); - } - } -} - std::string mapTypeStringRecursive(const std::string &cpp_type) { - auto [rule, subs] = search(types_, cpp_type, GetTypeMapKey(cpp_type)); + auto [rule, subs] = RuleRegistry::SearchType(cpp_type); if (!rule) { llvm::errs() << "cpp_type: " << cpp_type << '\n'; - assert(0 && "Type is not present in types_"); + assert(0 && "Type is not present in the registry"); } for (auto &ty : subs) { if (ty) { @@ -459,23 +67,21 @@ std::string mapTypeStringRecursive(const std::string &cpp_type) { } // namespace -PushASTContext::PushASTContext(clang::ASTContext &ctx) : prev_(ctx_) { - ctx_ = &ctx; +bool Contains(clang::ASTContext &ctx, clang::QualType qual_type) { + return RuleRegistry::Search(ctx, qual_type).first != nullptr; } -PushASTContext::~PushASTContext() { ctx_ = prev_; } -bool Contains(clang::QualType qual_type) { - return search(qual_type).first != nullptr; +bool Contains(clang::ASTContext &ctx, const clang::Expr *expr) { + return RuleRegistry::Search(ctx, expr) != nullptr; } -bool Contains(const clang::Expr *expr) { return search(expr) != nullptr; } - -const TranslationRule::ExprRule *GetExprRule(const clang::Expr *expr) { - return search(expr); +const TranslationRule::ExprRule *GetExprRule(clang::ASTContext &ctx, + const clang::Expr *expr) { + return RuleRegistry::Search(ctx, expr); } -bool IsLibcPassthrough(const clang::Expr *expr) { - const auto *tgt_ir = GetExprRule(expr); +bool IsLibcPassthrough(clang::ASTContext &ctx, const clang::Expr *expr) { + const auto *tgt_ir = GetExprRule(ctx, expr); if (tgt_ir == nullptr || !tgt_ir->body.empty() || !tgt_ir->is_extern) { return false; } @@ -487,19 +93,23 @@ bool IsLibcPassthrough(const clang::Expr *expr) { decl->getLocation()); } -std::string MapFunctionName(const clang::FunctionDecl *decl) { +std::string MapFunctionName(clang::ASTContext &ctx, + const clang::FunctionDecl *decl) { assert(decl); if (!IsUserDefinedDecl(decl) && - exprs_.contains(GetExprMapKey(Printer::ToString(*ctx_, decl)))) { + RuleRegistry::HasExprKey(Printer::ToString(ctx, decl))) { return std::format("libcc2rs::{}_{}", decl->getNameAsString(), - model_ == Model::kRefCount ? "refcount" : "unsafe"); + RuleRegistry::CurrentModel() == Model::kRefCount + ? "refcount" + : "unsafe"); } return GetNamedDeclAsString(decl->getCanonicalDecl()); } -std::string InstantiateTemplate(const clang::Expr *expr, unsigned n) { - auto expr_str = Printer::ToString(*ctx_, expr); - auto [rule, subs] = search(exprs_, expr_str, GetExprMapKey(expr_str)); +std::string InstantiateTemplate(clang::ASTContext &ctx, const clang::Expr *expr, + unsigned n) { + auto expr_str = Printer::ToString(ctx, expr); + auto [rule, subs] = RuleRegistry::SearchExpr(expr_str); auto text = std::format("T{}", n); if (!rule) { return text; @@ -511,8 +121,8 @@ std::string InstantiateTemplate(const clang::Expr *expr, unsigned n) { return instantiateTgt(subs, text); } -std::string Map(clang::QualType qual_type) { - auto [rule, subs] = search(qual_type); +std::string Map(clang::ASTContext &ctx, clang::QualType qual_type) { + auto [rule, subs] = RuleRegistry::Search(ctx, qual_type); if (rule) { for (auto &ty : subs) { if (ty) { @@ -524,8 +134,8 @@ std::string Map(clang::QualType qual_type) { return {}; } -std::string MapInitializer(clang::QualType qual_type) { - auto [rule, subs] = search(qual_type); +std::string MapInitializer(clang::ASTContext &ctx, clang::QualType qual_type) { + auto [rule, subs] = RuleRegistry::Search(ctx, qual_type); if (rule && !rule->initializer.empty()) { for (auto &ty : subs) { if (ty) { @@ -537,41 +147,44 @@ std::string MapInitializer(clang::QualType qual_type) { return {}; } -bool MapsToPointer(clang::QualType qual_type) { - auto rule = search(qual_type).first; +bool MapsToPointer(clang::ASTContext &ctx, clang::QualType qual_type) { + auto rule = RuleRegistry::Search(ctx, qual_type).first; return rule && rule->type_info.is_pointer(); } -bool MapsToRefcountPointer(clang::QualType qual_type) { - auto rule = search(qual_type).first; +bool MapsToRefcountPointer(clang::ASTContext &ctx, clang::QualType qual_type) { + auto rule = RuleRegistry::Search(ctx, qual_type).first; return rule && rule->type_info.is_refcount_pointer; } -const std::vector *MappedDerives(clang::QualType qual_type) { - auto rule = search(qual_type).first; +const std::vector *MappedDerives(clang::ASTContext &ctx, + clang::QualType qual_type) { + auto rule = RuleRegistry::Search(ctx, qual_type).first; return rule ? &rule->type_info.derives : nullptr; } -void SetDerives(clang::QualType qual_type, std::vector derives) { - if (auto *rule = search(qual_type).first) { +void SetDerives(clang::ASTContext &ctx, clang::QualType qual_type, + std::vector derives) { + if (auto *rule = RuleRegistry::Search(ctx, qual_type).first) { rule->type_info.derives = std::move(derives); } } -bool ReturnsPointer(const clang::Expr *expr) { - auto rule = search(expr); +bool ReturnsPointer(clang::ASTContext &ctx, const clang::Expr *expr) { + auto rule = RuleRegistry::Search(ctx, expr); return rule && rule->return_type.is_pointer(); } -const TranslationRule::TypeInfo &GetParamInfo(const clang::Expr *expr, - unsigned index) { - auto rule = search(expr); +const TranslationRule::TypeInfo & +GetParamInfo(clang::ASTContext &ctx, const clang::Expr *expr, unsigned index) { + auto rule = RuleRegistry::Search(ctx, expr); assert(rule && "expression must have a translation rule"); return rule->params.at(index); } -std::string GetParamType(const clang::Expr *expr, unsigned index) { - auto expr_str = Printer::ToString(*ctx_, expr); - auto [rule, subs] = search(exprs_, expr_str, GetExprMapKey(expr_str)); +std::string GetParamType(clang::ASTContext &ctx, const clang::Expr *expr, + unsigned index) { + auto expr_str = Printer::ToString(ctx, expr); + auto [rule, subs] = RuleRegistry::SearchExpr(expr_str); for (auto &ty : subs) { if (ty) { ty = mapTypeStringRecursive(*ty); @@ -580,76 +193,9 @@ std::string GetParamType(const clang::Expr *expr, unsigned index) { return instantiateTgt(subs, rule->params.at(index).type); } -bool ParamIsPointer(const clang::Expr *expr, unsigned index) { - return GetParamInfo(expr, index).is_pointer(); -} - -void AddRuleForUserDefinedType(clang::NamedDecl *decl) { - auto cpp_name = Printer::ToString(*ctx_, GetTypeForDecl(*ctx_, decl)); - auto rs_name = Printer::ToRustName(cpp_name); - - AddTypeRule(cpp_name, TranslationRule::TypeRule::Plain(rs_name)); - - if (auto record_decl = llvm::dyn_cast(decl)) { - // Forward declaration - if (!record_decl->isThisDeclarationADefinition()) { - return; - } - - if (auto cxx_decl = llvm::dyn_cast(record_decl)) { - if (cxx_decl->isAbstract()) { - switch (model_) { - case Model::kUnsafe: - AddTypeRule(cpp_name + " *", TranslationRule::TypeRule::UnsafePtr( - "*mut dyn " + rs_name)); - break; - case Model::kRefCount: - AddTypeRule(cpp_name + " *", TranslationRule::TypeRule::RefcountPtr( - "PtrDyn')); - break; - } - } else { - switch (model_) { - case Model::kUnsafe: - AddTypeRule(cpp_name + " *", - TranslationRule::TypeRule::UnsafePtr("*mut " + rs_name)); - break; - case Model::kRefCount: - AddTypeRule(cpp_name + " *", TranslationRule::TypeRule::RefcountPtr( - "Ptr<" + rs_name + '>')); - break; - } - } - - for (auto *nested : GetNestedStructs(cxx_decl)) { - AddRuleForUserDefinedType(nested); - } - } - } -} - -void LoadTranslationRules(Model model, clang::ASTContext &ctx, - const std::string &rules_dir) { - ctx_ = &ctx; - model_ = model; - - if (translation_rules_loaded_) { - return; - } - translation_rules_loaded_ = true; - - addRulesFromDirectory(rules_dir, model); - -#if 0 - for (auto &[src, rule] : exprs_) { - log() << "Expr key: " << src << '\n'; - rule.dump(); - } - for (auto &[src, rule] : types_) { - log() << "Type key: " << src << '\n'; - rule.dump(); - } -#endif +bool ParamIsPointer(clang::ASTContext &ctx, const clang::Expr *expr, + unsigned index) { + return GetParamInfo(ctx, expr, index).is_pointer(); } } // namespace cpp2rust::Mapper diff --git a/cpp2rust/converter/mapper.h b/cpp2rust/converter/mapper.h index 4864d299f..6aa0f4af5 100644 --- a/cpp2rust/converter/mapper.h +++ b/cpp2rust/converter/mapper.h @@ -9,39 +9,30 @@ #include -#include "converter/factory.h" #include "converter/translation_rule.h" namespace cpp2rust::Mapper { -class PushASTContext { -public: - explicit PushASTContext(clang::ASTContext &ctx); - ~PushASTContext(); - PushASTContext(const PushASTContext &) = delete; - PushASTContext &operator=(const PushASTContext &) = delete; - -private: - clang::ASTContext *prev_; -}; - -bool Contains(clang::QualType qual_type); -bool Contains(const clang::Expr *expr); - -std::string Map(clang::QualType qual_type); -std::string MapInitializer(clang::QualType qual_type); -const TranslationRule::ExprRule *GetExprRule(const clang::Expr *expr); -bool IsLibcPassthrough(const clang::Expr *expr); -std::string MapFunctionName(const clang::FunctionDecl *decl); -std::string InstantiateTemplate(const clang::Expr *expr, unsigned n); -bool ReturnsPointer(const clang::Expr *expr); -std::string GetParamType(const clang::Expr *expr, unsigned index); -bool ParamIsPointer(const clang::Expr *expr, unsigned index); -bool MapsToPointer(clang::QualType qual_type); -bool MapsToRefcountPointer(clang::QualType qual_type); -const std::vector *MappedDerives(clang::QualType qual_type); -void SetDerives(clang::QualType qual_type, std::vector derives); - -void LoadTranslationRules(Model model, clang::ASTContext &ctx, - const std::string &rules_dir); -void AddRuleForUserDefinedType(clang::NamedDecl *decl); +bool Contains(clang::ASTContext &ctx, clang::QualType qual_type); +bool Contains(clang::ASTContext &ctx, const clang::Expr *expr); + +std::string Map(clang::ASTContext &ctx, clang::QualType qual_type); +std::string MapInitializer(clang::ASTContext &ctx, clang::QualType qual_type); +const TranslationRule::ExprRule *GetExprRule(clang::ASTContext &ctx, + const clang::Expr *expr); +bool IsLibcPassthrough(clang::ASTContext &ctx, const clang::Expr *expr); +std::string MapFunctionName(clang::ASTContext &ctx, + const clang::FunctionDecl *decl); +std::string InstantiateTemplate(clang::ASTContext &ctx, const clang::Expr *expr, + unsigned n); +bool ReturnsPointer(clang::ASTContext &ctx, const clang::Expr *expr); +std::string GetParamType(clang::ASTContext &ctx, const clang::Expr *expr, + unsigned index); +bool ParamIsPointer(clang::ASTContext &ctx, const clang::Expr *expr, + unsigned index); +bool MapsToPointer(clang::ASTContext &ctx, clang::QualType qual_type); +bool MapsToRefcountPointer(clang::ASTContext &ctx, clang::QualType qual_type); +const std::vector *MappedDerives(clang::ASTContext &ctx, + clang::QualType qual_type); +void SetDerives(clang::ASTContext &ctx, clang::QualType qual_type, + std::vector derives); } // namespace cpp2rust::Mapper diff --git a/cpp2rust/converter/models/converter_refcount.cpp b/cpp2rust/converter/models/converter_refcount.cpp index c7a33aab1..84c8131ca 100644 --- a/cpp2rust/converter/models/converter_refcount.cpp +++ b/cpp2rust/converter/models/converter_refcount.cpp @@ -50,35 +50,37 @@ static bool IsBoxedType(std::string_view type) { return type.starts_with("Vec<") || type.starts_with("Box<"); } -static bool IsBoxedType(clang::QualType type) { - return IsBoxedType(Mapper::Map(type.getUnqualifiedType())); +static bool IsBoxedType(clang::ASTContext &ctx, clang::QualType type) { + return IsBoxedType(Mapper::Map(ctx, type.getUnqualifiedType())); } -static bool NeedsMutAccess(const clang::CXXMethodDecl *method, +static bool NeedsMutAccess(clang::ASTContext &ctx, + const clang::CXXMethodDecl *method, clang::QualType base_type) { - return !method->isConst() && IsBoxedType(base_type); + return !method->isConst() && IsBoxedType(ctx, base_type); } -static bool IsPointerType(clang::QualType type) { - return type->isPointerType() || - GetStrongestIteratorCategory(type) == IteratorCategory::Contiguous; +static bool IsPointerType(clang::ASTContext &ctx, clang::QualType type) { + return type->isPointerType() || GetStrongestIteratorCategory(ctx, type) == + IteratorCategory::Contiguous; } -bool ConverterRefCount::PendingDeref::compute_inner_boxed(clang::Expr *expr) { +bool ConverterRefCount::PendingDeref::compute_inner_boxed( + clang::Expr *expr) const { if (!expr) { return false; } - if (!IsBoxedType(expr->getType().getNonReferenceType())) { + if (!IsBoxedType(ctx, expr->getType().getNonReferenceType())) { return false; } if (auto *ase = clang::dyn_cast(expr)) { auto base_type = ase->getBase()->IgnoreCasts()->getType(); if (base_type->isPointerType()) - return IsBoxedType(base_type->getPointeeType()); - return IsBoxedType(base_type.getNonReferenceType()); + return IsBoxedType(ctx, base_type->getPointeeType()); + return IsBoxedType(ctx, base_type.getNonReferenceType()); } if (auto *oce = clang::dyn_cast(expr)) { - return IsBoxedType(oce->getArg(0)->getType().getNonReferenceType()); + return IsBoxedType(ctx, oce->getArg(0)->getType().getNonReferenceType()); } return false; } @@ -117,7 +119,7 @@ ConverterRefCount::PushUnboxedIfSimple::PushUnboxedIfSimple( // Vectors are boxed until the last element if (!unboxed && (outer == "Vec<%>" || outer == "Box<%>")) { - if (!IsBoxedType(inner_type)) { + if (!IsBoxedType(c.ctx_, inner_type)) { unboxed = true; } } @@ -171,7 +173,7 @@ bool ConverterRefCount::Convert(clang::QualType qual_type) { return false; } - if (!Mapper::Contains(qual_type)) + if (!Mapper::Contains(ctx_, qual_type)) qual_type = qual_type.getUnqualifiedType().getDesugaredType(ctx_); if (qual_type->isReferenceType() || qual_type->isIncompleteArrayType()) { @@ -305,7 +307,7 @@ std::string ConverterRefCount::ConvertObject(clang::Expr *expr, object_shape_ = saved_shape; if (shape == ObjectShape::Element && expr->getType()->isPointerType()) { auto pointee = expr->getType()->getPointeeType(); - if (IsBoxedType(pointee) || pointee->isArrayType()) { + if (IsBoxedType(ctx_, pointee) || pointee->isArrayType()) { computed_expr_type_ = ComputedExprType::FreshPointer; return std::format("Ptr::<{}>::decay(&({}))", ToString(pointee), std::move(str)); @@ -321,7 +323,7 @@ ConverterRefCount::ConvertFreshObject(clang::Expr *expr, if (!target_ptr_type.empty()) { auto type = expr->getType().getNonReferenceType(); auto pointee = type->isPointerType() ? type->getPointeeType() : type; - if (IsBoxedType(pointee) || pointee->isArrayType()) { + if (IsBoxedType(ctx_, pointee) || pointee->isArrayType()) { auto normalize = [](std::string s) { std::erase(s, ' '); for (size_t pos; (pos = s.find("::<")) != std::string::npos;) @@ -391,7 +393,7 @@ ConverterRefCount::MaterializeTemp(const std::string &binding_name, std::string ConverterRefCount::ConvertPtrType(clang::QualType type) { std::string str; // decays into Ptr; remove the outer type Vec<> - if (IsBoxedType(type)) { + if (IsBoxedType(ctx_, type)) { str = GetInnerType(type); } else { PushConversionKind push(*this, ConversionKind::Ptr); @@ -549,7 +551,7 @@ void ConverterRefCount::EmitRustUnion(clang::RecordDecl *decl) { auto name = GetRecordName(decl); auto attrs = GetStructAttributes(decl); - Mapper::SetDerives(ctx_.getCanonicalTagType(decl), + Mapper::SetDerives(ctx_, ctx_.getCanonicalTagType(decl), std::vector(attrs.begin(), attrs.end())); StrCat(std::format("pub struct {} {{ __bytes: Value> }}", name)); @@ -882,7 +884,7 @@ void ConverterRefCount::ConvertDeclRefValue(clang::Expr *expr, } if (auto pointee = ref->getPointeeType(); - isObject() && WantsElementPtr() && IsBoxedType(pointee)) { + isObject() && WantsElementPtr() && IsBoxedType(ctx_, pointee)) { StrCat(std::format("Ptr::<{}>::decay(&({}))", ToString(pointee), std::move(str))); computed_expr_type_ = ComputedExprType::FreshPointer; @@ -1121,7 +1123,8 @@ bool ConverterRefCount::VisitCallExpr(clang::CallExpr *expr) { return false; } - if (IsImplicitAssignmentCall(expr) && !Mapper::Contains(expr->getCallee())) { + if (IsImplicitAssignmentCall(expr) && + !Mapper::Contains(ctx_, expr->getCallee())) { auto *call = clang::cast(expr); ConvertAssignment(call->getImplicitObjectArgument(), call->getArg(0), "="); return false; @@ -1133,7 +1136,7 @@ bool ConverterRefCount::VisitCallExpr(clang::CallExpr *expr) { if (auto *opcall = clang::dyn_cast(expr); opcall && !IsUserOperatorCall(opcall) && - !Mapper::Contains(expr->getCallee())) { + !Mapper::Contains(ctx_, expr->getCallee())) { return ConvertCXXOperatorCallExpr(opcall); } @@ -1167,14 +1170,14 @@ bool ConverterRefCount::VisitCallExpr(clang::CallExpr *expr) { return false; } - if (isAddrOf() && !ty->isReferenceType() && !IsPointerType(ty)) { + if (isAddrOf() && !ty->isReferenceType() && !IsPointerType(ctx_, ty)) { PushConversionKind push(*this, ConversionKind::FullRefCount); StrCat(BoxValue(std::move(str)), ".as_pointer()"); return false; } if (isObject() && WantsElementPtr() && ref && - IsBoxedType(ref->getPointeeType())) { + IsBoxedType(ctx_, ref->getPointeeType())) { StrCat(std::format("Ptr::<{}>::decay(&({}))", ToString(ref->getPointeeType()), std::move(str))); computed_expr_type_ = ComputedExprType::FreshPointer; @@ -1188,7 +1191,7 @@ bool ConverterRefCount::VisitCallExpr(clang::CallExpr *expr) { if (IsPassThroughRule(expr)) { return false; } - if (IsPointerType(ty) || ty->isReferenceType()) { + if (IsPointerType(ctx_, ty) || ty->isReferenceType()) { computed_expr_type_ = ComputedExprType::FreshPointer; } else { computed_expr_type_ = ComputedExprType::FreshValue; @@ -1738,7 +1741,7 @@ bool ConverterRefCount::VisitMemberExpr(clang::MemberExpr *expr) { ConvertDeclRefValue(expr, member); return false; } - bool known = Mapper::Contains(expr); + bool known = Mapper::Contains(ctx_, expr); if (auto *method = clang::dyn_cast(member); method && !known) { @@ -1756,7 +1759,7 @@ bool ConverterRefCount::VisitMemberExpr(clang::MemberExpr *expr) { if (base_type->isPointerType()) { base_type = base_type->getPointeeType(); } - bool needs_mut = NeedsMutAccess(method, base_type); + bool needs_mut = NeedsMutAccess(ctx_, method, base_type); PushExprKind push(*this, needs_mut ? ExprKind::LValue : ExprKind::RValue); Converter::ConvertMemberExpr(expr); SetFreshType(expr->getType()); @@ -1884,7 +1887,7 @@ bool ConverterRefCount::VisitCXXForRangeStmtMap(clang::CXXForRangeStmt *stmt) { EmitByValueShadow( loop_var_name, loop_var->getType(), std::string(loop_var_name), - "Value<" + Mapper::Map(GetForRangeIteratorType(stmt)) + '>'); + "Value<" + Mapper::Map(ctx_, GetForRangeIteratorType(stmt)) + '>'); ConvertForRangeBody(stmt, loop_var); @@ -1906,7 +1909,7 @@ bool ConverterRefCount::VisitCXXForRangeStmtVector( PushBrace brace(*this); // handle multi-level types such as Vec>> - if (IsBoxedType(stmt->getRangeInit()->getType()) && + if (IsBoxedType(ctx_, stmt->getRangeInit()->getType()) && GetInnerType(stmt->getRangeInit()->getType()).starts_with("Value<")) { StrCat(keyword::kLet, loop_var_name, token::kColon); @@ -2110,7 +2113,7 @@ std::string ConverterRefCount::GetDefaultAsString(clang::QualType qual_type) { return BoxValue(std::move(arr)); } - if (auto init = Mapper::MapInitializer(qual_type); !init.empty()) { + if (auto init = Mapper::MapInitializer(ctx_, qual_type); !init.empty()) { computed_expr_type_ = ComputedExprType::FreshValue; return BoxValue(std::move(init)); } @@ -2342,19 +2345,19 @@ void ConverterRefCount::ConvertGenericBinaryOperator( if (may_cause_borrow_mut_err) { StrCat(std::format( "{{ let _lhs = {}; _lhs {} {} }}", - ConvertFreshRValue(lhs, - GetOperandImplicitConversionTarget(expr, lhs, rhs)), + ConvertFreshRValue( + lhs, GetOperandImplicitConversionTarget(ctx_, expr, lhs, rhs)), opcode, ConvertFreshRValue( - rhs, GetOperandImplicitConversionTarget(expr, rhs, lhs)))); + rhs, GetOperandImplicitConversionTarget(ctx_, expr, rhs, lhs)))); computed_expr_type_ = ComputedExprType::FreshValue; return; } PushParen outer(*this); - Convert(lhs, GetOperandImplicitConversionTarget(expr, lhs, rhs)); + Convert(lhs, GetOperandImplicitConversionTarget(ctx_, expr, lhs, rhs)); StrCat(opcode); - Convert(rhs, GetOperandImplicitConversionTarget(expr, rhs, lhs)); + Convert(rhs, GetOperandImplicitConversionTarget(ctx_, expr, rhs, lhs)); computed_expr_type_ = ComputedExprType::FreshValue; } @@ -2390,7 +2393,7 @@ bool ConverterRefCount::ConvertCXXOperatorCallExpr( break; } - if (GetStrongestIteratorCategory(expr->getArg(0)->getType()) == + if (GetStrongestIteratorCategory(ctx_, expr->getArg(0)->getType()) == IteratorCategory::Bidirectional) { Convert(expr->getArg(0)); break; @@ -2430,8 +2433,8 @@ bool ConverterRefCount::ConvertCXXOperatorCallExpr( } bool is_inner_boxed = - IsBoxedType(expr->getType().getNonReferenceType()) && - IsBoxedType(expr->getArg(0)->getType().getNonReferenceType()); + IsBoxedType(ctx_, expr->getType().getNonReferenceType()) && + IsBoxedType(ctx_, expr->getArg(0)->getType().getNonReferenceType()); if (isLValue()) { PushConversionKind push_ck(*this, ConversionKind::Unboxed); @@ -2673,7 +2676,7 @@ void ConverterRefCount::ConvertDeref(clang::Expr *expr) { } if (isObject() && WantsElementPtr() && - (IsBoxedType(pointee_type) || pointee_type->isArrayType())) { + (IsBoxedType(ctx_, pointee_type) || pointee_type->isArrayType())) { StrCat(std::format("Ptr::<{}>::decay(&({}))", ToString(pointee_type), std::move(str))); computed_expr_type_ = ComputedExprType::FreshPointer; @@ -2694,7 +2697,7 @@ void ConverterRefCount::ConvertArrow(clang::Expr *expr) { return; } - if (GetStrongestIteratorCategory(op->getArg(0)->getType()) == + if (GetStrongestIteratorCategory(ctx_, op->getArg(0)->getType()) == IteratorCategory::Bidirectional) { Convert(op->getArg(0)); return; @@ -2711,7 +2714,7 @@ std::string ConverterRefCount::AccessLValueObject(clang::MemberExpr *member) { if (member->isArrow()) { auto *op = clang::dyn_cast(object->IgnoreImplicit()); - if (op && GetStrongestIteratorCategory(op->getArg(0)->getType()) == + if (op && GetStrongestIteratorCategory(ctx_, op->getArg(0)->getType()) == IteratorCategory::Bidirectional) { return ConvertRValue(op->getArg(0)); } @@ -2781,7 +2784,7 @@ std::string ConverterRefCount::ConvertMappedMethodCall( return Converter::ConvertMappedMethodCall(expr, mc, args, num_args, ctx); } - auto param_type = Mapper::GetParamType(GetCalleeOrExpr(expr), arg_idx); + auto param_type = Mapper::GetParamType(ctx_, GetCalleeOrExpr(expr), arg_idx); if (arg->getType()->isPointerType()) { return std::format("{}.with_mut(|__v: {}| __v{})", ConvertPointer(arg), diff --git a/cpp2rust/converter/models/converter_refcount.h b/cpp2rust/converter/models/converter_refcount.h index 0b410404e..368a7d6f3 100644 --- a/cpp2rust/converter/models/converter_refcount.h +++ b/cpp2rust/converter/models/converter_refcount.h @@ -405,7 +405,8 @@ class ConverterRefCount final : public Converter { // emit ptr.write(rhs), or by ConvertMappedMethodCall to emit // ptr.with_mut(...). struct PendingDeref { - explicit PendingDeref(ComputedExprType &type) : type(type) {} + PendingDeref(ComputedExprType &type, clang::ASTContext &ctx) + : type(type), ctx(ctx) {} void set(std::string str, bool fresh, clang::Expr *expr = nullptr); void set_unchecked(std::string str, bool fresh, clang::Expr *expr = nullptr); @@ -424,11 +425,12 @@ class ConverterRefCount final : public Converter { } private: - static bool compute_inner_boxed(clang::Expr *expr); + bool compute_inner_boxed(clang::Expr *expr) const; ComputedExprType &type; + clang::ASTContext &ctx; std::string value; bool pointee_is_boxed = false; bool ptr_is_fresh = false; - } pending_deref_{computed_expr_type_}; + } pending_deref_{computed_expr_type_, ctx_}; }; } // namespace cpp2rust diff --git a/cpp2rust/converter/printer.cpp b/cpp2rust/converter/printer.cpp index 914ec5a9e..37e982960 100644 --- a/cpp2rust/converter/printer.cpp +++ b/cpp2rust/converter/printer.cpp @@ -133,7 +133,8 @@ std::string ToString(clang::ASTContext &ctx, clang::QualType qual_type, bool builtin_alias = canonical->isBuiltinType() && (pointee->getAs() || pointee->getAs()); - if (!builtin_alias && Mapper::Map(pointee) == Mapper::Map(canonical)) { + if (!builtin_alias && + Mapper::Map(ctx, pointee) == Mapper::Map(ctx, canonical)) { pointee = canonical; } std::string out; @@ -212,7 +213,7 @@ std::string ToString(clang::ASTContext &ctx, const clang::NamedDecl *decl) { if (const auto op = func_decl->getOverloadedOperator(); op >= clang::OverloadedOperatorKind::OO_LessLess && op <= clang::OverloadedOperatorKind::OO_GreaterGreaterEqual) { - // ensure matchTemplate does not consider these operator names when matching + // ensure MatchTemplate does not consider these operator names when matching func_decl->getQualifier().print(os, getPrintPolicy(ctx)); os << "operator "; switch (op) { diff --git a/cpp2rust/converter/rules/matching.cpp b/cpp2rust/converter/rules/matching.cpp new file mode 100644 index 000000000..81cb942d5 --- /dev/null +++ b/cpp2rust/converter/rules/matching.cpp @@ -0,0 +1,261 @@ +// Copyright (c) 2022-present INESC-ID. +// Distributed under the MIT license that can be found in the LICENSE file. + +#include "converter/rules/matching.h" + +#include +#include +#include +#include + +#include "converter/translation_rule.h" + +namespace cpp2rust::Matching { + +// Attempts to unify an instantiated C++ type or function signature with a +// corresponding template pattern. If the two match structurally, it returns +// a mapping from template parameter names (e.g., "T1") to their concrete +// instantiated types (e.g., "int"). If no match is possible, returns nullopt. +// +// Example: +// template_str = "std::vector::vector()" +// instantiated = "std::vector::vector()" +// result = { "int" } +std::optional MatchTemplate(const std::string &template_str, + const std::string &instantiated) { + auto matchLiteralAt = [&](const std::string &input_str, size_t pos, + std::string_view literal, size_t &end_pos) -> bool { + size_t i = pos; + size_t j = 0; + + while (true) { + while (i < input_str.size() && std::isspace(input_str[i])) { + i++; + } + + while (j < literal.size() && std::isspace(literal[j])) { + j++; + } + + if (j == literal.size()) { + end_pos = i; + return true; + } + + if (i >= input_str.size()) { + return false; + } + + if (input_str[i] != literal[j]) { + return false; + } + + i++; + j++; + } + }; + + auto findNextLiteralSameDepth = [&](const std::string &s, size_t start, + std::string_view lit) -> size_t { + int ang = 0; + int par = 0; + int sq = 0; + + for (size_t i = 0; i < s.size() && i < start; i++) { + switch (s[i]) { + case '<': { + ang++; + break; + } + case '>': { + ang--; + break; + } + case '(': { + par++; + break; + } + case ')': { + par--; + break; + } + case '[': { + sq++; + break; + } + case ']': { + sq--; + break; + } + default: + break; + } + assert(ang >= 0 && par >= 0 && sq >= 0 && "Unbalanced ang, par or sq"); + } + + int base_ang = ang; + int base_par = par; + int base_sq = sq; + + for (size_t i = start; i <= s.size(); i++) { + if (ang == base_ang && par == base_par && sq == base_sq) { + size_t end_i = 0; + if (matchLiteralAt(s, i, lit, end_i)) { + return i; + } + } + + if (i == s.size()) { + break; + } + + char c = s[i]; + switch (c) { + case '<': { + ang++; + break; + } + case '>': { + ang--; + break; + } + case '(': { + par++; + break; + } + case ')': { + par--; + break; + } + case '[': { + sq++; + break; + } + case ']': { + sq--; + break; + } + default: + break; + } + + if (ang < 0 || par < 0 || sq < 0) { + return std::string::npos; + } + } + + return std::string::npos; + }; + + Bindings captured; + + size_t ti = 0; + size_t si = 0; + + while (ti < template_str.size()) { + if (template_str[ti] == 'T' && ti + 1 < template_str.size() && + std::isdigit(template_str[ti + 1])) { + size_t tj = ti + 2; + while (tj < template_str.size() && std::isdigit(template_str[tj])) { + tj++; + } + + size_t type_idx = std::stoi(&template_str[ti + 1]) - 1; + assert(type_idx < TranslationRule::kMaxGenerics && + "template placeholder exceeds kMaxGenerics"); + ti = tj; + + std::string_view nextLit; + size_t scan = ti; + while (scan < template_str.size()) { + if (template_str[scan] == 'T' && scan + 1 < template_str.size() && + std::isdigit(template_str[scan + 1])) { + break; + } + scan++; + } + nextLit = std::string_view(template_str).substr(ti, scan - ti); + + captured.resize(std::max(captured.size(), type_idx + 1)); + auto &repl = captured[type_idx]; + if (repl.has_value()) { + size_t end_pos = 0; + if (!matchLiteralAt(instantiated, si, *repl, end_pos)) { + return std::nullopt; + } + si = end_pos; + } else { + if (!nextLit.empty()) { + size_t k = findNextLiteralSameDepth(instantiated, si, nextLit); + if (k == std::string::npos) { + return std::nullopt; + } + + size_t a = si; + size_t b = k; + + while (a < b && std::isspace(instantiated[a])) { + a++; + } + while (b > a && std::isspace(instantiated[b - 1])) { + b--; + } + + repl = instantiated.substr(a, b - a); + si = k; + } else { + size_t a = si; + size_t b = instantiated.size(); + + while (a < b && std::isspace(instantiated[a])) { + a++; + } + while (b > a && std::isspace(instantiated[b - 1])) { + b--; + } + + repl = instantiated.substr(a, b - a); + si = instantiated.size(); + } + } + + if (!nextLit.empty()) { + size_t end_pos = 0; + if (!matchLiteralAt(instantiated, si, nextLit, end_pos)) { + return std::nullopt; + } + si = end_pos; + ti += nextLit.size(); + } + } else { + size_t tj = ti; + while (tj < template_str.size()) { + if (template_str[tj] == 'T' && tj + 1 < template_str.size() && + std::isdigit(template_str[tj + 1])) { + break; + } + ++tj; + } + + auto lit = std::string_view(template_str).substr(ti, tj - ti); + size_t end_pos = 0; + if (!matchLiteralAt(instantiated, si, lit, end_pos)) { + return std::nullopt; + } + si = end_pos; + ti = tj; + } + } + + while (si < instantiated.size() && std::isspace(instantiated[si])) { + si++; + } + + if (si != instantiated.size()) { + return std::nullopt; + } + + return captured; +} + +} // namespace cpp2rust::Matching diff --git a/cpp2rust/converter/rules/matching.h b/cpp2rust/converter/rules/matching.h new file mode 100644 index 000000000..c5026b1b4 --- /dev/null +++ b/cpp2rust/converter/rules/matching.h @@ -0,0 +1,16 @@ +#pragma once + +// Copyright (c) 2022-present INESC-ID. +// Distributed under the MIT license that can be found in the LICENSE file. + +#include +#include +#include + +namespace cpp2rust::Matching { +// Concrete C++ types bound to the rule's T1, T2, ... placeholders. +using Bindings = std::vector>; + +std::optional MatchTemplate(const std::string &template_str, + const std::string &instantiated); +} // namespace cpp2rust::Matching diff --git a/cpp2rust/converter/rules/registry.cpp b/cpp2rust/converter/rules/registry.cpp new file mode 100644 index 000000000..f64a35be6 --- /dev/null +++ b/cpp2rust/converter/rules/registry.cpp @@ -0,0 +1,246 @@ +// Copyright (c) 2022-present INESC-ID. +// Distributed under the MIT license that can be found in the LICENSE file. + +#include "converter/rules/registry.h" + +#include +#include +#include +#include +#include + +#include "converter/converter_lib.h" +#include "converter/printer.h" + +namespace cpp2rust::RuleRegistry { + +namespace { + +using ExprRuleMap = + std::unordered_multimap; +using TypeRuleMap = + std::unordered_multimap; + +Model model_ = Model::kUnsafe; +bool translation_rules_loaded_ = false; + +ExprRuleMap exprs_; // key -> ExprRule +TypeRuleMap types_; // key -> TypeRule + +std::string ExprKey(const std::string &str) { + // Extract the function name from something like + // const T1 & std::foo::fn_name(args) + auto n = str.find_first_of('('); + if (n == std::string::npos) { + n = str.size(); + } + + // Walk backwards from '(' tracking <> depth: + // - skip characters inside template arguments (depth > 0) + // - stop at the first space outside all angle brackets + std::string result; + int depth = 0; + for (int i = (int)n - 1; i >= 0; --i) { + char c = str[i]; + if (c == '>') + ++depth; + else if (c == '<') + --depth; + else if (c == ' ' && depth == 0) + break; + else if (depth == 0) + result += c; + } + std::reverse(result.begin(), result.end()); + return result; +} + +std::string TypeKey(const std::string &str) { + auto n = str.find_first_of("<["); + if (n == std::string::npos || str[n] == '<') { + return str.substr(0, n); + } + // something like int[][] or T1[] -> [] + return str.substr(n + 1); +} + +void AddTypeRule(std::string src, TranslationRule::TypeRule &&rule) { + auto key = TypeKey(src); + rule.src = std::move(src); + types_.emplace(std::move(key), std::move(rule)); +} + +void addRulesFromDirectory(const std::filesystem::path &dir, Model model) { + namespace fs = std::filesystem; + for (const auto &entry : fs::directory_iterator(dir)) { + const auto &path = entry.path(); + assert(fs::exists(path / "ir_src.json") && + (fs::exists(path / "ir_unsafe.json") || + fs::exists(path / "ir_refcount.json"))); + auto [expr_rules, type_rules] = TranslationRule::Load(path, model); + if (expr_rules.empty() && type_rules.empty()) { + log() << "No rules found in " << path << '\n'; + continue; + } + for (auto &[_, rule] : expr_rules) { + exprs_.emplace(ExprKey(rule.src), std::move(rule)); + } + for (auto &[_, rule] : type_rules) { + auto key = TypeKey(rule.src); + auto [begin, end] = types_.equal_range(key); + for (auto it = begin; it != end; ++it) { + if (it->second.src == rule.src) { + llvm::errs() << "ERROR: duplicate type rule for C++ type '" + << rule.src << "': maps to both '" + << it->second.type_info.type << "' and '" + << rule.type_info.type << "'\n"; + std::exit(EXIT_FAILURE); + } + } + types_.emplace(std::move(key), std::move(rule)); + } + } +} + +template +Match search(std::unordered_multimap &map, + const std::string &txt, const std::string &key) { + auto [it, end] = map.equal_range(key); + T *rule = nullptr; + Matching::Bindings subs; + + for (; it != end; ++it) { + auto &this_rule = it->second; + auto this_subs = Matching::MatchTemplate(this_rule.src, txt); + if (!this_subs) { + continue; + } + // tie breaker: prefer more specific rules (usually the longer ones) + if (!rule || this_rule.src.size() > rule->src.size()) { + rule = &this_rule; + subs = *std::move(this_subs); + } + } + return {rule, std::move(subs)}; +} + +} // namespace + +Match SearchExpr(const std::string &str) { + return search(exprs_, str, ExprKey(str)); +} + +Match SearchType(const std::string &str) { + return search(types_, str, TypeKey(str)); +} + +TranslationRule::ExprRule *Search(clang::ASTContext &ctx, + const clang::Expr *expr) { + if (RefersToUserDefinedDecl(expr)) { + return nullptr; + } + auto qualified_name = Printer::ToString(ctx, expr); + auto [rule, subs] = SearchExpr(qualified_name); + log() << "search expr " << qualified_name << ", result:\n"; + if (rule) { + rule->dump(); + } else { + log() << "None\n"; + } + return rule; +} + +Match Search(clang::ASTContext &ctx, + clang::QualType qual_type) { + auto sugared = + Printer::ToString(ctx, qual_type, Printer::ScalarSugar::kPreserve); + if (auto res = SearchType(sugared); res.first) { + log() << "search type " << sugared + << ", result: " << res.first->type_info.type << '\n'; + return res; + } + auto type = Printer::ToString(ctx, qual_type); + if (type == sugared) { + log() << "search type " << type << ", result: None\n"; + return {}; + } + auto res = SearchType(type); + log() << "search type " << type + << ", result: " << (res.first ? res.first->type_info.type : "None") + << '\n'; + return res; +} + +bool HasExprKey(const std::string &str) { + return exprs_.contains(ExprKey(str)); +} + +Model CurrentModel() { return model_; } + +void AddRuleForUserDefinedType(clang::ASTContext &ctx, clang::NamedDecl *decl) { + auto cpp_name = Printer::ToString(ctx, GetTypeForDecl(ctx, decl)); + auto rs_name = Printer::ToRustName(cpp_name); + + AddTypeRule(cpp_name, TranslationRule::TypeRule::Plain(rs_name)); + + if (auto record_decl = llvm::dyn_cast(decl)) { + // Forward declaration + if (!record_decl->isThisDeclarationADefinition()) { + return; + } + + if (auto cxx_decl = llvm::dyn_cast(record_decl)) { + if (cxx_decl->isAbstract()) { + switch (model_) { + case Model::kUnsafe: + AddTypeRule(cpp_name + " *", TranslationRule::TypeRule::UnsafePtr( + "*mut dyn " + rs_name)); + break; + case Model::kRefCount: + AddTypeRule(cpp_name + " *", TranslationRule::TypeRule::RefcountPtr( + "PtrDyn')); + break; + } + } else { + switch (model_) { + case Model::kUnsafe: + AddTypeRule(cpp_name + " *", + TranslationRule::TypeRule::UnsafePtr("*mut " + rs_name)); + break; + case Model::kRefCount: + AddTypeRule(cpp_name + " *", TranslationRule::TypeRule::RefcountPtr( + "Ptr<" + rs_name + '>')); + break; + } + } + + for (auto *nested : GetNestedStructs(cxx_decl)) { + AddRuleForUserDefinedType(ctx, nested); + } + } + } +} + +void Load(Model model, const std::string &rules_dir) { + model_ = model; + + if (translation_rules_loaded_) { + return; + } + translation_rules_loaded_ = true; + + addRulesFromDirectory(rules_dir, model); + +#if 0 + for (auto &[src, rule] : exprs_) { + log() << "Expr key: " << src << '\n'; + rule.dump(); + } + for (auto &[src, rule] : types_) { + log() << "Type key: " << src << '\n'; + rule.dump(); + } +#endif +} + +} // namespace cpp2rust::RuleRegistry diff --git a/cpp2rust/converter/rules/registry.h b/cpp2rust/converter/rules/registry.h new file mode 100644 index 000000000..eba8a93be --- /dev/null +++ b/cpp2rust/converter/rules/registry.h @@ -0,0 +1,33 @@ +#pragma once + +// Copyright (c) 2022-present INESC-ID. +// Distributed under the MIT license that can be found in the LICENSE file. + +#include +#include +#include +#include + +#include +#include + +#include "converter/factory.h" +#include "converter/rules/matching.h" +#include "converter/translation_rule.h" + +namespace cpp2rust::RuleRegistry { +template using Match = std::pair; + +Match SearchExpr(const std::string &str); +Match SearchType(const std::string &str); +TranslationRule::ExprRule *Search(clang::ASTContext &ctx, + const clang::Expr *expr); +Match Search(clang::ASTContext &ctx, + clang::QualType qual_type); +bool HasExprKey(const std::string &str); + +Model CurrentModel(); + +void Load(Model model, const std::string &rules_dir); +void AddRuleForUserDefinedType(clang::ASTContext &ctx, clang::NamedDecl *decl); +} // namespace cpp2rust::RuleRegistry diff --git a/cpp2rust/cpp_rule_preprocessor.cpp b/cpp2rust/cpp_rule_preprocessor.cpp index f24dc6b6c..cff775b80 100644 --- a/cpp2rust/cpp_rule_preprocessor.cpp +++ b/cpp2rust/cpp_rule_preprocessor.cpp @@ -33,7 +33,6 @@ #include "compat/platform_flags.h" #include "converter/converter_lib.h" -#include "converter/mapper.h" #include "converter/printer.h" namespace fs = std::filesystem; @@ -98,7 +97,6 @@ class Callback : public clang::ast_matchers::MatchFinder::MatchCallback { void run(const clang::ast_matchers::MatchFinder::MatchResult &R) override { assert(sema_); - Mapper::PushASTContext scoped(*R.Context); if (auto func = R.Nodes.getNodeAs("validate_func")) { const char *err = nullptr; if (auto body = From 05e3d41166f6c1895653989a6856272572135f95 Mon Sep 17 00:00:00 2001 From: Lucian Popescu Date: Tue, 29 Sep 2026 14:37:53 +0100 Subject: [PATCH 3/7] Add matcher inteface --- cpp2rust/converter/mapper.cpp | 74 ++++----- cpp2rust/converter/mapper.h | 4 + cpp2rust/converter/printer.cpp | 2 +- cpp2rust/converter/rules/matching.h | 31 +++- cpp2rust/converter/rules/registry.cpp | 130 +++------------- cpp2rust/converter/rules/registry.h | 27 ++-- .../{matching.cpp => string_matcher.cpp} | 145 +++++++++++++++++- cpp2rust/converter/rules/string_matcher.h | 22 +++ 8 files changed, 265 insertions(+), 170 deletions(-) rename cpp2rust/converter/rules/{matching.cpp => string_matcher.cpp} (58%) create mode 100644 cpp2rust/converter/rules/string_matcher.h diff --git a/cpp2rust/converter/mapper.cpp b/cpp2rust/converter/mapper.cpp index 98067eddb..1b0e8b039 100644 --- a/cpp2rust/converter/mapper.cpp +++ b/cpp2rust/converter/mapper.cpp @@ -14,7 +14,6 @@ #include #include "converter/converter_lib.h" -#include "converter/printer.h" #include "converter/rules/registry.h" #include "converter/translation_rule.h" @@ -22,6 +21,19 @@ namespace cpp2rust::Mapper { namespace { +Matching::Bindings mapBindings(clang::ASTContext &ctx, + const Matching::Bindings &bindings) { + Matching::Bindings mapped(bindings.size()); + for (unsigned i = 0; i < bindings.size(); ++i) { + if (bindings[i]) { + mapped[i] = RuleRegistry::GetMatcher().MapBinding(ctx, bindings, i); + } + } + return mapped; +} + +} // namespace + // Substitutes concrete types into a target template string using the provided // type mapping. Each template parameter in `tgt_template` is replaced with its // corresponding instantiated type from `types`. @@ -30,7 +42,7 @@ namespace { // types = { {"i32"} } // tgt_template = "Vec" // result = "Vec" -std::string instantiateTgt(const Matching::Bindings &types, +std::string InstantiateTgt(const Matching::Bindings &types, const std::string &tgt_template) { assert(types.size() <= TranslationRule::kMaxGenerics && "template placeholder exceeds kMaxGenerics"); @@ -51,33 +63,17 @@ std::string instantiateTgt(const Matching::Bindings &types, return instantiated_template; } -std::string mapTypeStringRecursive(const std::string &cpp_type) { - auto [rule, subs] = RuleRegistry::SearchType(cpp_type); - if (!rule) { - llvm::errs() << "cpp_type: " << cpp_type << '\n'; - assert(0 && "Type is not present in the registry"); - } - for (auto &ty : subs) { - if (ty) { - ty = mapTypeStringRecursive(*ty); - } - } - return instantiateTgt(subs, rule->type_info.type); -} - -} // namespace - bool Contains(clang::ASTContext &ctx, clang::QualType qual_type) { return RuleRegistry::Search(ctx, qual_type).first != nullptr; } bool Contains(clang::ASTContext &ctx, const clang::Expr *expr) { - return RuleRegistry::Search(ctx, expr) != nullptr; + return RuleRegistry::Search(ctx, expr).first != nullptr; } const TranslationRule::ExprRule *GetExprRule(clang::ASTContext &ctx, const clang::Expr *expr) { - return RuleRegistry::Search(ctx, expr); + return RuleRegistry::Search(ctx, expr).first; } bool IsLibcPassthrough(clang::ASTContext &ctx, const clang::Expr *expr) { @@ -97,7 +93,7 @@ std::string MapFunctionName(clang::ASTContext &ctx, const clang::FunctionDecl *decl) { assert(decl); if (!IsUserDefinedDecl(decl) && - RuleRegistry::HasExprKey(Printer::ToString(ctx, decl))) { + RuleRegistry::GetMatcher().HasRuleNamed(ctx, decl)) { return std::format("libcc2rs::{}_{}", decl->getNameAsString(), RuleRegistry::CurrentModel() == Model::kRefCount ? "refcount" @@ -108,28 +104,22 @@ std::string MapFunctionName(clang::ASTContext &ctx, std::string InstantiateTemplate(clang::ASTContext &ctx, const clang::Expr *expr, unsigned n) { - auto expr_str = Printer::ToString(ctx, expr); - auto [rule, subs] = RuleRegistry::SearchExpr(expr_str); + auto [rule, subs] = RuleRegistry::Search(ctx, expr); auto text = std::format("T{}", n); if (!rule) { return text; } auto &ty = subs.at(n - 1); if (ty) { - ty = mapTypeStringRecursive(*ty); + ty = RuleRegistry::GetMatcher().MapBinding(ctx, subs, n - 1); } - return instantiateTgt(subs, text); + return InstantiateTgt(subs, text); } std::string Map(clang::ASTContext &ctx, clang::QualType qual_type) { auto [rule, subs] = RuleRegistry::Search(ctx, qual_type); if (rule) { - for (auto &ty : subs) { - if (ty) { - ty = mapTypeStringRecursive(*ty); - } - } - return instantiateTgt(subs, rule->type_info.type); + return InstantiateTgt(mapBindings(ctx, subs), rule->type_info.type); } return {}; } @@ -137,12 +127,7 @@ std::string Map(clang::ASTContext &ctx, clang::QualType qual_type) { std::string MapInitializer(clang::ASTContext &ctx, clang::QualType qual_type) { auto [rule, subs] = RuleRegistry::Search(ctx, qual_type); if (rule && !rule->initializer.empty()) { - for (auto &ty : subs) { - if (ty) { - ty = mapTypeStringRecursive(*ty); - } - } - return instantiateTgt(subs, rule->initializer); + return InstantiateTgt(mapBindings(ctx, subs), rule->initializer); } return {}; } @@ -169,28 +154,23 @@ void SetDerives(clang::ASTContext &ctx, clang::QualType qual_type, rule->type_info.derives = std::move(derives); } } + bool ReturnsPointer(clang::ASTContext &ctx, const clang::Expr *expr) { - auto rule = RuleRegistry::Search(ctx, expr); + auto rule = RuleRegistry::Search(ctx, expr).first; return rule && rule->return_type.is_pointer(); } const TranslationRule::TypeInfo & GetParamInfo(clang::ASTContext &ctx, const clang::Expr *expr, unsigned index) { - auto rule = RuleRegistry::Search(ctx, expr); + auto rule = RuleRegistry::Search(ctx, expr).first; assert(rule && "expression must have a translation rule"); return rule->params.at(index); } std::string GetParamType(clang::ASTContext &ctx, const clang::Expr *expr, unsigned index) { - auto expr_str = Printer::ToString(ctx, expr); - auto [rule, subs] = RuleRegistry::SearchExpr(expr_str); - for (auto &ty : subs) { - if (ty) { - ty = mapTypeStringRecursive(*ty); - } - } - return instantiateTgt(subs, rule->params.at(index).type); + auto [rule, subs] = RuleRegistry::Search(ctx, expr); + return InstantiateTgt(mapBindings(ctx, subs), rule->params.at(index).type); } bool ParamIsPointer(clang::ASTContext &ctx, const clang::Expr *expr, diff --git a/cpp2rust/converter/mapper.h b/cpp2rust/converter/mapper.h index 6aa0f4af5..a2f8e76c3 100644 --- a/cpp2rust/converter/mapper.h +++ b/cpp2rust/converter/mapper.h @@ -9,6 +9,7 @@ #include +#include "converter/rules/matching.h" #include "converter/translation_rule.h" namespace cpp2rust::Mapper { @@ -35,4 +36,7 @@ const std::vector *MappedDerives(clang::ASTContext &ctx, clang::QualType qual_type); void SetDerives(clang::ASTContext &ctx, clang::QualType qual_type, std::vector derives); + +std::string InstantiateTgt(const Matching::Bindings &types, + const std::string &tgt_template); } // namespace cpp2rust::Mapper diff --git a/cpp2rust/converter/printer.cpp b/cpp2rust/converter/printer.cpp index 37e982960..7475a4f5b 100644 --- a/cpp2rust/converter/printer.cpp +++ b/cpp2rust/converter/printer.cpp @@ -213,7 +213,7 @@ std::string ToString(clang::ASTContext &ctx, const clang::NamedDecl *decl) { if (const auto op = func_decl->getOverloadedOperator(); op >= clang::OverloadedOperatorKind::OO_LessLess && op <= clang::OverloadedOperatorKind::OO_GreaterGreaterEqual) { - // ensure MatchTemplate does not consider these operator names when matching + // ensure matchTemplate does not consider these operator names when matching func_decl->getQualifier().print(os, getPrintPolicy(ctx)); os << "operator "; switch (op) { diff --git a/cpp2rust/converter/rules/matching.h b/cpp2rust/converter/rules/matching.h index c5026b1b4..41a0a4f24 100644 --- a/cpp2rust/converter/rules/matching.h +++ b/cpp2rust/converter/rules/matching.h @@ -3,14 +3,39 @@ // Copyright (c) 2022-present INESC-ID. // Distributed under the MIT license that can be found in the LICENSE file. +#include +#include +#include +#include + #include #include +#include #include +#include "converter/translation_rule.h" + namespace cpp2rust::Matching { -// Concrete C++ types bound to the rule's T1, T2, ... placeholders. using Bindings = std::vector>; -std::optional MatchTemplate(const std::string &template_str, - const std::string &instantiated); +template using Match = std::pair; + +class Matcher { +public: + virtual ~Matcher() = default; + + virtual std::string Key(const TranslationRule::ExprRule &rule) const = 0; + virtual std::string Key(const TranslationRule::TypeRule &rule) const = 0; + + virtual Match Find(clang::ASTContext &ctx, + const clang::Expr *expr) = 0; + virtual Match Find(clang::ASTContext &ctx, + clang::QualType type) = 0; + + virtual bool HasRuleNamed(clang::ASTContext &ctx, + const clang::FunctionDecl *decl) = 0; + + virtual std::string MapBinding(clang::ASTContext &ctx, + const Bindings &bindings, unsigned n) = 0; +}; } // namespace cpp2rust::Matching diff --git a/cpp2rust/converter/rules/registry.cpp b/cpp2rust/converter/rules/registry.cpp index f64a35be6..f121a7edf 100644 --- a/cpp2rust/converter/rules/registry.cpp +++ b/cpp2rust/converter/rules/registry.cpp @@ -6,67 +6,28 @@ #include #include #include -#include +#include #include #include "converter/converter_lib.h" #include "converter/printer.h" +#include "converter/rules/string_matcher.h" namespace cpp2rust::RuleRegistry { namespace { -using ExprRuleMap = - std::unordered_multimap; -using TypeRuleMap = - std::unordered_multimap; - Model model_ = Model::kUnsafe; bool translation_rules_loaded_ = false; +std::unique_ptr matcher_; + ExprRuleMap exprs_; // key -> ExprRule TypeRuleMap types_; // key -> TypeRule -std::string ExprKey(const std::string &str) { - // Extract the function name from something like - // const T1 & std::foo::fn_name(args) - auto n = str.find_first_of('('); - if (n == std::string::npos) { - n = str.size(); - } - - // Walk backwards from '(' tracking <> depth: - // - skip characters inside template arguments (depth > 0) - // - stop at the first space outside all angle brackets - std::string result; - int depth = 0; - for (int i = (int)n - 1; i >= 0; --i) { - char c = str[i]; - if (c == '>') - ++depth; - else if (c == '<') - --depth; - else if (c == ' ' && depth == 0) - break; - else if (depth == 0) - result += c; - } - std::reverse(result.begin(), result.end()); - return result; -} - -std::string TypeKey(const std::string &str) { - auto n = str.find_first_of("<["); - if (n == std::string::npos || str[n] == '<') { - return str.substr(0, n); - } - // something like int[][] or T1[] -> [] - return str.substr(n + 1); -} - void AddTypeRule(std::string src, TranslationRule::TypeRule &&rule) { - auto key = TypeKey(src); rule.src = std::move(src); + auto key = matcher_->Key(rule); types_.emplace(std::move(key), std::move(rule)); } @@ -83,10 +44,10 @@ void addRulesFromDirectory(const std::filesystem::path &dir, Model model) { continue; } for (auto &[_, rule] : expr_rules) { - exprs_.emplace(ExprKey(rule.src), std::move(rule)); + exprs_.emplace(matcher_->Key(rule), std::move(rule)); } for (auto &[_, rule] : type_rules) { - auto key = TypeKey(rule.src); + auto key = matcher_->Key(rule); auto [begin, end] = types_.equal_range(key); for (auto it = begin; it != end; ++it) { if (it->second.src == rule.src) { @@ -102,79 +63,35 @@ void addRulesFromDirectory(const std::filesystem::path &dir, Model model) { } } -template -Match search(std::unordered_multimap &map, - const std::string &txt, const std::string &key) { - auto [it, end] = map.equal_range(key); - T *rule = nullptr; - Matching::Bindings subs; - - for (; it != end; ++it) { - auto &this_rule = it->second; - auto this_subs = Matching::MatchTemplate(this_rule.src, txt); - if (!this_subs) { - continue; - } - // tie breaker: prefer more specific rules (usually the longer ones) - if (!rule || this_rule.src.size() > rule->src.size()) { - rule = &this_rule; - subs = *std::move(this_subs); - } - } - return {rule, std::move(subs)}; -} - } // namespace -Match SearchExpr(const std::string &str) { - return search(exprs_, str, ExprKey(str)); +std::ranges::subrange +ExprCandidates(const std::string &key) { + auto [begin, end] = exprs_.equal_range(key); + return {begin, end}; } -Match SearchType(const std::string &str) { - return search(types_, str, TypeKey(str)); +std::ranges::subrange +TypeCandidates(const std::string &key) { + auto [begin, end] = types_.equal_range(key); + return {begin, end}; } -TranslationRule::ExprRule *Search(clang::ASTContext &ctx, - const clang::Expr *expr) { +Matching::Match Search(clang::ASTContext &ctx, + const clang::Expr *expr) { if (RefersToUserDefinedDecl(expr)) { - return nullptr; - } - auto qualified_name = Printer::ToString(ctx, expr); - auto [rule, subs] = SearchExpr(qualified_name); - log() << "search expr " << qualified_name << ", result:\n"; - if (rule) { - rule->dump(); - } else { - log() << "None\n"; - } - return rule; -} - -Match Search(clang::ASTContext &ctx, - clang::QualType qual_type) { - auto sugared = - Printer::ToString(ctx, qual_type, Printer::ScalarSugar::kPreserve); - if (auto res = SearchType(sugared); res.first) { - log() << "search type " << sugared - << ", result: " << res.first->type_info.type << '\n'; - return res; - } - auto type = Printer::ToString(ctx, qual_type); - if (type == sugared) { - log() << "search type " << type << ", result: None\n"; return {}; } - auto res = SearchType(type); - log() << "search type " << type - << ", result: " << (res.first ? res.first->type_info.type : "None") - << '\n'; - return res; + return matcher_->Find(ctx, expr); } -bool HasExprKey(const std::string &str) { - return exprs_.contains(ExprKey(str)); +Matching::Match Search(clang::ASTContext &ctx, + clang::QualType qual_type) { + return matcher_->Find(ctx, qual_type); } +Matching::Matcher &GetMatcher() { return *matcher_; } + Model CurrentModel() { return model_; } void AddRuleForUserDefinedType(clang::ASTContext &ctx, clang::NamedDecl *decl) { @@ -229,6 +146,7 @@ void Load(Model model, const std::string &rules_dir) { } translation_rules_loaded_ = true; + matcher_ = std::make_unique(); addRulesFromDirectory(rules_dir, model); #if 0 diff --git a/cpp2rust/converter/rules/registry.h b/cpp2rust/converter/rules/registry.h index eba8a93be..8b1a1e698 100644 --- a/cpp2rust/converter/rules/registry.h +++ b/cpp2rust/converter/rules/registry.h @@ -8,23 +8,30 @@ #include #include +#include #include -#include +#include #include "converter/factory.h" #include "converter/rules/matching.h" #include "converter/translation_rule.h" namespace cpp2rust::RuleRegistry { -template using Match = std::pair; - -Match SearchExpr(const std::string &str); -Match SearchType(const std::string &str); -TranslationRule::ExprRule *Search(clang::ASTContext &ctx, - const clang::Expr *expr); -Match Search(clang::ASTContext &ctx, - clang::QualType qual_type); -bool HasExprKey(const std::string &str); +using ExprRuleMap = + std::unordered_multimap; +using TypeRuleMap = + std::unordered_multimap; + +std::ranges::subrange +ExprCandidates(const std::string &key); +std::ranges::subrange +TypeCandidates(const std::string &key); + +Matching::Match Search(clang::ASTContext &ctx, + const clang::Expr *expr); +Matching::Match Search(clang::ASTContext &ctx, + clang::QualType qual_type); +Matching::Matcher &GetMatcher(); Model CurrentModel(); diff --git a/cpp2rust/converter/rules/matching.cpp b/cpp2rust/converter/rules/string_matcher.cpp similarity index 58% rename from cpp2rust/converter/rules/matching.cpp rename to cpp2rust/converter/rules/string_matcher.cpp index 81cb942d5..45e7edb8c 100644 --- a/cpp2rust/converter/rules/matching.cpp +++ b/cpp2rust/converter/rules/string_matcher.cpp @@ -1,17 +1,22 @@ // Copyright (c) 2022-present INESC-ID. // Distributed under the MIT license that can be found in the LICENSE file. -#include "converter/rules/matching.h" +#include "converter/rules/string_matcher.h" #include #include #include #include -#include "converter/translation_rule.h" +#include "converter/converter_lib.h" +#include "converter/mapper.h" +#include "converter/printer.h" +#include "converter/rules/registry.h" namespace cpp2rust::Matching { +namespace { + // Attempts to unify an instantiated C++ type or function signature with a // corresponding template pattern. If the two match structurally, it returns // a mapping from template parameter names (e.g., "T1") to their concrete @@ -21,7 +26,7 @@ namespace cpp2rust::Matching { // template_str = "std::vector::vector()" // instantiated = "std::vector::vector()" // result = { "int" } -std::optional MatchTemplate(const std::string &template_str, +std::optional matchTemplate(const std::string &template_str, const std::string &instantiated) { auto matchLiteralAt = [&](const std::string &input_str, size_t pos, std::string_view literal, size_t &end_pos) -> bool { @@ -258,4 +263,138 @@ std::optional MatchTemplate(const std::string &template_str, return captured; } +std::string exprKey(const std::string &str) { + // Extract the function name from something like + // const T1 & std::foo::fn_name(args) + auto n = str.find_first_of('('); + if (n == std::string::npos) { + n = str.size(); + } + + // Walk backwards from '(' tracking <> depth: + // - skip characters inside template arguments (depth > 0) + // - stop at the first space outside all angle brackets + std::string result; + int depth = 0; + for (int i = (int)n - 1; i >= 0; --i) { + char c = str[i]; + if (c == '>') + ++depth; + else if (c == '<') + --depth; + else if (c == ' ' && depth == 0) + break; + else if (depth == 0) + result += c; + } + std::reverse(result.begin(), result.end()); + return result; +} + +std::string typeKey(const std::string &str) { + auto n = str.find_first_of("<["); + if (n == std::string::npos || str[n] == '<') { + return str.substr(0, n); + } + // something like int[][] or T1[] -> [] + return str.substr(n + 1); +} + +template +Match search(Candidates candidates, const std::string &txt) { + Rule *rule = nullptr; + Bindings subs; + + for (auto &[_, this_rule] : candidates) { + auto this_subs = matchTemplate(this_rule.src, txt); + if (!this_subs) { + continue; + } + // tie breaker: prefer more specific rules (usually the longer ones) + if (!rule || this_rule.src.size() > rule->src.size()) { + rule = &this_rule; + subs = *std::move(this_subs); + } + } + return {rule, std::move(subs)}; +} + +Match searchExpr(const std::string &txt) { + return search( + RuleRegistry::ExprCandidates(exprKey(txt)), txt); +} + +Match searchType(const std::string &txt) { + return search( + RuleRegistry::TypeCandidates(typeKey(txt)), txt); +} + +std::string mapTypeString(const std::string &cpp_type) { + auto [rule, subs] = searchType(cpp_type); + if (!rule) { + llvm::errs() << "cpp_type: " << cpp_type << '\n'; + assert(0 && "Type is not present in the registry"); + } + for (auto &ty : subs) { + if (ty) { + ty = mapTypeString(*ty); + } + } + return Mapper::InstantiateTgt(subs, rule->type_info.type); +} + +} // namespace + +std::string StringMatcher::Key(const TranslationRule::ExprRule &rule) const { + return exprKey(rule.src); +} + +std::string StringMatcher::Key(const TranslationRule::TypeRule &rule) const { + return typeKey(rule.src); +} + +Match StringMatcher::Find(clang::ASTContext &ctx, + const clang::Expr *expr) { + auto qualified_name = Printer::ToString(ctx, expr); + auto res = searchExpr(qualified_name); + log() << "search expr " << qualified_name << ", result:\n"; + if (res.first) { + res.first->dump(); + } else { + log() << "None\n"; + } + return res; +} + +Match StringMatcher::Find(clang::ASTContext &ctx, + clang::QualType type) { + auto sugared = Printer::ToString(ctx, type, Printer::ScalarSugar::kPreserve); + if (auto res = searchType(sugared); res.first) { + log() << "search type " << sugared + << ", result: " << res.first->type_info.type << '\n'; + return res; + } + auto desugared = Printer::ToString(ctx, type); + if (desugared == sugared) { + log() << "search type " << desugared << ", result: None\n"; + return {}; + } + auto res = searchType(desugared); + log() << "search type " << desugared + << ", result: " << (res.first ? res.first->type_info.type : "None") + << '\n'; + return res; +} + +bool StringMatcher::HasRuleNamed(clang::ASTContext &ctx, + const clang::FunctionDecl *decl) { + return !RuleRegistry::ExprCandidates(exprKey(Printer::ToString(ctx, decl))) + .empty(); +} + +std::string StringMatcher::MapBinding(clang::ASTContext &ctx, + const Bindings &bindings, unsigned n) { + return mapTypeString(bindings.at(n).value()); +} + } // namespace cpp2rust::Matching diff --git a/cpp2rust/converter/rules/string_matcher.h b/cpp2rust/converter/rules/string_matcher.h new file mode 100644 index 000000000..34dfe9dc4 --- /dev/null +++ b/cpp2rust/converter/rules/string_matcher.h @@ -0,0 +1,22 @@ +#pragma once + +// Copyright (c) 2022-present INESC-ID. +// Distributed under the MIT license that can be found in the LICENSE file. + +#include "converter/rules/matching.h" + +namespace cpp2rust::Matching { +class StringMatcher final : public Matcher { +public: + std::string Key(const TranslationRule::ExprRule &rule) const override; + std::string Key(const TranslationRule::TypeRule &rule) const override; + Match Find(clang::ASTContext &ctx, + const clang::Expr *expr) override; + Match Find(clang::ASTContext &ctx, + clang::QualType type) override; + bool HasRuleNamed(clang::ASTContext &ctx, + const clang::FunctionDecl *decl) override; + std::string MapBinding(clang::ASTContext &ctx, const Bindings &bindings, + unsigned n) override; +}; +} // namespace cpp2rust::Matching From 8aff06e0ab5b4bcea6ad282e1ab5567c6a8c8e35 Mon Sep 17 00:00:00 2001 From: Lucian Popescu Date: Tue, 29 Sep 2026 14:42:00 +0100 Subject: [PATCH 4/7] Move instantiateTgt in matcher --- cpp2rust/converter/mapper.cpp | 40 ++++----------------- cpp2rust/converter/mapper.h | 4 --- cpp2rust/converter/rules/matcher.cpp | 40 +++++++++++++++++++++ cpp2rust/converter/rules/matching.h | 3 ++ cpp2rust/converter/rules/string_matcher.cpp | 3 +- 5 files changed, 50 insertions(+), 40 deletions(-) create mode 100644 cpp2rust/converter/rules/matcher.cpp diff --git a/cpp2rust/converter/mapper.cpp b/cpp2rust/converter/mapper.cpp index 1b0e8b039..fbb3dfebe 100644 --- a/cpp2rust/converter/mapper.cpp +++ b/cpp2rust/converter/mapper.cpp @@ -6,7 +6,6 @@ #include #include -#include #include #include #include @@ -34,35 +33,6 @@ Matching::Bindings mapBindings(clang::ASTContext &ctx, } // namespace -// Substitutes concrete types into a target template string using the provided -// type mapping. Each template parameter in `tgt_template` is replaced with its -// corresponding instantiated type from `types`. -// -// Example: -// types = { {"i32"} } -// tgt_template = "Vec" -// result = "Vec" -std::string InstantiateTgt(const Matching::Bindings &types, - const std::string &tgt_template) { - assert(types.size() <= TranslationRule::kMaxGenerics && - "template placeholder exceeds kMaxGenerics"); - std::string instantiated_template = tgt_template; - std::string::size_type pos = 0; - while ((pos = instantiated_template.find('T', pos)) != std::string::npos) { - if (pos + 1 >= instantiated_template.size()) { - break; - } - if (!std::isdigit(instantiated_template[pos + 1])) { - ++pos; - continue; - } - const auto &repl = types.at(instantiated_template[pos + 1] - '1').value(); - instantiated_template.replace(pos, 2, repl); - pos += repl.length(); - } - return instantiated_template; -} - bool Contains(clang::ASTContext &ctx, clang::QualType qual_type) { return RuleRegistry::Search(ctx, qual_type).first != nullptr; } @@ -113,13 +83,14 @@ std::string InstantiateTemplate(clang::ASTContext &ctx, const clang::Expr *expr, if (ty) { ty = RuleRegistry::GetMatcher().MapBinding(ctx, subs, n - 1); } - return InstantiateTgt(subs, text); + return Matching::InstantiateTgt(subs, text); } std::string Map(clang::ASTContext &ctx, clang::QualType qual_type) { auto [rule, subs] = RuleRegistry::Search(ctx, qual_type); if (rule) { - return InstantiateTgt(mapBindings(ctx, subs), rule->type_info.type); + return Matching::InstantiateTgt(mapBindings(ctx, subs), + rule->type_info.type); } return {}; } @@ -127,7 +98,7 @@ std::string Map(clang::ASTContext &ctx, clang::QualType qual_type) { std::string MapInitializer(clang::ASTContext &ctx, clang::QualType qual_type) { auto [rule, subs] = RuleRegistry::Search(ctx, qual_type); if (rule && !rule->initializer.empty()) { - return InstantiateTgt(mapBindings(ctx, subs), rule->initializer); + return Matching::InstantiateTgt(mapBindings(ctx, subs), rule->initializer); } return {}; } @@ -170,7 +141,8 @@ GetParamInfo(clang::ASTContext &ctx, const clang::Expr *expr, unsigned index) { std::string GetParamType(clang::ASTContext &ctx, const clang::Expr *expr, unsigned index) { auto [rule, subs] = RuleRegistry::Search(ctx, expr); - return InstantiateTgt(mapBindings(ctx, subs), rule->params.at(index).type); + return Matching::InstantiateTgt(mapBindings(ctx, subs), + rule->params.at(index).type); } bool ParamIsPointer(clang::ASTContext &ctx, const clang::Expr *expr, diff --git a/cpp2rust/converter/mapper.h b/cpp2rust/converter/mapper.h index a2f8e76c3..6aa0f4af5 100644 --- a/cpp2rust/converter/mapper.h +++ b/cpp2rust/converter/mapper.h @@ -9,7 +9,6 @@ #include -#include "converter/rules/matching.h" #include "converter/translation_rule.h" namespace cpp2rust::Mapper { @@ -36,7 +35,4 @@ const std::vector *MappedDerives(clang::ASTContext &ctx, clang::QualType qual_type); void SetDerives(clang::ASTContext &ctx, clang::QualType qual_type, std::vector derives); - -std::string InstantiateTgt(const Matching::Bindings &types, - const std::string &tgt_template); } // namespace cpp2rust::Mapper diff --git a/cpp2rust/converter/rules/matcher.cpp b/cpp2rust/converter/rules/matcher.cpp new file mode 100644 index 000000000..ad93454ee --- /dev/null +++ b/cpp2rust/converter/rules/matcher.cpp @@ -0,0 +1,40 @@ +// Copyright (c) 2022-present INESC-ID. +// Distributed under the MIT license that can be found in the LICENSE file. + +#include +#include + +#include "converter/rules/matching.h" + +namespace cpp2rust::Matching { + +// Substitutes concrete types into a target template string using the provided +// type mapping. Each template parameter in `tgt_template` is replaced with its +// corresponding instantiated type from `types`. +// +// Example: +// types = { {"i32"} } +// tgt_template = "Vec" +// result = "Vec" +std::string InstantiateTgt(const Bindings &types, + const std::string &tgt_template) { + assert(types.size() <= TranslationRule::kMaxGenerics && + "template placeholder exceeds kMaxGenerics"); + std::string instantiated_template = tgt_template; + std::string::size_type pos = 0; + while ((pos = instantiated_template.find('T', pos)) != std::string::npos) { + if (pos + 1 >= instantiated_template.size()) { + break; + } + if (!std::isdigit(instantiated_template[pos + 1])) { + ++pos; + continue; + } + const auto &repl = types.at(instantiated_template[pos + 1] - '1').value(); + instantiated_template.replace(pos, 2, repl); + pos += repl.length(); + } + return instantiated_template; +} + +} // namespace cpp2rust::Matching diff --git a/cpp2rust/converter/rules/matching.h b/cpp2rust/converter/rules/matching.h index 41a0a4f24..5dd4933c1 100644 --- a/cpp2rust/converter/rules/matching.h +++ b/cpp2rust/converter/rules/matching.h @@ -38,4 +38,7 @@ class Matcher { virtual std::string MapBinding(clang::ASTContext &ctx, const Bindings &bindings, unsigned n) = 0; }; + +std::string InstantiateTgt(const Bindings &types, + const std::string &tgt_template); } // namespace cpp2rust::Matching diff --git a/cpp2rust/converter/rules/string_matcher.cpp b/cpp2rust/converter/rules/string_matcher.cpp index 45e7edb8c..4daab7e39 100644 --- a/cpp2rust/converter/rules/string_matcher.cpp +++ b/cpp2rust/converter/rules/string_matcher.cpp @@ -9,7 +9,6 @@ #include #include "converter/converter_lib.h" -#include "converter/mapper.h" #include "converter/printer.h" #include "converter/rules/registry.h" @@ -340,7 +339,7 @@ std::string mapTypeString(const std::string &cpp_type) { ty = mapTypeString(*ty); } } - return Mapper::InstantiateTgt(subs, rule->type_info.type); + return InstantiateTgt(subs, rule->type_info.type); } } // namespace From b13230cef4227094a6d8fea200dfaaac91f3fead Mon Sep 17 00:00:00 2001 From: Lucian Popescu Date: Tue, 29 Sep 2026 14:46:35 +0100 Subject: [PATCH 5/7] Drop isntante of matcher --- cpp2rust/converter/mapper.cpp | 26 ++++++------ cpp2rust/converter/rules/matcher.cpp | 8 ++-- cpp2rust/converter/rules/matcher.h | 38 ++++++++++++++++++ cpp2rust/converter/rules/matching.h | 44 --------------------- cpp2rust/converter/rules/registry.cpp | 25 +++++------- cpp2rust/converter/rules/registry.h | 11 +++--- cpp2rust/converter/rules/string_matcher.cpp | 26 ++++++------ cpp2rust/converter/rules/string_matcher.h | 22 ----------- 8 files changed, 81 insertions(+), 119 deletions(-) create mode 100644 cpp2rust/converter/rules/matcher.h delete mode 100644 cpp2rust/converter/rules/matching.h delete mode 100644 cpp2rust/converter/rules/string_matcher.h diff --git a/cpp2rust/converter/mapper.cpp b/cpp2rust/converter/mapper.cpp index fbb3dfebe..85a3a36e2 100644 --- a/cpp2rust/converter/mapper.cpp +++ b/cpp2rust/converter/mapper.cpp @@ -13,6 +13,7 @@ #include #include "converter/converter_lib.h" +#include "converter/rules/matcher.h" #include "converter/rules/registry.h" #include "converter/translation_rule.h" @@ -20,12 +21,12 @@ namespace cpp2rust::Mapper { namespace { -Matching::Bindings mapBindings(clang::ASTContext &ctx, - const Matching::Bindings &bindings) { - Matching::Bindings mapped(bindings.size()); +Matcher::Bindings mapBindings(clang::ASTContext &ctx, + const Matcher::Bindings &bindings) { + Matcher::Bindings mapped(bindings.size()); for (unsigned i = 0; i < bindings.size(); ++i) { if (bindings[i]) { - mapped[i] = RuleRegistry::GetMatcher().MapBinding(ctx, bindings, i); + mapped[i] = Matcher::MapBinding(ctx, bindings, i); } } return mapped; @@ -62,8 +63,7 @@ bool IsLibcPassthrough(clang::ASTContext &ctx, const clang::Expr *expr) { std::string MapFunctionName(clang::ASTContext &ctx, const clang::FunctionDecl *decl) { assert(decl); - if (!IsUserDefinedDecl(decl) && - RuleRegistry::GetMatcher().HasRuleNamed(ctx, decl)) { + if (!IsUserDefinedDecl(decl) && Matcher::HasRuleNamed(ctx, decl)) { return std::format("libcc2rs::{}_{}", decl->getNameAsString(), RuleRegistry::CurrentModel() == Model::kRefCount ? "refcount" @@ -81,16 +81,16 @@ std::string InstantiateTemplate(clang::ASTContext &ctx, const clang::Expr *expr, } auto &ty = subs.at(n - 1); if (ty) { - ty = RuleRegistry::GetMatcher().MapBinding(ctx, subs, n - 1); + ty = Matcher::MapBinding(ctx, subs, n - 1); } - return Matching::InstantiateTgt(subs, text); + return Matcher::InstantiateTgt(subs, text); } std::string Map(clang::ASTContext &ctx, clang::QualType qual_type) { auto [rule, subs] = RuleRegistry::Search(ctx, qual_type); if (rule) { - return Matching::InstantiateTgt(mapBindings(ctx, subs), - rule->type_info.type); + return Matcher::InstantiateTgt(mapBindings(ctx, subs), + rule->type_info.type); } return {}; } @@ -98,7 +98,7 @@ std::string Map(clang::ASTContext &ctx, clang::QualType qual_type) { std::string MapInitializer(clang::ASTContext &ctx, clang::QualType qual_type) { auto [rule, subs] = RuleRegistry::Search(ctx, qual_type); if (rule && !rule->initializer.empty()) { - return Matching::InstantiateTgt(mapBindings(ctx, subs), rule->initializer); + return Matcher::InstantiateTgt(mapBindings(ctx, subs), rule->initializer); } return {}; } @@ -141,8 +141,8 @@ GetParamInfo(clang::ASTContext &ctx, const clang::Expr *expr, unsigned index) { std::string GetParamType(clang::ASTContext &ctx, const clang::Expr *expr, unsigned index) { auto [rule, subs] = RuleRegistry::Search(ctx, expr); - return Matching::InstantiateTgt(mapBindings(ctx, subs), - rule->params.at(index).type); + return Matcher::InstantiateTgt(mapBindings(ctx, subs), + rule->params.at(index).type); } bool ParamIsPointer(clang::ASTContext &ctx, const clang::Expr *expr, diff --git a/cpp2rust/converter/rules/matcher.cpp b/cpp2rust/converter/rules/matcher.cpp index ad93454ee..a00cad0de 100644 --- a/cpp2rust/converter/rules/matcher.cpp +++ b/cpp2rust/converter/rules/matcher.cpp @@ -1,12 +1,12 @@ // Copyright (c) 2022-present INESC-ID. // Distributed under the MIT license that can be found in the LICENSE file. +#include "converter/rules/matcher.h" + #include #include -#include "converter/rules/matching.h" - -namespace cpp2rust::Matching { +namespace cpp2rust::Matcher { // Substitutes concrete types into a target template string using the provided // type mapping. Each template parameter in `tgt_template` is replaced with its @@ -37,4 +37,4 @@ std::string InstantiateTgt(const Bindings &types, return instantiated_template; } -} // namespace cpp2rust::Matching +} // namespace cpp2rust::Matcher diff --git a/cpp2rust/converter/rules/matcher.h b/cpp2rust/converter/rules/matcher.h new file mode 100644 index 000000000..5676ee9be --- /dev/null +++ b/cpp2rust/converter/rules/matcher.h @@ -0,0 +1,38 @@ +#pragma once + +// Copyright (c) 2022-present INESC-ID. +// Distributed under the MIT license that can be found in the LICENSE file. + +#include +#include +#include +#include + +#include +#include +#include +#include + +#include "converter/translation_rule.h" + +namespace cpp2rust::Matcher { +using Bindings = std::vector>; + +template using Match = std::pair; + +std::string Key(const TranslationRule::ExprRule &rule); +std::string Key(const TranslationRule::TypeRule &rule); + +Match Find(clang::ASTContext &ctx, + const clang::Expr *expr); +Match Find(clang::ASTContext &ctx, + clang::QualType type); + +bool HasRuleNamed(clang::ASTContext &ctx, const clang::FunctionDecl *decl); + +std::string MapBinding(clang::ASTContext &ctx, const Bindings &bindings, + unsigned n); + +std::string InstantiateTgt(const Bindings &types, + const std::string &tgt_template); +} // namespace cpp2rust::Matcher diff --git a/cpp2rust/converter/rules/matching.h b/cpp2rust/converter/rules/matching.h deleted file mode 100644 index 5dd4933c1..000000000 --- a/cpp2rust/converter/rules/matching.h +++ /dev/null @@ -1,44 +0,0 @@ -#pragma once - -// Copyright (c) 2022-present INESC-ID. -// Distributed under the MIT license that can be found in the LICENSE file. - -#include -#include -#include -#include - -#include -#include -#include -#include - -#include "converter/translation_rule.h" - -namespace cpp2rust::Matching { -using Bindings = std::vector>; - -template using Match = std::pair; - -class Matcher { -public: - virtual ~Matcher() = default; - - virtual std::string Key(const TranslationRule::ExprRule &rule) const = 0; - virtual std::string Key(const TranslationRule::TypeRule &rule) const = 0; - - virtual Match Find(clang::ASTContext &ctx, - const clang::Expr *expr) = 0; - virtual Match Find(clang::ASTContext &ctx, - clang::QualType type) = 0; - - virtual bool HasRuleNamed(clang::ASTContext &ctx, - const clang::FunctionDecl *decl) = 0; - - virtual std::string MapBinding(clang::ASTContext &ctx, - const Bindings &bindings, unsigned n) = 0; -}; - -std::string InstantiateTgt(const Bindings &types, - const std::string &tgt_template); -} // namespace cpp2rust::Matching diff --git a/cpp2rust/converter/rules/registry.cpp b/cpp2rust/converter/rules/registry.cpp index f121a7edf..bf5b3a249 100644 --- a/cpp2rust/converter/rules/registry.cpp +++ b/cpp2rust/converter/rules/registry.cpp @@ -6,12 +6,10 @@ #include #include #include -#include #include #include "converter/converter_lib.h" #include "converter/printer.h" -#include "converter/rules/string_matcher.h" namespace cpp2rust::RuleRegistry { @@ -20,14 +18,12 @@ namespace { Model model_ = Model::kUnsafe; bool translation_rules_loaded_ = false; -std::unique_ptr matcher_; - ExprRuleMap exprs_; // key -> ExprRule TypeRuleMap types_; // key -> TypeRule void AddTypeRule(std::string src, TranslationRule::TypeRule &&rule) { rule.src = std::move(src); - auto key = matcher_->Key(rule); + auto key = Matcher::Key(rule); types_.emplace(std::move(key), std::move(rule)); } @@ -44,10 +40,10 @@ void addRulesFromDirectory(const std::filesystem::path &dir, Model model) { continue; } for (auto &[_, rule] : expr_rules) { - exprs_.emplace(matcher_->Key(rule), std::move(rule)); + exprs_.emplace(Matcher::Key(rule), std::move(rule)); } for (auto &[_, rule] : type_rules) { - auto key = matcher_->Key(rule); + auto key = Matcher::Key(rule); auto [begin, end] = types_.equal_range(key); for (auto it = begin; it != end; ++it) { if (it->second.src == rule.src) { @@ -77,21 +73,19 @@ TypeCandidates(const std::string &key) { return {begin, end}; } -Matching::Match Search(clang::ASTContext &ctx, - const clang::Expr *expr) { +Matcher::Match Search(clang::ASTContext &ctx, + const clang::Expr *expr) { if (RefersToUserDefinedDecl(expr)) { return {}; } - return matcher_->Find(ctx, expr); + return Matcher::Find(ctx, expr); } -Matching::Match Search(clang::ASTContext &ctx, - clang::QualType qual_type) { - return matcher_->Find(ctx, qual_type); +Matcher::Match Search(clang::ASTContext &ctx, + clang::QualType qual_type) { + return Matcher::Find(ctx, qual_type); } -Matching::Matcher &GetMatcher() { return *matcher_; } - Model CurrentModel() { return model_; } void AddRuleForUserDefinedType(clang::ASTContext &ctx, clang::NamedDecl *decl) { @@ -146,7 +140,6 @@ void Load(Model model, const std::string &rules_dir) { } translation_rules_loaded_ = true; - matcher_ = std::make_unique(); addRulesFromDirectory(rules_dir, model); #if 0 diff --git a/cpp2rust/converter/rules/registry.h b/cpp2rust/converter/rules/registry.h index 8b1a1e698..a4a084a2a 100644 --- a/cpp2rust/converter/rules/registry.h +++ b/cpp2rust/converter/rules/registry.h @@ -13,7 +13,7 @@ #include #include "converter/factory.h" -#include "converter/rules/matching.h" +#include "converter/rules/matcher.h" #include "converter/translation_rule.h" namespace cpp2rust::RuleRegistry { @@ -27,11 +27,10 @@ ExprCandidates(const std::string &key); std::ranges::subrange TypeCandidates(const std::string &key); -Matching::Match Search(clang::ASTContext &ctx, - const clang::Expr *expr); -Matching::Match Search(clang::ASTContext &ctx, - clang::QualType qual_type); -Matching::Matcher &GetMatcher(); +Matcher::Match Search(clang::ASTContext &ctx, + const clang::Expr *expr); +Matcher::Match Search(clang::ASTContext &ctx, + clang::QualType qual_type); Model CurrentModel(); diff --git a/cpp2rust/converter/rules/string_matcher.cpp b/cpp2rust/converter/rules/string_matcher.cpp index 4daab7e39..57e552de6 100644 --- a/cpp2rust/converter/rules/string_matcher.cpp +++ b/cpp2rust/converter/rules/string_matcher.cpp @@ -1,8 +1,6 @@ // Copyright (c) 2022-present INESC-ID. // Distributed under the MIT license that can be found in the LICENSE file. -#include "converter/rules/string_matcher.h" - #include #include #include @@ -10,9 +8,10 @@ #include "converter/converter_lib.h" #include "converter/printer.h" +#include "converter/rules/matcher.h" #include "converter/rules/registry.h" -namespace cpp2rust::Matching { +namespace cpp2rust::Matcher { namespace { @@ -344,16 +343,16 @@ std::string mapTypeString(const std::string &cpp_type) { } // namespace -std::string StringMatcher::Key(const TranslationRule::ExprRule &rule) const { +std::string Key(const TranslationRule::ExprRule &rule) { return exprKey(rule.src); } -std::string StringMatcher::Key(const TranslationRule::TypeRule &rule) const { +std::string Key(const TranslationRule::TypeRule &rule) { return typeKey(rule.src); } -Match StringMatcher::Find(clang::ASTContext &ctx, - const clang::Expr *expr) { +Match Find(clang::ASTContext &ctx, + const clang::Expr *expr) { auto qualified_name = Printer::ToString(ctx, expr); auto res = searchExpr(qualified_name); log() << "search expr " << qualified_name << ", result:\n"; @@ -365,8 +364,8 @@ Match StringMatcher::Find(clang::ASTContext &ctx, return res; } -Match StringMatcher::Find(clang::ASTContext &ctx, - clang::QualType type) { +Match Find(clang::ASTContext &ctx, + clang::QualType type) { auto sugared = Printer::ToString(ctx, type, Printer::ScalarSugar::kPreserve); if (auto res = searchType(sugared); res.first) { log() << "search type " << sugared @@ -385,15 +384,14 @@ Match StringMatcher::Find(clang::ASTContext &ctx, return res; } -bool StringMatcher::HasRuleNamed(clang::ASTContext &ctx, - const clang::FunctionDecl *decl) { +bool HasRuleNamed(clang::ASTContext &ctx, const clang::FunctionDecl *decl) { return !RuleRegistry::ExprCandidates(exprKey(Printer::ToString(ctx, decl))) .empty(); } -std::string StringMatcher::MapBinding(clang::ASTContext &ctx, - const Bindings &bindings, unsigned n) { +std::string MapBinding(clang::ASTContext &ctx, const Bindings &bindings, + unsigned n) { return mapTypeString(bindings.at(n).value()); } -} // namespace cpp2rust::Matching +} // namespace cpp2rust::Matcher diff --git a/cpp2rust/converter/rules/string_matcher.h b/cpp2rust/converter/rules/string_matcher.h deleted file mode 100644 index 34dfe9dc4..000000000 --- a/cpp2rust/converter/rules/string_matcher.h +++ /dev/null @@ -1,22 +0,0 @@ -#pragma once - -// Copyright (c) 2022-present INESC-ID. -// Distributed under the MIT license that can be found in the LICENSE file. - -#include "converter/rules/matching.h" - -namespace cpp2rust::Matching { -class StringMatcher final : public Matcher { -public: - std::string Key(const TranslationRule::ExprRule &rule) const override; - std::string Key(const TranslationRule::TypeRule &rule) const override; - Match Find(clang::ASTContext &ctx, - const clang::Expr *expr) override; - Match Find(clang::ASTContext &ctx, - clang::QualType type) override; - bool HasRuleNamed(clang::ASTContext &ctx, - const clang::FunctionDecl *decl) override; - std::string MapBinding(clang::ASTContext &ctx, const Bindings &bindings, - unsigned n) override; -}; -} // namespace cpp2rust::Matching From 4b8c9a01d395426cb90f5afdfcc432889a48faf7 Mon Sep 17 00:00:00 2001 From: Lucian Popescu Date: Tue, 29 Sep 2026 14:51:36 +0100 Subject: [PATCH 6/7] Move slim shims in RuleRegistry --- cpp2rust/converter/converter.cpp | 35 ++++---- cpp2rust/converter/converter_lib.cpp | 3 +- cpp2rust/converter/mapper.cpp | 81 +------------------ cpp2rust/converter/mapper.h | 12 --- .../converter/models/converter_refcount.cpp | 6 +- cpp2rust/converter/rules/matcher.cpp | 10 +++ cpp2rust/converter/rules/matcher.h | 1 + cpp2rust/converter/rules/registry.cpp | 60 ++++++++++++++ cpp2rust/converter/rules/registry.h | 14 ++++ 9 files changed, 114 insertions(+), 108 deletions(-) diff --git a/cpp2rust/converter/converter.cpp b/cpp2rust/converter/converter.cpp index 3e6f7221e..20837adc4 100644 --- a/cpp2rust/converter/converter.cpp +++ b/cpp2rust/converter/converter.cpp @@ -673,14 +673,15 @@ bool Converter::RecordDerivesDefault(const clang::RecordDecl *decl) { } bool Converter::IsPassThroughRule(clang::Expr *expr) const { - const auto *rule = Mapper::GetExprRule(ctx_, GetCalleeOrExpr(expr)); + const auto *rule = RuleRegistry::GetExprRule(ctx_, GetCalleeOrExpr(expr)); return rule && rule->body.size() == 1 && std::holds_alternative( rule->body[0]); } bool Converter::RecordDerivesCopy(const clang::RecordDecl *decl) const { - auto *derives = Mapper::MappedDerives(ctx_, ctx_.getCanonicalTagType(decl)); + auto *derives = + RuleRegistry::MappedDerives(ctx_, ctx_.getCanonicalTagType(decl)); return derives && std::find(derives->begin(), derives->end(), "Copy") != derives->end(); } @@ -803,8 +804,9 @@ void Converter::EmitRustStructOrUnion(clang::RecordDecl *decl) { EmitReprC(decl); } auto attrs = GetStructAttributes(decl); - Mapper::SetDerives(ctx_, ctx_.getCanonicalTagType(decl), - std::vector(attrs.begin(), attrs.end())); + RuleRegistry::SetDerives( + ctx_, ctx_.getCanonicalTagType(decl), + std::vector(attrs.begin(), attrs.end())); StrCat("#[derive("); for (auto *attr : attrs) { StrCat(attr, ','); @@ -902,8 +904,9 @@ void Converter::EmitReprC(clang::RecordDecl *decl) { void Converter::EmitRustUnion(clang::RecordDecl *decl) { EmitReprC(decl); auto attrs = GetStructAttributes(decl); - Mapper::SetDerives(ctx_, ctx_.getCanonicalTagType(decl), - std::vector(attrs.begin(), attrs.end())); + RuleRegistry::SetDerives( + ctx_, ctx_.getCanonicalTagType(decl), + std::vector(attrs.begin(), attrs.end())); StrCat("#[derive("); for (auto *attr : attrs) { StrCat(attr, ','); @@ -1797,7 +1800,7 @@ bool Converter::VisitCallExpr(clang::CallExpr *expr) { } if (Mapper::Contains(ctx_, expr->getCallee())) { - if (Mapper::IsLibcPassthrough(ctx_, GetCalleeOrExpr(expr))) { + if (RuleRegistry::IsLibcPassthrough(ctx_, GetCalleeOrExpr(expr))) { ConvertGenericCallExpr(expr); return false; } @@ -1937,7 +1940,7 @@ Converter::CallInfo Converter::CollectCallInfo(clang::CallExpr *expr) { info.is_variadic = function ? function->isVariadic() : proto->isVariadic(); info.is_fn_ptr_call = !function; info.is_libc_passthrough = - Mapper::IsLibcPassthrough(ctx_, GetCalleeOrExpr(expr)); + RuleRegistry::IsLibcPassthrough(ctx_, GetCalleeOrExpr(expr)); for (unsigned i = 0; i < num_named_params && i < num_args; ++i) { auto *arg = expr->getArg(i + arg_begin); @@ -2421,7 +2424,7 @@ void Converter::ConvertIntegralToBooleanCast(clang::ImplicitCastExpr *expr) { bool Converter::IsCastRedundantInRust(clang::Expr *expr, clang::QualType target_type) { auto target = GetUnsafeTypeAsString(target_type); - if (const auto *rule = Mapper::GetExprRule(ctx_, expr)) { + if (const auto *rule = RuleRegistry::GetExprRule(ctx_, expr)) { return rule->return_type.type == target; } return GetUnsafeTypeAsString(expr->getType()) == target; @@ -3278,7 +3281,7 @@ replaceNonUniformLibcField(clang::MemberExpr *expr) { void Converter::ConvertMemberExpr(clang::MemberExpr *expr) { if (auto mapped = GetMappedAsString(expr); !mapped.empty()) { - if (Mapper::ReturnsPointer(ctx_, expr)) { + if (RuleRegistry::ReturnsPointer(ctx_, expr)) { StrCat(token::kStar, mapped); } else { StrCat(mapped); @@ -4950,7 +4953,7 @@ std::string Converter::ConvertMappedMethodCall( std::string Converter::GetMappedAsString(clang::Expr *expr, clang::Expr **args, unsigned num_args, TempMaterializationCtx *ctx) { - auto *tgt_ir = Mapper::GetExprRule(ctx_, GetCalleeOrExpr(expr)); + auto *tgt_ir = RuleRegistry::GetExprRule(ctx_, GetCalleeOrExpr(expr)); if (!tgt_ir) return {}; @@ -4990,9 +4993,9 @@ std::string Converter::ConvertIRFragment( .access = ph->access, .is_receiver = is_receiver, .is_cpp_ptr = arg->getType()->isPointerType(), - .maps_to_rust_ptr = Mapper::MapsToPointer(ctx_, arg->getType()), - .declared_in_rule_as_rust_ptr = - Mapper::ParamIsPointer(ctx_, GetCalleeOrExpr(expr), arg_idx), + .maps_to_rust_ptr = RuleRegistry::MapsToPointer(ctx_, arg->getType()), + .declared_in_rule_as_rust_ptr = RuleRegistry::ParamIsPointer( + ctx_, GetCalleeOrExpr(expr), arg_idx), .is_index_base = ph->is_index_base, }; result += ConvertPlaceholder(expr, arg, ph_ctx); @@ -5012,7 +5015,7 @@ std::string Converter::ConvertIRFragment( std::string Converter::ConvertVariadicTail(clang::Expr *expr, const std::vector &all_args) { - const auto *tgt_ir = Mapper::GetExprRule(ctx_, GetCalleeOrExpr(expr)); + const auto *tgt_ir = RuleRegistry::GetExprRule(ctx_, GetCalleeOrExpr(expr)); unsigned fixed = tgt_ir ? tgt_ir->params.size() : 0; Buffer buf(*this); @@ -5031,7 +5034,7 @@ Converter::ConvertVariadicTail(clang::Expr *expr, std::string Converter::ConvertInitFragment(clang::Expr *expr, const std::vector &all_args) { - const auto *tgt_ir = Mapper::GetExprRule(ctx_, GetCalleeOrExpr(expr)); + const auto *tgt_ir = RuleRegistry::GetExprRule(ctx_, GetCalleeOrExpr(expr)); assert(tgt_ir && tgt_ir->init_type.valid()); auto *callee = clang::cast(expr)->getDirectCallee(); assert(callee); diff --git a/cpp2rust/converter/converter_lib.cpp b/cpp2rust/converter/converter_lib.cpp index c413a0cd9..9e3182bff 100644 --- a/cpp2rust/converter/converter_lib.cpp +++ b/cpp2rust/converter/converter_lib.cpp @@ -27,6 +27,7 @@ #include "converter/lex.h" #include "converter/mapper.h" #include "converter/printer.h" +#include "converter/rules/registry.h" // https://doc.rust-lang.org/reference/keywords.html static const char rust_keywords[][12] = { @@ -1480,7 +1481,7 @@ GetStrongestIteratorCategory(clang::ASTContext &ctx, clang::QualType type) { if (!Mapper::Contains(ctx, type)) { return std::nullopt; } - if (Mapper::MapsToRefcountPointer(ctx, type)) { + if (RuleRegistry::MapsToRefcountPointer(ctx, type)) { return IteratorCategory::Contiguous; } auto mapped = Mapper::Map(ctx, type); diff --git a/cpp2rust/converter/mapper.cpp b/cpp2rust/converter/mapper.cpp index 85a3a36e2..6067aeff3 100644 --- a/cpp2rust/converter/mapper.cpp +++ b/cpp2rust/converter/mapper.cpp @@ -4,7 +4,6 @@ #include "converter/mapper.h" #include -#include #include #include @@ -19,21 +18,6 @@ namespace cpp2rust::Mapper { -namespace { - -Matcher::Bindings mapBindings(clang::ASTContext &ctx, - const Matcher::Bindings &bindings) { - Matcher::Bindings mapped(bindings.size()); - for (unsigned i = 0; i < bindings.size(); ++i) { - if (bindings[i]) { - mapped[i] = Matcher::MapBinding(ctx, bindings, i); - } - } - return mapped; -} - -} // namespace - bool Contains(clang::ASTContext &ctx, clang::QualType qual_type) { return RuleRegistry::Search(ctx, qual_type).first != nullptr; } @@ -42,24 +26,6 @@ bool Contains(clang::ASTContext &ctx, const clang::Expr *expr) { return RuleRegistry::Search(ctx, expr).first != nullptr; } -const TranslationRule::ExprRule *GetExprRule(clang::ASTContext &ctx, - const clang::Expr *expr) { - return RuleRegistry::Search(ctx, expr).first; -} - -bool IsLibcPassthrough(clang::ASTContext &ctx, const clang::Expr *expr) { - const auto *tgt_ir = GetExprRule(ctx, expr); - if (tgt_ir == nullptr || !tgt_ir->body.empty() || !tgt_ir->is_extern) { - return false; - } - const auto *ref = - clang::dyn_cast(expr->IgnoreParenImpCasts()); - const auto *decl = ref != nullptr ? ref->getDecl() : nullptr; - return decl != nullptr && - decl->getASTContext().getSourceManager().isInSystemHeader( - decl->getLocation()); -} - std::string MapFunctionName(clang::ASTContext &ctx, const clang::FunctionDecl *decl) { assert(decl); @@ -89,7 +55,7 @@ std::string InstantiateTemplate(clang::ASTContext &ctx, const clang::Expr *expr, std::string Map(clang::ASTContext &ctx, clang::QualType qual_type) { auto [rule, subs] = RuleRegistry::Search(ctx, qual_type); if (rule) { - return Matcher::InstantiateTgt(mapBindings(ctx, subs), + return Matcher::InstantiateTgt(Matcher::MapBindings(ctx, subs), rule->type_info.type); } return {}; @@ -98,56 +64,17 @@ std::string Map(clang::ASTContext &ctx, clang::QualType qual_type) { std::string MapInitializer(clang::ASTContext &ctx, clang::QualType qual_type) { auto [rule, subs] = RuleRegistry::Search(ctx, qual_type); if (rule && !rule->initializer.empty()) { - return Matcher::InstantiateTgt(mapBindings(ctx, subs), rule->initializer); + return Matcher::InstantiateTgt(Matcher::MapBindings(ctx, subs), + rule->initializer); } return {}; } -bool MapsToPointer(clang::ASTContext &ctx, clang::QualType qual_type) { - auto rule = RuleRegistry::Search(ctx, qual_type).first; - return rule && rule->type_info.is_pointer(); -} - -bool MapsToRefcountPointer(clang::ASTContext &ctx, clang::QualType qual_type) { - auto rule = RuleRegistry::Search(ctx, qual_type).first; - return rule && rule->type_info.is_refcount_pointer; -} - -const std::vector *MappedDerives(clang::ASTContext &ctx, - clang::QualType qual_type) { - auto rule = RuleRegistry::Search(ctx, qual_type).first; - return rule ? &rule->type_info.derives : nullptr; -} - -void SetDerives(clang::ASTContext &ctx, clang::QualType qual_type, - std::vector derives) { - if (auto *rule = RuleRegistry::Search(ctx, qual_type).first) { - rule->type_info.derives = std::move(derives); - } -} - -bool ReturnsPointer(clang::ASTContext &ctx, const clang::Expr *expr) { - auto rule = RuleRegistry::Search(ctx, expr).first; - return rule && rule->return_type.is_pointer(); -} - -const TranslationRule::TypeInfo & -GetParamInfo(clang::ASTContext &ctx, const clang::Expr *expr, unsigned index) { - auto rule = RuleRegistry::Search(ctx, expr).first; - assert(rule && "expression must have a translation rule"); - return rule->params.at(index); -} - std::string GetParamType(clang::ASTContext &ctx, const clang::Expr *expr, unsigned index) { auto [rule, subs] = RuleRegistry::Search(ctx, expr); - return Matcher::InstantiateTgt(mapBindings(ctx, subs), + return Matcher::InstantiateTgt(Matcher::MapBindings(ctx, subs), rule->params.at(index).type); } -bool ParamIsPointer(clang::ASTContext &ctx, const clang::Expr *expr, - unsigned index) { - return GetParamInfo(ctx, expr, index).is_pointer(); -} - } // namespace cpp2rust::Mapper diff --git a/cpp2rust/converter/mapper.h b/cpp2rust/converter/mapper.h index 6aa0f4af5..8d4838492 100644 --- a/cpp2rust/converter/mapper.h +++ b/cpp2rust/converter/mapper.h @@ -17,22 +17,10 @@ bool Contains(clang::ASTContext &ctx, const clang::Expr *expr); std::string Map(clang::ASTContext &ctx, clang::QualType qual_type); std::string MapInitializer(clang::ASTContext &ctx, clang::QualType qual_type); -const TranslationRule::ExprRule *GetExprRule(clang::ASTContext &ctx, - const clang::Expr *expr); -bool IsLibcPassthrough(clang::ASTContext &ctx, const clang::Expr *expr); std::string MapFunctionName(clang::ASTContext &ctx, const clang::FunctionDecl *decl); std::string InstantiateTemplate(clang::ASTContext &ctx, const clang::Expr *expr, unsigned n); -bool ReturnsPointer(clang::ASTContext &ctx, const clang::Expr *expr); std::string GetParamType(clang::ASTContext &ctx, const clang::Expr *expr, unsigned index); -bool ParamIsPointer(clang::ASTContext &ctx, const clang::Expr *expr, - unsigned index); -bool MapsToPointer(clang::ASTContext &ctx, clang::QualType qual_type); -bool MapsToRefcountPointer(clang::ASTContext &ctx, clang::QualType qual_type); -const std::vector *MappedDerives(clang::ASTContext &ctx, - clang::QualType qual_type); -void SetDerives(clang::ASTContext &ctx, clang::QualType qual_type, - std::vector derives); } // namespace cpp2rust::Mapper diff --git a/cpp2rust/converter/models/converter_refcount.cpp b/cpp2rust/converter/models/converter_refcount.cpp index 84c8131ca..853ad1590 100644 --- a/cpp2rust/converter/models/converter_refcount.cpp +++ b/cpp2rust/converter/models/converter_refcount.cpp @@ -16,6 +16,7 @@ #include "converter/lex.h" #include "converter/mapper.h" #include "converter/printer.h" +#include "converter/rules/registry.h" namespace cpp2rust { std::map @@ -551,8 +552,9 @@ void ConverterRefCount::EmitRustUnion(clang::RecordDecl *decl) { auto name = GetRecordName(decl); auto attrs = GetStructAttributes(decl); - Mapper::SetDerives(ctx_, ctx_.getCanonicalTagType(decl), - std::vector(attrs.begin(), attrs.end())); + RuleRegistry::SetDerives( + ctx_, ctx_.getCanonicalTagType(decl), + std::vector(attrs.begin(), attrs.end())); StrCat(std::format("pub struct {} {{ __bytes: Value> }}", name)); diff --git a/cpp2rust/converter/rules/matcher.cpp b/cpp2rust/converter/rules/matcher.cpp index a00cad0de..a6ce29cc8 100644 --- a/cpp2rust/converter/rules/matcher.cpp +++ b/cpp2rust/converter/rules/matcher.cpp @@ -8,6 +8,16 @@ namespace cpp2rust::Matcher { +Bindings MapBindings(clang::ASTContext &ctx, const Bindings &bindings) { + Bindings mapped(bindings.size()); + for (unsigned i = 0; i < bindings.size(); ++i) { + if (bindings[i]) { + mapped[i] = MapBinding(ctx, bindings, i); + } + } + return mapped; +} + // Substitutes concrete types into a target template string using the provided // type mapping. Each template parameter in `tgt_template` is replaced with its // corresponding instantiated type from `types`. diff --git a/cpp2rust/converter/rules/matcher.h b/cpp2rust/converter/rules/matcher.h index 5676ee9be..d3f1448e0 100644 --- a/cpp2rust/converter/rules/matcher.h +++ b/cpp2rust/converter/rules/matcher.h @@ -32,6 +32,7 @@ bool HasRuleNamed(clang::ASTContext &ctx, const clang::FunctionDecl *decl); std::string MapBinding(clang::ASTContext &ctx, const Bindings &bindings, unsigned n); +Bindings MapBindings(clang::ASTContext &ctx, const Bindings &bindings); std::string InstantiateTgt(const Bindings &types, const std::string &tgt_template); diff --git a/cpp2rust/converter/rules/registry.cpp b/cpp2rust/converter/rules/registry.cpp index bf5b3a249..c0d048b42 100644 --- a/cpp2rust/converter/rules/registry.cpp +++ b/cpp2rust/converter/rules/registry.cpp @@ -3,6 +3,8 @@ #include "converter/rules/registry.h" +#include + #include #include #include @@ -59,6 +61,13 @@ void addRulesFromDirectory(const std::filesystem::path &dir, Model model) { } } +const TranslationRule::TypeInfo & +GetParamInfo(clang::ASTContext &ctx, const clang::Expr *expr, unsigned index) { + auto rule = Search(ctx, expr).first; + assert(rule && "expression must have a translation rule"); + return rule->params.at(index); +} + } // namespace std::ranges::subrange @@ -86,6 +95,57 @@ Matcher::Match Search(clang::ASTContext &ctx, return Matcher::Find(ctx, qual_type); } +const TranslationRule::ExprRule *GetExprRule(clang::ASTContext &ctx, + const clang::Expr *expr) { + return Search(ctx, expr).first; +} + +bool MapsToPointer(clang::ASTContext &ctx, clang::QualType qual_type) { + auto rule = Search(ctx, qual_type).first; + return rule && rule->type_info.is_pointer(); +} + +bool MapsToRefcountPointer(clang::ASTContext &ctx, clang::QualType qual_type) { + auto rule = Search(ctx, qual_type).first; + return rule && rule->type_info.is_refcount_pointer; +} + +const std::vector *MappedDerives(clang::ASTContext &ctx, + clang::QualType qual_type) { + auto rule = Search(ctx, qual_type).first; + return rule ? &rule->type_info.derives : nullptr; +} + +void SetDerives(clang::ASTContext &ctx, clang::QualType qual_type, + std::vector derives) { + if (auto *rule = Search(ctx, qual_type).first) { + rule->type_info.derives = std::move(derives); + } +} + +bool ReturnsPointer(clang::ASTContext &ctx, const clang::Expr *expr) { + auto rule = Search(ctx, expr).first; + return rule && rule->return_type.is_pointer(); +} + +bool ParamIsPointer(clang::ASTContext &ctx, const clang::Expr *expr, + unsigned index) { + return GetParamInfo(ctx, expr, index).is_pointer(); +} + +bool IsLibcPassthrough(clang::ASTContext &ctx, const clang::Expr *expr) { + const auto *tgt_ir = GetExprRule(ctx, expr); + if (tgt_ir == nullptr || !tgt_ir->body.empty() || !tgt_ir->is_extern) { + return false; + } + const auto *ref = + clang::dyn_cast(expr->IgnoreParenImpCasts()); + const auto *decl = ref != nullptr ? ref->getDecl() : nullptr; + return decl != nullptr && + decl->getASTContext().getSourceManager().isInSystemHeader( + decl->getLocation()); +} + Model CurrentModel() { return model_; } void AddRuleForUserDefinedType(clang::ASTContext &ctx, clang::NamedDecl *decl) { diff --git a/cpp2rust/converter/rules/registry.h b/cpp2rust/converter/rules/registry.h index a4a084a2a..aaa65d03b 100644 --- a/cpp2rust/converter/rules/registry.h +++ b/cpp2rust/converter/rules/registry.h @@ -11,6 +11,7 @@ #include #include #include +#include #include "converter/factory.h" #include "converter/rules/matcher.h" @@ -32,6 +33,19 @@ Matcher::Match Search(clang::ASTContext &ctx, Matcher::Match Search(clang::ASTContext &ctx, clang::QualType qual_type); +const TranslationRule::ExprRule *GetExprRule(clang::ASTContext &ctx, + const clang::Expr *expr); +bool IsLibcPassthrough(clang::ASTContext &ctx, const clang::Expr *expr); +bool ReturnsPointer(clang::ASTContext &ctx, const clang::Expr *expr); +bool ParamIsPointer(clang::ASTContext &ctx, const clang::Expr *expr, + unsigned index); +bool MapsToPointer(clang::ASTContext &ctx, clang::QualType qual_type); +bool MapsToRefcountPointer(clang::ASTContext &ctx, clang::QualType qual_type); +const std::vector *MappedDerives(clang::ASTContext &ctx, + clang::QualType qual_type); +void SetDerives(clang::ASTContext &ctx, clang::QualType qual_type, + std::vector derives); + Model CurrentModel(); void Load(Model model, const std::string &rules_dir); From 4026a0e1c08d548b1b8b8f12ef786bf1d0c07dee Mon Sep 17 00:00:00 2001 From: Lucian Popescu Date: Tue, 29 Sep 2026 15:14:36 +0100 Subject: [PATCH 7/7] Remove ctx from MapBinding --- cpp2rust/converter/mapper.cpp | 8 ++++---- cpp2rust/converter/rules/matcher.cpp | 4 ++-- cpp2rust/converter/rules/matcher.h | 5 ++--- cpp2rust/converter/rules/string_matcher.cpp | 3 +-- 4 files changed, 9 insertions(+), 11 deletions(-) diff --git a/cpp2rust/converter/mapper.cpp b/cpp2rust/converter/mapper.cpp index 6067aeff3..1fe6c4c0d 100644 --- a/cpp2rust/converter/mapper.cpp +++ b/cpp2rust/converter/mapper.cpp @@ -47,7 +47,7 @@ std::string InstantiateTemplate(clang::ASTContext &ctx, const clang::Expr *expr, } auto &ty = subs.at(n - 1); if (ty) { - ty = Matcher::MapBinding(ctx, subs, n - 1); + ty = Matcher::MapBinding(subs, n - 1); } return Matcher::InstantiateTgt(subs, text); } @@ -55,7 +55,7 @@ std::string InstantiateTemplate(clang::ASTContext &ctx, const clang::Expr *expr, std::string Map(clang::ASTContext &ctx, clang::QualType qual_type) { auto [rule, subs] = RuleRegistry::Search(ctx, qual_type); if (rule) { - return Matcher::InstantiateTgt(Matcher::MapBindings(ctx, subs), + return Matcher::InstantiateTgt(Matcher::MapBindings(subs), rule->type_info.type); } return {}; @@ -64,7 +64,7 @@ std::string Map(clang::ASTContext &ctx, clang::QualType qual_type) { std::string MapInitializer(clang::ASTContext &ctx, clang::QualType qual_type) { auto [rule, subs] = RuleRegistry::Search(ctx, qual_type); if (rule && !rule->initializer.empty()) { - return Matcher::InstantiateTgt(Matcher::MapBindings(ctx, subs), + return Matcher::InstantiateTgt(Matcher::MapBindings(subs), rule->initializer); } return {}; @@ -73,7 +73,7 @@ std::string MapInitializer(clang::ASTContext &ctx, clang::QualType qual_type) { std::string GetParamType(clang::ASTContext &ctx, const clang::Expr *expr, unsigned index) { auto [rule, subs] = RuleRegistry::Search(ctx, expr); - return Matcher::InstantiateTgt(Matcher::MapBindings(ctx, subs), + return Matcher::InstantiateTgt(Matcher::MapBindings(subs), rule->params.at(index).type); } diff --git a/cpp2rust/converter/rules/matcher.cpp b/cpp2rust/converter/rules/matcher.cpp index a6ce29cc8..1dd53121f 100644 --- a/cpp2rust/converter/rules/matcher.cpp +++ b/cpp2rust/converter/rules/matcher.cpp @@ -8,11 +8,11 @@ namespace cpp2rust::Matcher { -Bindings MapBindings(clang::ASTContext &ctx, const Bindings &bindings) { +Bindings MapBindings(const Bindings &bindings) { Bindings mapped(bindings.size()); for (unsigned i = 0; i < bindings.size(); ++i) { if (bindings[i]) { - mapped[i] = MapBinding(ctx, bindings, i); + mapped[i] = MapBinding(bindings, i); } } return mapped; diff --git a/cpp2rust/converter/rules/matcher.h b/cpp2rust/converter/rules/matcher.h index d3f1448e0..f5f9e0db4 100644 --- a/cpp2rust/converter/rules/matcher.h +++ b/cpp2rust/converter/rules/matcher.h @@ -30,9 +30,8 @@ Match Find(clang::ASTContext &ctx, bool HasRuleNamed(clang::ASTContext &ctx, const clang::FunctionDecl *decl); -std::string MapBinding(clang::ASTContext &ctx, const Bindings &bindings, - unsigned n); -Bindings MapBindings(clang::ASTContext &ctx, const Bindings &bindings); +std::string MapBinding(const Bindings &bindings, unsigned n); +Bindings MapBindings(const Bindings &bindings); std::string InstantiateTgt(const Bindings &types, const std::string &tgt_template); diff --git a/cpp2rust/converter/rules/string_matcher.cpp b/cpp2rust/converter/rules/string_matcher.cpp index 57e552de6..66cc3b980 100644 --- a/cpp2rust/converter/rules/string_matcher.cpp +++ b/cpp2rust/converter/rules/string_matcher.cpp @@ -389,8 +389,7 @@ bool HasRuleNamed(clang::ASTContext &ctx, const clang::FunctionDecl *decl) { .empty(); } -std::string MapBinding(clang::ASTContext &ctx, const Bindings &bindings, - unsigned n) { +std::string MapBinding(const Bindings &bindings, unsigned n) { return mapTypeString(bindings.at(n).value()); }