16#include "mlir/IR/Builders.h"
17#include "mlir/IR/ImplicitLocOpBuilder.h"
18#include "mlir/IR/SymbolTable.h"
19#include "llvm/ADT/STLExtras.h"
29 auto *
context = parser.getContext();
30 auto loc = parser.getCurrentLocation();
32 if (parser.parseString(&rawPath))
35 return parser.emitError(loc,
"invalid base path");
44 p << elt.module.getValue() <<
'/' << elt.
instance.getValue();
51 StringAttr &module, StringAttr &ref,
54 auto *
context = parser.getContext();
55 auto loc = parser.getCurrentLocation();
57 if (parser.parseString(&rawPath))
60 return parser.emitError(loc,
"invalid path");
65 StringAttr module, StringAttr ref,
68 for (
const auto &elt : path)
69 p << elt.module.getValue() <<
'/' << elt.instance.getValue() <<
':';
70 if (!module.getValue().empty())
71 p << module.getValue();
72 if (!ref.getValue().empty())
73 p <<
'>' << ref.getValue();
74 if (!field.getValue().empty())
75 p << field.getValue();
79static ParseResult
parseFieldLocs(OpAsmParser &parser, ArrayAttr &fieldLocs) {
80 if (parser.parseOptionalKeyword(
"field_locs"))
82 if (parser.parseLParen() || parser.parseAttribute(fieldLocs) ||
83 parser.parseRParen()) {
90 ArrayAttr fieldLocs) {
91 mlir::OpPrintingFlags flags;
92 if (!flags.shouldPrintDebugInfo() || !fieldLocs)
94 printer <<
"field_locs(";
95 printer.printAttribute(fieldLocs);
103 SmallVectorImpl<Attribute> &fieldNames,
104 SmallVectorImpl<Type> &fieldTypes) {
106 llvm::StringMap<SMLoc> nameLocMap;
107 auto parseElt = [&]() -> ParseResult {
109 std::string fieldName;
110 if (parser.parseKeywordOrString(&fieldName))
112 SMLoc currLoc = parser.getCurrentLocation();
113 if (nameLocMap.count(fieldName)) {
114 parser.emitError(currLoc,
"field \"")
115 << fieldName <<
"\" is defined twice";
116 parser.emitError(nameLocMap[fieldName]) <<
"previous definition is here";
119 nameLocMap[fieldName] = currLoc;
120 fieldNames.push_back(StringAttr::get(parser.getContext(), fieldName));
123 fieldTypes.emplace_back();
124 if (parser.parseColonType(fieldTypes.back()))
130 return parser.parseCommaSeparatedList(OpAsmParser::Delimiter::Paren,
134template <
typename ClassTy>
137 (void)mlir::impl::parseOptionalVisibilityKeyword(parser, state.attributes);
141 if (parser.parseSymbolName(symName, ClassTy::getSymNameAttrName(state.name),
146 SmallVector<OpAsmParser::Argument> args;
147 if (parser.parseArgumentList(args, OpAsmParser::Delimiter::Paren,
151 SmallVector<Type> fieldTypes;
152 SmallVector<Attribute> fieldNames;
153 if (succeeded(parser.parseOptionalArrow()))
157 SmallVector<NamedAttribute> fieldTypesMap;
158 if (!fieldNames.empty()) {
159 for (
auto [name, type] : zip(fieldNames, fieldTypes))
160 fieldTypesMap.push_back(
161 NamedAttribute(cast<StringAttr>(name), TypeAttr::get(type)));
163 auto *ctx = parser.getContext();
164 state.addAttribute(
"fieldNames", mlir::ArrayAttr::get(ctx, fieldNames));
165 state.addAttribute(
"fieldTypes",
166 mlir::DictionaryAttr::get(ctx, fieldTypesMap));
169 if (failed(parser.parseOptionalAttrDictWithKeyword(state.attributes)))
173 Region *region = state.addRegion();
174 if (parser.parseRegion(*region, args))
179 region->emplaceBlock();
182 auto argNames = llvm::map_range(args, [&](OpAsmParser::Argument arg) {
183 return StringAttr::get(parser.getContext(), arg.ssaName.name.drop_front());
187 ArrayAttr::get(parser.getContext(), SmallVector<Attribute>(argNames)));
196 StringRef visibilityAttrName =
197 mlir::SymbolOpInterface::getDefaultVisibilityAttrName();
198 if (
auto visibility =
199 classLike->getAttrOfType<StringAttr>(visibilityAttrName))
200 printer << visibility.getValue() <<
' ';
203 printer.printSymbolName(classLike.getSymName());
206 auto argNames = SmallVector<StringRef>(
207 classLike.getFormalParamNames().getAsValueRange<StringAttr>());
208 ArrayRef<BlockArgument> args = classLike.getBodyBlock()->getArguments();
212 for (
size_t i = 0, e = args.size(); i < e; ++i) {
213 printer <<
'%' << argNames[i] <<
": " << args[i].getType();
219 ArrayRef<Attribute> fieldNames =
220 cast<ArrayAttr>(classLike->getAttr(
"fieldNames")).getValue();
222 if (!fieldNames.empty()) {
224 for (
size_t i = 0, e = fieldNames.size(); i < e; ++i) {
227 StringAttr name = cast<StringAttr>(fieldNames[i]);
228 printer.printKeywordOrString(name.getValue());
230 Type type = classLike.getFieldType(name).value();
231 printer.printType(type);
237 SmallVector<StringRef> elidedAttrs{
238 classLike.getSymNameAttrName(), classLike.getFormalParamNamesAttrName(),
239 visibilityAttrName,
"fieldTypes",
"fieldNames"};
240 printer.printOptionalAttrDictWithKeyword(classLike.getOperation()->getAttrs(),
244 printer.printRegion(classLike.getBody(),
false,
250 if (classLike.getFormalParamNames().size() !=
251 classLike.getBodyBlock()->getArguments().size()) {
252 auto error = classLike.emitOpError(
253 "formal parameter name list doesn't match formal parameter value list");
254 error.attachNote(classLike.getLoc())
255 <<
"formal parameter names: " << classLike.getFormalParamNames();
256 error.attachNote(classLike.getLoc())
257 <<
"formal parameter values: "
258 << classLike.getBodyBlock()->getArguments();
268 auto argNames = SmallVector<StringRef>(
269 classLike.getFormalParamNames().getAsValueRange<StringAttr>());
270 ArrayRef<BlockArgument> args = classLike.getBodyBlock()->getArguments();
273 for (
size_t i = 0, e = args.size(); i < e; ++i)
274 setNameFn(args[i], argNames[i]);
278 return NamedAttribute(name, TypeAttr::get(type));
283 return NamedAttribute(StringAttr(name),
284 mlir::IntegerAttr::get(mlir::IndexType::get(ctx), i));
289 DictionaryAttr fieldTypes = mlir::cast<DictionaryAttr>(
290 classLike.getOperation()->getAttr(
"fieldTypes"));
291 Attribute type = fieldTypes.get(name);
292 if (
auto field = dyn_cast_or_null<TypeAttr>(type))
293 return field.getValue();
298 AttrTypeReplacer &replacer) {
299 classLike->setAttr(
"fieldTypes", cast<DictionaryAttr>(replacer.replace(
300 classLike.getFieldTypes())));
307ParseResult circt::om::ClassOp::parse(OpAsmParser &parser,
308 OperationState &state) {
309 return parseClassLike<ClassOp>(parser, state);
312circt::om::ClassOp circt::om::ClassOp::buildSimpleClassOp(
313 OpBuilder &odsBuilder, Location loc, Twine name,
314 ArrayRef<StringRef> formalParamNames, ArrayRef<StringRef> fieldNames,
315 ArrayRef<Type> fieldTypes) {
316 circt::om::ClassOp classOp = circt::om::ClassOp::create(
317 odsBuilder, loc, odsBuilder.getStringAttr(name), {},
318 odsBuilder.getStrArrayAttr(formalParamNames),
319 odsBuilder.getStrArrayAttr(fieldNames),
320 odsBuilder.getDictionaryAttr(llvm::map_to_vector(
321 llvm::zip(fieldNames, fieldTypes), [&](
auto field) -> NamedAttribute {
322 return NamedAttribute(odsBuilder.getStringAttr(std::get<0>(field)),
323 TypeAttr::get(std::get<1>(field)));
325 Block *body = &classOp.getRegion().emplaceBlock();
326 auto prevLoc = odsBuilder.saveInsertionPoint();
327 odsBuilder.setInsertionPointToEnd(body);
329 mlir::SmallVector<Attribute> locAttrs(fieldNames.size(), LocationAttr(loc));
331 ClassFieldsOp::create(odsBuilder, loc,
332 llvm::map_to_vector(fieldTypes,
333 [&](Type type) -> Value {
334 return body->addArgument(type,
337 odsBuilder.getArrayAttr(locAttrs));
339 odsBuilder.restoreInsertionPoint(prevLoc);
344void circt::om::ClassOp::print(OpAsmPrinter &printer) {
348LogicalResult circt::om::ClassOp::verify() {
return verifyClassLike(*
this); }
350LogicalResult circt::om::ClassOp::verifyRegions() {
352 dyn_cast_or_null<ClassFieldsOp>(this->
getBodyBlock()->getTerminator());
354 return this->emitOpError(
"expected terminator to be ClassFieldsOp");
357 if (fieldsOp.getNumOperands() != this->getFieldNames().size()) {
358 auto diag = this->emitOpError()
359 <<
"returns '" << this->getFieldNames().size()
360 <<
"' fields, but its terminator returned '"
361 << fieldsOp.getNumOperands() <<
"' fields";
362 return diag.attachNote(fieldsOp.getLoc()) <<
"see terminator:";
366 auto types = this->getFieldTypes();
367 for (
auto [fieldName, terminatorOperandType] :
368 llvm::zip(this->getFieldNames(), fieldsOp.getOperandTypes())) {
370 auto fieldNameAttr = dyn_cast_or_null<StringAttr>(fieldName);
372 return this->emitOpError(
"field name is not a StringAttr");
374 if (
auto fieldType = types.get(fieldNameAttr))
375 if (
auto typeAttr = dyn_cast<TypeAttr>(fieldType))
376 if (typeAttr.getValue() == terminatorOperandType)
379 auto diag = this->emitOpError()
380 <<
"returns different field types than its terminator";
381 return diag.attachNote(fieldsOp.getLoc()) <<
"see terminator:";
387void circt::om::ClassOp::getAsmBlockArgumentNames(
392std::optional<mlir::Type>
393circt::om::ClassOp::getFieldType(mlir::StringAttr field) {
397void circt::om::ClassOp::replaceFieldTypes(AttrTypeReplacer replacer) {
401void circt::om::ClassOp::updateFields(
402 mlir::ArrayRef<mlir::Location> newLocations,
403 mlir::ArrayRef<mlir::Value> newValues,
404 mlir::ArrayRef<mlir::Attribute> newNames) {
406 auto fieldsOp = getFieldsOp();
407 assert(fieldsOp &&
"The fields op should exist");
409 SmallVector<Attribute> names(getFieldNamesAttr().getAsRange<StringAttr>());
411 SmallVector<NamedAttribute> fieldTypes(getFieldTypesAttr().getValue());
413 SmallVector<Value> fieldVals(fieldsOp.getFields());
415 Location fieldOpLoc = fieldsOp->getLoc();
418 SmallVector<Location> locations;
419 if (
auto fl = dyn_cast<FusedLoc>(fieldOpLoc)) {
420 auto metadataArr = dyn_cast<ArrayAttr>(fl.getMetadata());
421 assert(metadataArr &&
"Expected the metadata for the fused location");
422 auto r = metadataArr.getAsRange<LocationAttr>();
423 locations.append(r.begin(), r.end());
426 locations.append(names.size(), fieldOpLoc);
430 names.append(newNames.begin(), newNames.end());
431 locations.append(newLocations.begin(), newLocations.end());
432 fieldVals.append(newValues.begin(), newValues.end());
435 for (
auto [v, n] :
llvm::zip(newValues, newNames))
436 fieldTypes.emplace_back(
437 NamedAttribute(
llvm::cast<StringAttr>(n), TypeAttr::
get(v.getType())));
440 SmallVector<Attribute> locationsAttr;
441 llvm::for_each(locations, [&](Location &l) {
442 locationsAttr.push_back(cast<Attribute>(l));
445 ImplicitLocOpBuilder builder(
getLoc(), *
this);
447 setFieldNamesAttr(builder.getArrayAttr(names));
449 setFieldTypesAttr(builder.getDictionaryAttr(fieldTypes));
450 fieldsOp.getFieldsMutable().assign(fieldVals);
452 fieldsOp->setLoc(builder.getFusedLoc(
453 locations, ArrayAttr::get(getContext(), locationsAttr)));
456void circt::om::ClassOp::addNewFieldsOp(mlir::OpBuilder &builder,
457 mlir::ArrayRef<Location> locs,
458 mlir::ArrayRef<Value> values) {
461 assert(locs.size() == values.size() &&
"Expected a location per value");
462 mlir::SmallVector<Attribute> locAttrs;
463 for (
auto loc : locs) {
464 locAttrs.push_back(cast<Attribute>(LocationAttr(loc)));
468 ClassFieldsOp::create(builder, builder.getFusedLoc(locs), values,
469 builder.getArrayAttr(locAttrs));
472mlir::Location circt::om::ClassOp::getFieldLocByIndex(
size_t i) {
473 auto fieldsOp = this->getFieldsOp();
474 auto fieldLocs = fieldsOp.getFieldLocs();
475 if (!fieldLocs.has_value())
476 return fieldsOp.getLoc();
477 assert(i < fieldLocs.value().size() &&
478 "field index too large for location array");
479 return cast<LocationAttr>(fieldLocs.value()[i]);
486ParseResult circt::om::ClassExternOp::parse(OpAsmParser &parser,
487 OperationState &state) {
488 return parseClassLike<ClassExternOp>(parser, state);
491void circt::om::ClassExternOp::print(OpAsmPrinter &printer) {
495LogicalResult circt::om::ClassExternOp::verify() {
501 return this->emitOpError(
"external class body should be empty");
507void circt::om::ClassExternOp::getAsmBlockArgumentNames(
512std::optional<mlir::Type>
513circt::om::ClassExternOp::getFieldType(mlir::StringAttr field) {
517void circt::om::ClassExternOp::replaceFieldTypes(AttrTypeReplacer replacer) {
525LogicalResult circt::om::ClassFieldsOp::verify() {
526 auto fieldLocs = this->getFieldLocs();
527 if (fieldLocs.has_value()) {
528 auto fieldLocsVal = fieldLocs.value();
529 if (fieldLocsVal.size() != this->getFields().size()) {
530 auto error = this->emitOpError(
"size of field_locs (")
531 << fieldLocsVal.size()
532 <<
") does not match number of fields ("
533 << this->getFields().size() <<
")";
543void circt::om::ObjectOp::build(::mlir::OpBuilder &odsBuilder,
544 ::mlir::OperationState &odsState,
546 ::mlir::ValueRange actualParams) {
547 return build(odsBuilder, odsState,
548 om::ClassType::get(odsBuilder.getContext(),
549 mlir::FlatSymbolRefAttr::get(classOp)),
550 mlir::FlatSymbolRefAttr::get(classOp.getNameAttr()),
554static FailureOr<ClassLike>
556 ClassType resultType, StringAttr className) {
557 StringAttr resultClassName = resultType.getClassName().getAttr();
558 if (resultClassName != className)
559 return op->emitOpError(
"result type (")
560 << resultClassName <<
") does not match referred to class ("
563 auto classDef = dyn_cast_or_null<ClassLike>(
564 symbolTable.lookupNearestSymbolFrom(op, className));
566 return op->emitOpError(
"refers to non-existant class (")
572circt::om::ObjectOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
575 getClassNameAttr().getAttr());
576 if (failed(classDef))
579 auto actualTypes = getActualParams().getTypes();
580 auto formalTypes = classDef->getBodyBlock()->getArgumentTypes();
583 if (actualTypes.size() != formalTypes.size()) {
584 auto error = emitOpError(
585 "actual parameter list doesn't match formal parameter list");
586 error.attachNote(classDef->getLoc())
587 <<
"formal parameters: " << classDef->getBodyBlock()->getArguments();
588 error.attachNote(
getLoc()) <<
"actual parameters: " << getActualParams();
593 for (
size_t i = 0, e = actualTypes.size(); i < e; ++i) {
594 if (actualTypes[i] != formalTypes[i]) {
595 return emitOpError(
"actual parameter type (")
596 << actualTypes[i] <<
") doesn't match formal parameter type ("
597 << formalTypes[i] <<
')';
609circt::om::ObjectFieldOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
610 auto classType = getObject().getType();
611 auto className = classType.getClassName().getAttr();
614 auto classDef = dyn_cast_or_null<ClassLike>(
615 symbolTable.lookupNearestSymbolFrom(*
this, className));
617 return emitOpError(
"class ") << className <<
" was not found";
620 auto fieldName = getFieldAttr();
621 std::optional<Type> fieldType = classDef.getFieldType(fieldName);
623 auto diag = emitOpError(
"referenced non-existent field ") << fieldName;
624 diag.attachNote(classDef.getLoc()) <<
"class defined here";
629 if (getResult().getType() != fieldType.value())
630 return emitOpError(
"expected type ")
631 << getResult().getType() <<
", but accessed field has type "
632 << fieldType.value();
640void circt::om::ElaboratedObjectOp::build(OpBuilder &odsBuilder,
641 OperationState &odsState,
642 om::ClassLike classOp,
643 ValueRange fieldValues) {
644 return build(odsBuilder, odsState,
646 odsBuilder.getContext(),
647 mlir::FlatSymbolRefAttr::get(classOp.getSymNameAttr())),
648 mlir::FlatSymbolRefAttr::get(classOp.getSymNameAttr()),
652LogicalResult circt::om::ElaboratedObjectOp::verifySymbolUses(
653 SymbolTableCollection &symbolTable) {
656 getClassNameAttr().getAttr());
657 if (failed(classDef))
660 auto fieldNames = classDef->getFieldNames();
661 auto fieldValues = getFieldValues();
662 if (fieldValues.size() != fieldNames.size())
663 return emitOpError(
"field value list doesn't match class field list, "
665 << fieldNames.size() <<
" values but got " << fieldValues.size();
667 for (
auto [fieldName, fieldValue] :
llvm::zip(fieldNames, fieldValues)) {
669 classDef->getFieldType(cast<StringAttr>(fieldName)).value();
670 if (fieldValue.getType() != expectedType)
671 return emitOpError(
"field value type for ")
672 << cast<StringAttr>(fieldName) <<
" (" << fieldValue.getType()
673 <<
") doesn't match class field type (" << expectedType <<
')';
683void circt::om::ConstantOp::build(::mlir::OpBuilder &odsBuilder,
684 ::mlir::OperationState &odsState,
685 ::mlir::TypedAttr constVal) {
686 return build(odsBuilder, odsState, constVal.getType(), constVal);
689OpFoldResult circt::om::ConstantOp::fold(FoldAdaptor adaptor) {
690 assert(adaptor.getOperands().empty() &&
"constant has no operands");
691 return getValueAttr();
698void circt::om::ListCreateOp::print(OpAsmPrinter &p) {
700 p.printOperands(getInputs());
701 p.printOptionalAttrDict((*this)->getAttrs());
702 p <<
" : " << getType().getElementType();
705ParseResult circt::om::ListCreateOp::parse(OpAsmParser &parser,
706 OperationState &result) {
707 llvm::SmallVector<OpAsmParser::UnresolvedOperand, 16> operands;
710 if (parser.parseOperandList(operands) ||
711 parser.parseOptionalAttrDict(result.attributes) || parser.parseColon() ||
712 parser.parseType(elemType))
714 result.addTypes({circt::om::ListType::get(elemType)});
716 for (
auto operand : operands)
717 if (parser.resolveOperand(operand, elemType, result.operands))
727BasePathCreateOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
728 auto hierPath = symbolTable.lookupNearestSymbolFrom<hw::HierPathOp>(
729 *
this, getTargetAttr());
731 return emitOpError(
"invalid symbol reference");
740PathCreateOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
741 auto hierPath = symbolTable.lookupNearestSymbolFrom<hw::HierPathOp>(
742 *
this, getTargetAttr());
744 return emitOpError(
"invalid symbol reference");
753 auto value = attr.getValue();
754 if (value.getType().isSignedInteger())
755 return value.getAPSInt();
760 return APSInt(value.getValue(),
false);
764 llvm::function_ref<FailureOr<APSInt>(
const APSInt &,
const APSInt &)>;
769 auto lhs = dyn_cast_or_null<circt::om::IntegerAttr>(lhsAttr);
770 auto rhs = dyn_cast_or_null<circt::om::IntegerAttr>(rhsAttr);
779 if (lhsVal.getBitWidth() > rhsVal.getBitWidth())
780 rhsVal = rhsVal.extend(lhsVal.getBitWidth());
781 else if (rhsVal.getBitWidth() > lhsVal.getBitWidth())
782 lhsVal = lhsVal.extend(rhsVal.getBitWidth());
785 auto result = evaluate(lhsVal, rhsVal);
790 auto *ctx = lhsAttr.getContext();
791 return circt::om::IntegerAttr::get(
792 ctx, mlir::IntegerAttr::get(ctx, result.value()));
799OpFoldResult IntegerAddOp::fold(FoldAdaptor adaptor) {
801 adaptor.getLhs(), adaptor.getRhs(),
802 [](
const APSInt &lhs,
const APSInt &rhs) { return success(lhs + rhs); });
809OpFoldResult IntegerMulOp::fold(FoldAdaptor adaptor) {
811 adaptor.getLhs(), adaptor.getRhs(),
812 [](
const APSInt &lhs,
const APSInt &rhs) { return success(lhs * rhs); });
819OpFoldResult IntegerShrOp::fold(FoldAdaptor adaptor) {
821 adaptor.getLhs(), adaptor.getRhs(),
822 [&](
const APSInt &lhs,
const APSInt &rhs) -> FailureOr<APSInt> {
824 if (!rhs.isNonNegative())
825 return (emitOpError(
"shift amount must be non-negative"), failure());
828 if (!rhs.isRepresentableByInt64())
829 return (emitOpError(
"shift amount must be representable in 64 bits"),
831 return success(lhs >> rhs.getExtValue());
839OpFoldResult IntegerShlOp::fold(FoldAdaptor adaptor) {
841 adaptor.getLhs(), adaptor.getRhs(),
842 [&](
const APSInt &lhs,
const APSInt &rhs) -> FailureOr<APSInt> {
844 if (!rhs.isNonNegative())
845 return (emitOpError(
"shift amount must be non-negative"), failure());
848 if (!rhs.isRepresentableByInt64())
849 return (emitOpError(
"shift amount must be representable in 64 bits"),
851 int64_t shiftAmt = rhs.getExtValue();
853 return success(lhs.extend(lhs.getBitWidth() + shiftAmt) << shiftAmt);
861OpFoldResult StringConcatOp::fold(FoldAdaptor adaptor) {
863 if (getStrings().size() == 1) {
864 if (
auto strAttr = adaptor.getStrings()[0])
867 return getStrings()[0];
871 if (!llvm::all_of(adaptor.getStrings(), [](Attribute operand) {
872 return isa_and_nonnull<StringAttr>(operand);
877 SmallString<64> result;
878 for (
auto operand : adaptor.getStrings())
879 result += cast<StringAttr>(operand).getValue();
881 return StringAttr::get(result, getResult().getType());
889 using OpRewritePattern::OpRewritePattern;
892 matchAndRewrite(StringConcatOp concat,
893 mlir::PatternRewriter &rewriter)
const override {
897 bool hasNestedConcat = llvm::any_of(concat.getStrings(), [](Value operand) {
898 auto nestedConcat = operand.getDefiningOp<StringConcatOp>();
899 return nestedConcat && operand.hasOneUse();
902 if (!hasNestedConcat)
906 SmallVector<Value> flatOperands;
907 for (
auto input : concat.getStrings()) {
908 if (
auto nestedConcat = input.getDefiningOp<StringConcatOp>();
909 nestedConcat && input.hasOneUse())
910 llvm::append_range(flatOperands, nestedConcat.getStrings());
912 flatOperands.push_back(input);
915 rewriter.modifyOpInPlace(concat,
916 [&]() { concat->setOperands(flatOperands); });
923class MergeAdjacentOMStringConstants
926 using OpRewritePattern::OpRewritePattern;
929 matchAndRewrite(StringConcatOp concat,
930 mlir::PatternRewriter &rewriter)
const override {
932 SmallVector<Value> newOperands;
933 SmallString<64> accumulatedLit;
934 SmallVector<ConstantOp> accumulatedOps;
935 bool changed =
false;
937 auto flushLiterals = [&]() {
938 if (accumulatedOps.empty())
942 if (accumulatedOps.size() == 1) {
943 newOperands.push_back(accumulatedOps[0]);
946 auto newLit = rewriter.createOrFold<ConstantOp>(
948 StringAttr::get(accumulatedLit, concat.getResult().getType()));
949 newOperands.push_back(newLit);
952 accumulatedLit.clear();
953 accumulatedOps.clear();
956 for (
auto operand : concat.getStrings()) {
957 if (
auto litOp = operand.getDefiningOp<ConstantOp>()) {
958 if (
auto strAttr = dyn_cast<StringAttr>(litOp.getValue())) {
960 if (strAttr.getValue().empty()) {
964 accumulatedLit += strAttr.getValue();
965 accumulatedOps.push_back(litOp);
971 newOperands.push_back(operand);
981 if (newOperands.empty())
982 return rewriter.replaceOpWithNewOp<ConstantOp>(
983 concat, StringAttr::get(
"", concat.getResult().getType())),
987 rewriter.modifyOpInPlace(concat,
988 [&]() { concat->setOperands(newOperands); });
995void StringConcatOp::getCanonicalizationPatterns(RewritePatternSet &results,
997 results.insert<FlattenOMStringConcat, MergeAdjacentOMStringConstants>(
1005static FailureOr<mlir::Attribute>
1007 auto resultType = mlir::IntegerType::get(lhsAttr.getContext(), 1);
1010 if (
auto lhs = dyn_cast<mlir::StringAttr>(lhsAttr))
1011 if (
auto rhs = dyn_cast<mlir::StringAttr>(rhsAttr))
1012 return mlir::Attribute(
1013 mlir::IntegerAttr::get(resultType, lhs == rhs ? 1 : 0));
1016 if (
auto lhs = dyn_cast<circt::om::IntegerAttr>(lhsAttr))
1017 if (
auto rhs = dyn_cast<circt::om::IntegerAttr>(rhsAttr)) {
1020 if (lhsVal.getBitWidth() > rhsVal.getBitWidth())
1021 rhsVal = rhsVal.extend(lhsVal.getBitWidth());
1022 else if (rhsVal.getBitWidth() > lhsVal.getBitWidth())
1023 lhsVal = lhsVal.extend(rhsVal.getBitWidth());
1024 return mlir::Attribute(
1025 mlir::IntegerAttr::get(resultType, lhsVal == rhsVal ? 1 : 0));
1029 if (
auto lhs = dyn_cast<mlir::IntegerAttr>(lhsAttr))
1030 if (
auto rhs = dyn_cast<mlir::IntegerAttr>(rhsAttr))
1031 return mlir::Attribute(
1032 mlir::IntegerAttr::get(resultType, lhs == rhs ? 1 : 0));
1037OpFoldResult PropEqOp::fold(FoldAdaptor adaptor) {
1038 auto lhsAttr = adaptor.getLhs();
1039 auto rhsAttr = adaptor.getRhs();
1040 if (!lhsAttr || !rhsAttr)
1056 auto lhsInt = dyn_cast_or_null<mlir::IntegerAttr>(lhsAttr);
1057 auto rhsInt = dyn_cast_or_null<mlir::IntegerAttr>(rhsAttr);
1058 if (!lhsInt || !rhsInt)
1060 APSInt lhsVal(lhsInt.getValue());
1061 APSInt rhsVal(rhsInt.getValue());
1062 auto result = evaluate(lhsVal, rhsVal);
1065 return mlir::IntegerAttr::get(
1066 lhsInt.getType(), result->extOrTrunc(lhsInt.getValue().getBitWidth()));
1071 auto i = dyn_cast_or_null<mlir::IntegerAttr>(a);
1072 return i && i.getValue().isZero();
1077 auto i = dyn_cast_or_null<mlir::IntegerAttr>(a);
1078 return i && i.getValue().isAllOnes();
1081OpFoldResult IntegerAndOp::fold(FoldAdaptor adaptor) {
1083 adaptor.getLhs(), adaptor.getRhs(),
1084 [](
const APSInt &lhs,
const APSInt &rhs) {
1085 return success(APSInt(lhs & rhs, false));
1090 return mlir::IntegerAttr::get(getResult().getType(),
1091 APInt::getZero(getType().
getWidth()));
1100OpFoldResult IntegerOrOp::fold(FoldAdaptor adaptor) {
1102 adaptor.getLhs(), adaptor.getRhs(),
1103 [](
const APSInt &lhs,
const APSInt &rhs) {
1104 return success(APSInt(lhs | rhs, false));
1109 return mlir::IntegerAttr::get(getResult().getType(),
1110 APInt::getAllOnes(getType().
getWidth()));
1119OpFoldResult IntegerXorOp::fold(FoldAdaptor adaptor) {
1121 adaptor.getLhs(), adaptor.getRhs(),
1122 [](
const APSInt &lhs,
const APSInt &rhs) {
1123 return success(APSInt(lhs ^ rhs, false));
1138LogicalResult circt::om::UnknownValueOp::verifySymbolUses(
1139 SymbolTableCollection &symbolTable) {
1142 auto classType = dyn_cast<ClassType>(getType());
1147 auto className = classType.getClassName();
1148 if (symbolTable.lookupNearestSymbolFrom<ClassLike>(*
this, className))
1151 return emitOpError() <<
"refers to non-existant class (\""
1152 << className.getValue() <<
"\")";
1159#define GET_OP_CLASSES
1160#include "circt/Dialect/OM/OM.cpp.inc"
assert(baseType &&"element must be base type")
static std::unique_ptr< Context > context
static Location getLoc(DefSlot slot)
static OpFoldResult foldIntegerBitwise(Attribute lhsAttr, Attribute rhsAttr, IntegerBinaryFn evaluate)
LogicalResult verifyClassLike(ClassLike classLike)
std::optional< Type > getClassLikeFieldType(ClassLike classLike, StringAttr name)
void getClassLikeAsmBlockArgumentNames(ClassLike classLike, Region ®ion, OpAsmSetValueNameFn setNameFn)
static ParseResult parseBasePathString(OpAsmParser &parser, PathAttr &path)
static ParseResult parsePathString(OpAsmParser &parser, PathAttr &path, StringAttr &module, StringAttr &ref, StringAttr &field)
static APSInt getAPSIntForOMIntegerAttr(circt::om::IntegerAttr attr)
static void printBasePathString(OpAsmPrinter &p, Operation *op, PathAttr path)
llvm::function_ref< FailureOr< APSInt >(const APSInt &, const APSInt &)> IntegerBinaryFn
static FailureOr< ClassLike > verifyClassLikeSymbolUser(Operation *op, SymbolTableCollection &symbolTable, ClassType resultType, StringAttr className)
static void printFieldLocs(OpAsmPrinter &printer, Operation *op, ArrayAttr fieldLocs)
static bool isZeroInt(Attribute a)
static ParseResult parseFieldLocs(OpAsmParser &parser, ArrayAttr &fieldLocs)
static ParseResult parseClassFieldsList(OpAsmParser &parser, SmallVectorImpl< Attribute > &fieldNames, SmallVectorImpl< Type > &fieldTypes)
static FailureOr< mlir::Attribute > evaluateBinaryEquality(mlir::Attribute lhsAttr, mlir::Attribute rhsAttr)
static void printClassLike(ClassLike classLike, OpAsmPrinter &printer)
void replaceClassLikeFieldTypes(ClassLike classLike, AttrTypeReplacer &replacer)
NamedAttribute makeFieldType(StringAttr name, Type type)
NamedAttribute makeFieldIdx(MLIRContext *ctx, mlir::StringAttr name, unsigned i)
static void printPathString(OpAsmPrinter &p, Operation *op, PathAttr path, StringAttr module, StringAttr ref, StringAttr field)
static ParseResult parseClassLike(OpAsmParser &parser, OperationState &state)
static OpFoldResult foldIntegerBinaryArithmetic(Attribute lhsAttr, Attribute rhsAttr, IntegerBinaryFn evaluate)
static bool isAllOnesInt(Attribute a)
static Block * getBodyBlock(FModuleLike mod)
static InstancePath empty
Direction get(bool isOutput)
Returns an output direction if isOutput is true, otherwise returns an input direction.
uint64_t getWidth(Type t)
void error(Twine message)
ParseResult parsePath(MLIRContext *context, StringRef spelling, PathAttr &path, StringAttr &module, StringAttr &ref, StringAttr &field)
Parse a target string in to a path.
ParseResult parseBasePath(MLIRContext *context, StringRef spelling, PathAttr &path)
Parse a target string of the form "Foo/bar:Bar/baz" in to a base path.
function_ref< void(Value, StringRef)> OpAsmSetValueNameFn
A module name, and the name of an instance inside that module.
mlir::StringAttr mlir::StringAttr instance