[flang-commits] [flang] [flang] Enumeration Type: (PR 4/5) Lowering (PR #193571)

via flang-commits flang-commits at lists.llvm.org
Tue Oct 6 08:54:03 PDT 2026


================
@@ -1727,12 +1728,316 @@ class HlfirBuilder {
   gen(const Fortran::evaluate::FunctionRef<T> &expr) {
     mlir::Type resType =
         Fortran::lower::TypeBuilder<T>::genType(getConverter(), expr);
+
+    // Intercept enumeration-type intrinsics (NEXT, PREVIOUS, HUGE) that return
+    // SomeDerived but lower to i32 operations.
+    if constexpr (std::is_same_v<T, Fortran::evaluate::SomeDerived>) {
+      if (const auto *intrinsic = expr.proc().GetSpecificIntrinsic()) {
+        if (intrinsic->name == "next") {
+          return genEnumerationNext(expr, resType);
+        }
+        if (intrinsic->name == "previous") {
+          return genEnumerationPrevious(expr, resType);
+        }
+        if (intrinsic->name == "huge") {
+          return genEnumerationHuge(expr, resType);
+        }
+      }
+    }
+
     auto result = Fortran::lower::convertCallToHLFIR(
         getLoc(), getConverter(), expr, resType, getSymMap(), getStmtCtx());
     assert(result.has_value());
     return *result;
   }
 
+  // Helper to extract enumeration type info from a NEXT/PREVIOUS/HUGE
+  // intrinsic's result type.
+  std::pair<const Fortran::semantics::DerivedTypeSpec *, int>
+  getEnumerationTypeInfo(
+      const Fortran::evaluate::FunctionRef<Fortran::evaluate::SomeDerived>
+          &expr,
+      llvm::StringRef name) {
+    auto resultDynType = expr.proc().GetType();
+    assert(resultDynType && (name + " must have a result type").str().c_str());
+    const auto *derived = Fortran::evaluate::GetDerivedTypeSpec(*resultDynType);
+    assert(derived &&
+           derived->typeSymbol()
+               .detailsIf<Fortran::semantics::DerivedTypeDetails>() &&
+           derived->typeSymbol()
+               .detailsIf<Fortran::semantics::DerivedTypeDetails>()
+               ->isEnumerationType() &&
+           (name + " result must be enumeration type").str().c_str());
+    int count = derived->typeSymbol()
+                    .GetUltimate()
+                    .get<Fortran::semantics::DerivedTypeDetails>()
+                    .enumeratorCount();
+    return {derived, count};
+  }
+
+  // Return the syntactically supplied STAT expression of NEXT/PREVIOUS, or
+  // nullptr if none was written.
+  const Fortran::lower::SomeExpr *getEnumerationStatExpr(
+      const Fortran::evaluate::FunctionRef<Fortran::evaluate::SomeDerived>
+          &expr) {
+    if (expr.arguments().size() < 2 || !expr.arguments()[1])
+      return nullptr;
+    const auto *statExpr = expr.arguments()[1]->UnwrapExpr();
+    assert(statExpr && "STAT argument must be an expression");
+    return statExpr;
+  }
+
+  // Return an i1 telling whether STAT is present at runtime, or a null value
+  // if it is always present. An absent optional dummy or an unallocated or
+  // disassociated allocatable/pointer actual makes STAT not present.
+  mlir::Value
+  genEnumerationStatIsPresent(const Fortran::lower::SomeExpr &statExpr,
+                              hlfir::Entity stat) {
+    if (!Fortran::evaluate::MayBePassedAsAbsentOptional(statExpr))
+      return {};
+    mlir::Location loc = getLoc();
+    fir::FirOpBuilder &builder = getBuilder();
+    if (Fortran::evaluate::IsAllocatableOrPointerObject(statExpr))
+      return builder.genIsNotNullAddr(
+          loc, hlfir::genVariableRawAddress(loc, builder, stat));
+    return fir::IsPresentOp::create(builder, loc, builder.getI1Type(), stat)
+        .getResult();
+  }
+
+  // Emit genPresent() when STAT is present at runtime and genAbsent()
+  // otherwise. A null isPresent means STAT is always present.
+  template <typename PresentFn, typename AbsentFn>
+  void genIfEnumerationStatPresent(mlir::Value isPresent, PresentFn genPresent,
+                                   AbsentFn genAbsent) {
+    if (!isPresent) {
+      genPresent();
+      return;
+    }
+    fir::FirOpBuilder &builder = getBuilder();
+    auto ifOp = fir::IfOp::create(builder, getLoc(), {}, isPresent,
+                                  /*withElseRegion=*/true);
+    builder.setInsertionPointToStart(&ifOp.getThenRegion().front());
+    genPresent();
+    builder.setInsertionPointToStart(&ifOp.getElseRegion().front());
+    genAbsent();
+    builder.setInsertionPointAfter(ifOp);
+  }
+
+  // Error termination if cond (i1) is true; used when STAT is not present.
+  void genEnumerationBoundaryFatal(mlir::Value cond) {
+    mlir::Location loc = getLoc();
+    fir::FirOpBuilder &builder = getBuilder();
+    auto ifOp = fir::IfOp::create(builder, loc, {}, cond,
+                                  /*withElseRegion=*/false);
+    builder.setInsertionPointToStart(&ifOp.getThenRegion().front());
+    fir::runtime::genReportFatalUserError(
+        builder, loc,
+        "NEXT or PREVIOUS of enumeration type at boundary without STAT=");
+    builder.setInsertionPointAfter(ifOp);
+  }
+
+  // Helper to lower STAT argument handling for NEXT/PREVIOUS.
+  // atBoundary is a boolean indicating whether the boundary condition was hit.
+  void genEnumerationStatHandling(
+      const Fortran::evaluate::FunctionRef<Fortran::evaluate::SomeDerived>
+          &expr,
+      mlir::Value atBoundary, mlir::Type resType) {
+    mlir::Location loc = getLoc();
+    fir::FirOpBuilder &builder = getBuilder();
+    const Fortran::lower::SomeExpr *statExpr = getEnumerationStatExpr(expr);
+    if (!statExpr) {
+      genEnumerationBoundaryFatal(atBoundary);
+      return;
+    }
+    hlfir::Entity statAddr = Fortran::lower::convertExprToHLFIR(
+        loc, converter, *statExpr, getSymMap(), getStmtCtx());
+    auto genStatAssign = [&]() {
+      // Assign 0 or FORTRAN_RUNTIME_STAT_ENUM_BOUNDARY (112).
+      mlir::Type statType = statAddr.getFortranElementType();
+      mlir::Value boundaryConst = builder.createIntegerConstant(
+          loc, statType, 112 /* FORTRAN_RUNTIME_STAT_ENUM_BOUNDARY */);
+      mlir::Value zeroConst = builder.createIntegerConstant(loc, statType, 0);
+      mlir::Value statVal = mlir::arith::SelectOp::create(
+          builder, loc, atBoundary, boundaryConst, zeroConst);
+      hlfir::AssignOp::create(builder, loc, statVal, statAddr);
+    };
+    genIfEnumerationStatPresent(
+        genEnumerationStatIsPresent(*statExpr, statAddr), genStatAssign,
+        [&]() { genEnumerationBoundaryFatal(atBoundary); });
+  }
+
+  // Compute the per-element NEXT/PREVIOUS result and boundary flag from a
+  // scalar ordinal.  isNext selects NEXT (min(ordinal+1, count)) versus
+  // PREVIOUS (max(ordinal-1, 1)); the returned atBoundary is an i1.
+  std::pair<mlir::Value, mlir::Value>
+  genEnumOrdinalStep(fir::FirOpBuilder &builder, mlir::Location loc,
+                     mlir::Value ordinal, mlir::Type resType, int count,
+                     bool isNext) {
+    mlir::Value one = builder.createIntegerConstant(loc, resType, 1);
+    if (isNext) {
+      mlir::Value maxVal = builder.createIntegerConstant(loc, resType, count);
+      mlir::Value incremented =
+          mlir::arith::AddIOp::create(builder, loc, ordinal, one);
+      mlir::Value cmp = mlir::arith::CmpIOp::create(
+          builder, loc, mlir::arith::CmpIPredicate::sle, incremented, maxVal);
+      mlir::Value result =
+          mlir::arith::SelectOp::create(builder, loc, cmp, incremented, maxVal);
+      mlir::Value atBoundary = mlir::arith::CmpIOp::create(
+          builder, loc, mlir::arith::CmpIPredicate::eq, ordinal, maxVal);
+      return {result, atBoundary};
+    }
+    mlir::Value decremented =
+        mlir::arith::SubIOp::create(builder, loc, ordinal, one);
+    mlir::Value cmp = mlir::arith::CmpIOp::create(
+        builder, loc, mlir::arith::CmpIPredicate::sge, decremented, one);
+    mlir::Value result =
+        mlir::arith::SelectOp::create(builder, loc, cmp, decremented, one);
+    mlir::Value atBoundary = mlir::arith::CmpIOp::create(
+        builder, loc, mlir::arith::CmpIPredicate::eq, ordinal, one);
+    return {result, atBoundary};
+  }
+
+  // Lower NEXT/PREVIOUS applied to a whole-array (elemental) enumeration
+  // argument.  The result is a pure hlfir.elemental; STAT/error-termination
+  // side effects are handled after it, per element.
+  hlfir::EntityWithAttributes genEnumerationArray(
+      const Fortran::evaluate::FunctionRef<Fortran::evaluate::SomeDerived>
+          &expr,
+      hlfir::Entity arg, mlir::Type resType, int count, bool isNext) {
+    mlir::Location loc = getLoc();
+    fir::FirOpBuilder &builder = getBuilder();
+    mlir::Value shape = hlfir::genShape(loc, builder, arg);
+    // resType is the whole array type here; the ordinal arithmetic and result
+    // element type need the scalar (i32) element type.
+    mlir::Type eleTy = hlfir::getFortranElementType(resType);
+
+    auto resultKernel = [&](mlir::Location l, fir::FirOpBuilder &b,
+                            mlir::ValueRange idx) -> hlfir::Entity {
+      mlir::Value ordinal =
+          hlfir::loadTrivialScalar(l, b, hlfir::getElementAt(l, b, arg, idx));
+      auto [result, atBoundary] =
+          genEnumOrdinalStep(b, l, ordinal, eleTy, count, isNext);
+      (void)atBoundary;
+      return hlfir::Entity{result};
+    };
+    mlir::Value resultElem =
+        hlfir::genElementalOp(loc, builder, eleTy, shape, /*typeParams=*/{},
+                              resultKernel, /*isUnordered=*/true);
+    fir::FirOpBuilder *bldr = &builder;
+    getStmtCtx().attachCleanup(
+        [=]() { hlfir::DestroyOp::create(*bldr, loc, resultElem); });
+
+    // STAT not present: error termination if any element is at a boundary.
+    auto genBoundaryFatal = [&]() {
+      mlir::Type logType = fir::LogicalType::get(builder.getContext(), 4);
+      auto maskKernel = [&](mlir::Location l, fir::FirOpBuilder &b,
+                            mlir::ValueRange idx) -> hlfir::Entity {
+        mlir::Value ordinal =
+            hlfir::loadTrivialScalar(l, b, hlfir::getElementAt(l, b, arg, idx));
+        auto [result, atBoundary] =
+            genEnumOrdinalStep(b, l, ordinal, eleTy, count, isNext);
+        (void)result;
+        return hlfir::Entity{b.createConvert(l, logType, atBoundary)};
+      };
+      mlir::Value mask =
+          hlfir::genElementalOp(loc, builder, logType, shape, /*typeParams=*/{},
+                                maskKernel, /*isUnordered=*/true);
+      mlir::Value anyBoundary =
+          hlfir::AnyOp::create(builder, loc, logType, mask,
+                               /*dim=*/mlir::Value{});
+      genEnumerationBoundaryFatal(
+          builder.createConvert(loc, builder.getI1Type(), anyBoundary));
+      hlfir::DestroyOp::create(builder, loc, mask);
+    };
+
+    const Fortran::lower::SomeExpr *statExpr = getEnumerationStatExpr(expr);
+    if (!statExpr) {
+      genBoundaryFatal();
+      return hlfir::EntityWithAttributes{resultElem};
+    }
+    hlfir::Entity statEntity = Fortran::lower::convertExprToHLFIR(
+        loc, converter, *statExpr, getSymMap(), getStmtCtx());
+    // STAT present: elementwise 0/112 into the conformable STAT array. The
+    // temporary is destroyed inline since it may live inside a fir.if region.
+    auto genStatAssign = [&]() {
+      mlir::Type statType = statEntity.getFortranElementType();
+      auto statKernel = [&](mlir::Location l, fir::FirOpBuilder &b,
+                            mlir::ValueRange idx) -> hlfir::Entity {
+        mlir::Value ordinal =
+            hlfir::loadTrivialScalar(l, b, hlfir::getElementAt(l, b, arg, idx));
+        auto [result, atBoundary] =
+            genEnumOrdinalStep(b, l, ordinal, eleTy, count, isNext);
+        (void)result;
+        mlir::Value boundaryConst = b.createIntegerConstant(l, statType, 112);
+        mlir::Value zeroConst = b.createIntegerConstant(l, statType, 0);
+        return hlfir::Entity{mlir::arith::SelectOp::create(
+            b, l, atBoundary, boundaryConst, zeroConst)};
+      };
+      mlir::Value statElem = hlfir::genElementalOp(
+          loc, builder, statType, shape, /*typeParams=*/{}, statKernel,
+          /*isUnordered=*/true);
+      hlfir::AssignOp::create(builder, loc, statElem, statEntity);
----------------
kwyatt-ext wrote:

Resolved through conversation to add this note to the code:
"NOTE: STAT is listed as "scalar", but also as INTENT(OUT) and is in an elemental function.  Taking them together, this means that, by F2023 15.9.1 ΒΆ4, it should be a conforming argument to A."

https://github.com/llvm/llvm-project/pull/193571


More information about the flang-commits mailing list