[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
Fri Oct 2 16:47:51 PDT 2026


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

>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/6] [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/6] [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

>From 199346c560af33877e58a4c6f8a938ad22b44d0e Mon Sep 17 00:00:00 2001
From: Andre Kuhlenschmidt <akuhlenschmi at nvidia.com>
Date: Thu, 1 Oct 2026 11:52:38 -0700
Subject: [PATCH 3/6] [flang][CUDA] Handle module interfaces and expression
 call contexts

Recover CUDA callee attributes from separate module interfaces. Check call contexts in SELECT CASE selectors, I/O and allocation operands, and evaluated specification expressions. Preserve guarded-call allowances and existing I/O warnings.
---
 flang/lib/Semantics/check-cuda.cpp            |  73 +++-
 .../cuf-hostdevice-expression-context.cuf     | 317 ++++++++++++++++++
 .../CUDA/cuf-hostdevice-module-calls.cuf      |  52 +++
 3 files changed, 440 insertions(+), 2 deletions(-)
 create mode 100644 flang/test/Semantics/CUDA/cuf-hostdevice-expression-context.cuf
 create mode 100644 flang/test/Semantics/CUDA/cuf-hostdevice-module-calls.cuf

diff --git a/flang/lib/Semantics/check-cuda.cpp b/flang/lib/Semantics/check-cuda.cpp
index 13545803428ee7b..8234d8f5d55089c 100644
--- a/flang/lib/Semantics/check-cuda.cpp
+++ b/flang/lib/Semantics/check-cuda.cpp
@@ -116,6 +116,11 @@ struct DeviceExprChecker
     if (const Symbol * sym{x.GetInterfaceSymbol()}) {
       const Symbol &ultimate{sym->GetUltimate()};
       const auto *subp{ultimate.detailsIf<semantics::SubprogramDetails>()};
+      if (subp && subp->moduleInterface()) {
+        subp = subp->moduleInterface()
+                   ->GetUltimate()
+                   .detailsIf<semantics::SubprogramDetails>();
+      }
       if (subp) {
         if (const auto &stmtFunction{subp->stmtFunction()};
             stmtFunction && IsCUDADeviceContext(&ultimate.owner())) {
@@ -401,10 +406,41 @@ template <bool CUF_KERNEL> struct ActionStmtChecker {
   }
 };
 
+// Check analyzed expressions in statements and specification parts that do
+// not have a single typed expression for the entire construct. Do not descend
+// into an expression already checked, or into declarations of other entities
+// whose expressions are not evaluated on entry to this scope.
+struct DeviceExprVisitor {
+  SemanticsContext &context;
+  CallContext callContext;
+  template <typename A> bool Pre(const A &) { return true; }
+  template <typename A> void Post(const A &) {}
+  bool Pre(const parser::Expr &x) { return Check(x.typedExpr, x.source); }
+  bool Pre(const parser::Variable &x) {
+    return Check(x.typedExpr, x.GetSource());
+  }
+  bool Pre(const parser::InterfaceBlock &) { return false; }
+  bool Pre(const parser::DerivedTypeDef &) { return false; }
+  bool Pre(const parser::StmtFunctionStmt &) { return false; }
+
+private:
+  bool Check(const parser::TypedExpr &expr, parser::CharBlock source) {
+    if (expr) {
+      if (auto msg{DeviceExprChecker{context, callContext}(expr)}) {
+        context.Say(source, std::move(*msg));
+      }
+      return false;
+    }
+    return true;
+  }
+};
+
 template <bool IsCUFKernelDo> class DeviceContextChecker {
 public:
   explicit DeviceContextChecker(SemanticsContext &c) : context_{c} {}
-  void CheckSubprogram(const parser::Name &name, const parser::Block &body) {
+  template <typename A>
+  void CheckSubprogram(const parser::Name &name, const A &header,
+      const parser::SpecificationPart &spec, const parser::Block &body) {
     if (name.symbol) {
       const auto *subp{
           name.symbol->GetUltimate().detailsIf<SubprogramDetails>()};
@@ -422,6 +458,8 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
                 common::CUDASubprogramAttrs::HostDevice;
         callContext_ =
             isHostDevice ? CallContext::HostDevice : CallContext::Device;
+        CheckExpressions(header);
+        CheckExpressions(spec);
         Check(body);
       }
     }
@@ -433,6 +471,10 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
   }
 
 private:
+  template <typename A> void CheckExpressions(const A &x) {
+    DeviceExprVisitor visitor{context_, callContext_};
+    parser::Walk(x, visitor);
+  }
   void Check(const parser::ExecutionPartConstruct &epc) {
     common::visit(
         common::visitors{
@@ -466,12 +508,17 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
               Check(std::get<parser::Block>(x.value().t));
             },
             [&](const common::Indirection<parser::BlockConstruct> &x) {
+              CheckExpressions(
+                  std::get<parser::BlockSpecificationPart>(x.value().t));
               Check(std::get<parser::Block>(x.value().t));
             },
             [&](const common::Indirection<parser::IfConstruct> &x) {
               Check(x.value());
             },
             [&](const common::Indirection<parser::CaseConstruct> &x) {
+              CheckExpressions(
+                  std::get<parser::Statement<parser::SelectCaseStmt>>(
+                      x.value().t));
               const auto &caseList{
                   std::get<std::list<parser::CaseConstruct::Case>>(
                       x.value().t)};
@@ -551,8 +598,11 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
               ErrorInCUFKernel(source);
             },
             [&](const common::Indirection<parser::StopStmt> &) { return; },
-            [&](const common::Indirection<parser::PrintStmt> &) {},
+            [&](const common::Indirection<parser::PrintStmt> &x) {
+              CheckExpressions(x.value());
+            },
             [&](const common::Indirection<parser::WriteStmt> &x) {
+              CheckExpressions(x.value());
               if (x.value().format) { // Formatted write to '*' or '6'
                 if (std::holds_alternative<Fortran::parser::Star>(
                         x.value().format->u)) {
@@ -567,26 +617,39 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
               WarnIfNotInternal(x.value(), source);
             },
             [&](const common::Indirection<parser::CloseStmt> &x) {
+              CheckExpressions(x.value());
               WarnOnIoStmt(source);
             },
             [&](const common::Indirection<parser::EndfileStmt> &x) {
+              CheckExpressions(x.value());
               WarnOnIoStmt(source);
             },
             [&](const common::Indirection<parser::OpenStmt> &x) {
+              CheckExpressions(x.value());
               WarnOnIoStmt(source);
             },
             [&](const common::Indirection<parser::ReadStmt> &x) {
+              CheckExpressions(x.value());
               WarnIfNotInternal(x.value(), source);
             },
             [&](const common::Indirection<parser::InquireStmt> &x) {
+              CheckExpressions(x.value());
               WarnOnIoStmt(source);
             },
             [&](const common::Indirection<parser::RewindStmt> &x) {
+              CheckExpressions(x.value());
               WarnOnIoStmt(source);
             },
             [&](const common::Indirection<parser::BackspaceStmt> &x) {
+              CheckExpressions(x.value());
               WarnOnIoStmt(source);
             },
+            [&](const common::Indirection<parser::AllocateStmt> &x) {
+              CheckExpressions(x.value());
+            },
+            [&](const common::Indirection<parser::DeallocateStmt> &x) {
+              CheckExpressions(x.value());
+            },
             [&](const common::Indirection<parser::IfStmt> &x) {
               Check(x.value());
             },
@@ -697,6 +760,8 @@ void CUDAChecker::Enter(const parser::SubroutineSubprogram &x) {
   DeviceContextChecker<false>{context_}.CheckSubprogram(
       std::get<parser::Name>(
           std::get<parser::Statement<parser::SubroutineStmt>>(x.t).statement.t),
+      std::get<parser::Statement<parser::SubroutineStmt>>(x.t).statement,
+      std::get<parser::SpecificationPart>(x.t),
       std::get<parser::ExecutionPart>(x.t).v);
 }
 
@@ -704,12 +769,16 @@ void CUDAChecker::Enter(const parser::FunctionSubprogram &x) {
   DeviceContextChecker<false>{context_}.CheckSubprogram(
       std::get<parser::Name>(
           std::get<parser::Statement<parser::FunctionStmt>>(x.t).statement.t),
+      std::get<parser::Statement<parser::FunctionStmt>>(x.t).statement,
+      std::get<parser::SpecificationPart>(x.t),
       std::get<parser::ExecutionPart>(x.t).v);
 }
 
 void CUDAChecker::Enter(const parser::SeparateModuleSubprogram &x) {
   DeviceContextChecker<false>{context_}.CheckSubprogram(
       std::get<parser::Statement<parser::MpSubprogramStmt>>(x.t).statement.v,
+      std::get<parser::Statement<parser::MpSubprogramStmt>>(x.t).statement,
+      std::get<parser::SpecificationPart>(x.t),
       std::get<parser::ExecutionPart>(x.t).v);
 }
 
diff --git a/flang/test/Semantics/CUDA/cuf-hostdevice-expression-context.cuf b/flang/test/Semantics/CUDA/cuf-hostdevice-expression-context.cuf
new file mode 100644
index 000000000000000..4dec3bb845c9af7
--- /dev/null
+++ b/flang/test/Semantics/CUDA/cuf-hostdevice-expression-context.cuf
@@ -0,0 +1,317 @@
+! RUN: %python %S/../test_errors.py %s %flang_fc1
+! Check call targets in selectors, I/O, allocation operands, and evaluated
+! specification expressions. Guards apply to executable and BLOCK expressions.
+module expression_context
+  use cudadevice, only: on_device
+contains
+  pure integer function h()
+    h = 1
+  end function
+  attributes(device) pure integer function d()
+    d = 1
+  end function
+  attributes(host,device) pure integer function b()
+    b = 1
+  end function
+  attributes(host,device) subroutine select_h()
+    integer :: x
+    !ERROR: 'h' may not be called in device code
+    select case(h())
+    case(1)
+      x = 1
+    end select
+  end subroutine
+  attributes(host,device) subroutine select_d()
+    integer :: x
+    !ERROR: 'd' may not be called in host code
+    select case(d())
+    case(1)
+      x = 1
+    end select
+  end subroutine
+  attributes(host,device) subroutine select_b()
+    integer :: x
+    select case(b())
+    case(1)
+      x = 1
+    end select
+  end subroutine
+  attributes(host,device) subroutine print_h()
+    !ERROR: 'h' may not be called in device code
+    print *, h()
+  end subroutine
+  attributes(host,device) subroutine print_d()
+    !ERROR: 'd' may not be called in host code
+    print *, d()
+  end subroutine
+  attributes(host,device) subroutine print_b()
+    print *, b()
+  end subroutine
+  attributes(host,device) subroutine write_item_h()
+    !ERROR: 'h' may not be called in device code
+    write(*,*) h()
+  end subroutine
+  attributes(host,device) subroutine write_item_d()
+    !ERROR: 'd' may not be called in host code
+    write(*,*) d()
+  end subroutine
+  attributes(host,device) subroutine write_item_b()
+    write(*,*) b()
+  end subroutine
+  attributes(host,device) subroutine write_unit_h()
+    !ERROR: 'h' may not be called in device code
+    !WARNING: I/O statement might not be supported on device [-Wcuda-usage]
+    write(h(),*) 1
+  end subroutine
+  attributes(host,device) subroutine write_unit_d()
+    !ERROR: 'd' may not be called in host code
+    !WARNING: I/O statement might not be supported on device [-Wcuda-usage]
+    write(d(),*) 1
+  end subroutine
+  attributes(host,device) subroutine write_unit_b()
+    !WARNING: I/O statement might not be supported on device [-Wcuda-usage]
+    write(b(),*) 1
+  end subroutine
+  attributes(host,device) subroutine write_implied_do_h()
+    integer :: i
+    !ERROR: 'h' may not be called in device code
+    write(*,*) (i, i=1,h())
+  end subroutine
+  attributes(host,device) subroutine write_implied_do_d()
+    integer :: i
+    !ERROR: 'd' may not be called in host code
+    write(*,*) (i, i=1,d())
+  end subroutine
+  attributes(host,device) subroutine write_implied_do_b()
+    integer :: i
+    write(*,*) (i, i=1,b())
+  end subroutine
+  attributes(host,device) subroutine write_subscript_h()
+    integer :: x(4)
+    !ERROR: 'h' may not be called in device code
+    write(*,*) x(h())
+  end subroutine
+  attributes(host,device) subroutine write_subscript_d()
+    integer :: x(4)
+    !ERROR: 'd' may not be called in host code
+    write(*,*) x(d())
+  end subroutine
+  attributes(host,device) subroutine write_subscript_b()
+    integer :: x(4)
+    write(*,*) x(b())
+  end subroutine
+  attributes(host,device) subroutine read_subscript_h()
+    integer :: x(4)
+    !ERROR: 'h' may not be called in device code
+    !WARNING: I/O statement might not be supported on device [-Wcuda-usage]
+    read(*,*) x(h())
+  end subroutine
+  attributes(host,device) subroutine read_subscript_d()
+    integer :: x(4)
+    !ERROR: 'd' may not be called in host code
+    !WARNING: I/O statement might not be supported on device [-Wcuda-usage]
+    read(*,*) x(d())
+  end subroutine
+  attributes(host,device) subroutine read_subscript_b()
+    integer :: x(4)
+    !WARNING: I/O statement might not be supported on device [-Wcuda-usage]
+    read(*,*) x(b())
+  end subroutine
+  attributes(host,device) subroutine allocate_bound_h()
+    integer, allocatable :: a(:)
+    !ERROR: 'h' may not be called in device code
+    allocate(a(h()))
+  end subroutine
+  attributes(host,device) subroutine allocate_bound_d()
+    integer, allocatable :: a(:)
+    !ERROR: 'd' may not be called in host code
+    allocate(a(d()))
+  end subroutine
+  attributes(host,device) subroutine allocate_bound_b()
+    integer, allocatable :: a(:)
+    allocate(a(b()))
+  end subroutine
+  attributes(host,device) subroutine allocate_source_h()
+    integer, allocatable :: a(:)
+    !ERROR: 'h' may not be called in device code
+    allocate(a,source=[h()])
+  end subroutine
+  attributes(host,device) subroutine allocate_source_d()
+    integer, allocatable :: a(:)
+    !ERROR: 'd' may not be called in host code
+    allocate(a,source=[d()])
+  end subroutine
+  attributes(host,device) subroutine allocate_source_b()
+    integer, allocatable :: a(:)
+    allocate(a,source=[b()])
+  end subroutine
+  attributes(host,device) subroutine allocate_mold_h()
+    integer, allocatable :: a(:)
+    !ERROR: 'h' may not be called in device code
+    allocate(a,mold=[h()])
+  end subroutine
+  attributes(host,device) subroutine allocate_mold_d()
+    integer, allocatable :: a(:)
+    !ERROR: 'd' may not be called in host code
+    allocate(a,mold=[d()])
+  end subroutine
+  attributes(host,device) subroutine allocate_mold_b()
+    integer, allocatable :: a(:)
+    allocate(a,mold=[b()])
+  end subroutine
+  attributes(host,device) subroutine allocate_length_h()
+    character(:), allocatable :: text
+    !ERROR: 'h' may not be called in device code
+    allocate(character(h()) :: text)
+  end subroutine
+  attributes(host,device) subroutine allocate_length_d()
+    character(:), allocatable :: text
+    !ERROR: 'd' may not be called in host code
+    allocate(character(d()) :: text)
+  end subroutine
+  attributes(host,device) subroutine allocate_length_b()
+    character(:), allocatable :: text
+    allocate(character(b()) :: text)
+  end subroutine
+  attributes(host,device) subroutine deallocate_stat_h()
+    integer, allocatable :: a(:)
+    integer :: status(4)
+    !ERROR: 'h' may not be called in device code
+    deallocate(a,stat=status(h()))
+  end subroutine
+  attributes(host,device) subroutine deallocate_stat_d()
+    integer, allocatable :: a(:)
+    integer :: status(4)
+    !ERROR: 'd' may not be called in host code
+    deallocate(a,stat=status(d()))
+  end subroutine
+  attributes(host,device) subroutine deallocate_stat_b()
+    integer, allocatable :: a(:)
+    integer :: status(4)
+    deallocate(a,stat=status(b()))
+  end subroutine
+  attributes(host,device) subroutine automatic_bound_h()
+    !ERROR: 'h' may not be called in device code
+    integer :: a(h())
+    a = 1
+  end subroutine
+  attributes(host,device) subroutine automatic_bound_d()
+    !ERROR: 'd' may not be called in host code
+    integer :: a(d())
+    a = 1
+  end subroutine
+  attributes(host,device) subroutine automatic_bound_b()
+    integer :: a(b())
+    a = 1
+  end subroutine
+  attributes(host,device) subroutine character_length_h()
+    !ERROR: 'h' may not be called in device code
+    character(h()) :: text
+    text = 'x'
+  end subroutine
+  attributes(host,device) subroutine character_length_d()
+    !ERROR: 'd' may not be called in host code
+    character(d()) :: text
+    text = 'x'
+  end subroutine
+  attributes(host,device) subroutine character_length_b()
+    character(b()) :: text
+    text = 'x'
+  end subroutine
+  attributes(host,device) subroutine guarded()
+    integer :: x(4), i, status(4)
+    integer, allocatable :: a(:)
+    character(:), allocatable :: text
+    if (on_device()) then
+      select case(h())
+      case(1)
+        x = 1
+      end select
+      print *, h()
+      write(*,*) h()
+      !WARNING: I/O statement might not be supported on device [-Wcuda-usage]
+      write(h(),*) 1
+      write(*,*) (i, i=1,h())
+      write(*,*) x(h())
+      !WARNING: I/O statement might not be supported on device [-Wcuda-usage]
+      read(*,*) x(h())
+      allocate(a(h()))
+      allocate(a,source=[h()])
+      allocate(a,mold=[h()])
+      allocate(character(h()) :: text)
+      deallocate(a,stat=status(h()))
+      block
+        integer :: local(h())
+        character(h()) :: local_text
+        local = 1
+        local_text = 'x'
+      end block
+    else
+      select case(d())
+      case(1)
+        x = 1
+      end select
+      print *, d()
+      write(*,*) d()
+      !WARNING: I/O statement might not be supported on device [-Wcuda-usage]
+      write(d(),*) 1
+      write(*,*) (i, i=1,d())
+      write(*,*) x(d())
+      !WARNING: I/O statement might not be supported on device [-Wcuda-usage]
+      read(*,*) x(d())
+      allocate(a(d()))
+      allocate(a,source=[d()])
+      allocate(a,mold=[d()])
+      allocate(character(d()) :: text)
+      deallocate(a,stat=status(d()))
+      block
+        integer :: local(d())
+        character(d()) :: local_text
+        local = 1
+        local_text = 'x'
+      end block
+    end if
+  end subroutine
+  attributes(host,device) subroutine declarations_only()
+    integer :: unused
+    interface
+      subroutine host_interface(n,a)
+        import h
+        integer, intent(in) :: n
+        integer :: a(h())
+      end subroutine
+    end interface
+    unused() = h()
+  end subroutine
+  attributes(host,device) subroutine block_bound_h()
+    block
+      !ERROR: 'h' may not be called in device code
+      integer :: a(h())
+      a = 1
+    end block
+  end subroutine
+  !ERROR: 'h' may not be called in device code
+  attributes(host,device) character(h()) function header_h()
+    header_h = 'x'
+  end function
+  attributes(host,device) subroutine block_bound_d()
+    block
+      !ERROR: 'd' may not be called in host code
+      integer :: a(d())
+      a = 1
+    end block
+  end subroutine
+  !ERROR: 'd' may not be called in host code
+  attributes(host,device) character(d()) function header_d()
+    header_d = 'x'
+  end function
+  attributes(host,device) subroutine block_bound_b()
+    block
+      integer :: a(b())
+      a = 1
+    end block
+  end subroutine
+  attributes(host,device) character(b()) function header_b()
+    header_b = 'x'
+  end function
+end module
diff --git a/flang/test/Semantics/CUDA/cuf-hostdevice-module-calls.cuf b/flang/test/Semantics/CUDA/cuf-hostdevice-module-calls.cuf
new file mode 100644
index 000000000000000..bbfb1dd655aeb38
--- /dev/null
+++ b/flang/test/Semantics/CUDA/cuf-hostdevice-module-calls.cuf
@@ -0,0 +1,52 @@
+! RUN: %python %S/../test_errors.py %s %flang_fc1
+! Separate module procedure implementations inherit CUDA attributes from
+! their module interfaces, for both function references and CALL statements.
+module module_calls
+  use cudadevice, only: on_device
+  interface
+    attributes(host,device) module integer function both(n)
+      integer, intent(in) :: n
+    end function
+    module integer function host_only(n)
+      integer, intent(in) :: n
+    end function
+    attributes(device) module integer function device_only(n)
+      integer, intent(in) :: n
+    end function
+    attributes(host,device) module subroutine both_subroutine(n)
+      integer, intent(inout) :: n
+    end subroutine
+    attributes(host,device) module integer function caller(n)
+      integer, intent(inout) :: n
+    end function
+  end interface
+end module
+
+submodule(module_calls) implementations
+contains
+  module procedure both
+    both = n
+  end procedure
+  module procedure host_only
+    host_only = n
+  end procedure
+  module procedure device_only
+    device_only = n
+  end procedure
+  module procedure both_subroutine
+    n = n + 1
+  end procedure
+  module procedure caller
+    caller = both(n)
+    call both_subroutine(n)
+    !ERROR: 'host_only' may not be called in device code
+    caller = host_only(n)
+    !ERROR: 'device_only' may not be called in host code
+    caller = device_only(n)
+    if (on_device()) then
+      caller = device_only(n)
+    else
+      caller = host_only(n)
+    end if
+  end procedure
+end submodule

>From 44ce6faccebf0b3915a48a388aa580da7aa2a3b6 Mon Sep 17 00:00:00 2001
From: Andre Kuhlenschmidt <akuhlenschmi at nvidia.com>
Date: Thu, 1 Oct 2026 14:54:33 -0700
Subject: [PATCH 4/6] [flang][CUDA] Allow calls after ON_DEVICE return guards

---
 flang/lib/Semantics/check-cuda.cpp            |  37 ++++-
 .../CUDA/cuf-hostdevice-return-guard.cuf      | 157 ++++++++++++++++++
 2 files changed, 189 insertions(+), 5 deletions(-)
 create mode 100644 flang/test/Semantics/CUDA/cuf-hostdevice-return-guard.cuf

diff --git a/flang/lib/Semantics/check-cuda.cpp b/flang/lib/Semantics/check-cuda.cpp
index 8234d8f5d55089c..340667294b8e16d 100644
--- a/flang/lib/Semantics/check-cuda.cpp
+++ b/flang/lib/Semantics/check-cuda.cpp
@@ -7,6 +7,7 @@
 //===----------------------------------------------------------------------===//
 
 #include "check-cuda.h"
+#include "flang/Common/restorer.h"
 #include "flang/Common/template.h"
 #include "flang/Evaluate/fold.h"
 #include "flang/Evaluate/tools.h"
@@ -465,6 +466,8 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
     }
   }
   void Check(const parser::Block &block) {
+    // A guard established by a return affects only this block's continuation.
+    auto restore{common::Restorer{callContext_, callContext_}};
     for (const auto &epc : block) {
       Check(epc);
     }
@@ -680,8 +683,9 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
     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 &elseIfBlocks{
+        std::get<std::list<parser::IfConstruct::ElseIfBlock>>(ic.t)};
+    for (const auto &eib : elseIfBlocks) {
       const auto &eIfS{std::get<parser::Statement<parser::ElseIfStmt>>(eib.t)};
       const auto &elseIfCondition{
           std::get<parser::ScalarLogicalExpr>(eIfS.statement.t)};
@@ -689,11 +693,18 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
       AllowGuardedCalls(elseIfCondition);
       Check(std::get<parser::Block>(eib.t));
     }
-    if (const auto &eb{
-            std::get<std::optional<parser::IfConstruct::ElseBlock>>(ic.t)}) {
+    const auto &eb{
+        std::get<std::optional<parser::IfConstruct::ElseBlock>>(ic.t)};
+    if (eb) {
       Check(std::get<parser::Block>(eb->t));
     }
     callContext_ = incoming;
+    if (elseIfBlocks.empty() &&
+        (HasUnconditionalReturn(std::get<parser::Block>(ic.t)) ||
+            (eb && HasUnconditionalReturn(std::get<parser::Block>(eb->t))))) {
+      // The remaining statements are an implicit opposite branch arm.
+      AllowGuardedCalls(condition);
+    }
   }
   void Check(const parser::IfStmt &is) {
     const CallContext incoming{callContext_};
@@ -703,7 +714,23 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
     CheckUnwrappedExpr(context_, uS.source, condition, callContext_);
     AllowGuardedCalls(condition);
     Check(uS.statement, uS.source);
-    callContext_ = incoming;
+    if (!IsReturn(uS.statement)) {
+      callContext_ = incoming;
+    }
+  }
+  static bool IsReturn(const parser::ActionStmt &stmt) {
+    const auto *ret{parser::Unwrap<parser::ReturnStmt>(stmt)};
+    return ret && !ret->v;
+  }
+  static bool HasUnconditionalReturn(const parser::Block &block) {
+    for (const auto &epc : block) {
+      if (const auto *stmt{parser::Unwrap<parser::ActionStmt>(epc)}) {
+        if (IsReturn(*stmt)) {
+          return true;
+        }
+      }
+    }
+    return false;
   }
   void AllowGuardedCalls(const parser::ScalarLogicalExpr &condition) {
     // Accept either kind of callee in either arm of an ON_DEVICE guard.
diff --git a/flang/test/Semantics/CUDA/cuf-hostdevice-return-guard.cuf b/flang/test/Semantics/CUDA/cuf-hostdevice-return-guard.cuf
new file mode 100644
index 000000000000000..6d3ca577cb588e5
--- /dev/null
+++ b/flang/test/Semantics/CUDA/cuf-hostdevice-return-guard.cuf
@@ -0,0 +1,157 @@
+! RUN: %python %S/../test_errors.py %s %flang_fc1
+! The remainder of a block is guarded when an ON_DEVICE branch returns.
+! As for explicit branch arms, checking the appropriate target is deferred.
+module return_guard
+  use cudadevice, only: on_device, on_gpu => on_device
+contains
+  subroutine host_only()
+  end subroutine
+  attributes(device) subroutine device_only()
+  end subroutine
+  attributes(host,device) subroutine both()
+    ! Review example: the device version returns before the host-only call.
+    if (on_device()) return
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine negated()
+    if (.not. on_gpu()) return
+    call device_only()
+  end subroutine
+  attributes(host,device) subroutine block_return()
+    if (on_device()) then
+      return
+    end if
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine return_before_last_statement()
+    if (on_device()) then
+      call device_only()
+      return
+      call device_only()
+    end if
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine else_return_before_last_statement()
+    if (on_device()) then
+      call device_only()
+    else
+      return
+      call host_only()
+    end if
+    call device_only()
+  end subroutine
+  attributes(host,device) subroutine nested_conditional_return(choose)
+    logical, intent(in) :: choose
+    if (on_device()) then
+      if (choose) then
+        return
+      end if
+      call device_only()
+    end if
+    ! A nested conditional return does not eliminate fallthrough.
+    !ERROR: 'host_only' may not be called in device code
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine else_return()
+    if (on_device()) then
+      call device_only()
+    else
+      return
+    end if
+    call device_only()
+  end subroutine
+  attributes(host,device) subroutine nested(choose)
+    logical, intent(in) :: choose
+    if (choose) then
+      if (on_device()) return
+      call host_only()
+    end if
+    ! The outer IF can bypass the return guard.
+    !ERROR: 'host_only' may not be called in device code
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine loop_return(choose)
+    logical, intent(in) :: choose
+    do while (choose)
+      if (on_device()) return
+      call host_only()
+    end do
+    ! The loop need not execute.
+    !ERROR: 'host_only' may not be called in device code
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine ordinary_return(choose)
+    logical, intent(in) :: choose
+    if (choose) return
+    !ERROR: 'host_only' may not be called in device code
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine conditional_return(choose)
+    logical, intent(in) :: choose
+    if (on_device()) then
+      if (choose) return
+    end if
+    ! The device branch can fall through.
+    !ERROR: 'host_only' may not be called in device code
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine unrelated_label()
+10  continue
+    if (on_device()) return
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine bypass()
+    goto 10
+    if (on_device()) return
+    ! Accept the continuation guard; analyzing jumps is deferred.
+10  call host_only()
+  end subroutine
+  attributes(host,device) subroutine bypass_block()
+    !WARNING: Label '10' is in a construct that should not be used as a branch target here [-Wbranch-into-construct]
+    goto 10
+    if (on_device()) then
+      return
+10  end if
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine else_if_bypass(choose)
+    logical, intent(in) :: choose
+    if (choose) then
+      !ERROR: 'host_only' may not be called in device code
+      call host_only()
+    else if (on_device()) then
+      return
+    end if
+    ! The first branch can bypass the ON_DEVICE condition.
+    !ERROR: 'host_only' may not be called in device code
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine select_dispatch(result)
+    integer, intent(out) :: result
+    ! The selector example from the review is already diagnosed.
+    !ERROR: 'device_value' may not be called in host code
+    select case (device_value())
+    case (1)
+      result = 10
+    case default
+      result = 20
+    end select
+  end subroutine
+  attributes(device) integer function device_value()
+    device_value = 1
+  end function
+end module
+
+module fake_return_guard
+contains
+  attributes(host,device) logical function on_device()
+    on_device = .false.
+  end function
+  subroutine host_only()
+  end subroutine
+  attributes(host,device) subroutine caller()
+    if (on_device()) return
+    !ERROR: 'host_only' may not be called in device code
+    call host_only()
+  end subroutine
+end module

>From 39f4e8af7a20a286ec8278088bb314bb1fe3b4e0 Mon Sep 17 00:00:00 2001
From: Andre Kuhlenschmidt <akuhlenschmi at nvidia.com>
Date: Thu, 1 Oct 2026 16:01:49 -0700
Subject: [PATCH 5/6] [flang][CUDA] Extend return guards through blocks and IF
 arms

---
 flang/lib/Semantics/check-cuda.cpp            | 39 ++++++--
 .../CUDA/cuf-hostdevice-return-guard.cuf      | 90 ++++++++++++++++++-
 2 files changed, 117 insertions(+), 12 deletions(-)

diff --git a/flang/lib/Semantics/check-cuda.cpp b/flang/lib/Semantics/check-cuda.cpp
index 340667294b8e16d..6729fdbf683a819 100644
--- a/flang/lib/Semantics/check-cuda.cpp
+++ b/flang/lib/Semantics/check-cuda.cpp
@@ -468,12 +468,15 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
   void Check(const parser::Block &block) {
     // A guard established by a return affects only this block's continuation.
     auto restore{common::Restorer{callContext_, callContext_}};
+    CheckStatements(block);
+  }
+
+private:
+  void CheckStatements(const parser::Block &block) {
     for (const auto &epc : block) {
       Check(epc);
     }
   }
-
-private:
   template <typename A> void CheckExpressions(const A &x) {
     DeviceExprVisitor visitor{context_, callContext_};
     parser::Walk(x, visitor);
@@ -513,7 +516,9 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
             [&](const common::Indirection<parser::BlockConstruct> &x) {
               CheckExpressions(
                   std::get<parser::BlockSpecificationPart>(x.value().t));
-              Check(std::get<parser::Block>(x.value().t));
+              // BLOCK executes unconditionally, so its guard also applies to
+              // the enclosing continuation.
+              CheckStatements(std::get<parser::Block>(x.value().t));
             },
             [&](const common::Indirection<parser::IfConstruct> &x) {
               Check(x.value());
@@ -698,12 +703,10 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
     if (eb) {
       Check(std::get<parser::Block>(eb->t));
     }
-    callContext_ = incoming;
-    if (elseIfBlocks.empty() &&
-        (HasUnconditionalReturn(std::get<parser::Block>(ic.t)) ||
-            (eb && HasUnconditionalReturn(std::get<parser::Block>(eb->t))))) {
-      // The remaining statements are an implicit opposite branch arm.
-      AllowGuardedCalls(condition);
+    // An unconditional return in any arm permits continuation under this IF's
+    // guard.
+    if (!HasReturningArm(ic)) {
+      callContext_ = incoming;
     }
   }
   void Check(const parser::IfStmt &is) {
@@ -728,10 +731,28 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
         if (IsReturn(*stmt)) {
           return true;
         }
+      } else if (const auto *bc{parser::Unwrap<parser::BlockConstruct>(epc)}) {
+        if (HasUnconditionalReturn(std::get<parser::Block>(bc->t))) {
+          return true;
+        }
       }
     }
     return false;
   }
+  static bool HasReturningArm(const parser::IfConstruct &ic) {
+    if (HasUnconditionalReturn(std::get<parser::Block>(ic.t))) {
+      return true;
+    }
+    for (const auto &eib :
+        std::get<std::list<parser::IfConstruct::ElseIfBlock>>(ic.t)) {
+      if (HasUnconditionalReturn(std::get<parser::Block>(eib.t))) {
+        return true;
+      }
+    }
+    const auto &eb{
+        std::get<std::optional<parser::IfConstruct::ElseBlock>>(ic.t)};
+    return eb && HasUnconditionalReturn(std::get<parser::Block>(eb->t));
+  }
   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.
diff --git a/flang/test/Semantics/CUDA/cuf-hostdevice-return-guard.cuf b/flang/test/Semantics/CUDA/cuf-hostdevice-return-guard.cuf
index 6d3ca577cb588e5..0f20d0bef880403 100644
--- a/flang/test/Semantics/CUDA/cuf-hostdevice-return-guard.cuf
+++ b/flang/test/Semantics/CUDA/cuf-hostdevice-return-guard.cuf
@@ -1,5 +1,6 @@
 ! RUN: %python %S/../test_errors.py %s %flang_fc1
-! The remainder of a block is guarded when an ON_DEVICE branch returns.
+! The remainder of a block is guarded when an ON_DEVICE branch has a RETURN.
+! Only direct RETURNs and RETURNs in unconditional BLOCKs qualify.
 ! As for explicit branch arms, checking the appropriate target is deferred.
 module return_guard
   use cudadevice, only: on_device, on_gpu => on_device
@@ -23,6 +24,90 @@ contains
     end if
     call host_only()
   end subroutine
+  attributes(host,device) subroutine unconditional_block()
+    block
+      if (on_device()) return
+    end block
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine nested_block_return()
+    if (on_device()) then
+      block
+        block
+          return
+        end block
+      end block
+    end if
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine return_with_else_if(choose)
+    logical, intent(in) :: choose
+    if (on_device()) then
+      return
+    else if (choose) then
+      call host_only()
+    end if
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine else_if_guard_return(choose)
+    logical, intent(in) :: choose
+    if (choose) then
+      call both()
+    else if (on_device()) then
+      return
+    end if
+    ! Accept this permissively; the first arm can bypass the guard.
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine nested_else_if_return(choose)
+    logical, intent(in) :: choose
+    if (on_device()) then
+      if (choose) then
+        call device_only()
+      else if (.not. choose) then
+        return
+      end if
+    end if
+    ! The nested IF does not unconditionally return.
+    !ERROR: 'host_only' may not be called in device code
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine nested_else_return(choose)
+    logical, intent(in) :: choose
+    if (on_device()) then
+      if (choose) then
+        call device_only()
+      else
+        return
+      end if
+    end if
+    ! The nested ELSE does not make the enclosing arm return unconditionally.
+    !ERROR: 'host_only' may not be called in device code
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine guarded_loop_return(choose)
+    logical, intent(in) :: choose
+    if (on_device()) then
+      do while (choose)
+        return
+      end do
+    end if
+    ! Loops may not execute and are not traversed by the return scan.
+    !ERROR: 'host_only' may not be called in device code
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine conditional_block_guard(choose)
+    logical, intent(in) :: choose
+    if (choose) then
+      block
+        if (on_device()) return
+      end block
+      call host_only()
+    end if
+    ! The enclosing conditional can bypass the whole BLOCK.
+    !ERROR: 'host_only' may not be called in device code
+    call host_only()
+  end subroutine
   attributes(host,device) subroutine return_before_last_statement()
     if (on_device()) then
       call device_only()
@@ -122,8 +207,7 @@ contains
     else if (on_device()) then
       return
     end if
-    ! The first branch can bypass the ON_DEVICE condition.
-    !ERROR: 'host_only' may not be called in device code
+    ! The continuation is accepted even if the first arm bypasses the guard.
     call host_only()
   end subroutine
   attributes(host,device) subroutine select_dispatch(result)

>From 961db8feb00dff386b2f1504b34740bac6846631 Mon Sep 17 00:00:00 2001
From: Andre Kuhlenschmidt <akuhlenschmi at nvidia.com>
Date: Fri, 2 Oct 2026 16:43:09 -0700
Subject: [PATCH 6/6] [flang][CUDA] Address host-device call checking review
 feedback

---
 flang/include/flang/Evaluate/call.h           |   9 +
 flang/lib/Semantics/check-cuda.cpp            | 147 ++++++++--------
 flang/lib/Semantics/expression.cpp            |  22 ++-
 .../CUDA/cuf-hostdevice-c-pointers.cuf        |  47 +++++
 .../CUDA/cuf-hostdevice-call-context.cuf      |   4 +-
 .../cuf-hostdevice-expression-context.cuf     |  30 ++++
 .../cuf-hostdevice-generic-diagnostics.cuf    | 162 ++++++++++++++++++
 .../CUDA/cuf-hostdevice-return-guard.cuf      | 130 ++++++++++++++
 .../CUDA/cuf-hostdevice-stop-calls.cuf        |  57 ++++++
 9 files changed, 528 insertions(+), 80 deletions(-)
 create mode 100644 flang/test/Semantics/CUDA/cuf-hostdevice-c-pointers.cuf
 create mode 100644 flang/test/Semantics/CUDA/cuf-hostdevice-generic-diagnostics.cuf
 create mode 100644 flang/test/Semantics/CUDA/cuf-hostdevice-stop-calls.cuf

diff --git a/flang/include/flang/Evaluate/call.h b/flang/include/flang/Evaluate/call.h
index f04ec6c373c88d1..6bdccc5c98a1193 100644
--- a/flang/include/flang/Evaluate/call.h
+++ b/flang/include/flang/Evaluate/call.h
@@ -316,6 +316,11 @@ struct ProcedureDesignator {
   const Symbol *GetInterfaceSymbol() const;
 
   std::string GetName() const;
+  std::optional<parser::CharBlock> genericName() const { return genericName_; }
+  ProcedureDesignator &set_genericName(parser::CharBlock name) {
+    genericName_ = name;
+    return *this;
+  }
   std::optional<DynamicType> GetType() const;
   int Rank() const;
   bool IsElemental() const;
@@ -327,6 +332,10 @@ struct ProcedureDesignator {
   std::variant<SpecificIntrinsic, SymbolRef,
       common::CopyableIndirection<Component>>
       u;
+
+private:
+  // Diagnostic metadata only; equality and lowering use the resolved specific.
+  std::optional<parser::CharBlock> genericName_;
 };
 
 using Chevrons = std::vector<Expr<SomeType>>;
diff --git a/flang/lib/Semantics/check-cuda.cpp b/flang/lib/Semantics/check-cuda.cpp
index 6729fdbf683a819..fdd0a4b64f7bcb6 100644
--- a/flang/lib/Semantics/check-cuda.cpp
+++ b/flang/lib/Semantics/check-cuda.cpp
@@ -68,6 +68,11 @@ static const llvm::StringSet<> warpFunctions_ = {"match_all_syncjj",
     "match_any_syncjj", "match_any_syncjx", "match_any_syncjf",
     "match_any_syncjd"};
 
+// These builtin procedures lower to inline pointer operations on either target.
+static const llvm::StringSet<> inlinePointerFunctions_ = {"c_associated_c_ptr",
+    "c_associated_c_funptr", "__builtin_c_ptr_eq", "__builtin_c_ptr_ne",
+    "__builtin_c_devptr_eq", "__builtin_c_devptr_ne"};
+
 enum class CallContext { Device, HostDevice, GuardedHostDevice };
 
 // Match the BIND(C) intrinsic-module procedure intercepted by CUDA lowering,
@@ -138,8 +143,7 @@ struct DeviceExprChecker
             }
             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 CallError(x, true);
             }
             return {};
           }
@@ -160,6 +164,11 @@ struct DeviceExprChecker
       if (mod && mod->name() == "ieee_arithmetic") {
         return {};
       }
+      if (mod && mod->attrs().test(Attr::INTRINSIC) &&
+          mod->name() == "__fortran_builtins" &&
+          inlinePointerFunctions_.contains(ultimate.name().ToString())) {
+        return {};
+      }
     } else if (x.GetSpecificIntrinsic()) {
       // TODO(CUDA): Check for unsupported intrinsics here
       return {};
@@ -168,12 +177,32 @@ struct DeviceExprChecker
     if (callContext_ == CallContext::GuardedHostDevice) {
       return {};
     }
-    return parser::MessageFormattedText(
-        "'%s' may not be called in device code"_err_en_US, x.GetName());
+    return CallError(x, false);
   }
 
   SemanticsContext &context_;
   CallContext callContext_{CallContext::Device};
+
+private:
+  static Result CallError(
+      const evaluate::ProcedureDesignator &proc, bool inHostCode) {
+    if (auto generic{proc.genericName()}) {
+      const Symbol *specific{proc.GetInterfaceSymbol()};
+      const std::string specificName{specific
+              ? specific->GetUltimate().name().ToString()
+              : proc.GetName()};
+      if (*generic != specificName) {
+        return parser::MessageFormattedText(inHostCode
+                ? "'%s' (specific procedure '%s') may not be called in host code"_err_en_US
+                : "'%s' (specific procedure '%s') may not be called in device code"_err_en_US,
+            *generic, specificName);
+      }
+    }
+    return parser::MessageFormattedText(inHostCode
+            ? "'%s' may not be called in host code"_err_en_US
+            : "'%s' may not be called in device code"_err_en_US,
+        proc.GetName());
+  }
 };
 
 static bool IsHostArray(const Symbol &symbol) {
@@ -250,18 +279,9 @@ struct FindHostArray
   }
 };
 
-template <typename A>
-static MaybeMsg CheckUnwrappedExpr(SemanticsContext &context, const A &x,
-    CallContext callContext = CallContext::Device) {
-  if (const auto *expr{parser::Unwrap<parser::Expr>(x)}) {
-    return DeviceExprChecker{context, callContext}(expr->typedExpr);
-  }
-  return {};
-}
-
 template <typename A>
 static void CheckUnwrappedExpr(SemanticsContext &context, SourceName at,
-    const A &x, CallContext callContext = CallContext::Device) {
+    const A &x, CallContext callContext) {
   if (const auto *expr{parser::Unwrap<parser::Expr>(x)}) {
     if (auto msg{DeviceExprChecker{context, callContext}(expr->typedExpr)}) {
       context.Say(at, std::move(*msg));
@@ -271,8 +291,8 @@ static void CheckUnwrappedExpr(SemanticsContext &context, SourceName at,
 
 template <bool CUF_KERNEL> struct ActionStmtChecker {
   template <typename A>
-  static MaybeMsg WhyNotOk(SemanticsContext &context, const A &x,
-      CallContext callContext = CallContext::Device) {
+  static MaybeMsg WhyNotOk(
+      SemanticsContext &context, const A &x, CallContext callContext) {
     if constexpr (ConstraintTrait<A>) {
       return WhyNotOk(context, x.thing, callContext);
     } else if constexpr (WrapperTrait<A>) {
@@ -288,14 +308,12 @@ template <bool CUF_KERNEL> struct ActionStmtChecker {
   }
   template <typename A>
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const common::Indirection<A> &x,
-      CallContext callContext = CallContext::Device) {
+      const common::Indirection<A> &x, CallContext callContext) {
     return WhyNotOk(context, x.value(), callContext);
   }
   template <typename... As>
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const std::variant<As...> &x,
-      CallContext callContext = CallContext::Device) {
+      const std::variant<As...> &x, CallContext callContext) {
     return common::visit(
         [&context, callContext](
             const auto &x) { return WhyNotOk(context, x, callContext); },
@@ -303,8 +321,7 @@ template <bool CUF_KERNEL> struct ActionStmtChecker {
   }
   template <std::size_t J = 0, typename... As>
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const std::tuple<As...> &x,
-      CallContext callContext = CallContext::Device) {
+      const std::tuple<As...> &x, CallContext callContext) {
     if constexpr (J == sizeof...(As)) {
       return {};
     } else if (auto msg{WhyNotOk(context, std::get<J>(x), callContext)}) {
@@ -315,7 +332,7 @@ template <bool CUF_KERNEL> struct ActionStmtChecker {
   }
   template <typename A>
   static MaybeMsg WhyNotOk(SemanticsContext &context, const std::list<A> &x,
-      CallContext callContext = CallContext::Device) {
+      CallContext callContext) {
     for (const auto &y : x) {
       if (MaybeMsg result{WhyNotOk(context, y, callContext)}) {
         return result;
@@ -325,7 +342,7 @@ template <bool CUF_KERNEL> struct ActionStmtChecker {
   }
   template <typename A>
   static MaybeMsg WhyNotOk(SemanticsContext &context, const std::optional<A> &x,
-      CallContext callContext = CallContext::Device) {
+      CallContext callContext) {
     if (x) {
       return WhyNotOk(context, *x, callContext);
     } else {
@@ -334,64 +351,33 @@ template <bool CUF_KERNEL> struct ActionStmtChecker {
   }
   template <typename A>
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::UnlabeledStatement<A> &x,
-      CallContext callContext = CallContext::Device) {
+      const parser::UnlabeledStatement<A> &x, CallContext callContext) {
     return WhyNotOk(context, x.statement, callContext);
   }
   template <typename A>
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::Statement<A> &x,
-      CallContext callContext = CallContext::Device) {
+      const parser::Statement<A> &x, CallContext callContext) {
     return WhyNotOk(context, x.statement, callContext);
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::AllocateStmt &,
-      CallContext callContext = CallContext::Device) {
-    return {}; // AllocateObjects are checked elsewhere
-  }
-  static MaybeMsg WhyNotOk(SemanticsContext &context,
-      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 &,
-      CallContext callContext = CallContext::Device) {
-    return {}; // AllocateObjects are checked elsewhere
-  }
-  static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::AssignmentStmt &x,
-      CallContext callContext = CallContext::Device) {
+      const parser::AssignmentStmt &x, CallContext callContext) {
     return DeviceExprChecker{context, callContext}(x.typedAssignment);
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context, const parser::CallStmt &x,
-      CallContext callContext = CallContext::Device) {
+      CallContext callContext) {
     return DeviceExprChecker{context, callContext}(x.typedCall);
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::ContinueStmt &,
-      CallContext callContext = CallContext::Device) {
+      const parser::ContinueStmt &, CallContext callContext) {
     return {};
   }
-  static MaybeMsg WhyNotOk(SemanticsContext &, const parser::PauseStmt &,
-      CallContext callContext = CallContext::Device) {
+  static MaybeMsg WhyNotOk(
+      SemanticsContext &, const parser::PauseStmt &, CallContext callContext) {
     return parser::MessageFormattedText{
         "device subprograms may not contain PAUSE statements"_err_en_US};
   }
-  static MaybeMsg WhyNotOk(SemanticsContext &context, const parser::IfStmt &x,
-      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,
-        callContext);
-  }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::NullifyStmt &x,
-      CallContext callContext = CallContext::Device) {
+      const parser::NullifyStmt &x, CallContext callContext) {
     for (const auto &y : x.v) {
       if (MaybeMsg result{
               DeviceExprChecker{context, callContext}(y.typedExpr)}) {
@@ -401,8 +387,7 @@ template <bool CUF_KERNEL> struct ActionStmtChecker {
     return {};
   }
   static MaybeMsg WhyNotOk(SemanticsContext &context,
-      const parser::PointerAssignmentStmt &x,
-      CallContext callContext = CallContext::Device) {
+      const parser::PointerAssignmentStmt &x, CallContext callContext) {
     return DeviceExprChecker{context, callContext}(x.typedAssignment);
   }
 };
@@ -605,7 +590,9 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
             [&](const common::Indirection<parser::GotoStmt> &) {
               ErrorInCUFKernel(source);
             },
-            [&](const common::Indirection<parser::StopStmt> &) { return; },
+            [&](const common::Indirection<parser::StopStmt> &x) {
+              CheckExpressions(x.value());
+            },
             [&](const common::Indirection<parser::PrintStmt> &x) {
               CheckExpressions(x.value());
             },
@@ -687,10 +674,14 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
     const auto &condition{std::get<parser::ScalarLogicalExpr>(ifS.statement.t)};
     CheckUnwrappedExpr(context_, ifS.source, condition, callContext_);
     AllowGuardedCalls(condition);
+    const CallContext branchContext{callContext_};
     Check(std::get<parser::Block>(ic.t));
     const auto &elseIfBlocks{
         std::get<std::list<parser::IfConstruct::ElseIfBlock>>(ic.t)};
     for (const auto &eib : elseIfBlocks) {
+      // An ELSEIF guard applies to its own arm and its plain ELSE, not to a
+      // later unrelated ELSEIF. Preserve a guard on the initial IF, if any.
+      callContext_ = branchContext;
       const auto &eIfS{std::get<parser::Statement<parser::ElseIfStmt>>(eib.t)};
       const auto &elseIfCondition{
           std::get<parser::ScalarLogicalExpr>(eIfS.statement.t)};
@@ -703,9 +694,8 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
     if (eb) {
       Check(std::get<parser::Block>(eb->t));
     }
-    // An unconditional return in any arm permits continuation under this IF's
-    // guard.
-    if (!HasReturningArm(ic)) {
+    // A return only extends a guard when its arm is associated with ON_DEVICE.
+    if (!HasReturningOnDeviceArm(ic)) {
       callContext_ = incoming;
     }
   }
@@ -739,19 +729,30 @@ template <bool IsCUFKernelDo> class DeviceContextChecker {
     }
     return false;
   }
-  static bool HasReturningArm(const parser::IfConstruct &ic) {
-    if (HasUnconditionalReturn(std::get<parser::Block>(ic.t))) {
+  bool HasReturningOnDeviceArm(const parser::IfConstruct &ic) {
+    const auto &ifS{std::get<parser::Statement<parser::IfThenStmt>>(ic.t)};
+    bool checksOnDevice{ChecksOnDevice(
+        context_, std::get<parser::ScalarLogicalExpr>(ifS.statement.t))};
+    if (checksOnDevice &&
+        HasUnconditionalReturn(std::get<parser::Block>(ic.t))) {
       return true;
     }
     for (const auto &eib :
         std::get<std::list<parser::IfConstruct::ElseIfBlock>>(ic.t)) {
-      if (HasUnconditionalReturn(std::get<parser::Block>(eib.t))) {
+      const auto &elseIfS{
+          std::get<parser::Statement<parser::ElseIfStmt>>(eib.t)};
+      checksOnDevice = ChecksOnDevice(
+          context_, std::get<parser::ScalarLogicalExpr>(elseIfS.statement.t));
+      if (checksOnDevice &&
+          HasUnconditionalReturn(std::get<parser::Block>(eib.t))) {
         return true;
       }
     }
     const auto &eb{
         std::get<std::optional<parser::IfConstruct::ElseBlock>>(ic.t)};
-    return eb && HasUnconditionalReturn(std::get<parser::Block>(eb->t));
+    // A plain ELSE is associated with the preceding IF or ELSEIF condition.
+    return checksOnDevice && eb &&
+        HasUnconditionalReturn(std::get<parser::Block>(eb->t));
   }
   void AllowGuardedCalls(const parser::ScalarLogicalExpr &condition) {
     // Accept either kind of callee in either arm of an ON_DEVICE guard.
diff --git a/flang/lib/Semantics/expression.cpp b/flang/lib/Semantics/expression.cpp
index f514adc0f476d02..e74e6a50f386154 100644
--- a/flang/lib/Semantics/expression.cpp
+++ b/flang/lib/Semantics/expression.cpp
@@ -2784,6 +2784,13 @@ auto ExpressionAnalyzer::AnalyzeProcedureComponentRef(
             *sym);
         return std::nullopt;
       }
+      const bool isGeneric{sym->has<semantics::GenericDetails>()};
+      auto withGenericName{[&](ProcedureDesignator proc) {
+        if (isGeneric) {
+          proc.set_genericName(sc.Component().source);
+        }
+        return proc;
+      }};
       if (auto *dtExpr{UnwrapExpr<Expr<SomeDerived>>(*base)}) {
         if (sym->has<semantics::GenericDetails>()) {
           const Symbol &generic{*sym};
@@ -2867,7 +2874,8 @@ auto ExpressionAnalyzer::AnalyzeProcedureComponentRef(
                 GetBindingResolution(dtExpr->GetType(), *sym)}) {
           AddPassArg(arguments, std::move(*dtExpr), *sym, false);
           return CalleeAndArguments{
-              ProcedureDesignator{*resolution}, std::move(arguments)};
+              withGenericName(ProcedureDesignator{*resolution}),
+              std::move(arguments)};
         } else if (dataRef.has_value()) {
           if (ExtractCoarrayRef(*dataRef)) {
             if (IsProcedurePointer(*sym)) {
@@ -2884,7 +2892,7 @@ auto ExpressionAnalyzer::AnalyzeProcedureComponentRef(
               if (auto component{CreateComponent(std::move(*dataRef), *sym,
                       *dtSpec->scope(), /*C919bAlreadyEnforced=*/true)}) {
                 return CalleeAndArguments{
-                    ProcedureDesignator{std::move(*component)},
+                    withGenericName(ProcedureDesignator{std::move(*component)}),
                     std::move(arguments)};
               }
             }
@@ -2896,7 +2904,8 @@ auto ExpressionAnalyzer::AnalyzeProcedureComponentRef(
                 Expr<SomeDerived>{Designator<SomeDerived>{std::move(*dataRef)}},
                 *sym);
             return CalleeAndArguments{
-                ProcedureDesignator{*sym}, std::move(arguments)};
+                withGenericName(ProcedureDesignator{*sym}),
+                std::move(arguments)};
           }
         }
       }
@@ -3617,8 +3626,11 @@ auto ExpressionAnalyzer::GetCalleeAndArguments(const parser::Name &name,
           semantics::SymbolRef{*resolution}, std::move(arguments)};
     }
   } else if (IsProcedure(*resolution)) {
-    return CalleeAndArguments{
-        ProcedureDesignator{*resolution}, std::move(arguments)};
+    ProcedureDesignator proc{*resolution};
+    if (isGenericInterface) {
+      proc.set_genericName(name.source);
+    }
+    return CalleeAndArguments{std::move(proc), std::move(arguments)};
   }
   if (!context_.HasError(*resolution)) {
     AttachDeclaration(
diff --git a/flang/test/Semantics/CUDA/cuf-hostdevice-c-pointers.cuf b/flang/test/Semantics/CUDA/cuf-hostdevice-c-pointers.cuf
new file mode 100644
index 000000000000000..57be393c63af60a
--- /dev/null
+++ b/flang/test/Semantics/CUDA/cuf-hostdevice-c-pointers.cuf
@@ -0,0 +1,47 @@
+! RUN: %python %S/../test_errors.py %s %flang_fc1
+! These intrinsic-module procedures lower inline on both host and device.
+module c_pointer_calls
+  use iso_c_binding, only: c_ptr, c_funptr, c_associated, operator(==), operator(/=)
+  use __fortran_builtins, only: c_devptr => __builtin_c_devptr
+contains
+  attributes(host,device) subroutine both(p, q, f, g, d, e, result)
+    type(c_ptr), value :: p, q
+    type(c_funptr), value :: f, g
+    type(c_devptr), value :: d, e
+    logical :: result(8)
+    result(1) = c_associated(p)
+    result(2) = c_associated(p, q)
+    result(3) = c_associated(f)
+    result(4) = c_associated(f, g)
+    result(5) = p == q
+    result(6) = p /= q
+    result(7) = d == e
+    result(8) = d /= e
+  end subroutine
+  attributes(device) subroutine device(p, q, f, g, d, e, result)
+    type(c_ptr), value :: p, q
+    type(c_funptr), value :: f, g
+    type(c_devptr), value :: d, e
+    logical :: result(8)
+    result(1) = c_associated(p)
+    result(2) = c_associated(p, q)
+    result(3) = c_associated(f)
+    result(4) = c_associated(f, g)
+    result(5) = p == q
+    result(6) = p /= q
+    result(7) = d == e
+    result(8) = d /= e
+  end subroutine
+end module
+
+module ordinary_named_builtin
+contains
+  logical function c_associated_c_ptr()
+    c_associated_c_ptr = .true.
+  end function
+  attributes(host,device) subroutine caller(result)
+    logical :: result
+    !ERROR: 'c_associated_c_ptr' may not be called in device code
+    result = c_associated_c_ptr()
+  end subroutine
+end module
diff --git a/flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf b/flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf
index 606ced0689b7de3..e32476a67faa6c1 100644
--- a/flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf
+++ b/flang/test/Semantics/CUDA/cuf-hostdevice-call-context.cuf
@@ -232,9 +232,9 @@ contains
 
   attributes(host,device) subroutine generic_subroutines(n)
     integer, intent(inout) :: n
-    !ERROR: 'host_subroutine' may not be called in device code
+    !ERROR: 'host_generic' (specific procedure 'host_subroutine') may not be called in device code
     call host_generic(n)
-    !ERROR: 'device_subroutine' may not be called in host code
+    !ERROR: 'device_generic' (specific procedure 'device_subroutine') may not be called in host code
     call device_generic(n)
   end subroutine
 
diff --git a/flang/test/Semantics/CUDA/cuf-hostdevice-expression-context.cuf b/flang/test/Semantics/CUDA/cuf-hostdevice-expression-context.cuf
index 4dec3bb845c9af7..00ca8f6502c6a41 100644
--- a/flang/test/Semantics/CUDA/cuf-hostdevice-expression-context.cuf
+++ b/flang/test/Semantics/CUDA/cuf-hostdevice-expression-context.cuf
@@ -315,3 +315,33 @@ contains
     header_b = 'x'
   end function
 end module
+
+! Pin the expression checks in device, global, and CUF kernel contexts.
+module device_contexts
+contains
+  pure integer function h(i)
+    integer, value :: i
+    h = i
+  end function
+  attributes(device) subroutine device_bound(n)
+    integer, value :: n
+    !ERROR: 'h' may not be called in device code
+    real :: a(h(n))
+    a = 0.0
+  end subroutine
+  attributes(global) subroutine global_print(n)
+    integer, value :: n
+    !ERROR: 'h' may not be called in device code
+    print *, h(n)
+  end subroutine
+  subroutine kernel_print(a, n)
+    integer :: n, i
+    integer, device :: a(n)
+    !$cuf kernel do <<<*,*>>>
+    do i = 1, n
+      !ERROR: 'h' may not be called in device code
+      print *, h(i)
+      a(i) = i
+    end do
+  end subroutine
+end module
diff --git a/flang/test/Semantics/CUDA/cuf-hostdevice-generic-diagnostics.cuf b/flang/test/Semantics/CUDA/cuf-hostdevice-generic-diagnostics.cuf
new file mode 100644
index 000000000000000..599775338dd4d57
--- /dev/null
+++ b/flang/test/Semantics/CUDA/cuf-hostdevice-generic-diagnostics.cuf
@@ -0,0 +1,162 @@
+! RUN: %python %S/../test_errors.py %s %flang_fc1
+! Name the referenced generic and the selected specific in call diagnostics.
+module generic_targets
+  type :: bound_targets
+    contains
+    procedure, nopass :: host_binding => host_integer
+    procedure, nopass :: device_binding => device_integer
+    generic :: host_bound => host_binding
+    generic :: device_bound => device_binding
+  end type
+  interface host_generic
+    module procedure host_integer, host_real
+  end interface
+  interface device_generic
+    module procedure device_integer
+  end interface
+  interface both_generic
+    module procedure both_integer
+  end interface
+  interface host_sub_generic
+    module procedure host_sub
+  end interface
+  interface device_sub_generic
+    module procedure device_sub
+  end interface
+contains
+  integer function host_integer(n)
+    integer, value :: n
+    host_integer = n
+  end function
+  real function host_real(n)
+    real, value :: n
+    host_real = n
+  end function
+  attributes(device) integer function device_integer(n)
+    integer, value :: n
+    device_integer = n
+  end function
+  attributes(host,device) integer function both_integer(n)
+    integer, value :: n
+    both_integer = n
+  end function
+  subroutine host_sub(n)
+    integer, value :: n
+    n = 1
+  end subroutine
+  attributes(device) subroutine device_sub(n)
+    integer, value :: n
+    n = 1
+  end subroutine
+end module
+
+module generic_callers
+  use generic_targets
+  use cudadevice, only: on_device
+contains
+  attributes(host,device) subroutine unguarded(n, x)
+    integer :: n
+    real :: x
+    !ERROR: 'host_generic' (specific procedure 'host_integer') may not be called in device code
+    n = host_generic(n)
+    !ERROR: 'host_generic' (specific procedure 'host_real') may not be called in device code
+    x = host_generic(x)
+    !ERROR: 'device_generic' (specific procedure 'device_integer') may not be called in host code
+    n = device_generic(n)
+    !ERROR: 'host_sub_generic' (specific procedure 'host_sub') may not be called in device code
+    call host_sub_generic(n)
+    !ERROR: 'device_sub_generic' (specific procedure 'device_sub') may not be called in host code
+    call device_sub_generic(n)
+    ! Ordinary calls keep their single-name diagnostics.
+    !ERROR: 'host_integer' may not be called in device code
+    n = host_integer(n)
+    !ERROR: 'device_integer' may not be called in host code
+    n = device_integer(n)
+    n = both_generic(n)
+  end subroutine
+  attributes(device) subroutine device_caller(n)
+    integer :: n
+    !ERROR: 'host_generic' (specific procedure 'host_integer') may not be called in device code
+    n = host_generic(n)
+    n = device_generic(n)
+  end subroutine
+  attributes(host,device) subroutine guarded(n)
+    integer :: n
+    if (on_device()) then
+      n = device_generic(n)
+      call device_sub_generic(n)
+    else
+      n = host_generic(n)
+      call host_sub_generic(n)
+    end if
+  end subroutine
+end module
+
+module generic_alias_callers
+  use generic_targets, only: renamed_host => host_generic, another_host => host_generic
+  use generic_targets, only: renamed_device => device_generic, renamed_sub => host_sub_generic
+  use generic_targets, only: both_generic, specific_alias => host_integer
+contains
+  attributes(host,device) subroutine aliases(n)
+    integer :: n
+    !ERROR: 'renamed_host' (specific procedure 'host_integer') may not be called in device code
+    n = renamed_host(n)
+    !ERROR: 'another_host' (specific procedure 'host_integer') may not be called in device code
+    n = another_host(n)
+    !ERROR: 'renamed_device' (specific procedure 'device_integer') may not be called in host code
+    n = renamed_device(n)
+    !ERROR: 'renamed_sub' (specific procedure 'host_sub') may not be called in device code
+    call renamed_sub(n)
+    !ERROR: 'renamed_host' (specific procedure 'host_integer') may not be called in device code
+    n = both_generic(renamed_host(n))
+    ! A USE alias of an ordinary specific is not a generic reference.
+    !ERROR: 'specific_alias' may not be called in device code
+    n = specific_alias(n)
+  end subroutine
+end module
+
+module bound_generic_callers
+  use generic_targets
+contains
+  attributes(host,device) subroutine bound_calls(t, poly, n)
+    type(bound_targets) :: t
+    class(bound_targets) :: poly
+    integer :: n
+    !ERROR: 'host_bound' (specific procedure 'host_integer') may not be called in device code
+    n = t%host_bound(n)
+    !ERROR: 'device_bound' (specific procedure 'device_integer') may not be called in host code
+    n = t%device_bound(n)
+    !ERROR: 'host_bound' (specific procedure 'host_integer') may not be called in device code
+    n = poly%host_bound(n)
+    !ERROR: 'device_bound' (specific procedure 'device_integer') may not be called in host code
+    n = poly%device_bound(n)
+  end subroutine
+end module
+
+module native_cuda_generic
+  use cudadevice, only: atomicadd
+contains
+  attributes(host,device) subroutine atomic_call(a, n, old)
+    integer, device :: a(*)
+    integer :: n, old
+    !ERROR: 'atomicadd' (specific procedure 'atomicaddi') may not be called in host code
+    old = atomicadd(a(1), n)
+  end subroutine
+end module
+
+module identical_generic_name
+  interface same_name
+    module procedure same_name
+  end interface
+contains
+  integer function same_name(n)
+    integer, value :: n
+    same_name = n
+  end function
+  attributes(host,device) subroutine caller(n)
+    integer :: n
+    ! Do not repeat the name when generic and specific have the same spelling.
+    !ERROR: 'same_name' may not be called in device code
+    n = same_name(n)
+  end subroutine
+end module
diff --git a/flang/test/Semantics/CUDA/cuf-hostdevice-return-guard.cuf b/flang/test/Semantics/CUDA/cuf-hostdevice-return-guard.cuf
index 0f20d0bef880403..2632d59531e5791 100644
--- a/flang/test/Semantics/CUDA/cuf-hostdevice-return-guard.cuf
+++ b/flang/test/Semantics/CUDA/cuf-hostdevice-return-guard.cuf
@@ -165,6 +165,136 @@ contains
     !ERROR: 'host_only' may not be called in device code
     call host_only()
   end subroutine
+  attributes(host,device) subroutine reset_between_branches(choose, other)
+    logical :: choose, other
+    if (choose) then
+      if (on_device()) return
+      call host_only()
+    else if (other) then
+      !ERROR: 'host_only' may not be called in device code
+      call host_only()
+    else
+      !ERROR: 'device_only' may not be called in host code
+      call device_only()
+    end if
+    !ERROR: 'host_only' may not be called in device code
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine reset_after_guarded_else_if(choose, other)
+    logical :: choose, other
+    if (choose) then
+      call both()
+    else if (on_device()) then
+      call device_only()
+    else if (other) then
+      !ERROR: 'host_only' may not be called in device code
+      call host_only()
+      return
+    else
+      !ERROR: 'device_only' may not be called in host code
+      call device_only()
+    end if
+    !ERROR: 'host_only' may not be called in device code
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine preserve_inherited_guard(choose)
+    logical :: choose
+    if (on_device()) return
+    if (choose) then
+      call host_only()
+    else
+      call device_only()
+    end if
+    call host_only()
+    call device_only()
+  end subroutine
+  attributes(host,device) subroutine ordinary_block_return(choose)
+    logical, intent(in) :: choose
+    if (choose) then
+      return
+    else
+      call both()
+    end if
+    !ERROR: 'host_only' may not be called in device code
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine ordinary_else_return(choose)
+    logical, intent(in) :: choose
+    if (choose) then
+      call both()
+    else
+      return
+    end if
+    !ERROR: 'device_only' may not be called in host code
+    call device_only()
+  end subroutine
+  attributes(host,device) subroutine unrelated_arm_returns(choose)
+    logical, intent(in) :: choose
+    if (on_device()) then
+      call device_only()
+    else if (choose) then
+      return
+    end if
+    ! The returning arm does not test ON_DEVICE.
+    !ERROR: 'host_only' may not be called in device code
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine return_before_guard_arm(choose)
+    logical, intent(in) :: choose
+    if (choose) then
+      return
+    else if (on_device()) then
+      call device_only()
+    end if
+    !ERROR: 'host_only' may not be called in device code
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine else_if_guard_else_returns(choose)
+    logical, intent(in) :: choose
+    if (choose) then
+      call both()
+    else if (on_device()) then
+      call device_only()
+    else
+      return
+    end if
+    call host_only()
+    call device_only()
+  end subroutine
+  attributes(host,device) subroutine deferred_direction()
+    ! TODO: Direction checking belongs to the separate guard-analysis extension.
+    if (on_device()) return
+    call device_only()
+  end subroutine
+  attributes(host,device) subroutine deferred_negated_direction()
+    ! TODO: The device copy reaches this host-only call.
+    if (.not. on_device()) return
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine deferred_compound_guard(n)
+    integer :: n
+    ! TODO: The device copy can reach the call when n <= 0.
+    if (on_device() .and. n > 0) return
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine deferred_exit_bypass(skip)
+    logical :: skip
+    ! TODO: EXIT can bypass the RETURN; control-flow checking is deferred.
+    guard: if (on_device()) then
+      if (skip) exit guard
+      return
+    end if guard
+    call host_only()
+  end subroutine
+  attributes(host,device) subroutine deferred_block_exit(skip)
+    logical :: skip
+    ! TODO: EXIT can bypass the guard; control-flow checking is deferred.
+    outer: block
+      if (skip) exit outer
+      if (on_device()) return
+    end block outer
+    call host_only()
+  end subroutine
   attributes(host,device) subroutine ordinary_return(choose)
     logical, intent(in) :: choose
     if (choose) return
diff --git a/flang/test/Semantics/CUDA/cuf-hostdevice-stop-calls.cuf b/flang/test/Semantics/CUDA/cuf-hostdevice-stop-calls.cuf
new file mode 100644
index 000000000000000..a08448af736d5c1
--- /dev/null
+++ b/flang/test/Semantics/CUDA/cuf-hostdevice-stop-calls.cuf
@@ -0,0 +1,57 @@
+! RUN: %python %S/../test_errors.py %s %flang_fc1
+! Check calls in STOP and ERROR STOP codes and QUIET expressions.
+module stop_call_test
+  use cudadevice, only: on_device
+contains
+  attributes(device) integer function d()
+    d = 1
+  end function
+  integer function h()
+    h = 1
+  end function
+  attributes(host,device) integer function both()
+    both = 1
+  end function
+  attributes(device) logical function device_quiet()
+    device_quiet = .true.
+  end function
+  logical function host_quiet()
+    host_quiet = .true.
+  end function
+
+  attributes(host,device) subroutine caller()
+    ! Review example: the STOP code expression runs on both targets.
+    !ERROR: 'd' may not be called in host code
+    stop d()
+  end subroutine
+  attributes(host,device) subroutine host_code()
+    !ERROR: 'h' may not be called in device code
+    stop h()
+  end subroutine
+  attributes(host,device) subroutine error_code()
+    !ERROR: 'd' may not be called in host code
+    error stop d()
+  end subroutine
+  attributes(host,device) subroutine quiet_device()
+    !ERROR: 'device_quiet' may not be called in host code
+    stop 1, quiet=device_quiet()
+  end subroutine
+  attributes(host,device) subroutine quiet_host()
+    !ERROR: 'host_quiet' may not be called in device code
+    error stop 'failed', quiet=host_quiet()
+  end subroutine
+  attributes(host,device) subroutine both_code()
+    stop both()
+  end subroutine
+  attributes(host,device) subroutine guarded()
+    if (on_device()) then
+      stop d(), quiet=device_quiet()
+    else
+      error stop h(), quiet=host_quiet()
+    end if
+  end subroutine
+  attributes(device) subroutine device_caller()
+    !ERROR: 'h' may not be called in device code
+    stop h()
+  end subroutine
+end module



More information about the flang-commits mailing list