[flang-commits] [flang] [flang][cuda][semantics] Reject unguarded host-only and device-only calls in CUDA HOST, DEVICE procedures (PR #228176)

via flang-commits flang-commits at lists.llvm.org
Thu Oct 1 11:03:49 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-flang-semantics

Author: Andre Kuhlenschmidt (akuhlens)

<details>
<summary>Changes</summary>

 Previously, semantic checking allowed these calls throughout a HOST,DEVICE
  procedure, although each unguarded call must be valid in both its host and device
  versions. The change requires unguarded callees to support both targets.

  An ON_DEVICE() condition permits host-only and device-only calls in either branch
  arm, including nested branches, one-line conditions, and calls through an
  intrinsic alias. Determining whether a call appears in the correct arm is deferred
  to a separate extension.

  Regression tests cover function and subroutine calls, guarded and unguarded cases,
  and user-defined procedures named on_device that must not enable the allowance.
  Build and full Flang checks passed with no unexpected failures.

---

Patch is 26.01 KiB, truncated to 20.00 KiB below, full version: https://github.com/llvm/llvm-project/pull/228176.diff


2 Files Affected:

- (modified) flang/lib/Semantics/check-cuda.cpp (+127-61) 
- (added) flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf (+335) 


``````````diff
diff --git a/flang/lib/Semantics/check-cuda.cpp b/flang/lib/Semantics/check-cuda.cpp
index 8922e0eb559e8..13545803428ee 100644
--- a/flang/lib/Semantics/check-cuda.cpp
+++ b/flang/lib/Semantics/check-cuda.cpp
@@ -67,6 +67,37 @@ static const llvm::StringSet<> warpFunctions_ = {"match_all_syncjj",
     "match_any_syncjj", "match_any_syncjx", "match_any_syncjf",
     "match_any_syncjd"};
 
+enum class CallContext { Device, HostDevice, GuardedHostDevice };
+
+// Match the BIND(C) intrinsic-module procedure intercepted by CUDA lowering,
+// including calls through a USE rename. A user procedure named ON_DEVICE is
+// an ordinary call and does not allow host-only or device-only callees.
+static bool IsOnDevice(const evaluate::ProcedureDesignator &proc) {
+  const Symbol *sym{proc.GetSymbol()};
+  if (!sym) {
+    return false;
+  }
+  const Symbol &ultimate{sym->GetUltimate()};
+  const Symbol *module{ultimate.owner().GetSymbol()};
+  return ultimate.name() == "on_device" && IsBindCProcedure(ultimate) &&
+      module && module->attrs().test(Attr::INTRINSIC);
+}
+
+struct FindOnDevice : public evaluate::AnyTraverse<FindOnDevice, bool> {
+  using Base = evaluate::AnyTraverse<FindOnDevice, bool>;
+  FindOnDevice() : Base(*this) {}
+  using Base::operator();
+  bool operator()(const evaluate::ProcedureDesignator &proc) const {
+    return IsOnDevice(proc);
+  }
+};
+
+static bool ChecksOnDevice(
+    SemanticsContext &context, const parser::ScalarLogicalExpr &condition) {
+  const auto *expr{GetExpr(context, condition)};
+  return expr && FindOnDevice{}(*expr);
+}
+
 // Traverses an evaluate::Expr<> in search of unsupported operations
 // on the device.
 
@@ -74,10 +105,14 @@ struct DeviceExprChecker
     : public evaluate::AnyTraverse<DeviceExprChecker, MaybeMsg> {
   using Result = MaybeMsg;
   using Base = evaluate::AnyTraverse<DeviceExprChecker, Result>;
-  explicit DeviceExprChecker(SemanticsContext &c, bool allowHostCallees = false)
-      : Base(*this), context_{c}, allowHostCallees_{allowHostCallees} {}
+  explicit DeviceExprChecker(
+      SemanticsContext &c, CallContext callContext = CallContext::Device)
+      : Base(*this), context_{c}, callContext_{callContext} {}
   using Base::operator();
   Result operator()(const evaluate::ProcedureDesignator &x) const {
+    if (IsOnDevice(x)) {
+      return {};
+    }
     if (const Symbol * sym{x.GetInterfaceSymbol()}) {
       const Symbol &ultimate{sym->GetUltimate()};
       const auto *subp{ultimate.detailsIf<semantics::SubprogramDetails>()};
@@ -95,6 +130,11 @@ struct DeviceExprChecker
               return parser::MessageFormattedText(
                   "warp match function disabled"_err_en_US);
             }
+            if (*attrs == common::CUDASubprogramAttrs::Device &&
+                callContext_ == CallContext::HostDevice) {
+              return parser::MessageFormattedText(
+                  "'%s' may not be called in host code"_err_en_US, x.GetName());
+            }
             return {};
           }
           if (*attrs == common::CUDASubprogramAttrs::Global) {
@@ -119,10 +159,7 @@ struct DeviceExprChecker
       return {};
     }
 
-    // A host,device subprogram is compiled for the host as well as the device,
-    // so a call to a host procedure (typically guarded at run time by a test
-    // such as ON_DEVICE()) is legitimate in its host compilation.
-    if (allowHostCallees_) {
+    if (callContext_ == CallContext::GuardedHostDevice) {
       return {};
     }
     return parser::MessageFormattedText(
@@ -130,7 +167,7 @@ struct DeviceExprChecker
   }
 
   SemanticsContext &context_;
-  bool allowHostCallees_{false};
+  CallContext callContext_{CallContext::Device};
 };
 
 static bool IsHostArray(const Symbol &symbol) {
@@ -208,20 +245,19 @@ struct FindHostArray
 };
 
 template <typename A>
-static MaybeMsg CheckUnwrappedExpr(
-    SemanticsContext &context, const A &x, bool allowHostCallees = false) {
+static MaybeMsg CheckUnwrappedExpr(SemanticsContext &context, const A &x,
+    CallContext callContext = CallContext::Device) {
   if (const auto *expr{parser::Unwrap<parser::Expr>(x)}) {
-    return DeviceExprChecker{context, allowHostCallees}(expr->typedExpr);
+    return DeviceExprChecker{context, callContext}(expr->typedExpr);
   }
   return {};
 }
 
 template <typename A>
 static void CheckUnwrappedExpr(SemanticsContext &context, SourceName at,
-    const A &x, bool allowHostCallees = false) {
+    const A &x, CallContext callContext = CallContext::Device) {
   if (const auto *expr{parser::Unwrap<parser::Expr>(x)}) {
-    if (auto msg{
-            DeviceExprChecker{context, allowHostCallees}(expr->typedExpr)}) {
+    if (auto msg{DeviceExprChecker{context, callContext}(expr->typedExpr)}) {
       context.Say(at, std::move(*msg));
     }
   }
@@ -229,16 +265,16 @@ static void CheckUnwrappedExpr(SemanticsContext &context, SourceName at,
 
 template <bool CUF_KERNEL> struct ActionStmtChecker {
   template <typename A>
-  static MaybeMsg WhyNotOk(
-      SemanticsContext &context, const A &x, bool allowHostCallees = false) {
+  static MaybeMsg WhyNotOk(SemanticsContext &context, const A &x,
+      CallContext callContext = CallContext::Device) {
     if constexpr (ConstraintTrait<A>) {
-      return WhyNotOk(context, x.thing, allowHostCallees);
+      return WhyNotOk(context, x.thing, callContext);
     } else if constexpr (WrapperTrait<A>) {
-      return WhyNotOk(context, x.v, allowHostCallees);
+      return WhyNotOk(context, x.v, callContext);
     } else if constexpr (UnionTrait<A>) {
-      return WhyNotOk(context, x.u, allowHostCallees);
+      return WhyNotOk(context, x.u, callContext);
     } else if constexpr (TupleTrait<A>) {
-      return WhyNotOk(context, x.t, allowHostCallees);
+      return WhyNotOk(context, x.t, callContext);
     } else {
       return parser::MessageFormattedText{
           "Statement may not appear in device code"_err_en_US};
@@ -246,33 +282,36 @@ template <bool CUF_KERNEL> struct ActionStmtChecker {
   }
   template <typename A>
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const common::Indirection<A> &x, bool allowHostCallees = false) {
-    return WhyNotOk(context, x.value(), allowHostCallees);
+      const common::Indirection<A> &x,
+      CallContext callContext = CallContext::Device) {
+    return WhyNotOk(context, x.value(), callContext);
   }
   template <typename... As>
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const std::variant<As...> &x, bool allowHostCallees = false) {
+      const std::variant<As...> &x,
+      CallContext callContext = CallContext::Device) {
     return common::visit(
-        [&context, allowHostCallees](
-            const auto &x) { return WhyNotOk(context, x, allowHostCallees); },
+        [&context, callContext](
+            const auto &x) { return WhyNotOk(context, x, callContext); },
         x);
   }
   template <std::size_t J = 0, typename... As>
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const std::tuple<As...> &x, bool allowHostCallees = false) {
+      const std::tuple<As...> &x,
+      CallContext callContext = CallContext::Device) {
     if constexpr (J == sizeof...(As)) {
       return {};
-    } else if (auto msg{WhyNotOk(context, std::get<J>(x), allowHostCallees)}) {
+    } else if (auto msg{WhyNotOk(context, std::get<J>(x), callContext)}) {
       return msg;
     } else {
-      return WhyNotOk<(J + 1)>(context, x, allowHostCallees);
+      return WhyNotOk<(J + 1)>(context, x, callContext);
     }
   }
   template <typename A>
   static MaybeMsg WhyNotOk(SemanticsContext &context, const std::list<A> &x,
-      bool allowHostCallees = false) {
+      CallContext callContext = CallContext::Device) {
     for (const auto &y : x) {
-      if (MaybeMsg result{WhyNotOk(context, y, allowHostCallees)}) {
+      if (MaybeMsg result{WhyNotOk(context, y, callContext)}) {
         return result;
       }
     }
@@ -280,76 +319,85 @@ template <bool CUF_KERNEL> struct ActionStmtChecker {
   }
   template <typename A>
   static MaybeMsg WhyNotOk(SemanticsContext &context, const std::optional<A> &x,
-      bool allowHostCallees = false) {
+      CallContext callContext = CallContext::Device) {
     if (x) {
-      return WhyNotOk(context, *x, allowHostCallees);
+      return WhyNotOk(context, *x, callContext);
     } else {
       return {};
     }
   }
   template <typename A>
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::UnlabeledStatement<A> &x, bool allowHostCallees = false) {
-    return WhyNotOk(context, x.statement, allowHostCallees);
+      const parser::UnlabeledStatement<A> &x,
+      CallContext callContext = CallContext::Device) {
+    return WhyNotOk(context, x.statement, callContext);
   }
   template <typename A>
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::Statement<A> &x, bool allowHostCallees = false) {
-    return WhyNotOk(context, x.statement, allowHostCallees);
+      const parser::Statement<A> &x,
+      CallContext callContext = CallContext::Device) {
+    return WhyNotOk(context, x.statement, callContext);
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::AllocateStmt &, bool allowHostCallees = false) {
+      const parser::AllocateStmt &,
+      CallContext callContext = CallContext::Device) {
     return {}; // AllocateObjects are checked elsewhere
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::AllocateCoarraySpec &, bool allowHostCallees = false) {
+      const parser::AllocateCoarraySpec &,
+      CallContext callContext = CallContext::Device) {
     return parser::MessageFormattedText(
         "A coarray may not be allocated on the device"_err_en_US);
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::DeallocateStmt &, bool allowHostCallees = false) {
+      const parser::DeallocateStmt &,
+      CallContext callContext = CallContext::Device) {
     return {}; // AllocateObjects are checked elsewhere
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::AssignmentStmt &x, bool allowHostCallees = false) {
-    return DeviceExprChecker{context, allowHostCallees}(x.typedAssignment);
+      const parser::AssignmentStmt &x,
+      CallContext callContext = CallContext::Device) {
+    return DeviceExprChecker{context, callContext}(x.typedAssignment);
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context, const parser::CallStmt &x,
-      bool allowHostCallees = false) {
-    return DeviceExprChecker{context, allowHostCallees}(x.typedCall);
+      CallContext callContext = CallContext::Device) {
+    return DeviceExprChecker{context, callContext}(x.typedCall);
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::ContinueStmt &, bool allowHostCallees = false) {
+      const parser::ContinueStmt &,
+      CallContext callContext = CallContext::Device) {
     return {};
   }
   static MaybeMsg WhyNotOk(SemanticsContext &, const parser::PauseStmt &,
-      bool allowHostCallees = false) {
+      CallContext callContext = CallContext::Device) {
     return parser::MessageFormattedText{
         "device subprograms may not contain PAUSE statements"_err_en_US};
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context, const parser::IfStmt &x,
-      bool allowHostCallees = false) {
-    if (auto result{CheckUnwrappedExpr(context,
-            std::get<parser::ScalarLogicalExpr>(x.t), allowHostCallees)}) {
+      CallContext callContext = CallContext::Device) {
+    if (auto result{CheckUnwrappedExpr(
+            context, std::get<parser::ScalarLogicalExpr>(x.t), callContext)}) {
       return result;
     }
     return WhyNotOk(context,
         std::get<parser::UnlabeledStatement<parser::ActionStmt>>(x.t).statement,
-        allowHostCallees);
+        callContext);
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::NullifyStmt &x, bool allowHostCallees = false) {
+      const parser::NullifyStmt &x,
+      CallContext callContext = CallContext::Device) {
     for (const auto &y : x.v) {
       if (MaybeMsg result{
-              DeviceExprChecker{context, allowHostCallees}(y.typedExpr)}) {
+              DeviceExprChecker{context, callContext}(y.typedExpr)}) {
         return result;
       }
     }
     return {};
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::PointerAssignmentStmt &x, bool allowHostCallees = false) {
-    return DeviceExprChecker{context, allowHostCallees}(x.typedAssignment);
+      const parser::PointerAssignmentStmt &x,
+      CallContext callContext = CallContext::Device) {
+    return DeviceExprChecker{context, callContext}(x.typedAssignment);
   }
 };
 
@@ -372,6 +420,8 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
         isHostDevice = subp->cudaSubprogramAttrs() &&
             subp->cudaSubprogramAttrs() ==
                 common::CUDASubprogramAttrs::HostDevice;
+        callContext_ =
+            isHostDevice ? CallContext::HostDevice : CallContext::Device;
         Check(body);
       }
     }
@@ -547,13 +597,13 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
                 ErrorIfHostSymbol(assign->rhs, source);
               }
               if (auto msg{ActionStmtChecker<IsCUFKernelDo>::WhyNotOk(
-                      context_, x, isHostDevice)}) {
+                      context_, x, callContext_)}) {
                 context_.Say(source, std::move(*msg));
               }
             },
             [&](const auto &x) {
               if (auto msg{ActionStmtChecker<IsCUFKernelDo>::WhyNotOk(
-                      context_, x, isHostDevice)}) {
+                      context_, x, callContext_)}) {
                 context_.Say(source, std::move(*msg));
               }
             },
@@ -561,28 +611,43 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
         stmt.u);
   }
   void Check(const parser::IfConstruct &ic) {
+    const CallContext incoming{callContext_};
     const auto &ifS{std::get<parser::Statement<parser::IfThenStmt>>(ic.t)};
-    CheckUnwrappedExpr(context_, ifS.source,
-        std::get<parser::ScalarLogicalExpr>(ifS.statement.t), isHostDevice);
+    const auto &condition{std::get<parser::ScalarLogicalExpr>(ifS.statement.t)};
+    CheckUnwrappedExpr(context_, ifS.source, condition, callContext_);
+    AllowGuardedCalls(condition);
     Check(std::get<parser::Block>(ic.t));
     for (const auto &eib :
         std::get<std::list<parser::IfConstruct::ElseIfBlock>>(ic.t)) {
       const auto &eIfS{std::get<parser::Statement<parser::ElseIfStmt>>(eib.t)};
-      CheckUnwrappedExpr(context_, eIfS.source,
-          std::get<parser::ScalarLogicalExpr>(eIfS.statement.t), isHostDevice);
+      const auto &elseIfCondition{
+          std::get<parser::ScalarLogicalExpr>(eIfS.statement.t)};
+      CheckUnwrappedExpr(context_, eIfS.source, elseIfCondition, callContext_);
+      AllowGuardedCalls(elseIfCondition);
       Check(std::get<parser::Block>(eib.t));
     }
     if (const auto &eb{
             std::get<std::optional<parser::IfConstruct::ElseBlock>>(ic.t)}) {
       Check(std::get<parser::Block>(eb->t));
     }
+    callContext_ = incoming;
   }
   void Check(const parser::IfStmt &is) {
+    const CallContext incoming{callContext_};
     const auto &uS{
         std::get<parser::UnlabeledStatement<parser::ActionStmt>>(is.t)};
-    CheckUnwrappedExpr(context_, uS.source,
-        std::get<parser::ScalarLogicalExpr>(is.t), isHostDevice);
+    const auto &condition{std::get<parser::ScalarLogicalExpr>(is.t)};
+    CheckUnwrappedExpr(context_, uS.source, condition, callContext_);
+    AllowGuardedCalls(condition);
     Check(uS.statement, uS.source);
+    callContext_ = incoming;
+  }
+  void AllowGuardedCalls(const parser::ScalarLogicalExpr &condition) {
+    // Accept either kind of callee in either arm of an ON_DEVICE guard.
+    // Determining whether a call is on the appropriate side is deferred.
+    if (isHostDevice && ChecksOnDevice(context_, condition)) {
+      callContext_ = CallContext::GuardedHostDevice;
+    }
   }
   void Check(const parser::LoopControl::Bounds &bounds) {
     Check(bounds.Lower());
@@ -618,13 +683,14 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
   }
   void Check(const parser::Expr &expr) {
     if (MaybeMsg msg{
-            DeviceExprChecker{context_, isHostDevice}(expr.typedExpr)}) {
+            DeviceExprChecker{context_, callContext_}(expr.typedExpr)}) {
       context_.Say(expr.source, std::move(*msg));
     }
   }
 
   SemanticsContext &context_;
   bool isHostDevice{false};
+  CallContext callContext_{CallContext::Device};
 };
 
 void CUDAChecker::Enter(const parser::SubroutineSubprogram &x) {
diff --git a/flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf b/flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf
new file mode 100644
index 0000000000000..606ced0689b7d
--- /dev/null
+++ b/flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf
@@ -0,0 +1,335 @@
+! RUN: %python %S/../test_errors.py %s %flang_fc1
+!
+! A host,device procedure has a host copy and a device copy. An unguarded
+! call must be valid in both. An ON_DEVICE() guard allows host and device
+! callees in either arm; checking the appropriate arm is deferred.
+
+module call_context
+  interface host_generic
+    module procedure host_subroutine
+  end interface
+  interface device_generic
+    module procedure device_subroutine
+  end interface
+contains
+  integer function host_only(n)
+    integer, value :: n
+    host_only = n + 1
+  end function
+
+  attributes(device) integer function device_only(n)
+    integer, value :: n
+    device_only = n + 2
+  end function
+
+  attributes(host,device) integer function both(n)
+    integer, value :: n
+    both = n + 3
+  end function
+
+  attributes(host,device) integer function unguarded_both(n)
+    integer, value :: n
+    unguarded_both = both(n)
+  end function
+
+  attributes(host,device) integer function unguarded_host(n)
+    integer, value :: n
+    !ERROR: 'host_only' may not be called in device code
+    unguarded_host = host_only(n)
+  end function
+
+  attributes(host,device) integer function unguarded_device(n)
+    integer, value :: n
+    !ERROR: 'device_only' may not be called in host code
+    unguarded_device = device_only(n)
+  end function
+
+  attributes(host,device) integer function guarded(n)
+    integer, value :: n
+    if (on_device()) then
+      guarded = device_only(n) + both(n)
+    else
+      guarded = host_only(n) + both(n)
+    end if
+  end function
+
+  attributes(host,device) integer function either_side(n)
+    integer, value :: n
+    if (on_device()) then
+      either_side = host_only(n)
+    else
+      either_side = device_only(n)
+    end if
+  end function
+
+  attributes(host,device) integer function guarded_negated(n)
+    integer, value :: n
+    if (.not. on_device()) then
+      guarded_negated = host_only(n)
+    else
+      guarded_negated = device_only(n)
+    end if
+  end function
+
+  attributes(host,device) integer function guarded_else_if(n, choose_first)
+    integer, value :: n
+    logical, value :: choose_first
+    if (choose_first) then
+      guarded_else_if = both(n)
+    else if (on_device()) then
+      guarded_else_if = device_only(n)
+    else
+      guarded_else_if = host_only(n)
+    end if
+  end function
+
+  attributes(host,device) integer function guarded_nested(n)
+    integer, value :: n
+    if (on_device()) then
+      if (n > 0) then
+        guarded_nested = device_only(n)
+      else
+        guarded_nested = both(n)
+      end if
+    else
+      if (n > 0) then
+        guarded_nested = host_only(n)
+      else
+        guarded_nested = both(n)
+      end if
+    end if
+  end function
+
+  attributes(host,device) integer function guarded_one_line(n)
+    integer, value :: n
+    guarded_one_line = both(n)
+    if (on_device()) guarded_one_line = device_only(n)
+    if (.not. on_device()) guarded_one_line = host_only(n)
+    ! Either callee is allowed under an ON_DEVICE check, regardless of polarity.
+    if (on_device()) guarded_one_line = host_only(n)
+    if (.not. on_device()) guarded_one_line = device_only(n)
+  end function
+
+  attributes(host,device) integer function ordinary_branch(n)
+    integer, value :: n
+    if...
[truncated]

``````````

</details>


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


More information about the flang-commits mailing list