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);
299 if (index < this->getNumArguments()) {
300 DictionaryAttr res = function_interface_impl::getArgAttrDict(*
this, index);
301 return res ? res.contains(PublicAttr::name) :
false;
315 assert(index < getNumArguments() &&
"argument index out of range");
330 assert(index < getNumResults() &&
"result index out of range");
342 return emitErrorFunc() <<
'\'' <<
ARG_NAME_ATTR_NAME <<
"' is only valid on function arguments";
345 return emitErrorFunc() <<
'\'' <<
RES_NAME_ATTR_NAME <<
"' is only valid on function results";
348 if (failed(verifyArgOrResNameAttrs(
355 if (failed(verifyArgOrResNameAttrs(
367 for (Type t : type.getInputs()) {
372 return emitErrorFunc().append(
373 "\"@", getName(),
"\" parameters cannot contain affine map attributes but found ", t
377 for (Type t : type.getResults()) {
385 WalkResult res = this->walk<WalkOrder::PreOrder>([
this](ModuleOp nestedMod) {
387 "cannot be nested within '", getOperation()->getName(),
"' operations"
389 return WalkResult::interrupt();
391 if (res.wasInterrupted()) {
403 llvm::ArrayRef<Type> resTypes = funcType.getResults();
405 if (resTypes.size() != 1) {
406 return origin.emitOpError().append(
410 if (failed(
checkSelfType(tables, parent, resTypes.front(), origin,
"return"))) {
421verifyFuncTypeProduct(
FuncDefOp &origin, SymbolTableCollection &tables, StructDefOp &parent) {
423 return verifyFuncTypeCompute(origin, tables, parent);
427verifyFuncTypeConstrain(
FuncDefOp &origin, SymbolTableCollection &tables, StructDefOp &parent) {
430 if (funcType.getResults().size() != 0) {
431 return origin.emitOpError() <<
"\"@" <<
FUNC_NAME_CONSTRAIN <<
"\" must have no return type";
435 llvm::ArrayRef<Type> inputTypes = funcType.getInputs();
436 if (inputTypes.size() < 1) {
438 <<
"\" must have at least one input type";
440 if (failed(
checkSelfType(tables, parent, inputTypes.front(), origin,
"first input"))) {
458 return verifyFuncTypeCompute(*
this, tables, parentStructOpt);
460 return verifyFuncTypeConstrain(*
this, tables, parentStructOpt);
462 return verifyFuncTypeProduct(*
this, tables, parentStructOpt);
477 assert(!body.empty() &&
"compute() function body is empty");
478 Block &block = body.back();
481 Operation *terminator = block.getTerminator();
482 assert(terminator &&
"compute() function has no terminator");
483 auto retOp = llvm::dyn_cast<ReturnOp>(terminator);
486 << terminator->getName() <<
"'\n";
487 llvm_unreachable(
"compute() function must end with ReturnOp");
489 return retOp.getOperands().front();
494 return getArguments().front();
507 auto function = getParentOp<FuncDefOp>();
510 const auto results =
function.getFunctionType().getResults();
511 if (getNumOperands() != results.size()) {
512 return emitOpError(
"has ") << getNumOperands() <<
" operands, but enclosing function (@"
513 <<
function.getName() <<
") returns " << results.size();
516 for (
unsigned i = 0, e = results.size(); i != e; ++i) {
517 if (!
typesUnify(getOperand(i).getType(), results[i])) {
518 return emitError() <<
"type of return operand " << i <<
" (" << getOperand(i).getType()
519 <<
") doesn't match function result type (" << results[i] <<
")"
520 <<
" in function @" <<
function.getName();
534 auto &prop = state.getOrAddProperties<
Properties>();
535 if (failed(reader.readAttribute(prop.callee)) ||
536 failed(reader.readAttribute(prop.mapOpGroupSizes)) ||
537 failed(reader.readOptionalAttribute(prop.numDimsPerMap))) {
541 if (reader.getBytecodeVersion() < 6) {
542 auto &propStorage = prop.operandSegmentSizes;
543 DenseI32ArrayAttr attr;
544 if (failed(reader.readAttribute(attr))) {
547 if (attr.size() >
static_cast<int64_t
>(
sizeof(propStorage) /
sizeof(int32_t))) {
548 reader.emitError(
"size mismatch for operand/result_segment_size");
551 llvm::copy(ArrayRef<int32_t>(attr), propStorage.begin());
556 if (succeeded(versionOpt)) {
558 if (ver.majorVersion >= 2) {
559 if (failed(reader.readOptionalAttribute(prop.templateParams))) {
565 if (reader.getBytecodeVersion() >= 6) {
566 return reader.readSparseArray(MutableArrayRef(prop.operandSegmentSizes));
573 auto &prop = getProperties();
574 writer.writeAttribute(prop.callee);
575 writer.writeAttribute(prop.mapOpGroupSizes);
576 writer.writeOptionalAttribute(prop.numDimsPerMap);
578 if (writer.getBytecodeVersion() < 6) {
579 auto &propStorage = prop.operandSegmentSizes;
580 writer.writeAttribute(DenseI32ArrayAttr::get(this->getContext(), propStorage));
583 writer.writeOptionalAttribute(prop.templateParams);
585 auto &propStorage = prop.operandSegmentSizes;
586 if (writer.getBytecodeVersion() >= 6) {
587 writer.writeSparseArray(ArrayRef(propStorage));
592 OpBuilder &odsBuilder, OperationState &odsState, TypeRange resultTypes, SymbolRefAttr callee,
593 ValueRange argOperands, ArrayRef<Attribute> templateParams
595 odsState.addTypes(resultTypes);
596 odsState.addOperands(argOperands);
600 props.setCallee(callee);
605 OpBuilder &odsBuilder, OperationState &odsState, TypeRange resultTypes, SymbolRefAttr callee,
606 ArrayRef<ValueRange> mapOperands, DenseI32ArrayAttr numDimsPerMap, ValueRange argOperands,
607 ArrayRef<Attribute> templateParams
609 odsState.addTypes(resultTypes);
610 odsState.addOperands(argOperands);
612 odsBuilder, odsState, mapOperands, numDimsPerMap,
615 props.setCallee(callee);
623 if (
auto intAttr = llvm::dyn_cast<IntegerAttr>(paramFromCallOp)) {
625 std::optional<Type> declaredType = targetParam.
getTypeOpt();
626 if (!declaredType || !llvm::isa<TypeVarType>(*declaredType)) {
627 auto diag = this->emitOpError().append(
628 "wildcard `?` can only be used for template parameters with `!poly.tvar` "
629 "type restriction, but parameter \"@",
630 targetParam.getName(),
"\" has "
633 diag.append(
"type restriction ", *declaredType);
635 diag.append(
"no type restriction");
642 if (std::optional<Type> declaredType = targetParam.
getTypeOpt()) {
643 bool compatible =
false;
644 if (
auto sym = llvm::dyn_cast<SymbolRefAttr>(paramFromCallOp)) {
645 if (sym.getNestedReferences().empty()) {
646 SymbolTableCollection tables;
648 if (failed(parentTemplate)) {
651 if (TemplateOp p = *parentTemplate) {
652 auto binding = p.getConstNamed<TemplateSymbolBindingOpInterface>(sym.getRootReference());
656 if (std::optional<Type> actualType = binding.getTypeOpt()) {
657 compatible =
typesUnify(*actualType, *declaredType);
664 }
else if (llvm::isa<TypeVarType>(*declaredType)) {
665 compatible = llvm::isa<TypeAttr>(paramFromCallOp);
666 }
else if (llvm::isa<FeltType>(*declaredType)) {
667 compatible = llvm::isa<FeltConstAttr, IntegerAttr>(paramFromCallOp) &&
669 }
else if (llvm::isa<IndexType, IntegerType>(*declaredType)) {
673 compatible = llvm::isa<IntegerAttr>(paramFromCallOp) &&
677 llvm_unreachable(
"inconsistent with `isValidConstReadType()`");
681 return this->emitOpError().append(
682 "instantiation value '", paramFromCallOp,
"' is not compatible with parameter \"@",
683 targetParam.getName(),
"\" type restriction ", *declaredType
691 llvm::iterator_range<Region::op_iterator<TemplateParamOp>> targetParamDefs
695 assert((callParams.size() == llvm::range_size(targetParamDefs)) &&
"pre-condition");
697 for (
auto [paramOp, attr] : llvm::zip_equal(targetParamDefs, callParams.getValue())) {
706 llvm::iterator_range<Region::op_iterator<TemplateParamOp>> targetParamDefs,
711 assert((callParams.size() == llvm::range_size(targetParamDefs)) &&
"pre-condition");
713 for (
auto [paramOp, attr] : llvm::zip_equal(targetParamDefs, callParams.getValue())) {
715 if (
auto intAttr = llvm::dyn_cast<IntegerAttr>(attr)) {
720 auto it = unifications.find({FlatSymbolRefAttr::get(paramOp.getNameAttr()),
Side::RHS});
721 if (it != unifications.end() && !
typeParamsUnify({attr}, {it->second})) {
723 return this->emitOpError().append(
724 "template instantiation value '", attr,
"' for parameter \"@", paramOp.getName(),
725 "\" conflicts with value '", it->second,
"' inferred from function type signature"
734struct CallOpVerifier {
735 CallOpVerifier(
CallOp *c,
FunctionKind tgtFuncKind) : callOp(c), tgtKind(tgtFuncKind) {}
736 CallOpVerifier(
CallOp *c, StringRef tgtName) : CallOpVerifier(c,
fnNameToKind(tgtName)) {}
737 virtual ~CallOpVerifier() =
default;
739 LogicalResult verify() {
742 LogicalResult aggregateResult = success();
743 if (failed(verifyTargetAttributes())) {
744 aggregateResult = failure();
746 if (failed(verifyInputs())) {
747 aggregateResult = failure();
749 if (failed(verifyOutputs())) {
750 aggregateResult = failure();
752 if (failed(verifyTemplateParams())) {
753 aggregateResult = failure();
755 if (failed(verifyAffineMapParams())) {
756 aggregateResult = failure();
758 return aggregateResult;
765 virtual LogicalResult verifyTargetAttributes() = 0;
766 virtual LogicalResult verifyInputs() = 0;
767 virtual LogicalResult verifyOutputs() = 0;
768 virtual LogicalResult verifyTemplateParams() = 0;
769 virtual LogicalResult verifyAffineMapParams() = 0;
772 LogicalResult verifyTargetAttributesMatch(FuncDefOp target) {
773 LogicalResult aggregateRes = success();
774 if (FuncDefOp caller = (*callOp)->getParentOfType<FuncDefOp>()) {
775 auto emitAttrErr = [&](StringLiteral attrName) {
776 aggregateRes = callOp->emitOpError()
777 <<
"target '@" << target.getName() <<
"' has '" << attrName
778 <<
"' attribute, which is not specified by the caller '@" << caller.getName()
783 emitAttrErr(AllowConstraintAttr::name);
786 emitAttrErr(AllowWitnessAttr::name);
789 emitAttrErr(AllowNonNativeFieldOpsAttr::name);
795 LogicalResult verifyNoTemplateInstantiations() {
798 return callOp->emitOpError().append(
799 "can only have template instantiations when targeting a templated free function"
805 LogicalResult verifyNoAffineMapInstantiations() {
808 return callOp->emitOpError().append(
809 "can only have affine map instantiations when targeting a \"@",
FUNC_NAME_COMPUTE,
815 assert(callOp->getMapOperands().empty());
820struct KnownTargetVerifier :
public CallOpVerifier {
821 KnownTargetVerifier(CallOp *c, SymbolLookupResult<FuncDefOp> &&tgtRes)
822 : CallOpVerifier(c, tgtRes.get().getSymName()), tgt(*tgtRes), tgtType(tgt.getFunctionType()),
823 includeSymNames(tgtRes.getNamespace()) {}
825 LogicalResult verifyTargetAttributes()
override {
826 return CallOpVerifier::verifyTargetAttributesMatch(tgt);
829 LogicalResult verifyInputs()
override {
830 return verifyTypesMatch(callOp->
getArgOperands().getTypes(), tgtType.getInputs(),
"operand");
833 LogicalResult verifyOutputs()
override {
834 return verifyTypesMatch(callOp->getResultTypes(), tgtType.getResults(),
"result");
837 LogicalResult verifyTemplateParams()
override {
838 Operation *tgtOp = tgt.getOperation();
841 return verifyNoTemplateInstantiations();
851 auto realParams = tgtOpParent.getConstOps<TemplateParamOp>();
856 llvm::SmallDenseSet<SymbolRefAttr> referencedInSignature;
860 bool allParamsReferenced = llvm::all_of(realParams, [&](TemplateParamOp p) {
861 return referencedInSignature.contains(FlatSymbolRefAttr::get(p.getNameAttr()));
863 if (allParamsReferenced) {
867 return callOp->emitOpError().append(
868 "must provide template instantiation parameters when calling \"@", tgt.getSymName(),
869 "\" because not all template parameters of \"@", tgtOpParent.getSymName(),
870 "\" appear in the function type signature"
876 return llzk::InFlightDiagnosticWrapper(this->callOp->emitOpError());
882 size_t numTemplateParams = llvm::range_size(realParams);
883 if (callParams.size() != numTemplateParams) {
885 return callOp->emitOpError().append(
886 "template instantiation has ", callParams.size(),
" parameter(s) but \"@",
887 tgtOpParent.getSymName(),
"\" expects ", numTemplateParams,
" template parameter(s)"
899 assert(succeeded(unifyResult) &&
"already checked by `verifyInputs()` and `verifyOutputs()`");
903 return verifyNoTemplateInstantiations();
907 LogicalResult verifyAffineMapParams()
override {
915 if (ArrayAttr params = retTy.getParams()) {
917 SmallVector<AffineMapAttr> mapAttrs;
918 for (Attribute a : params) {
919 if (AffineMapAttr m = dyn_cast<AffineMapAttr>(a)) {
920 mapAttrs.push_back(m);
931 return verifyNoAffineMapInstantiations();
936 template <
typename T>
938 verifyTypesMatch(ValueTypeRange<T> callOpTypes, ArrayRef<Type> tgtTypes,
const char *aspect) {
939 if (tgtTypes.size() != callOpTypes.size()) {
940 return callOp->emitOpError()
941 .append(
"incorrect number of ", aspect,
"s for callee, expected ", tgtTypes.size())
942 .attachNote(tgt.getLoc())
943 .append(
"callee defined here");
945 for (
unsigned i = 0, e = tgtTypes.size(); i != e; ++i) {
946 if (!
typesUnify(callOpTypes[i], tgtTypes[i], includeSymNames)) {
947 return callOp->emitOpError().append(
948 aspect,
" type mismatch: expected type ", tgtTypes[i],
", but found ", callOpTypes[i],
949 " for ", aspect,
" number ", i
957 FunctionType tgtType;
958 std::vector<llvm::StringRef> includeSymNames;
963LogicalResult checkSelfTypeUnknownTarget(
964 StringAttr expectedParamName, Type actualType,
CallOp *origin,
const char *aspect
966 if (!llvm::isa<TypeVarType>(actualType) ||
967 llvm::cast<TypeVarType>(actualType).getRefName() != expectedParamName) {
973 return origin->emitOpError().append(
974 "target \"@", origin->
getCallee().getLeafReference().getValue(),
"\" expected ", aspect,
975 " type '!",
TypeVarType::name,
"<@", expectedParamName.getValue(),
">' but found ",
991struct UnknownTargetVerifier :
public CallOpVerifier {
992 UnknownTargetVerifier(CallOp *c,
FunctionKind tgtFuncKind, SymbolRefAttr callee)
993 : CallOpVerifier(c, tgtFuncKind), calleeAttr(callee) {
1000 LogicalResult verifyTargetAttributes()
override {
1003 LogicalResult aggregateRes = success();
1004 if (FuncDefOp caller = (*callOp)->getParentOfType<FuncDefOp>()) {
1005 auto emitAttrErr = [&](StringLiteral attrName) {
1006 aggregateRes = callOp->emitOpError()
1007 <<
"target '" << calleeAttr <<
"' has '" << attrName
1008 <<
"' attribute, which is not specified by the caller '@" << caller.getName()
1014 if (!caller.hasAllowConstraintAttr()) {
1015 emitAttrErr(AllowConstraintAttr::name);
1019 if (!caller.hasAllowWitnessAttr()) {
1020 emitAttrErr(AllowWitnessAttr::name);
1024 if (!caller.hasAllowWitnessAttr()) {
1025 emitAttrErr(AllowWitnessAttr::name);
1027 if (!caller.hasAllowConstraintAttr()) {
1028 emitAttrErr(AllowConstraintAttr::name);
1035 return aggregateRes;
1038 LogicalResult verifyInputs()
override {
1044 Operation::operand_type_range inputTypes = callOp->
getArgOperands().getTypes();
1045 if (inputTypes.size() < 1) {
1047 return callOp->emitOpError()
1050 return checkSelfTypeUnknownTarget(
1051 calleeAttr.getRootReference(), inputTypes.front(), callOp,
"first input"
1057 LogicalResult verifyOutputs()
override {
1061 Operation::result_type_range resTypes = callOp->getResultTypes();
1062 if (resTypes.size() != 1) {
1064 return callOp->emitOpError().append(
1068 return checkSelfTypeUnknownTarget(
1069 calleeAttr.getRootReference(), resTypes.front(), callOp,
"return"
1073 if (callOp->getNumResults() != 0) {
1075 return callOp->emitOpError()
1082 LogicalResult verifyTemplateParams()
override {
1084 return verifyNoTemplateInstantiations();
1087 LogicalResult verifyAffineMapParams()
override {
1092 return verifyNoAffineMapInstantiations();
1098 SymbolRefAttr calleeAttr;
1112 return emitOpError(
"requires a 'callee' symbol reference attribute");
1117 if (calleeAttr.getNestedReferences().size() == 1) {
1119 if (parent.hasConstNamed<
TemplateParamOp>(calleeAttr.getRootReference())) {
1122 return UnknownTargetVerifier(
this, tgtKind, calleeAttr).verify();
1124 return this->emitError(
"expected parameterized callee to target a struct function")
1136 if (failed(tgtOpt)) {
1138 << calleeAttr <<
'"';
1140 return KnownTargetVerifier(
this, std::move(*tgtOpt)).verify();
1144 return FunctionType::get(getContext(),
getArgOperands().getTypes(), getResultTypes());
1150 return unifications;
1158bool calleeIsStructFunctionImpl(
1159 const char *funcName, SymbolRefAttr callee, llvm::function_ref<
StructType()> getType
1161 if (callee.getLeafReference() == funcName) {
1195 return getResults().front();
1204 Operation *thisOp = this->getOperation();
1206 assert(succeeded(root));
1229 llvm::SmallVector<ValueRange, 4> output;
1230 output.reserve(input.size());
1231 for (OperandRange r : input) {
1232 output.push_back(r);
1238 FailureOr<SymbolLookupResult<FuncDefOp>> res =
1240 if (failed(res) || res->isManaged()) {
1248 SymbolTableCollection tables;