574 lines · cpp
1//===- Nodes.cpp ----------------------------------------------------------===//2//3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.4// See https://llvm.org/LICENSE.txt for license information.5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception6//7//===----------------------------------------------------------------------===//8 9#include "mlir/Tools/PDLL/AST/Nodes.h"10#include "mlir/Tools/PDLL/AST/Context.h"11#include "llvm/ADT/SmallPtrSet.h"12#include "llvm/ADT/TypeSwitch.h"13#include <optional>14 15using namespace mlir;16using namespace mlir::pdll::ast;17 18/// Copy a string reference into the context with a null terminator.19static StringRef copyStringWithNull(Context &ctx, StringRef str) {20 if (str.empty())21 return str;22 23 char *data = ctx.getAllocator().Allocate<char>(str.size() + 1);24 llvm::copy(str, data);25 data[str.size()] = 0;26 return StringRef(data, str.size());27}28 29//===----------------------------------------------------------------------===//30// Name31//===----------------------------------------------------------------------===//32 33const Name &Name::create(Context &ctx, StringRef name, SMRange location) {34 return *new (ctx.getAllocator().Allocate<Name>())35 Name(copyStringWithNull(ctx, name), location);36}37 38//===----------------------------------------------------------------------===//39// Node40//===----------------------------------------------------------------------===//41 42namespace {43class NodeVisitor {44public:45 explicit NodeVisitor(function_ref<void(const Node *)> visitFn)46 : visitFn(visitFn) {}47 48 void visit(const Node *node) {49 if (!node || !alreadyVisited.insert(node).second)50 return;51 52 visitFn(node);53 TypeSwitch<const Node *>(node)54 .Case<55 // Statements.56 const CompoundStmt, const EraseStmt, const LetStmt,57 const ReplaceStmt, const ReturnStmt, const RewriteStmt,58 59 // Expressions.60 const AttributeExpr, const CallExpr, const DeclRefExpr,61 const MemberAccessExpr, const OperationExpr, const RangeExpr,62 const TupleExpr, const TypeExpr,63 64 // Core Constraint Decls.65 const AttrConstraintDecl, const OpConstraintDecl,66 const TypeConstraintDecl, const TypeRangeConstraintDecl,67 const ValueConstraintDecl, const ValueRangeConstraintDecl,68 69 // Decls.70 const NamedAttributeDecl, const OpNameDecl, const PatternDecl,71 const UserConstraintDecl, const UserRewriteDecl, const VariableDecl,72 73 const Module>(74 [&](auto derivedNode) { this->visitImpl(derivedNode); })75 .DefaultUnreachable("unknown AST node");76 }77 78private:79 void visitImpl(const CompoundStmt *stmt) {80 for (const Node *child : stmt->getChildren())81 visit(child);82 }83 void visitImpl(const EraseStmt *stmt) { visit(stmt->getRootOpExpr()); }84 void visitImpl(const LetStmt *stmt) { visit(stmt->getVarDecl()); }85 void visitImpl(const ReplaceStmt *stmt) {86 visit(stmt->getRootOpExpr());87 for (const Node *child : stmt->getReplExprs())88 visit(child);89 }90 void visitImpl(const ReturnStmt *stmt) { visit(stmt->getResultExpr()); }91 void visitImpl(const RewriteStmt *stmt) {92 visit(stmt->getRootOpExpr());93 visit(stmt->getRewriteBody());94 }95 96 void visitImpl(const AttributeExpr *expr) {}97 void visitImpl(const CallExpr *expr) {98 visit(expr->getCallableExpr());99 for (const Node *child : expr->getArguments())100 visit(child);101 }102 void visitImpl(const DeclRefExpr *expr) { visit(expr->getDecl()); }103 void visitImpl(const MemberAccessExpr *expr) { visit(expr->getParentExpr()); }104 void visitImpl(const OperationExpr *expr) {105 visit(expr->getNameDecl());106 for (const Node *child : expr->getOperands())107 visit(child);108 for (const Node *child : expr->getResultTypes())109 visit(child);110 for (const Node *child : expr->getAttributes())111 visit(child);112 }113 void visitImpl(const RangeExpr *expr) {114 for (const Node *child : expr->getElements())115 visit(child);116 }117 void visitImpl(const TupleExpr *expr) {118 for (const Node *child : expr->getElements())119 visit(child);120 }121 void visitImpl(const TypeExpr *expr) {}122 123 void visitImpl(const AttrConstraintDecl *decl) { visit(decl->getTypeExpr()); }124 void visitImpl(const OpConstraintDecl *decl) { visit(decl->getNameDecl()); }125 void visitImpl(const TypeConstraintDecl *decl) {}126 void visitImpl(const TypeRangeConstraintDecl *decl) {}127 void visitImpl(const ValueConstraintDecl *decl) {128 visit(decl->getTypeExpr());129 }130 void visitImpl(const ValueRangeConstraintDecl *decl) {131 visit(decl->getTypeExpr());132 }133 134 void visitImpl(const NamedAttributeDecl *decl) { visit(decl->getValue()); }135 void visitImpl(const OpNameDecl *decl) {}136 void visitImpl(const PatternDecl *decl) { visit(decl->getBody()); }137 void visitImpl(const UserConstraintDecl *decl) {138 for (const Node *child : decl->getInputs())139 visit(child);140 for (const Node *child : decl->getResults())141 visit(child);142 visit(decl->getBody());143 }144 void visitImpl(const UserRewriteDecl *decl) {145 for (const Node *child : decl->getInputs())146 visit(child);147 for (const Node *child : decl->getResults())148 visit(child);149 visit(decl->getBody());150 }151 void visitImpl(const VariableDecl *decl) {152 visit(decl->getInitExpr());153 for (const ConstraintRef &child : decl->getConstraints())154 visit(child.constraint);155 }156 157 void visitImpl(const Module *module) {158 for (const Node *child : module->getChildren())159 visit(child);160 }161 162 function_ref<void(const Node *)> visitFn;163 SmallPtrSet<const Node *, 16> alreadyVisited;164};165} // namespace166 167void Node::walk(function_ref<void(const Node *)> walkFn) const {168 return NodeVisitor(walkFn).visit(this);169}170 171//===----------------------------------------------------------------------===//172// DeclScope173//===----------------------------------------------------------------------===//174 175void DeclScope::add(Decl *decl) {176 const Name *name = decl->getName();177 assert(name && "expected a named decl");178 assert(!decls.count(name->getName()) && "decl with this name already exists");179 decls.try_emplace(name->getName(), decl);180}181 182Decl *DeclScope::lookup(StringRef name) {183 if (Decl *decl = decls.lookup(name))184 return decl;185 return parent ? parent->lookup(name) : nullptr;186}187 188//===----------------------------------------------------------------------===//189// CompoundStmt190//===----------------------------------------------------------------------===//191 192CompoundStmt *CompoundStmt::create(Context &ctx, SMRange loc,193 ArrayRef<Stmt *> children) {194 unsigned allocSize = CompoundStmt::totalSizeToAlloc<Stmt *>(children.size());195 void *rawData = ctx.getAllocator().Allocate(allocSize, alignof(CompoundStmt));196 197 CompoundStmt *stmt = new (rawData) CompoundStmt(loc, children.size());198 llvm::uninitialized_copy(children, stmt->getChildren().begin());199 return stmt;200}201 202//===----------------------------------------------------------------------===//203// LetStmt204//===----------------------------------------------------------------------===//205 206LetStmt *LetStmt::create(Context &ctx, SMRange loc, VariableDecl *varDecl) {207 return new (ctx.getAllocator().Allocate<LetStmt>()) LetStmt(loc, varDecl);208}209 210//===----------------------------------------------------------------------===//211// OpRewriteStmt212//===----------------------------------------------------------------------===//213 214//===----------------------------------------------------------------------===//215// EraseStmt216//===----------------------------------------------------------------------===//217 218EraseStmt *EraseStmt::create(Context &ctx, SMRange loc, Expr *rootOp) {219 return new (ctx.getAllocator().Allocate<EraseStmt>()) EraseStmt(loc, rootOp);220}221 222//===----------------------------------------------------------------------===//223// ReplaceStmt224//===----------------------------------------------------------------------===//225 226ReplaceStmt *ReplaceStmt::create(Context &ctx, SMRange loc, Expr *rootOp,227 ArrayRef<Expr *> replExprs) {228 unsigned allocSize = ReplaceStmt::totalSizeToAlloc<Expr *>(replExprs.size());229 void *rawData = ctx.getAllocator().Allocate(allocSize, alignof(ReplaceStmt));230 231 ReplaceStmt *stmt = new (rawData) ReplaceStmt(loc, rootOp, replExprs.size());232 llvm::uninitialized_copy(replExprs, stmt->getReplExprs().begin());233 return stmt;234}235 236//===----------------------------------------------------------------------===//237// RewriteStmt238//===----------------------------------------------------------------------===//239 240RewriteStmt *RewriteStmt::create(Context &ctx, SMRange loc, Expr *rootOp,241 CompoundStmt *rewriteBody) {242 return new (ctx.getAllocator().Allocate<RewriteStmt>())243 RewriteStmt(loc, rootOp, rewriteBody);244}245 246//===----------------------------------------------------------------------===//247// ReturnStmt248//===----------------------------------------------------------------------===//249 250ReturnStmt *ReturnStmt::create(Context &ctx, SMRange loc, Expr *resultExpr) {251 return new (ctx.getAllocator().Allocate<ReturnStmt>())252 ReturnStmt(loc, resultExpr);253}254 255//===----------------------------------------------------------------------===//256// AttributeExpr257//===----------------------------------------------------------------------===//258 259AttributeExpr *AttributeExpr::create(Context &ctx, SMRange loc,260 StringRef value) {261 return new (ctx.getAllocator().Allocate<AttributeExpr>())262 AttributeExpr(ctx, loc, copyStringWithNull(ctx, value));263}264 265//===----------------------------------------------------------------------===//266// CallExpr267//===----------------------------------------------------------------------===//268 269CallExpr *CallExpr::create(Context &ctx, SMRange loc, Expr *callable,270 ArrayRef<Expr *> arguments, Type resultType,271 bool isNegated) {272 unsigned allocSize = CallExpr::totalSizeToAlloc<Expr *>(arguments.size());273 void *rawData = ctx.getAllocator().Allocate(allocSize, alignof(CallExpr));274 275 CallExpr *expr = new (rawData)276 CallExpr(loc, resultType, callable, arguments.size(), isNegated);277 llvm::uninitialized_copy(arguments, expr->getArguments().begin());278 return expr;279}280 281//===----------------------------------------------------------------------===//282// DeclRefExpr283//===----------------------------------------------------------------------===//284 285DeclRefExpr *DeclRefExpr::create(Context &ctx, SMRange loc, Decl *decl,286 Type type) {287 return new (ctx.getAllocator().Allocate<DeclRefExpr>())288 DeclRefExpr(loc, decl, type);289}290 291//===----------------------------------------------------------------------===//292// MemberAccessExpr293//===----------------------------------------------------------------------===//294 295MemberAccessExpr *MemberAccessExpr::create(Context &ctx, SMRange loc,296 const Expr *parentExpr,297 StringRef memberName, Type type) {298 return new (ctx.getAllocator().Allocate<MemberAccessExpr>()) MemberAccessExpr(299 loc, parentExpr, memberName.copy(ctx.getAllocator()), type);300}301 302//===----------------------------------------------------------------------===//303// OperationExpr304//===----------------------------------------------------------------------===//305 306OperationExpr *307OperationExpr::create(Context &ctx, SMRange loc, const ods::Operation *odsOp,308 const OpNameDecl *name, ArrayRef<Expr *> operands,309 ArrayRef<Expr *> resultTypes,310 ArrayRef<NamedAttributeDecl *> attributes) {311 unsigned allocSize =312 OperationExpr::totalSizeToAlloc<Expr *, NamedAttributeDecl *>(313 operands.size() + resultTypes.size(), attributes.size());314 void *rawData =315 ctx.getAllocator().Allocate(allocSize, alignof(OperationExpr));316 317 Type resultType = OperationType::get(ctx, name->getName(), odsOp);318 OperationExpr *opExpr = new (rawData)319 OperationExpr(loc, resultType, name, operands.size(), resultTypes.size(),320 attributes.size(), name->getLoc());321 llvm::uninitialized_copy(operands, opExpr->getOperands().begin());322 llvm::uninitialized_copy(resultTypes, opExpr->getResultTypes().begin());323 llvm::uninitialized_copy(attributes, opExpr->getAttributes().begin());324 return opExpr;325}326 327std::optional<StringRef> OperationExpr::getName() const {328 return getNameDecl()->getName();329}330 331//===----------------------------------------------------------------------===//332// RangeExpr333//===----------------------------------------------------------------------===//334 335RangeExpr *RangeExpr::create(Context &ctx, SMRange loc,336 ArrayRef<Expr *> elements, RangeType type) {337 unsigned allocSize = RangeExpr::totalSizeToAlloc<Expr *>(elements.size());338 void *rawData = ctx.getAllocator().Allocate(allocSize, alignof(TupleExpr));339 340 RangeExpr *expr = new (rawData) RangeExpr(loc, type, elements.size());341 llvm::uninitialized_copy(elements, expr->getElements().begin());342 return expr;343}344 345//===----------------------------------------------------------------------===//346// TupleExpr347//===----------------------------------------------------------------------===//348 349TupleExpr *TupleExpr::create(Context &ctx, SMRange loc,350 ArrayRef<Expr *> elements,351 ArrayRef<StringRef> names) {352 unsigned allocSize = TupleExpr::totalSizeToAlloc<Expr *>(elements.size());353 void *rawData = ctx.getAllocator().Allocate(allocSize, alignof(TupleExpr));354 355 auto elementTypes = llvm::map_range(356 elements, [](const Expr *expr) { return expr->getType(); });357 TupleType type = TupleType::get(ctx, llvm::to_vector(elementTypes), names);358 359 TupleExpr *expr = new (rawData) TupleExpr(loc, type);360 llvm::uninitialized_copy(elements, expr->getElements().begin());361 return expr;362}363 364//===----------------------------------------------------------------------===//365// TypeExpr366//===----------------------------------------------------------------------===//367 368TypeExpr *TypeExpr::create(Context &ctx, SMRange loc, StringRef value) {369 return new (ctx.getAllocator().Allocate<TypeExpr>())370 TypeExpr(ctx, loc, copyStringWithNull(ctx, value));371}372 373//===----------------------------------------------------------------------===//374// Decl375//===----------------------------------------------------------------------===//376 377void Decl::setDocComment(Context &ctx, StringRef comment) {378 docComment = comment.copy(ctx.getAllocator());379}380 381//===----------------------------------------------------------------------===//382// AttrConstraintDecl383//===----------------------------------------------------------------------===//384 385AttrConstraintDecl *AttrConstraintDecl::create(Context &ctx, SMRange loc,386 Expr *typeExpr) {387 return new (ctx.getAllocator().Allocate<AttrConstraintDecl>())388 AttrConstraintDecl(loc, typeExpr);389}390 391//===----------------------------------------------------------------------===//392// OpConstraintDecl393//===----------------------------------------------------------------------===//394 395OpConstraintDecl *OpConstraintDecl::create(Context &ctx, SMRange loc,396 const OpNameDecl *nameDecl) {397 if (!nameDecl)398 nameDecl = OpNameDecl::create(ctx, SMRange());399 400 return new (ctx.getAllocator().Allocate<OpConstraintDecl>())401 OpConstraintDecl(loc, nameDecl);402}403 404std::optional<StringRef> OpConstraintDecl::getName() const {405 return getNameDecl()->getName();406}407 408//===----------------------------------------------------------------------===//409// TypeConstraintDecl410//===----------------------------------------------------------------------===//411 412TypeConstraintDecl *TypeConstraintDecl::create(Context &ctx, SMRange loc) {413 return new (ctx.getAllocator().Allocate<TypeConstraintDecl>())414 TypeConstraintDecl(loc);415}416 417//===----------------------------------------------------------------------===//418// TypeRangeConstraintDecl419//===----------------------------------------------------------------------===//420 421TypeRangeConstraintDecl *TypeRangeConstraintDecl::create(Context &ctx,422 SMRange loc) {423 return new (ctx.getAllocator().Allocate<TypeRangeConstraintDecl>())424 TypeRangeConstraintDecl(loc);425}426 427//===----------------------------------------------------------------------===//428// ValueConstraintDecl429//===----------------------------------------------------------------------===//430 431ValueConstraintDecl *ValueConstraintDecl::create(Context &ctx, SMRange loc,432 Expr *typeExpr) {433 return new (ctx.getAllocator().Allocate<ValueConstraintDecl>())434 ValueConstraintDecl(loc, typeExpr);435}436 437//===----------------------------------------------------------------------===//438// ValueRangeConstraintDecl439//===----------------------------------------------------------------------===//440 441ValueRangeConstraintDecl *442ValueRangeConstraintDecl::create(Context &ctx, SMRange loc, Expr *typeExpr) {443 return new (ctx.getAllocator().Allocate<ValueRangeConstraintDecl>())444 ValueRangeConstraintDecl(loc, typeExpr);445}446 447//===----------------------------------------------------------------------===//448// UserConstraintDecl449//===----------------------------------------------------------------------===//450 451std::optional<StringRef>452UserConstraintDecl::getNativeInputType(unsigned index) const {453 return hasNativeInputTypes ? getTrailingObjects<StringRef>()[index]454 : std::optional<StringRef>();455}456 457UserConstraintDecl *UserConstraintDecl::createImpl(458 Context &ctx, const Name &name, ArrayRef<VariableDecl *> inputs,459 ArrayRef<StringRef> nativeInputTypes, ArrayRef<VariableDecl *> results,460 std::optional<StringRef> codeBlock, const CompoundStmt *body,461 Type resultType) {462 bool hasNativeInputTypes = !nativeInputTypes.empty();463 assert(!hasNativeInputTypes || nativeInputTypes.size() == inputs.size());464 465 unsigned allocSize =466 UserConstraintDecl::totalSizeToAlloc<VariableDecl *, StringRef>(467 inputs.size() + results.size(),468 hasNativeInputTypes ? inputs.size() : 0);469 void *rawData =470 ctx.getAllocator().Allocate(allocSize, alignof(UserConstraintDecl));471 if (codeBlock)472 codeBlock = codeBlock->copy(ctx.getAllocator());473 474 UserConstraintDecl *decl = new (rawData)475 UserConstraintDecl(name, inputs.size(), hasNativeInputTypes,476 results.size(), codeBlock, body, resultType);477 llvm::uninitialized_copy(inputs, decl->getInputs().begin());478 llvm::uninitialized_copy(results, decl->getResults().begin());479 if (hasNativeInputTypes) {480 StringRef *nativeInputTypesPtr = decl->getTrailingObjects<StringRef>();481 for (unsigned i = 0, e = inputs.size(); i < e; ++i)482 nativeInputTypesPtr[i] = nativeInputTypes[i].copy(ctx.getAllocator());483 }484 485 return decl;486}487 488//===----------------------------------------------------------------------===//489// NamedAttributeDecl490//===----------------------------------------------------------------------===//491 492NamedAttributeDecl *NamedAttributeDecl::create(Context &ctx, const Name &name,493 Expr *value) {494 return new (ctx.getAllocator().Allocate<NamedAttributeDecl>())495 NamedAttributeDecl(name, value);496}497 498//===----------------------------------------------------------------------===//499// OpNameDecl500//===----------------------------------------------------------------------===//501 502OpNameDecl *OpNameDecl::create(Context &ctx, const Name &name) {503 return new (ctx.getAllocator().Allocate<OpNameDecl>()) OpNameDecl(name);504}505OpNameDecl *OpNameDecl::create(Context &ctx, SMRange loc) {506 return new (ctx.getAllocator().Allocate<OpNameDecl>()) OpNameDecl(loc);507}508 509//===----------------------------------------------------------------------===//510// PatternDecl511//===----------------------------------------------------------------------===//512 513PatternDecl *PatternDecl::create(Context &ctx, SMRange loc, const Name *name,514 std::optional<uint16_t> benefit,515 bool hasBoundedRecursion,516 const CompoundStmt *body) {517 return new (ctx.getAllocator().Allocate<PatternDecl>())518 PatternDecl(loc, name, benefit, hasBoundedRecursion, body);519}520 521//===----------------------------------------------------------------------===//522// UserRewriteDecl523//===----------------------------------------------------------------------===//524 525UserRewriteDecl *UserRewriteDecl::createImpl(Context &ctx, const Name &name,526 ArrayRef<VariableDecl *> inputs,527 ArrayRef<VariableDecl *> results,528 std::optional<StringRef> codeBlock,529 const CompoundStmt *body,530 Type resultType) {531 unsigned allocSize = UserRewriteDecl::totalSizeToAlloc<VariableDecl *>(532 inputs.size() + results.size());533 void *rawData =534 ctx.getAllocator().Allocate(allocSize, alignof(UserRewriteDecl));535 if (codeBlock)536 codeBlock = codeBlock->copy(ctx.getAllocator());537 538 UserRewriteDecl *decl = new (rawData) UserRewriteDecl(539 name, inputs.size(), results.size(), codeBlock, body, resultType);540 llvm::uninitialized_copy(inputs, decl->getInputs().begin());541 llvm::uninitialized_copy(results, decl->getResults().begin());542 return decl;543}544 545//===----------------------------------------------------------------------===//546// VariableDecl547//===----------------------------------------------------------------------===//548 549VariableDecl *VariableDecl::create(Context &ctx, const Name &name, Type type,550 Expr *initExpr,551 ArrayRef<ConstraintRef> constraints) {552 unsigned allocSize =553 VariableDecl::totalSizeToAlloc<ConstraintRef>(constraints.size());554 void *rawData = ctx.getAllocator().Allocate(allocSize, alignof(VariableDecl));555 556 VariableDecl *varDecl =557 new (rawData) VariableDecl(name, type, initExpr, constraints.size());558 llvm::uninitialized_copy(constraints, varDecl->getConstraints().begin());559 return varDecl;560}561 562//===----------------------------------------------------------------------===//563// Module564//===----------------------------------------------------------------------===//565 566Module *Module::create(Context &ctx, SMLoc loc, ArrayRef<Decl *> children) {567 unsigned allocSize = Module::totalSizeToAlloc<Decl *>(children.size());568 void *rawData = ctx.getAllocator().Allocate(allocSize, alignof(Module));569 570 Module *module = new (rawData) Module(loc, children.size());571 llvm::uninitialized_copy(children, module->getChildren().begin());572 return module;573}574