69 tables, origin, ArrayRef<ArrayRef<Type>> {funcType.getInputs(), funcType.getResults()}
77static LogicalResult verifyArgOrResNameAttrs(
78 ArrayAttr attrs, StringRef ownAttrName, StringRef crossAttrName, StringRef ownLabel,
84 llvm::DenseSet<StringAttr> seenNames;
85 for (
auto [i, attr] : llvm::enumerate(attrs)) {
86 auto dictAttr = llvm::dyn_cast<DictionaryAttr>(attr);
90 if (dictAttr.contains(crossAttrName)) {
91 return emitFn().append(
92 '\'', crossAttrName,
"' is only valid on function ", crossLabel,
"s but found on ",
96 Attribute nameAttr = dictAttr.get(ownAttrName);
100 auto name = llvm::dyn_cast<StringAttr>(nameAttr);
102 return emitFn().append(
103 '\'', ownAttrName,
"' on ", ownLabel,
' ', i,
" must be a string attribute"
106 if (!llvm::isa<NoneType>(name.getType())) {
107 return emitFn().append(
108 '\'', ownAttrName,
"' on ", ownLabel,
' ', i,
" must not have an explicit type"
111 if (name.getValue().empty()) {
112 return emitFn().append(
'\'', ownAttrName,
"' on ", ownLabel,
' ', i,
" must not be empty");
114 if (!seenNames.insert(name).second) {
115 return emitFn().append(
116 "duplicate '", ownAttrName,
"' value \"", name.getValue(),
"\" on ", ownLabel,
' ', i
123static std::optional<StringAttr>
124getFunctionNameAttrAtIndex(ArrayAttr attrs,
unsigned index, StringRef attrName) {
125 if (!attrs || index >= attrs.size()) {
128 if (
auto dictAttr = llvm::dyn_cast<DictionaryAttr>(attrs[index])) {
129 if (
auto nameAttr = llvm::dyn_cast_if_present<StringAttr>(dictAttr.get(attrName))) {
142 Location location, StringRef name, FunctionType type, ArrayRef<NamedAttribute> attrs
148 Location location, StringRef name, FunctionType type, Operation::dialect_attr_range attrs
150 SmallVector<NamedAttribute, 8> attrRef(attrs);
151 return create(location, name, type, llvm::ArrayRef(attrRef));
155 Location location, StringRef name, FunctionType type, ArrayRef<NamedAttribute> attrs,
156 ArrayRef<DictionaryAttr> argAttrs
158 FuncDefOp func =
create(location, name, type, attrs);
159 func.setAllArgAttrs(argAttrs);
164 OpBuilder &builder, OperationState &state, StringRef name, FunctionType type,
165 ArrayRef<NamedAttribute> attrs, ArrayRef<DictionaryAttr> argAttrs
167 state.addAttribute(SymbolTable::getSymbolAttrName(), builder.getStringAttr(name));
169 state.attributes.append(attrs.begin(), attrs.end());
172 if (argAttrs.empty()) {
175 assert(type.getNumInputs() == argAttrs.size());
176 function_interface_impl::addArgAndResultAttrs(
183 auto buildFuncType = [](Builder &builder, ArrayRef<Type> argTypes, ArrayRef<Type> results,
184 function_interface_impl::VariadicFlag,
185 std::string &) {
return builder.getFunctionType(argTypes, results); };
187 return function_interface_impl::parseFunctionOp(
194 function_interface_impl::printFunctionOp(
204 llvm::MapVector<StringAttr, Attribute> newAttrMap;
205 for (
const auto &attr : dest->getAttrs()) {
206 newAttrMap.insert({attr.getName(), attr.getValue()});
208 for (
const auto &attr : (*this)->getAttrs()) {
209 newAttrMap.insert({attr.getName(), attr.getValue()});
213 llvm::to_vector(llvm::map_range(newAttrMap, [](std::pair<StringAttr, Attribute> attrPair) {
214 return NamedAttribute(attrPair.first, attrPair.second);
216 dest->setAttrs(DictionaryAttr::get(getContext(), newAttrs));
229 FuncDefOp newFunc = llvm::cast<FuncDefOp>(getOperation()->cloneWithoutRegions());
237 unsigned oldNumArgs = oldType.getNumInputs();
238 SmallVector<Type, 4> newInputs;
239 newInputs.reserve(oldNumArgs);
240 for (
unsigned i = 0; i != oldNumArgs; ++i) {
241 if (!mapper.contains(getArgument(i))) {
242 newInputs.push_back(oldType.getInput(i));
248 if (newInputs.size() != oldNumArgs) {
249 newFunc.setType(FunctionType::get(oldType.getContext(), newInputs, oldType.getResults()));
251 if (ArrayAttr argAttrs = getAllArgAttrs()) {
252 SmallVector<Attribute> newArgAttrs;
253 newArgAttrs.reserve(newInputs.size());
254 for (
unsigned i = 0; i != oldNumArgs; ++i) {
255 if (!mapper.contains(getArgument(i))) {
256 newArgAttrs.push_back(argAttrs[i]);
259 newFunc.setAllArgAttrs(newArgAttrs);
271 return clone(mapper);
276 getOperation()->setAttr(AllowConstraintAttr::name, UnitAttr::get(getContext()));
278 getOperation()->removeAttr(AllowConstraintAttr::name);
284 getOperation()->setAttr(AllowWitnessAttr::name, UnitAttr::get(getContext()));
286 getOperation()->removeAttr(AllowWitnessAttr::name);
292 getOperation()->setAttr(AllowNonNativeFieldOpsAttr::name, UnitAttr::get(getContext()));
294 getOperation()->removeAttr(AllowNonNativeFieldOpsAttr::name);
300 getOperation()->setAttr(AllowVerifOpsAttr::name, UnitAttr::get(getContext()));
302 getOperation()->removeAttr(AllowVerifOpsAttr::name);
307 if (index < this->getNumArguments()) {
308 DictionaryAttr res = function_interface_impl::getArgAttrDict(*
this, index);
309 return res ? res.contains(PublicAttr::name) :
false;
323 assert(index < getNumArguments() &&
"argument index out of range");
338 assert(index < getNumResults() &&
"result index out of range");
350 return emitErrorFunc() <<
'\'' <<
ARG_NAME_ATTR_NAME <<
"' is only valid on function arguments";
353 return emitErrorFunc() <<
'\'' <<
RES_NAME_ATTR_NAME <<
"' is only valid on function results";
356 if (failed(verifyArgOrResNameAttrs(
363 if (failed(verifyArgOrResNameAttrs(
375 for (Type t : type.getInputs()) {
380 return emitErrorFunc().append(
381 "\"@", getName(),
"\" parameters cannot contain affine map attributes but found ", t
385 for (Type t : type.getResults()) {
393 WalkResult res = this->walk<WalkOrder::PreOrder>([
this](ModuleOp nestedMod) {
395 "cannot be nested within '", getOperation()->getName(),
"' operations"
397 return WalkResult::interrupt();
399 if (res.wasInterrupted()) {
411 llvm::ArrayRef<Type> resTypes = funcType.getResults();
413 if (resTypes.size() != 1) {
414 return origin.emitOpError().append(
418 if (failed(
checkSelfType(tables, parent, resTypes.front(), origin,
"return"))) {
429verifyFuncTypeProduct(
FuncDefOp &origin, SymbolTableCollection &tables, StructDefOp &parent) {
431 return verifyFuncTypeCompute(origin, tables, parent);
435verifyFuncTypeConstrain(
FuncDefOp &origin, SymbolTableCollection &tables, StructDefOp &parent) {
438 if (funcType.getResults().size() != 0) {
439 return origin.emitOpError() <<
"\"@" <<
FUNC_NAME_CONSTRAIN <<
"\" must have no return type";
443 llvm::ArrayRef<Type> inputTypes = funcType.getInputs();
444 if (inputTypes.size() < 1) {
446 <<
"\" must have at least one input type";
448 if (failed(
checkSelfType(tables, parent, inputTypes.front(), origin,
"first input"))) {
466 return verifyFuncTypeCompute(*
this, tables, parentStructOpt);
468 return verifyFuncTypeConstrain(*
this, tables, parentStructOpt);
470 return verifyFuncTypeProduct(*
this, tables, parentStructOpt);
485 assert(!body.empty() &&
"compute() function body is empty");
486 Block &block = body.back();
489 Operation *terminator = block.getTerminator();
490 assert(terminator &&
"compute() function has no terminator");
491 auto retOp = llvm::dyn_cast<ReturnOp>(terminator);
494 << terminator->getName() <<
"'\n";
495 llvm_unreachable(
"compute() function must end with ReturnOp");
497 return retOp.getOperands().front();
502 return getArguments().front();
515 auto function = getParentOp<FuncDefOp>();
518 const auto results =
function.getFunctionType().getResults();
519 if (getNumOperands() != results.size()) {
520 return emitOpError(
"has ") << getNumOperands() <<
" operands, but enclosing function (@"
521 <<
function.getName() <<
") returns " << results.size();
524 for (
unsigned i = 0, e = results.size(); i != e; ++i) {
525 if (!
typesUnify(getOperand(i).getType(), results[i])) {
526 return emitError() <<
"type of return operand " << i <<
" (" << getOperand(i).getType()
527 <<
") doesn't match function result type (" << results[i] <<
')'
528 <<
" in function @" <<
function.getName();
542 auto &prop = state.getOrAddProperties<
Properties>();
543 if (failed(reader.readAttribute(prop.callee)) ||
544 failed(reader.readAttribute(prop.mapOpGroupSizes)) ||
545 failed(reader.readOptionalAttribute(prop.numDimsPerMap))) {
549 if (reader.getBytecodeVersion() < 6) {
550 auto &propStorage = prop.operandSegmentSizes;
551 DenseI32ArrayAttr attr;
552 if (failed(reader.readAttribute(attr))) {
555 if (attr.size() >
static_cast<int64_t
>(
sizeof(propStorage) /
sizeof(int32_t))) {
556 reader.emitError(
"size mismatch for operand/result_segment_size");
559 llvm::copy(ArrayRef<int32_t>(attr), propStorage.begin());
564 if (succeeded(versionOpt)) {
566 if (ver.majorVersion >= 2) {
567 if (failed(reader.readOptionalAttribute(prop.templateParams))) {
573 if (reader.getBytecodeVersion() >= 6) {
574 return reader.readSparseArray(MutableArrayRef(prop.operandSegmentSizes));
581 auto &prop = getProperties();
582 writer.writeAttribute(prop.callee);
583 writer.writeAttribute(prop.mapOpGroupSizes);
584 writer.writeOptionalAttribute(prop.numDimsPerMap);
586 if (writer.getBytecodeVersion() < 6) {
587 auto &propStorage = prop.operandSegmentSizes;
588 writer.writeAttribute(DenseI32ArrayAttr::get(this->getContext(), propStorage));
591 writer.writeOptionalAttribute(prop.templateParams);
593 auto &propStorage = prop.operandSegmentSizes;
594 if (writer.getBytecodeVersion() >= 6) {
595 writer.writeSparseArray(ArrayRef(propStorage));
600 OpBuilder &odsBuilder, OperationState &odsState, TypeRange resultTypes, SymbolRefAttr callee,
601 ValueRange argOperands, ArrayRef<Attribute> templateParams
603 odsState.addTypes(resultTypes);
604 odsState.addOperands(argOperands);
608 props.setCallee(callee);
613 OpBuilder &odsBuilder, OperationState &odsState, TypeRange resultTypes, SymbolRefAttr callee,
614 ArrayRef<ValueRange> mapOperands, DenseI32ArrayAttr numDimsPerMap, ValueRange argOperands,
615 ArrayRef<Attribute> templateParams
617 odsState.addTypes(resultTypes);
618 odsState.addOperands(argOperands);
620 odsBuilder, odsState, mapOperands, numDimsPerMap,
623 props.setCallee(callee);
631 if (
auto intAttr = llvm::dyn_cast<IntegerAttr>(paramFromCallOp)) {
633 std::optional<Type> declaredType = targetParam.
getTypeOpt();
634 if (!declaredType || !llvm::isa<TypeVarType>(*declaredType)) {
635 auto diag = this->emitOpError().append(
636 "wildcard `?` can only be used for template parameters with `!poly.tvar` "
637 "type restriction, but parameter \"@",
638 targetParam.getName(),
"\" has "
641 diag.append(
"type restriction ", *declaredType);
643 diag.append(
"no type restriction");
650 if (std::optional<Type> declaredType = targetParam.
getTypeOpt()) {
651 bool compatible =
false;
652 if (
auto sym = llvm::dyn_cast<SymbolRefAttr>(paramFromCallOp)) {
653 if (sym.getNestedReferences().empty()) {
654 SymbolTableCollection tables;
656 if (failed(parentTemplate)) {
659 if (TemplateOp p = *parentTemplate) {
660 auto binding = p.getConstNamed<TemplateSymbolBindingOpInterface>(sym.getRootReference());
664 if (std::optional<Type> actualType = binding.getTypeOpt()) {
665 compatible =
typesUnify(*actualType, *declaredType);
672 }
else if (llvm::isa<TypeVarType>(*declaredType)) {
673 compatible = llvm::isa<TypeAttr>(paramFromCallOp);
674 }
else if (llvm::isa<FeltType>(*declaredType)) {
675 compatible = llvm::isa<FeltConstAttr, IntegerAttr>(paramFromCallOp) &&
677 }
else if (llvm::isa<IndexType, IntegerType>(*declaredType)) {
681 compatible = llvm::isa<IntegerAttr>(paramFromCallOp) &&
685 llvm_unreachable(
"inconsistent with `isValidConstReadType()`");
689 return this->emitOpError().append(
690 "instantiation value '", paramFromCallOp,
"' is not compatible with parameter \"@",
691 targetParam.getName(),
"\" type restriction ", *declaredType
699 llvm::iterator_range<Region::op_iterator<TemplateParamOp>> targetParamDefs
703 assert((callParams.size() == llvm::range_size(targetParamDefs)) &&
"pre-condition");
705 for (
auto [paramOp, attr] : llvm::zip_equal(targetParamDefs, callParams.getValue())) {
714 llvm::iterator_range<Region::op_iterator<TemplateParamOp>> targetParamDefs,
719 assert((callParams.size() == llvm::range_size(targetParamDefs)) &&
"pre-condition");
721 for (
auto [paramOp, attr] : llvm::zip_equal(targetParamDefs, callParams.getValue())) {
723 if (
auto intAttr = llvm::dyn_cast<IntegerAttr>(attr)) {
728 auto it = unifications.find({FlatSymbolRefAttr::get(paramOp.getNameAttr()),
Side::RHS});
729 if (it != unifications.end() && !
typeParamsUnify({attr}, {it->second})) {
731 return this->emitOpError().append(
732 "template instantiation value '", attr,
"' for parameter \"@", paramOp.getName(),
733 "\" conflicts with value '", it->second,
"' inferred from function type signature"
742struct CallOpVerifier {
743 CallOpVerifier(
CallOp *c,
FunctionKind tgtFuncKind) : callOp(c), tgtKind(tgtFuncKind) {}
744 CallOpVerifier(
CallOp *c, StringRef tgtName) : CallOpVerifier(c,
fnNameToKind(tgtName)) {}
745 virtual ~CallOpVerifier() =
default;
747 LogicalResult verify() {
750 LogicalResult aggregateResult = success();
751 if (failed(verifyTargetAttributes())) {
752 aggregateResult = failure();
754 if (failed(verifyInputs())) {
755 aggregateResult = failure();
757 if (failed(verifyOutputs())) {
758 aggregateResult = failure();
760 if (failed(verifyTemplateParams())) {
761 aggregateResult = failure();
763 if (failed(verifyAffineMapParams())) {
764 aggregateResult = failure();
766 return aggregateResult;
773 virtual LogicalResult verifyTargetAttributes() = 0;
774 virtual LogicalResult verifyInputs() = 0;
775 virtual LogicalResult verifyOutputs() = 0;
776 virtual LogicalResult verifyTemplateParams() = 0;
777 virtual LogicalResult verifyAffineMapParams() = 0;
780 LogicalResult verifyTargetAttributesMatch(FuncDefOp target) {
781 LogicalResult aggregateRes = success();
782 if (FuncDefOp caller = (*callOp)->getParentOfType<FuncDefOp>()) {
783 auto emitAttrErr = [&](StringLiteral attrName) {
784 aggregateRes = callOp->emitOpError()
785 <<
"target '@" << target.getName() <<
"' has '" << attrName
786 <<
"' attribute, which is not specified by the caller '@" << caller.getName()
791 emitAttrErr(AllowConstraintAttr::name);
794 emitAttrErr(AllowWitnessAttr::name);
797 emitAttrErr(AllowNonNativeFieldOpsAttr::name);
803 LogicalResult verifyNoTemplateInstantiations() {
806 return callOp->emitOpError().append(
807 "can only have template instantiations when targeting a templated free function"
813 LogicalResult verifyNoAffineMapInstantiations() {
816 return callOp->emitOpError().append(
817 "can only have affine map instantiations when targeting a \"@",
FUNC_NAME_COMPUTE,
823 assert(callOp->getMapOperands().empty());
828struct KnownTargetVerifier :
public CallOpVerifier {
829 KnownTargetVerifier(CallOp *c, SymbolLookupResult<FuncDefOp> &&tgtRes)
830 : CallOpVerifier(c, tgtRes.get().getSymName()), tgt(*tgtRes), tgtType(tgt.getFunctionType()),
831 includeSymNames(tgtRes.getNamespace()) {}
833 LogicalResult verifyTargetAttributes()
override {
834 return CallOpVerifier::verifyTargetAttributesMatch(tgt);
837 LogicalResult verifyInputs()
override {
838 return verifyTypesMatch(callOp->
getArgOperands().getTypes(), tgtType.getInputs(),
"operand");
841 LogicalResult verifyOutputs()
override {
842 return verifyTypesMatch(callOp->getResultTypes(), tgtType.getResults(),
"result");
845 LogicalResult verifyTemplateParams()
override {
846 Operation *tgtOp = tgt.getOperation();
849 return verifyNoTemplateInstantiations();
859 auto realParams = tgtOpParent.getConstOps<TemplateParamOp>();
864 llvm::SmallDenseSet<SymbolRefAttr> referencedInSignature;
868 bool allParamsReferenced = llvm::all_of(realParams, [&](TemplateParamOp p) {
869 return referencedInSignature.contains(FlatSymbolRefAttr::get(p.getNameAttr()));
871 if (allParamsReferenced) {
875 return callOp->emitOpError().append(
876 "must provide template instantiation parameters when calling \"@", tgt.getSymName(),
877 "\" because not all template parameters of \"@", tgtOpParent.getSymName(),
878 "\" appear in the function type signature"
884 return llzk::InFlightDiagnosticWrapper(this->callOp->emitOpError());
890 size_t numTemplateParams = llvm::range_size(realParams);
891 if (callParams.size() != numTemplateParams) {
893 return callOp->emitOpError().append(
894 "template instantiation has ", callParams.size(),
" parameter(s) but \"@",
895 tgtOpParent.getSymName(),
"\" expects ", numTemplateParams,
" template parameter(s)"
907 assert(succeeded(unifyResult) &&
"already checked by `verifyInputs()` and `verifyOutputs()`");
911 return verifyNoTemplateInstantiations();
915 LogicalResult verifyAffineMapParams()
override {
923 if (ArrayAttr params = retTy.getParams()) {
925 SmallVector<AffineMapAttr> mapAttrs;
926 for (Attribute a : params) {
927 if (AffineMapAttr m = dyn_cast<AffineMapAttr>(a)) {
928 mapAttrs.push_back(m);
939 return verifyNoAffineMapInstantiations();
944 template <
typename T>
946 verifyTypesMatch(ValueTypeRange<T> callOpTypes, ArrayRef<Type> tgtTypes,
const char *aspect) {
947 if (tgtTypes.size() != callOpTypes.size()) {
948 return callOp->emitOpError()
949 .append(
"incorrect number of ", aspect,
"s for callee, expected ", tgtTypes.size())
950 .attachNote(tgt.getLoc())
951 .append(
"callee defined here");
953 for (
unsigned i = 0, e = tgtTypes.size(); i != e; ++i) {
954 if (!
typesUnify(callOpTypes[i], tgtTypes[i], includeSymNames)) {
955 return callOp->emitOpError().append(
956 aspect,
" type mismatch: expected type ", tgtTypes[i],
", but found ", callOpTypes[i],
957 " for ", aspect,
" number ", i
965 FunctionType tgtType;
966 std::vector<llvm::StringRef> includeSymNames;
971LogicalResult checkSelfTypeUnknownTarget(
972 StringAttr expectedParamName, Type actualType,
CallOp *origin,
const char *aspect
974 if (!llvm::isa<TypeVarType>(actualType) ||
975 llvm::cast<TypeVarType>(actualType).getRefName() != expectedParamName) {
981 return origin->emitOpError().append(
982 "target \"@", origin->
getCallee().getLeafReference().getValue(),
"\" expected ", aspect,
983 " type '!",
TypeVarType::name,
"<@", expectedParamName.getValue(),
">' but found ",
999struct UnknownTargetVerifier :
public CallOpVerifier {
1000 UnknownTargetVerifier(CallOp *c,
FunctionKind tgtFuncKind, SymbolRefAttr callee)
1001 : CallOpVerifier(c, tgtFuncKind), calleeAttr(callee) {
1008 LogicalResult verifyTargetAttributes()
override {
1011 LogicalResult aggregateRes = success();
1012 if (FuncDefOp caller = (*callOp)->getParentOfType<FuncDefOp>()) {
1013 auto emitAttrErr = [&](StringLiteral attrName) {
1014 aggregateRes = callOp->emitOpError()
1015 <<
"target '" << calleeAttr <<
"' has '" << attrName
1016 <<
"' attribute, which is not specified by the caller '@" << caller.getName()
1022 if (!caller.hasAllowConstraintAttr()) {
1023 emitAttrErr(AllowConstraintAttr::name);
1027 if (!caller.hasAllowWitnessAttr()) {
1028 emitAttrErr(AllowWitnessAttr::name);
1032 if (!caller.hasAllowWitnessAttr()) {
1033 emitAttrErr(AllowWitnessAttr::name);
1035 if (!caller.hasAllowConstraintAttr()) {
1036 emitAttrErr(AllowConstraintAttr::name);
1043 return aggregateRes;
1046 LogicalResult verifyInputs()
override {
1052 Operation::operand_type_range inputTypes = callOp->
getArgOperands().getTypes();
1053 if (inputTypes.size() < 1) {
1055 return callOp->emitOpError()
1058 return checkSelfTypeUnknownTarget(
1059 calleeAttr.getRootReference(), inputTypes.front(), callOp,
"first input"
1065 LogicalResult verifyOutputs()
override {
1069 Operation::result_type_range resTypes = callOp->getResultTypes();
1070 if (resTypes.size() != 1) {
1072 return callOp->emitOpError().append(
1076 return checkSelfTypeUnknownTarget(
1077 calleeAttr.getRootReference(), resTypes.front(), callOp,
"return"
1081 if (callOp->getNumResults() != 0) {
1083 return callOp->emitOpError()
1090 LogicalResult verifyTemplateParams()
override {
1092 return verifyNoTemplateInstantiations();
1095 LogicalResult verifyAffineMapParams()
override {
1100 return verifyNoAffineMapInstantiations();
1106 SymbolRefAttr calleeAttr;
1120 return emitOpError(
"requires a 'callee' symbol reference attribute");
1125 if (calleeAttr.getNestedReferences().size() == 1) {
1127 if (parent.hasConstNamed<
TemplateParamOp>(calleeAttr.getRootReference())) {
1130 return UnknownTargetVerifier(
this, tgtKind, calleeAttr).verify();
1132 return this->emitError(
"expected parameterized callee to target a struct function")
1144 if (failed(tgtOpt)) {
1146 << calleeAttr <<
'"';
1148 return KnownTargetVerifier(
this, std::move(*tgtOpt)).verify();
1152 return FunctionType::get(getContext(),
getArgOperands().getTypes(), getResultTypes());
1158 return unifications;
1166bool calleeIsStructFunctionImpl(
1167 const char *funcName, SymbolRefAttr callee, llvm::function_ref<
StructType()> getType
1169 if (callee.getLeafReference() == funcName) {
1203 return getResults().front();
1212 Operation *thisOp = this->getOperation();
1214 assert(succeeded(root));
1237 llvm::SmallVector<ValueRange, 4> output;
1238 output.reserve(input.size());
1239 for (OperandRange r : input) {
1240 output.push_back(r);
1246 FailureOr<SymbolLookupResult<FuncDefOp>> res =
1248 if (failed(res) || res->isManaged()) {
1256 SymbolTableCollection tables;