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

Andre Kuhlenschmidt via flang-commits flang-commits at lists.llvm.org
Thu Oct 1 11:02:15 PDT 2026


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

 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.

>From 9cb35a8cd4f20b1c5936fbbf79eae19a26af3e20 Mon Sep 17 00:00:00 2001
From: Andre Kuhlenschmidt <akuhlenschmi at nvidia.com>
Date: Tue, 29 Sep 2026 14:56:53 -0700
Subject: [PATCH 1/2] [flang][CUDA] Check calls in host-device procedures by
 target

---
 flang/lib/Semantics/check-cuda.cpp            | 266 +++++++++++----
 .../CUDA/cuf-hostdevice-call-context.cuf      | 323 ++++++++++++++++++
 2 files changed, 529 insertions(+), 60 deletions(-)
 create mode 100644 flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf

diff --git a/flang/lib/Semantics/check-cuda.cpp b/flang/lib/Semantics/check-cuda.cpp
index 8922e0eb559e8e2..c6e616a6f8e851f 100644
--- a/flang/lib/Semantics/check-cuda.cpp
+++ b/flang/lib/Semantics/check-cuda.cpp
@@ -18,6 +18,7 @@
 #include "flang/Semantics/symbol.h"
 #include "flang/Semantics/tools.h"
 #include "llvm/ADT/StringSet.h"
+#include <utility>
 
 // Once labeled DO constructs have been canonicalized and their parse subtrees
 // transformed into parser::DoConstructs, scan the parser::Blocks of the program
@@ -67,6 +68,116 @@ static const llvm::StringSet<> warpFunctions_ = {"match_all_syncjj",
     "match_any_syncjj", "match_any_syncjx", "match_any_syncjf",
     "match_any_syncjd"};
 
+static constexpr unsigned HostTarget{1};
+static constexpr unsigned DeviceTarget{2};
+static constexpr unsigned BothTargets{HostTarget | DeviceTarget};
+
+// 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 refine the execution target.
+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);
+}
+
+static constexpr unsigned FalseResult{1};
+static constexpr unsigned TrueResult{2};
+static constexpr unsigned EitherResult{FalseResult | TrueResult};
+
+// Compute possible values of a logical condition in one copy of a procedure.
+// Unknown expressions can be true or false; this keeps target refinement
+// conservative without losing the useful implications of AND, OR, and NOT.
+static unsigned PossibleTruth(
+    const evaluate::Expr<evaluate::LogicalResult> &expr, bool onDevice) {
+  if (const auto *call{
+          evaluate::UnwrapExpr<evaluate::FunctionRef<evaluate::LogicalResult>>(
+              expr)}) {
+    if (IsOnDevice(call->proc())) {
+      return onDevice ? TrueResult : FalseResult;
+    }
+  }
+  if (const auto *negation{
+          evaluate::UnwrapExpr<evaluate::Not<evaluate::LogicalResult::kind>>(
+              expr)}) {
+    unsigned result{PossibleTruth(negation->left(), onDevice)};
+    return ((result & FalseResult) ? TrueResult : 0) |
+        ((result & TrueResult) ? FalseResult : 0);
+  }
+  if (const auto *parens{
+          evaluate::UnwrapExpr<evaluate::Parentheses<evaluate::LogicalResult>>(
+              expr)}) {
+    return PossibleTruth(parens->left(), onDevice);
+  }
+  if (const auto *binary{evaluate::UnwrapExpr<
+          evaluate::LogicalOperation<evaluate::LogicalResult::kind>>(expr)}) {
+    unsigned left{PossibleTruth(binary->left(), onDevice)};
+    unsigned right{PossibleTruth(binary->right(), onDevice)};
+    unsigned result{0};
+    for (bool a : {false, true}) {
+      if (!(left & (a ? TrueResult : FalseResult))) {
+        continue;
+      }
+      for (bool b : {false, true}) {
+        if (!(right & (b ? TrueResult : FalseResult))) {
+          continue;
+        }
+        bool value;
+        switch (binary->logicalOperator) {
+        case common::LogicalOperator::And:
+          value = a && b;
+          break;
+        case common::LogicalOperator::Or:
+          value = a || b;
+          break;
+        case common::LogicalOperator::Eqv:
+          value = a == b;
+          break;
+        case common::LogicalOperator::Neqv:
+          value = a != b;
+          break;
+        default:
+          return EitherResult;
+        }
+        result |= value ? TrueResult : FalseResult;
+      }
+    }
+    return result;
+  }
+  return EitherResult;
+}
+
+static std::pair<unsigned, unsigned> BranchTargets(SemanticsContext &context,
+    const parser::ScalarLogicalExpr &condition, unsigned incoming) {
+  const auto *analyzed{GetExpr(context, condition)};
+  const auto *logical{analyzed
+          ? evaluate::UnwrapExpr<evaluate::Expr<evaluate::LogicalResult>>(
+                *analyzed)
+          : nullptr};
+  if (!logical) {
+    return {incoming, incoming};
+  }
+  unsigned trueTargets{0};
+  unsigned falseTargets{0};
+  for (unsigned target : {HostTarget, DeviceTarget}) {
+    if (incoming & target) {
+      unsigned possible{PossibleTruth(*logical, target == DeviceTarget)};
+      if (possible & TrueResult) {
+        trueTargets |= target;
+      }
+      if (possible & FalseResult) {
+        falseTargets |= target;
+      }
+    }
+  }
+  return {trueTargets, falseTargets};
+}
+
 // Traverses an evaluate::Expr<> in search of unsupported operations
 // on the device.
 
@@ -74,10 +185,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, unsigned targets = DeviceTarget)
+      : Base(*this), context_{c}, targets_{targets} {}
   using Base::operator();
   Result operator()(const evaluate::ProcedureDesignator &x) const {
+    if (targets_ == 0 || IsOnDevice(x)) {
+      return {};
+    }
     if (const Symbol * sym{x.GetInterfaceSymbol()}) {
       const Symbol &ultimate{sym->GetUltimate()};
       const auto *subp{ultimate.detailsIf<semantics::SubprogramDetails>()};
@@ -95,6 +210,11 @@ struct DeviceExprChecker
               return parser::MessageFormattedText(
                   "warp match function disabled"_err_en_US);
             }
+            if (*attrs == common::CUDASubprogramAttrs::Device &&
+                (targets_ & HostTarget)) {
+              return parser::MessageFormattedText(
+                  "'%s' may not be called in host code"_err_en_US, x.GetName());
+            }
             return {};
           }
           if (*attrs == common::CUDASubprogramAttrs::Global) {
@@ -119,10 +239,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 (!(targets_ & DeviceTarget)) {
       return {};
     }
     return parser::MessageFormattedText(
@@ -130,7 +247,7 @@ struct DeviceExprChecker
   }
 
   SemanticsContext &context_;
-  bool allowHostCallees_{false};
+  unsigned targets_{DeviceTarget};
 };
 
 static bool IsHostArray(const Symbol &symbol) {
@@ -209,19 +326,18 @@ struct FindHostArray
 
 template <typename A>
 static MaybeMsg CheckUnwrappedExpr(
-    SemanticsContext &context, const A &x, bool allowHostCallees = false) {
+    SemanticsContext &context, const A &x, unsigned targets = DeviceTarget) {
   if (const auto *expr{parser::Unwrap<parser::Expr>(x)}) {
-    return DeviceExprChecker{context, allowHostCallees}(expr->typedExpr);
+    return DeviceExprChecker{context, targets}(expr->typedExpr);
   }
   return {};
 }
 
 template <typename A>
 static void CheckUnwrappedExpr(SemanticsContext &context, SourceName at,
-    const A &x, bool allowHostCallees = false) {
+    const A &x, unsigned targets = DeviceTarget) {
   if (const auto *expr{parser::Unwrap<parser::Expr>(x)}) {
-    if (auto msg{
-            DeviceExprChecker{context, allowHostCallees}(expr->typedExpr)}) {
+    if (auto msg{DeviceExprChecker{context, targets}(expr->typedExpr)}) {
       context.Say(at, std::move(*msg));
     }
   }
@@ -230,15 +346,15 @@ 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) {
+      SemanticsContext &context, const A &x, unsigned targets = DeviceTarget) {
     if constexpr (ConstraintTrait<A>) {
-      return WhyNotOk(context, x.thing, allowHostCallees);
+      return WhyNotOk(context, x.thing, targets);
     } else if constexpr (WrapperTrait<A>) {
-      return WhyNotOk(context, x.v, allowHostCallees);
+      return WhyNotOk(context, x.v, targets);
     } else if constexpr (UnionTrait<A>) {
-      return WhyNotOk(context, x.u, allowHostCallees);
+      return WhyNotOk(context, x.u, targets);
     } else if constexpr (TupleTrait<A>) {
-      return WhyNotOk(context, x.t, allowHostCallees);
+      return WhyNotOk(context, x.t, targets);
     } else {
       return parser::MessageFormattedText{
           "Statement may not appear in device code"_err_en_US};
@@ -246,33 +362,33 @@ 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, unsigned targets = DeviceTarget) {
+    return WhyNotOk(context, x.value(), targets);
   }
   template <typename... As>
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const std::variant<As...> &x, bool allowHostCallees = false) {
+      const std::variant<As...> &x, unsigned targets = DeviceTarget) {
     return common::visit(
-        [&context, allowHostCallees](
-            const auto &x) { return WhyNotOk(context, x, allowHostCallees); },
+        [&context, targets](
+            const auto &x) { return WhyNotOk(context, x, targets); },
         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, unsigned targets = DeviceTarget) {
     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), targets)}) {
       return msg;
     } else {
-      return WhyNotOk<(J + 1)>(context, x, allowHostCallees);
+      return WhyNotOk<(J + 1)>(context, x, targets);
     }
   }
   template <typename A>
   static MaybeMsg WhyNotOk(SemanticsContext &context, const std::list<A> &x,
-      bool allowHostCallees = false) {
+      unsigned targets = DeviceTarget) {
     for (const auto &y : x) {
-      if (MaybeMsg result{WhyNotOk(context, y, allowHostCallees)}) {
+      if (MaybeMsg result{WhyNotOk(context, y, targets)}) {
         return result;
       }
     }
@@ -280,76 +396,75 @@ template <bool CUF_KERNEL> struct ActionStmtChecker {
   }
   template <typename A>
   static MaybeMsg WhyNotOk(SemanticsContext &context, const std::optional<A> &x,
-      bool allowHostCallees = false) {
+      unsigned targets = DeviceTarget) {
     if (x) {
-      return WhyNotOk(context, *x, allowHostCallees);
+      return WhyNotOk(context, *x, targets);
     } 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, unsigned targets = DeviceTarget) {
+    return WhyNotOk(context, x.statement, targets);
   }
   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, unsigned targets = DeviceTarget) {
+    return WhyNotOk(context, x.statement, targets);
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::AllocateStmt &, bool allowHostCallees = false) {
+      const parser::AllocateStmt &, unsigned targets = DeviceTarget) {
     return {}; // AllocateObjects are checked elsewhere
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::AllocateCoarraySpec &, bool allowHostCallees = false) {
+      const parser::AllocateCoarraySpec &, unsigned targets = DeviceTarget) {
     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 &, unsigned targets = DeviceTarget) {
     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, unsigned targets = DeviceTarget) {
+    return DeviceExprChecker{context, targets}(x.typedAssignment);
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context, const parser::CallStmt &x,
-      bool allowHostCallees = false) {
-    return DeviceExprChecker{context, allowHostCallees}(x.typedCall);
+      unsigned targets = DeviceTarget) {
+    return DeviceExprChecker{context, targets}(x.typedCall);
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::ContinueStmt &, bool allowHostCallees = false) {
+      const parser::ContinueStmt &, unsigned targets = DeviceTarget) {
     return {};
   }
   static MaybeMsg WhyNotOk(SemanticsContext &, const parser::PauseStmt &,
-      bool allowHostCallees = false) {
+      unsigned targets = DeviceTarget) {
     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)}) {
+      unsigned targets = DeviceTarget) {
+    if (auto result{CheckUnwrappedExpr(
+            context, std::get<parser::ScalarLogicalExpr>(x.t), targets)}) {
       return result;
     }
     return WhyNotOk(context,
         std::get<parser::UnlabeledStatement<parser::ActionStmt>>(x.t).statement,
-        allowHostCallees);
+        targets);
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::NullifyStmt &x, bool allowHostCallees = false) {
+      const parser::NullifyStmt &x, unsigned targets = DeviceTarget) {
     for (const auto &y : x.v) {
-      if (MaybeMsg result{
-              DeviceExprChecker{context, allowHostCallees}(y.typedExpr)}) {
+      if (MaybeMsg result{DeviceExprChecker{context, targets}(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, unsigned targets = DeviceTarget) {
+    return DeviceExprChecker{context, targets}(x.typedAssignment);
   }
 };
 
@@ -372,11 +487,15 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
         isHostDevice = subp->cudaSubprogramAttrs() &&
             subp->cudaSubprogramAttrs() ==
                 common::CUDASubprogramAttrs::HostDevice;
+        currentTargets_ = isHostDevice ? BothTargets : DeviceTarget;
         Check(body);
       }
     }
   }
   void Check(const parser::Block &block) {
+    if (currentTargets_ == 0) {
+      return;
+    }
     for (const auto &epc : block) {
       Check(epc);
     }
@@ -489,6 +608,9 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
     }
   }
   void Check(const parser::ActionStmt &stmt, const parser::CharBlock &source) {
+    if (currentTargets_ == 0) {
+      return;
+    }
     common::visit(
         common::visitors{
             [&](const common::Indirection<parser::CycleStmt> &) {
@@ -547,13 +669,13 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
                 ErrorIfHostSymbol(assign->rhs, source);
               }
               if (auto msg{ActionStmtChecker<IsCUFKernelDo>::WhyNotOk(
-                      context_, x, isHostDevice)}) {
+                      context_, x, currentTargets_)}) {
                 context_.Say(source, std::move(*msg));
               }
             },
             [&](const auto &x) {
               if (auto msg{ActionStmtChecker<IsCUFKernelDo>::WhyNotOk(
-                      context_, x, isHostDevice)}) {
+                      context_, x, currentTargets_)}) {
                 context_.Say(source, std::move(*msg));
               }
             },
@@ -561,28 +683,51 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
         stmt.u);
   }
   void Check(const parser::IfConstruct &ic) {
+    const unsigned incoming{currentTargets_};
     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, incoming);
+    auto [thenTargets, remainingTargets]{isHostDevice
+            ? BranchTargets(context_, condition, incoming)
+            : std::pair<unsigned, unsigned>{incoming, incoming}};
+    currentTargets_ = thenTargets;
     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)};
+      currentTargets_ = remainingTargets;
+      if (remainingTargets != 0) {
+        CheckUnwrappedExpr(
+            context_, eIfS.source, elseIfCondition, remainingTargets);
+      }
+      auto [elseIfTargets, nextTargets]{isHostDevice
+              ? BranchTargets(context_, elseIfCondition, remainingTargets)
+              : std::pair<unsigned, unsigned>{
+                    remainingTargets, remainingTargets}};
+      currentTargets_ = elseIfTargets;
       Check(std::get<parser::Block>(eib.t));
+      remainingTargets = nextTargets;
     }
     if (const auto &eb{
             std::get<std::optional<parser::IfConstruct::ElseBlock>>(ic.t)}) {
+      currentTargets_ = remainingTargets;
       Check(std::get<parser::Block>(eb->t));
     }
+    currentTargets_ = incoming;
   }
   void Check(const parser::IfStmt &is) {
+    const unsigned incoming{currentTargets_};
     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, incoming);
+    currentTargets_ = isHostDevice
+        ? BranchTargets(context_, condition, incoming).first
+        : incoming;
     Check(uS.statement, uS.source);
+    currentTargets_ = incoming;
   }
   void Check(const parser::LoopControl::Bounds &bounds) {
     Check(bounds.Lower());
@@ -618,13 +763,14 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
   }
   void Check(const parser::Expr &expr) {
     if (MaybeMsg msg{
-            DeviceExprChecker{context_, isHostDevice}(expr.typedExpr)}) {
+            DeviceExprChecker{context_, currentTargets_}(expr.typedExpr)}) {
       context_.Say(expr.source, std::move(*msg));
     }
   }
 
   SemanticsContext &context_;
   bool isHostDevice{false};
+  unsigned currentTargets_{DeviceTarget};
 };
 
 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 000000000000000..088b0ffa1011f32
--- /dev/null
+++ b/flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf
@@ -0,0 +1,323 @@
+! 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. ON_DEVICE() selects the appropriate target for
+! calls that are only valid in one copy.
+
+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 wrong_side(n)
+    integer, value :: n
+    if (on_device()) then
+      !ERROR: 'host_only' may not be called in device code
+      wrong_side = host_only(n)
+    else
+      !ERROR: 'device_only' may not be called in host code
+      wrong_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)
+  end function
+
+  attributes(host,device) integer function compound_and(n)
+    integer, value :: n
+    if (on_device() .and. n > 0) then
+      compound_and = both(n)
+    else
+      ! A device copy reaches ELSE when n <= 0.
+      !ERROR: 'host_only' may not be called in device code
+      compound_and = host_only(n)
+    end if
+  end function
+
+  attributes(host,device) integer function compound_or(n)
+    integer, value :: n
+    if (on_device() .or. n > 0) then
+      ! A host copy reaches THEN when n > 0.
+      !ERROR: 'device_only' may not be called in host code
+      compound_or = device_only(n)
+    else
+      compound_or = both(n)
+    end if
+  end function
+
+  attributes(host,device) integer function safe_compound(n)
+    integer, value :: n
+    safe_compound = both(n)
+    if (on_device() .and. n > 0) safe_compound = device_only(n)
+    if (on_device() .or. n > 0) then
+      safe_compound = safe_compound + both(n)
+    else
+      safe_compound = safe_compound + host_only(n)
+    end if
+  end function
+
+  logical function host_condition(n)
+    integer, value :: n
+    host_condition = n > 0
+  end function
+
+  attributes(device) logical function device_condition(n)
+    integer, value :: n
+    device_condition = n > 0
+  end function
+
+  attributes(host,device) integer function else_if_condition(n)
+    integer, value :: n
+    if (on_device()) then
+      else_if_condition = device_only(n)
+    else if (host_condition(n)) then
+      else_if_condition = host_only(n)
+    else
+      else_if_condition = both(n)
+    end if
+  end function
+
+  attributes(host,device) integer function else_if_condition_negated(n)
+    integer, value :: n
+    if (.not. on_device()) then
+      else_if_condition_negated = host_only(n)
+    else if (device_condition(n)) then
+      else_if_condition_negated = device_only(n)
+    else
+      else_if_condition_negated = both(n)
+    end if
+  end function
+
+  attributes(host,device) integer function after_guard(n)
+    integer, value :: n
+    if (on_device()) then
+      after_guard = device_only(n)
+    else
+      after_guard = host_only(n)
+    end if
+    ! Both copies reach the statement after END IF.
+    !ERROR: 'host_only' may not be called in device code
+    after_guard = after_guard + host_only(n)
+  end function
+
+  subroutine host_subroutine(n)
+    integer, intent(inout) :: n
+    n = n + 1
+  end subroutine
+
+  attributes(device) subroutine device_subroutine(n)
+    integer, value :: n
+    integer :: scratch
+    scratch = n + 2
+  end subroutine
+
+  attributes(host,device) subroutine both_subroutine(n)
+    integer, intent(inout) :: n
+    n = n + 3
+  end subroutine
+
+  attributes(host,device) subroutine guarded_subroutine(n)
+    integer, intent(inout) :: n
+    call both_subroutine(n)
+    if (on_device()) then
+      call device_subroutine(n)
+    else
+      call host_subroutine(n)
+    end if
+  end subroutine
+
+  attributes(host,device) subroutine unguarded_subroutines(n)
+    integer, intent(inout) :: n
+    !ERROR: 'host_subroutine' may not be called in device code
+    call host_subroutine(n)
+    !ERROR: 'device_subroutine' may not be called in host code
+    call device_subroutine(n)
+  end subroutine
+
+  attributes(host,device) subroutine generic_subroutines(n)
+    integer, intent(inout) :: n
+    !ERROR: 'host_subroutine' may not be called in device code
+    call host_generic(n)
+    !ERROR: 'device_subroutine' may not be called in host code
+    call device_generic(n)
+  end subroutine
+
+  subroutine host_entry(x)
+    integer, intent(out) :: x
+    x = unguarded_both(1) + guarded(1) + guarded_negated(1) + &
+        guarded_else_if(1, .false.) + guarded_nested(1) + guarded_one_line(1)
+    call guarded_subroutine(x)
+  end subroutine
+
+  attributes(global) subroutine device_entry(x)
+    integer, device :: x(*)
+    x(1) = unguarded_both(1) + guarded(1) + guarded_negated(1) + &
+        guarded_else_if(1, .false.) + guarded_nested(1) + guarded_one_line(1)
+    call guarded_subroutine(x(1))
+  end subroutine
+end module
+
+module fake_bindc_guard
+contains
+  attributes(host,device) logical function on_device() bind(c, name="user_on_device")
+    on_device = .false.
+  end function
+
+  subroutine host_only(n)
+    integer, intent(inout) :: n
+    n = n + 1
+  end subroutine
+
+  attributes(device) subroutine device_only(n)
+    integer, intent(inout) :: n
+    n = n + 2
+  end subroutine
+
+  attributes(host,device) subroutine misleading_bindc_guard(n)
+    integer, intent(inout) :: n
+    if (on_device()) then
+      !ERROR: 'device_only' may not be called in host code
+      call device_only(n)
+    else
+      !ERROR: 'host_only' may not be called in device code
+      call host_only(n)
+    end if
+  end subroutine
+end module
+
+module aliased_guard
+  use cudadevice, only: on_gpu => on_device
+contains
+  subroutine host_only(n)
+    integer, intent(inout) :: n
+    n = n + 1
+  end subroutine
+
+  attributes(device) subroutine device_only(n)
+    integer, intent(inout) :: n
+    n = n + 2
+  end subroutine
+
+  attributes(host,device) subroutine real_aliased_guard(n)
+    integer, intent(inout) :: n
+    if (on_gpu()) then
+      call device_only(n)
+    else
+      call host_only(n)
+    end if
+  end subroutine
+end module
+
+module fake_guard
+contains
+  attributes(host,device) logical function on_device()
+    on_device = .false.
+  end function
+
+  subroutine host_only(n)
+    integer, intent(inout) :: n
+    n = n + 1
+  end subroutine
+
+  attributes(device) subroutine device_only(n)
+    integer, intent(inout) :: n
+    n = n + 2
+  end subroutine
+
+  attributes(host,device) subroutine misleading_guard(n)
+    integer, intent(inout) :: n
+    if (on_device()) then
+      !ERROR: 'device_only' may not be called in host code
+      call device_only(n)
+    else
+      !ERROR: 'host_only' may not be called in device code
+      call host_only(n)
+    end if
+  end subroutine
+end module

>From c9b59fe18d734c4d4540ce3b948a3b3ac10fd746 Mon Sep 17 00:00:00 2001
From: Andre Kuhlenschmidt <akuhlenschmi at nvidia.com>
Date: Thu, 1 Oct 2026 10:43:26 -0700
Subject: [PATCH 2/2] [flang][CUDA] Defer branch target analysis for ON_DEVICE
 guards

---
 flang/lib/Semantics/check-cuda.cpp            | 276 +++++++-----------
 .../CUDA/cuf-hostdevice-call-context.cuf      |  34 ++-
 2 files changed, 121 insertions(+), 189 deletions(-)

diff --git a/flang/lib/Semantics/check-cuda.cpp b/flang/lib/Semantics/check-cuda.cpp
index c6e616a6f8e851f..13545803428ee7b 100644
--- a/flang/lib/Semantics/check-cuda.cpp
+++ b/flang/lib/Semantics/check-cuda.cpp
@@ -18,7 +18,6 @@
 #include "flang/Semantics/symbol.h"
 #include "flang/Semantics/tools.h"
 #include "llvm/ADT/StringSet.h"
-#include <utility>
 
 // Once labeled DO constructs have been canonicalized and their parse subtrees
 // transformed into parser::DoConstructs, scan the parser::Blocks of the program
@@ -68,13 +67,11 @@ static const llvm::StringSet<> warpFunctions_ = {"match_all_syncjj",
     "match_any_syncjj", "match_any_syncjx", "match_any_syncjf",
     "match_any_syncjd"};
 
-static constexpr unsigned HostTarget{1};
-static constexpr unsigned DeviceTarget{2};
-static constexpr unsigned BothTargets{HostTarget | DeviceTarget};
+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 refine the execution target.
+// 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) {
@@ -86,96 +83,19 @@ static bool IsOnDevice(const evaluate::ProcedureDesignator &proc) {
       module && module->attrs().test(Attr::INTRINSIC);
 }
 
-static constexpr unsigned FalseResult{1};
-static constexpr unsigned TrueResult{2};
-static constexpr unsigned EitherResult{FalseResult | TrueResult};
-
-// Compute possible values of a logical condition in one copy of a procedure.
-// Unknown expressions can be true or false; this keeps target refinement
-// conservative without losing the useful implications of AND, OR, and NOT.
-static unsigned PossibleTruth(
-    const evaluate::Expr<evaluate::LogicalResult> &expr, bool onDevice) {
-  if (const auto *call{
-          evaluate::UnwrapExpr<evaluate::FunctionRef<evaluate::LogicalResult>>(
-              expr)}) {
-    if (IsOnDevice(call->proc())) {
-      return onDevice ? TrueResult : FalseResult;
-    }
-  }
-  if (const auto *negation{
-          evaluate::UnwrapExpr<evaluate::Not<evaluate::LogicalResult::kind>>(
-              expr)}) {
-    unsigned result{PossibleTruth(negation->left(), onDevice)};
-    return ((result & FalseResult) ? TrueResult : 0) |
-        ((result & TrueResult) ? FalseResult : 0);
-  }
-  if (const auto *parens{
-          evaluate::UnwrapExpr<evaluate::Parentheses<evaluate::LogicalResult>>(
-              expr)}) {
-    return PossibleTruth(parens->left(), onDevice);
-  }
-  if (const auto *binary{evaluate::UnwrapExpr<
-          evaluate::LogicalOperation<evaluate::LogicalResult::kind>>(expr)}) {
-    unsigned left{PossibleTruth(binary->left(), onDevice)};
-    unsigned right{PossibleTruth(binary->right(), onDevice)};
-    unsigned result{0};
-    for (bool a : {false, true}) {
-      if (!(left & (a ? TrueResult : FalseResult))) {
-        continue;
-      }
-      for (bool b : {false, true}) {
-        if (!(right & (b ? TrueResult : FalseResult))) {
-          continue;
-        }
-        bool value;
-        switch (binary->logicalOperator) {
-        case common::LogicalOperator::And:
-          value = a && b;
-          break;
-        case common::LogicalOperator::Or:
-          value = a || b;
-          break;
-        case common::LogicalOperator::Eqv:
-          value = a == b;
-          break;
-        case common::LogicalOperator::Neqv:
-          value = a != b;
-          break;
-        default:
-          return EitherResult;
-        }
-        result |= value ? TrueResult : FalseResult;
-      }
-    }
-    return result;
+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);
   }
-  return EitherResult;
-}
+};
 
-static std::pair<unsigned, unsigned> BranchTargets(SemanticsContext &context,
-    const parser::ScalarLogicalExpr &condition, unsigned incoming) {
-  const auto *analyzed{GetExpr(context, condition)};
-  const auto *logical{analyzed
-          ? evaluate::UnwrapExpr<evaluate::Expr<evaluate::LogicalResult>>(
-                *analyzed)
-          : nullptr};
-  if (!logical) {
-    return {incoming, incoming};
-  }
-  unsigned trueTargets{0};
-  unsigned falseTargets{0};
-  for (unsigned target : {HostTarget, DeviceTarget}) {
-    if (incoming & target) {
-      unsigned possible{PossibleTruth(*logical, target == DeviceTarget)};
-      if (possible & TrueResult) {
-        trueTargets |= target;
-      }
-      if (possible & FalseResult) {
-        falseTargets |= target;
-      }
-    }
-  }
-  return {trueTargets, falseTargets};
+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
@@ -186,11 +106,11 @@ struct DeviceExprChecker
   using Result = MaybeMsg;
   using Base = evaluate::AnyTraverse<DeviceExprChecker, Result>;
   explicit DeviceExprChecker(
-      SemanticsContext &c, unsigned targets = DeviceTarget)
-      : Base(*this), context_{c}, targets_{targets} {}
+      SemanticsContext &c, CallContext callContext = CallContext::Device)
+      : Base(*this), context_{c}, callContext_{callContext} {}
   using Base::operator();
   Result operator()(const evaluate::ProcedureDesignator &x) const {
-    if (targets_ == 0 || IsOnDevice(x)) {
+    if (IsOnDevice(x)) {
       return {};
     }
     if (const Symbol * sym{x.GetInterfaceSymbol()}) {
@@ -211,7 +131,7 @@ struct DeviceExprChecker
                   "warp match function disabled"_err_en_US);
             }
             if (*attrs == common::CUDASubprogramAttrs::Device &&
-                (targets_ & HostTarget)) {
+                callContext_ == CallContext::HostDevice) {
               return parser::MessageFormattedText(
                   "'%s' may not be called in host code"_err_en_US, x.GetName());
             }
@@ -239,7 +159,7 @@ struct DeviceExprChecker
       return {};
     }
 
-    if (!(targets_ & DeviceTarget)) {
+    if (callContext_ == CallContext::GuardedHostDevice) {
       return {};
     }
     return parser::MessageFormattedText(
@@ -247,7 +167,7 @@ struct DeviceExprChecker
   }
 
   SemanticsContext &context_;
-  unsigned targets_{DeviceTarget};
+  CallContext callContext_{CallContext::Device};
 };
 
 static bool IsHostArray(const Symbol &symbol) {
@@ -325,19 +245,19 @@ struct FindHostArray
 };
 
 template <typename A>
-static MaybeMsg CheckUnwrappedExpr(
-    SemanticsContext &context, const A &x, unsigned targets = DeviceTarget) {
+static MaybeMsg CheckUnwrappedExpr(SemanticsContext &context, const A &x,
+    CallContext callContext = CallContext::Device) {
   if (const auto *expr{parser::Unwrap<parser::Expr>(x)}) {
-    return DeviceExprChecker{context, targets}(expr->typedExpr);
+    return DeviceExprChecker{context, callContext}(expr->typedExpr);
   }
   return {};
 }
 
 template <typename A>
 static void CheckUnwrappedExpr(SemanticsContext &context, SourceName at,
-    const A &x, unsigned targets = DeviceTarget) {
+    const A &x, CallContext callContext = CallContext::Device) {
   if (const auto *expr{parser::Unwrap<parser::Expr>(x)}) {
-    if (auto msg{DeviceExprChecker{context, targets}(expr->typedExpr)}) {
+    if (auto msg{DeviceExprChecker{context, callContext}(expr->typedExpr)}) {
       context.Say(at, std::move(*msg));
     }
   }
@@ -345,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, unsigned targets = DeviceTarget) {
+  static MaybeMsg WhyNotOk(SemanticsContext &context, const A &x,
+      CallContext callContext = CallContext::Device) {
     if constexpr (ConstraintTrait<A>) {
-      return WhyNotOk(context, x.thing, targets);
+      return WhyNotOk(context, x.thing, callContext);
     } else if constexpr (WrapperTrait<A>) {
-      return WhyNotOk(context, x.v, targets);
+      return WhyNotOk(context, x.v, callContext);
     } else if constexpr (UnionTrait<A>) {
-      return WhyNotOk(context, x.u, targets);
+      return WhyNotOk(context, x.u, callContext);
     } else if constexpr (TupleTrait<A>) {
-      return WhyNotOk(context, x.t, targets);
+      return WhyNotOk(context, x.t, callContext);
     } else {
       return parser::MessageFormattedText{
           "Statement may not appear in device code"_err_en_US};
@@ -362,33 +282,36 @@ template <bool CUF_KERNEL> struct ActionStmtChecker {
   }
   template <typename A>
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const common::Indirection<A> &x, unsigned targets = DeviceTarget) {
-    return WhyNotOk(context, x.value(), targets);
+      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, unsigned targets = DeviceTarget) {
+      const std::variant<As...> &x,
+      CallContext callContext = CallContext::Device) {
     return common::visit(
-        [&context, targets](
-            const auto &x) { return WhyNotOk(context, x, targets); },
+        [&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, unsigned targets = DeviceTarget) {
+      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), targets)}) {
+    } else if (auto msg{WhyNotOk(context, std::get<J>(x), callContext)}) {
       return msg;
     } else {
-      return WhyNotOk<(J + 1)>(context, x, targets);
+      return WhyNotOk<(J + 1)>(context, x, callContext);
     }
   }
   template <typename A>
   static MaybeMsg WhyNotOk(SemanticsContext &context, const std::list<A> &x,
-      unsigned targets = DeviceTarget) {
+      CallContext callContext = CallContext::Device) {
     for (const auto &y : x) {
-      if (MaybeMsg result{WhyNotOk(context, y, targets)}) {
+      if (MaybeMsg result{WhyNotOk(context, y, callContext)}) {
         return result;
       }
     }
@@ -396,75 +319,85 @@ template <bool CUF_KERNEL> struct ActionStmtChecker {
   }
   template <typename A>
   static MaybeMsg WhyNotOk(SemanticsContext &context, const std::optional<A> &x,
-      unsigned targets = DeviceTarget) {
+      CallContext callContext = CallContext::Device) {
     if (x) {
-      return WhyNotOk(context, *x, targets);
+      return WhyNotOk(context, *x, callContext);
     } else {
       return {};
     }
   }
   template <typename A>
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::UnlabeledStatement<A> &x, unsigned targets = DeviceTarget) {
-    return WhyNotOk(context, x.statement, targets);
+      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, unsigned targets = DeviceTarget) {
-    return WhyNotOk(context, x.statement, targets);
+      const parser::Statement<A> &x,
+      CallContext callContext = CallContext::Device) {
+    return WhyNotOk(context, x.statement, callContext);
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::AllocateStmt &, unsigned targets = DeviceTarget) {
+      const parser::AllocateStmt &,
+      CallContext callContext = CallContext::Device) {
     return {}; // AllocateObjects are checked elsewhere
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::AllocateCoarraySpec &, unsigned targets = DeviceTarget) {
+      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 &, unsigned targets = DeviceTarget) {
+      const parser::DeallocateStmt &,
+      CallContext callContext = CallContext::Device) {
     return {}; // AllocateObjects are checked elsewhere
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::AssignmentStmt &x, unsigned targets = DeviceTarget) {
-    return DeviceExprChecker{context, targets}(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,
-      unsigned targets = DeviceTarget) {
-    return DeviceExprChecker{context, targets}(x.typedCall);
+      CallContext callContext = CallContext::Device) {
+    return DeviceExprChecker{context, callContext}(x.typedCall);
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::ContinueStmt &, unsigned targets = DeviceTarget) {
+      const parser::ContinueStmt &,
+      CallContext callContext = CallContext::Device) {
     return {};
   }
   static MaybeMsg WhyNotOk(SemanticsContext &, const parser::PauseStmt &,
-      unsigned targets = DeviceTarget) {
+      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,
-      unsigned targets = DeviceTarget) {
+      CallContext callContext = CallContext::Device) {
     if (auto result{CheckUnwrappedExpr(
-            context, std::get<parser::ScalarLogicalExpr>(x.t), targets)}) {
+            context, std::get<parser::ScalarLogicalExpr>(x.t), callContext)}) {
       return result;
     }
     return WhyNotOk(context,
         std::get<parser::UnlabeledStatement<parser::ActionStmt>>(x.t).statement,
-        targets);
+        callContext);
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::NullifyStmt &x, unsigned targets = DeviceTarget) {
+      const parser::NullifyStmt &x,
+      CallContext callContext = CallContext::Device) {
     for (const auto &y : x.v) {
-      if (MaybeMsg result{DeviceExprChecker{context, targets}(y.typedExpr)}) {
+      if (MaybeMsg result{
+              DeviceExprChecker{context, callContext}(y.typedExpr)}) {
         return result;
       }
     }
     return {};
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::PointerAssignmentStmt &x, unsigned targets = DeviceTarget) {
-    return DeviceExprChecker{context, targets}(x.typedAssignment);
+      const parser::PointerAssignmentStmt &x,
+      CallContext callContext = CallContext::Device) {
+    return DeviceExprChecker{context, callContext}(x.typedAssignment);
   }
 };
 
@@ -487,15 +420,13 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
         isHostDevice = subp->cudaSubprogramAttrs() &&
             subp->cudaSubprogramAttrs() ==
                 common::CUDASubprogramAttrs::HostDevice;
-        currentTargets_ = isHostDevice ? BothTargets : DeviceTarget;
+        callContext_ =
+            isHostDevice ? CallContext::HostDevice : CallContext::Device;
         Check(body);
       }
     }
   }
   void Check(const parser::Block &block) {
-    if (currentTargets_ == 0) {
-      return;
-    }
     for (const auto &epc : block) {
       Check(epc);
     }
@@ -608,9 +539,6 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
     }
   }
   void Check(const parser::ActionStmt &stmt, const parser::CharBlock &source) {
-    if (currentTargets_ == 0) {
-      return;
-    }
     common::visit(
         common::visitors{
             [&](const common::Indirection<parser::CycleStmt> &) {
@@ -669,13 +597,13 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
                 ErrorIfHostSymbol(assign->rhs, source);
               }
               if (auto msg{ActionStmtChecker<IsCUFKernelDo>::WhyNotOk(
-                      context_, x, currentTargets_)}) {
+                      context_, x, callContext_)}) {
                 context_.Say(source, std::move(*msg));
               }
             },
             [&](const auto &x) {
               if (auto msg{ActionStmtChecker<IsCUFKernelDo>::WhyNotOk(
-                      context_, x, currentTargets_)}) {
+                      context_, x, callContext_)}) {
                 context_.Say(source, std::move(*msg));
               }
             },
@@ -683,51 +611,43 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
         stmt.u);
   }
   void Check(const parser::IfConstruct &ic) {
-    const unsigned incoming{currentTargets_};
+    const CallContext incoming{callContext_};
     const auto &ifS{std::get<parser::Statement<parser::IfThenStmt>>(ic.t)};
     const auto &condition{std::get<parser::ScalarLogicalExpr>(ifS.statement.t)};
-    CheckUnwrappedExpr(context_, ifS.source, condition, incoming);
-    auto [thenTargets, remainingTargets]{isHostDevice
-            ? BranchTargets(context_, condition, incoming)
-            : std::pair<unsigned, unsigned>{incoming, incoming}};
-    currentTargets_ = thenTargets;
+    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)};
       const auto &elseIfCondition{
           std::get<parser::ScalarLogicalExpr>(eIfS.statement.t)};
-      currentTargets_ = remainingTargets;
-      if (remainingTargets != 0) {
-        CheckUnwrappedExpr(
-            context_, eIfS.source, elseIfCondition, remainingTargets);
-      }
-      auto [elseIfTargets, nextTargets]{isHostDevice
-              ? BranchTargets(context_, elseIfCondition, remainingTargets)
-              : std::pair<unsigned, unsigned>{
-                    remainingTargets, remainingTargets}};
-      currentTargets_ = elseIfTargets;
+      CheckUnwrappedExpr(context_, eIfS.source, elseIfCondition, callContext_);
+      AllowGuardedCalls(elseIfCondition);
       Check(std::get<parser::Block>(eib.t));
-      remainingTargets = nextTargets;
     }
     if (const auto &eb{
             std::get<std::optional<parser::IfConstruct::ElseBlock>>(ic.t)}) {
-      currentTargets_ = remainingTargets;
       Check(std::get<parser::Block>(eb->t));
     }
-    currentTargets_ = incoming;
+    callContext_ = incoming;
   }
   void Check(const parser::IfStmt &is) {
-    const unsigned incoming{currentTargets_};
+    const CallContext incoming{callContext_};
     const auto &uS{
         std::get<parser::UnlabeledStatement<parser::ActionStmt>>(is.t)};
     const auto &condition{std::get<parser::ScalarLogicalExpr>(is.t)};
-    CheckUnwrappedExpr(context_, uS.source, condition, incoming);
-    currentTargets_ = isHostDevice
-        ? BranchTargets(context_, condition, incoming).first
-        : incoming;
+    CheckUnwrappedExpr(context_, uS.source, condition, callContext_);
+    AllowGuardedCalls(condition);
     Check(uS.statement, uS.source);
-    currentTargets_ = incoming;
+    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());
@@ -763,14 +683,14 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
   }
   void Check(const parser::Expr &expr) {
     if (MaybeMsg msg{
-            DeviceExprChecker{context_, currentTargets_}(expr.typedExpr)}) {
+            DeviceExprChecker{context_, callContext_}(expr.typedExpr)}) {
       context_.Say(expr.source, std::move(*msg));
     }
   }
 
   SemanticsContext &context_;
   bool isHostDevice{false};
-  unsigned currentTargets_{DeviceTarget};
+  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
index 088b0ffa1011f32..606ced0689b7de3 100644
--- a/flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf
+++ b/flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf
@@ -1,8 +1,8 @@
 ! 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. ON_DEVICE() selects the appropriate target for
-! calls that are only valid in one copy.
+! 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
@@ -53,14 +53,12 @@ contains
     end if
   end function
 
-  attributes(host,device) integer function wrong_side(n)
+  attributes(host,device) integer function either_side(n)
     integer, value :: n
     if (on_device()) then
-      !ERROR: 'host_only' may not be called in device code
-      wrong_side = host_only(n)
+      either_side = host_only(n)
     else
-      !ERROR: 'device_only' may not be called in host code
-      wrong_side = device_only(n)
+      either_side = device_only(n)
     end if
   end function
 
@@ -107,6 +105,20 @@ contains
     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 (n > 0) then
+      !ERROR: 'host_only' may not be called in device code
+      ordinary_branch = host_only(n)
+    else
+      !ERROR: 'device_only' may not be called in host code
+      ordinary_branch = device_only(n)
+    end if
   end function
 
   attributes(host,device) integer function compound_and(n)
@@ -114,8 +126,6 @@ contains
     if (on_device() .and. n > 0) then
       compound_and = both(n)
     else
-      ! A device copy reaches ELSE when n <= 0.
-      !ERROR: 'host_only' may not be called in device code
       compound_and = host_only(n)
     end if
   end function
@@ -123,8 +133,6 @@ contains
   attributes(host,device) integer function compound_or(n)
     integer, value :: n
     if (on_device() .or. n > 0) then
-      ! A host copy reaches THEN when n > 0.
-      !ERROR: 'device_only' may not be called in host code
       compound_or = device_only(n)
     else
       compound_or = both(n)
@@ -207,8 +215,10 @@ contains
     call both_subroutine(n)
     if (on_device()) then
       call device_subroutine(n)
+      call host_subroutine(n)
     else
       call host_subroutine(n)
+      call device_subroutine(n)
     end if
   end subroutine
 
@@ -288,8 +298,10 @@ contains
     integer, intent(inout) :: n
     if (on_gpu()) then
       call device_only(n)
+      call host_only(n)
     else
       call host_only(n)
+      call device_only(n)
     end if
   end subroutine
 end module



More information about the flang-commits mailing list