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

via flang-commits flang-commits at lists.llvm.org
Thu Oct 1 19:49:57 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));
----------------
MattPD wrote:

NEXT without STAT now works with an allocated enumeration array through FIR lowering and execution.

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


More information about the flang-commits mailing list