19#include "mlir/IR/Builders.h"
20#include "mlir/IR/Diagnostics.h"
21#include "mlir/IR/IRMapping.h"
22#include "mlir/IR/Matchers.h"
23#include "mlir/IR/Operation.h"
24#include "mlir/IR/SymbolTable.h"
25#include "mlir/Interfaces/SideEffectInterfaces.h"
26#include "mlir/Support/WalkResult.h"
27#include "mlir/Transforms/GreedyPatternRewriteDriver.h"
28#include "llvm/ADT/STLExtras.h"
29#include "llvm/Support/LogicalResult.h"
33#define GEN_PASS_DEF_ELABORATEOBJECT
34#include "circt/Dialect/OM/OMPasses.h.inc"
44using FieldIndex = DenseMap<std::pair<StringAttr, StringAttr>,
unsigned>;
49 ObjectOpInliningPattern(MLIRContext *
context, SymbolTable &symTable,
50 bool replaceExternalWithUnknown)
52 replaceExternalWithUnknown(replaceExternalWithUnknown) {}
54 LogicalResult matchAndRewrite(ObjectOp objOp,
55 PatternRewriter &rewriter)
const override {
57 symTable.lookup<ClassLike>(objOp.getClassNameAttr().getAttr());
61 if (isa<ClassExternOp>(classLike)) {
62 if (!replaceExternalWithUnknown)
64 rewriter.replaceOpWithNewOp<UnknownValueOp>(objOp, objOp.getType());
68 auto classOp = dyn_cast<ClassOp>(classLike.getOperation());
73 for (
auto [formal, actual] :
llvm::zip(
74 classOp.
getBodyBlock()->getArguments(), objOp.getActualParams()))
75 mapper.map(formal, actual);
79 classOp.getBody().cloneInto(&clonedRegion, mapper);
80 Block *clonedBlock = &clonedRegion.front();
82 auto clonedFields = cast<ClassFieldsOp>(clonedBlock->getTerminator());
83 SmallVector<Value> fieldValues(clonedFields.getFields());
86 if (
auto classOp = dyn_cast<ClassOp>(classLike.getOperation()))
87 for (
auto [i, v] :
llvm::enumerate(fieldValues)) {
89 rewriter.getFusedLoc({classOp.getFieldLocByIndex(i), v.getLoc()});
90 if (
auto *fieldOp = v.getDefiningOp()) {
91 rewriter.modifyOpInPlace(fieldOp, [&] { fieldOp->setLoc(fieldLoc); });
94 rewriter.modifyOpInPlace(
95 cast<BlockArgument>(v).getOwner()->getParentOp(),
96 [&] { val.setLoc(fieldLoc); });
101 rewriter.eraseOp(clonedFields);
102 rewriter.inlineBlockBefore(clonedBlock, objOp);
104 rewriter.replaceOpWithNewOp<ElaboratedObjectOp>(objOp, classLike,
110 const SymbolTable &symTable;
111 bool replaceExternalWithUnknown;
117 EvaluateObjectField(MLIRContext *
context,
const SymbolTable &symTable,
118 const FieldIndex &fieldIndexes)
120 fieldIndexes(fieldIndexes) {}
122 LogicalResult matchAndRewrite(ObjectFieldOp op,
123 PatternRewriter &rewriter)
const override {
125 auto elaboratedOp = op.getObject().getDefiningOp<ElaboratedObjectOp>();
130 symTable.lookup<ClassLike>(elaboratedOp.getClassNameAttr().getAttr());
135 fieldIndexes.at({classLike.getSymNameAttr(), op.getFieldAttr()});
136 auto result = elaboratedOp.getFieldValues()[index];
140 if (op.getResult() == result)
143 rewriter.replaceOp(op, result);
147 const SymbolTable &symTable;
148 const FieldIndex &fieldIndexes;
153struct UnknownPropagationPattern : RewritePattern {
154 UnknownPropagationPattern(MLIRContext *
context)
155 : RewritePattern(MatchAnyOpTypeTag(), 1,
context) {}
157 LogicalResult matchAndRewrite(Operation *op,
158 PatternRewriter &rewriter)
const override {
162 if (!isa_and_nonnull<OMDialect>(op->getDialect()) || !isPure(op) ||
163 op->getNumResults() == 0)
170 if (!llvm::any_of(op->getOperands(), [](Value operand) {
171 return operand.getDefiningOp<UnknownValueOp>();
176 SmallVector<Value> unknowns;
177 for (Type resultType : op->getResultTypes())
179 UnknownValueOp::create(rewriter, op->
getLoc(), resultType));
181 rewriter.replaceOp(op, unknowns);
188bool isFullyEvaluated(Operation *op) {
191 ClassOp, ClassFieldsOp, ElaboratedObjectOp, AnyCastOp,
193 ConstantOp, UnknownValueOp,
195 FrozenBasePathCreateOp, FrozenPathCreateOp, FrozenEmptyPathOp,
197 ListCreateOp, ListConcatOp>(op);
200LogicalResult verifyResult(ClassOp module,
bool allowUnevaluated) {
201 auto isLegal = [allowUnevaluated](Operation *op) -> LogicalResult {
203 if (
auto assertOp = dyn_cast<PropertyAssertOp>(op)) {
206 auto *defOp = assertOp.getCondition().getDefiningOp();
208 auto checkAssert = [&](
bool cond) -> LogicalResult {
219 dyn_cast_or_null<ConstantOp>(assertOp.getMessage().getDefiningOp());
221 if (allowUnevaluated)
222 return op->emitError(
"OM property assertion failed: <unevaluated>");
224 auto diag = emitError(op->getLoc(),
225 "OM property assertion failed, but no message "
226 "is available as the message is unevaluated");
227 diag.attachNote(assertOp.getMessage().getLoc())
228 <<
"unevaluated message operation is here";
233 if (!matchPattern(assertOp.getMessage(), m_Constant(&message)))
234 return op->emitError()
235 <<
"OM property assertion failed, but no message is available "
236 "because the message is not a constant string";
237 return op->emitError(
"OM property assertion failed: ")
238 << message.getValue();
242 if (matchPattern(assertOp.getCondition(), m_ConstantInt(&value)))
243 return checkAssert(!value.isZero());
246 if (
auto unknownOp = dyn_cast_or_null<UnknownValueOp>(defOp))
247 return checkAssert(
true);
250 if (allowUnevaluated)
252 return emitError(op->getLoc(),
"failed to evaluate assertion condition");
255 if (!isFullyEvaluated(op)) {
256 if (allowUnevaluated)
258 return emitError(op->getLoc()) <<
"failed to evaluate " << op->getName();
263 bool encounteredError =
false;
264 module.walk([&](Operation *op) { encounteredError |= failed(isLegal(op)); });
266 return failure(encounteredError);
269struct ElaborateObjectPass
270 :
public circt::om::impl::ElaborateObjectBase<ElaborateObjectPass> {
273 static LogicalResult elaborateClass(ClassOp classOp, SymbolTable &symTable,
274 FieldIndex &fieldIndexes,
275 bool allowUnevaluated =
false) {
280 RewritePatternSet
patterns(classOp.getContext());
281 patterns.add<ObjectOpInliningPattern>(classOp.getContext(), symTable,
283 patterns.add<EvaluateObjectField>(classOp.getContext(), symTable,
285 patterns.add<UnknownPropagationPattern>(classOp.getContext());
286 GreedyRewriteConfig config;
288 config.setMaxIterations(GreedyRewriteConfig::kNoLimit);
289 if (failed(applyPatternsGreedily(classOp, std::move(
patterns), config)))
293 return verifyResult(classOp, allowUnevaluated);
296 LogicalResult initialize(MLIRContext *
context)
override {
298 allPublicClasses.getValue() + !targetClass.getValue().empty();
300 return emitError(UnknownLoc::get(
context))
301 <<
"exactly one of 'target-class' or 'all-public-classes' must "
306 void runOnOperation()
override {
307 auto module = getOperation();
308 auto &symTable = getAnalysis<SymbolTable>();
312 FieldIndex fieldIndexes;
313 for (
auto classOp : module.getOps<ClassLike>()) {
314 auto name = classOp.getSymNameAttr();
315 for (
auto [idx, fieldName] :
316 llvm::enumerate(classOp.getFieldNames().getAsRange<StringAttr>()))
317 fieldIndexes[{name, fieldName}] = idx;
321 if (allPublicClasses) {
322 for (
auto classOp : module.getOps<ClassOp>()) {
323 if (!classOp.isPublic())
325 if (failed(elaborateClass(classOp, symTable, fieldIndexes,
327 return signalPassFailure();
333 auto classOp = symTable.lookup<ClassOp>(targetClass);
335 emitError(module.getLoc())
336 <<
"target class '" << targetClass <<
"' was not found";
337 return signalPassFailure();
341 elaborateClass(classOp, symTable, fieldIndexes, allowUnevaluated)))
342 return signalPassFailure();
assert(baseType &&"element must be base type")
static std::unique_ptr< Context > context
static Location getLoc(DefSlot slot)
static Block * getBodyBlock(FModuleLike mod)
The InstanceGraph op interface, see InstanceGraphInterface.td for more details.