[clang] [clang][SYCL] Diagnose reference kernel arguments (PR #192957)
Mariya Podchishchaeva via cfe-commits
cfe-commits at lists.llvm.org
Mon Apr 20 05:07:29 PDT 2026
https://github.com/Fznamznon created https://github.com/llvm/llvm-project/pull/192957
Per SYCL 2020 spec: Reference types are not trivially copyable, so they may not be passed as kernel parameters.
This PR adds infrastructure for kernel object visiting and implements diagnostics for reference kernel parameters.
The infrastructure will be also used for other kernel parameter restrictions and functional code transformations that will be done in separate PRs.
>From cd2a8069daacb37f781a36c5bd21029de6427cc8 Mon Sep 17 00:00:00 2001
From: "Podchishchaeva, Mariya" <mariya.podchishchaeva at intel.com>
Date: Mon, 20 Apr 2026 04:59:25 -0700
Subject: [PATCH] [clang][SYCL] Diagnose reference kernel arguments
Per SYCL 2020 spec: Reference types are not trivially copyable, so they may not
be passed as kernel parameters.
This PR adds infrastructure for kernel object visiting and implements
diagnostics for reference kernel parameters.
The infrastructure will be also used for other kernel parameter restrictions and
functional code transformations that will be done in separate PRs.
---
clang/lib/Sema/SemaSYCL.cpp | 278 ++++++++++++++++++
.../SemaSYCL/sycl-kernel-arg-restrictions.cpp | 85 ++++++
2 files changed, 363 insertions(+)
create mode 100644 clang/test/SemaSYCL/sycl-kernel-arg-restrictions.cpp
diff --git a/clang/lib/Sema/SemaSYCL.cpp b/clang/lib/Sema/SemaSYCL.cpp
index 112a6e4416df2..0d22676aced29 100644
--- a/clang/lib/Sema/SemaSYCL.cpp
+++ b/clang/lib/Sema/SemaSYCL.cpp
@@ -294,6 +294,268 @@ void SemaSYCL::CheckSYCLExternalFunctionDecl(FunctionDecl *FD) {
}
}
+namespace {
+/// A special visitor to visit subobjects within a type, i.e. fields of a
+/// class or elements of an array. Useful for SYCl because in SYCL kernels are
+/// defined via lambda expressions or named callable objects and kernel
+/// parameters are fields of these. These visitors will be used for diagnosing
+/// invalid kernel arugments as well as for functional transformations.
+class SubobjectVisitor {
+ ASTContext &Ctx;
+
+ // These enable handler execution only when previous Handlers succeed.
+ template <typename... Tn>
+ bool handleField(FieldDecl *FD, QualType FDTy, Tn &&...tn) {
+ bool result = true;
+ (void)std::initializer_list<int>{(result = result && tn(FD, FDTy), 0)...};
+ return result;
+ }
+ template <typename... Tn>
+ bool handleField(const CXXBaseSpecifier &BD, QualType BDTy, Tn &&...tn) {
+ bool result = true;
+ std::initializer_list<int>{(result = result && tn(BD, BDTy), 0)...};
+ return result;
+ }
+
+#define KF_FOR_EACH(FUNC, Item, Qt) \
+ handleField(Item, Qt, ([&](FieldDecl *FD, QualType FDTy) { \
+ return Handlers.FUNC(FD, FDTy); \
+ })...)
+
+ // Parent contains the FieldDecl or CXXBaseSpecifier that was used to enter
+ // the Wrapper structure that we're currently visiting. Owner is the parent
+ // type (which doesn't exist in cases where it is a FieldDecl in the
+ // 'root'), and Wrapper is the current struct being unwrapped.
+ template <typename ParentTy, typename... HandlerTys>
+ void visitComplexRecord(const CXXRecordDecl *Owner, ParentTy &Parent,
+ const CXXRecordDecl *Wrapper, QualType RecordTy,
+ HandlerTys &...Handlers) {
+ (void)std::initializer_list<int>{
+ (Handlers.enterStruct(Owner, Parent, RecordTy), 0)...};
+ visitRecordHelper(Wrapper, Wrapper->bases(), Handlers...);
+ visitRecordHelper(Wrapper, Wrapper->fields(), Handlers...);
+ (void)std::initializer_list<int>{
+ (Handlers.leaveStruct(Owner, Parent, RecordTy), 0)...};
+ }
+
+ template <typename... HandlerTys>
+ void visitArray(const CXXRecordDecl *Owner, FieldDecl *Field,
+ QualType ArrayTy, HandlerTys &...Handlers) {
+ // TODO add support for simple array visiting, i.e. without entering array
+ // elements.
+ visitComplexArray(Owner, Field, ArrayTy, Handlers...);
+ }
+
+ template <typename ParentTy, typename... HandlerTys>
+ void visitRecord(const CXXRecordDecl *Owner, ParentTy &Parent,
+ const CXXRecordDecl *Wrapper, QualType RecordTy,
+ HandlerTys &...Handlers) {
+ // TODO add support for simple record visiting, i.e. without entering record
+ // fields.
+ visitComplexRecord(Owner, Parent, Wrapper, RecordTy, Handlers...);
+ }
+
+ template <typename... HandlerTys>
+ void visitRecordHelper(const CXXRecordDecl *Owner,
+ clang::CXXRecordDecl::base_class_const_range Range,
+ HandlerTys &...Handlers) {
+ for (const auto &Base : Range) {
+ QualType BaseTy = Base.getType();
+ visitRecord(Owner, Base, BaseTy->getAsCXXRecordDecl(), BaseTy,
+ Handlers...);
+ }
+ }
+
+ template <typename... HandlerTys>
+ void visitRecordHelper(const CXXRecordDecl *Owner, RecordDecl::field_range,
+ HandlerTys &...Handlers) {
+ visitRecordFields(Owner, Handlers...);
+ }
+
+ template <typename... HandlerTys>
+ void visitArrayElementImpl(const CXXRecordDecl *Owner, FieldDecl *ArrayField,
+ QualType ElementTy, uint64_t Index,
+ HandlerTys &...Handlers) {
+ visitField(Owner, ArrayField, ElementTy, Handlers...);
+ }
+
+ template <typename... HandlerTys>
+ void visitNthArrayElement(const CXXRecordDecl *Owner, FieldDecl *ArrayField,
+ QualType ElementTy, uint64_t Index,
+ HandlerTys &...Handlers) {
+ visitArrayElementImpl(Owner, ArrayField, ElementTy, Index, Handlers...);
+ }
+
+ template <typename... HandlerTys>
+ void visitComplexArray(const CXXRecordDecl *Owner, FieldDecl *Field,
+ QualType ArrayTy, HandlerTys &...Handlers) {
+ // Array workflow is:
+ // handleArrayType
+ // enterArray
+ // visitField (same as before, note that The FieldDecl is the of array
+ // itself, not the element)
+ // ... repeat per element, opt-out for duplicates.
+ // leaveArray
+
+ if (!KF_FOR_EACH(handleArrayType, Field, ArrayTy))
+ return;
+
+ const ConstantArrayType *CAT = Ctx.getAsConstantArrayType(ArrayTy);
+ assert(CAT && "Should only be called on constant-size array.");
+ QualType ET = CAT->getElementType();
+ uint64_t ElemCount = CAT->getSize().getZExtValue();
+
+ (void)std::initializer_list<int>{
+ (Handlers.enterArray(Field, ArrayTy, ET), 0)...};
+
+ for (uint64_t Index = 0; Index < ElemCount; ++Index)
+ visitNthArrayElement(Owner, Field, ET, Index, Handlers...);
+
+ (void)std::initializer_list<int>{
+ (Handlers.leaveArray(Field, ArrayTy, ET), 0)...};
+ }
+
+ template <typename... HandlerTys>
+ void visitField(const CXXRecordDecl *Owner, FieldDecl *Field,
+ QualType FieldTy, HandlerTys &...Handlers) {
+ if (FieldTy->isStructureOrClassType()) {
+ if (KF_FOR_EACH(handleStructType, Field, FieldTy)) {
+ CXXRecordDecl *RD = FieldTy->getAsCXXRecordDecl();
+ visitRecord(Owner, Field, RD, FieldTy, Handlers...);
+ }
+ } else if (FieldTy->isUnionType())
+ KF_FOR_EACH(handleUnionType, Field, FieldTy);
+ else if (FieldTy->isReferenceType())
+ KF_FOR_EACH(handleReferenceType, Field, FieldTy);
+ else if (FieldTy->isPointerType())
+ KF_FOR_EACH(handlePointerType, Field, FieldTy);
+ else if (FieldTy->isArrayType())
+ visitArray(Owner, Field, FieldTy, Handlers...);
+ else if (FieldTy->isScalarType() || FieldTy->isVectorType())
+ KF_FOR_EACH(handleScalarType, Field, FieldTy);
+ else
+ KF_FOR_EACH(handleOtherType, Field, FieldTy);
+ }
+
+public:
+ SubobjectVisitor(ASTContext &C) : Ctx(C) {}
+
+ template <typename... HandlerTys>
+ void visitRecordBases(const CXXRecordDecl *KernelFunctor,
+ HandlerTys &...Handlers) {
+ visitRecordHelper(KernelFunctor, KernelFunctor->bases(), Handlers...);
+ }
+
+ template <typename... HandlerTys>
+ void visitRecordFields(const CXXRecordDecl *Owner, HandlerTys &...Handlers) {
+ for (const auto Field : Owner->fields())
+ visitField(Owner, Field, Field->getType(), Handlers...);
+ }
+
+#undef KF_FOR_EACH
+};
+
+class SyclKernelFieldHandlerBase {
+public:
+ virtual bool handleStructType(FieldDecl *, QualType) { return true; }
+ virtual bool handleUnionType(FieldDecl *, QualType) { return true; }
+ virtual bool handleReferenceType(FieldDecl *, QualType) { return true; }
+ virtual bool handlePointerType(FieldDecl *, QualType) { return true; }
+ virtual bool handleArrayType(FieldDecl *, QualType) { return true; }
+ virtual bool handleScalarType(FieldDecl *, QualType) { return true; }
+ // Most handlers shouldn't be handling this, just the field checker.
+ virtual bool handleOtherType(FieldDecl *, QualType) { return true; }
+
+ virtual bool enterStruct(const CXXRecordDecl *, FieldDecl *, QualType) {
+ return true;
+ }
+ virtual bool leaveStruct(const CXXRecordDecl *, FieldDecl *, QualType) {
+ return true;
+ }
+ virtual bool enterStruct(const CXXRecordDecl *, const CXXBaseSpecifier &,
+ QualType) {
+ return true;
+ }
+ virtual bool leaveStruct(const CXXRecordDecl *, const CXXBaseSpecifier &,
+ QualType) {
+ return true;
+ }
+ // The following are used for stepping through array elements.
+ virtual bool enterArray(FieldDecl *, QualType, QualType) { return true; }
+ virtual bool leaveArray(FieldDecl *, QualType, QualType) { return true; }
+
+ virtual ~SyclKernelFieldHandlerBase() = default;
+};
+
+// A class to act as the direct base for all the SYCL Kernel related
+// tasks that contains a reference to Sema (and potentially any other
+// universally required data).
+class SyclKernelFieldHandler : public SyclKernelFieldHandlerBase {
+protected:
+ SemaSYCL &SemaSYCLRef;
+ SyclKernelFieldHandler(SemaSYCL &S) : SemaSYCLRef(S) {}
+};
+
+// A type to check the validity of all of the argument types.
+class SyclKernelFieldChecker : public SyclKernelFieldHandler {
+ bool IsInvalid = false;
+
+ bool checkArrayType(QualType Ty, SourceLocation Loc) {
+ if (Ty->isArrayType()) {
+ if (const auto *CAT =
+ SemaSYCLRef.getASTContext().getAsConstantArrayType(Ty)) {
+ QualType ET = CAT->getElementType();
+ return checkArrayType(ET, Loc);
+ }
+ return SemaSYCLRef.Diag(Loc, diag::err_bad_kernel_param_type) << Ty;
+ }
+ return false;
+ }
+
+public:
+ /// Constructor for the SyclKernelFieldChecker
+ /// \param S The SemaSYCL reference used for diagnostics and context.
+ explicit SyclKernelFieldChecker(SemaSYCL &S) : SyclKernelFieldHandler(S) {}
+ bool isValid() { return !IsInvalid; }
+
+ bool handleReferenceType(FieldDecl *FD, QualType FieldTy) final {
+ SemaSYCLRef.Diag(FD->getLocation(), diag::err_bad_kernel_param_type)
+ << FieldTy;
+ IsInvalid = true;
+ return isValid();
+ }
+
+ bool handlePointerType(FieldDecl *FD, QualType FieldTy) final {
+ while (FieldTy->isAnyPointerType()) {
+ FieldTy = QualType{FieldTy->getPointeeOrArrayElementType(), 0};
+ if (FieldTy->isVariableArrayType()) {
+ SemaSYCLRef.Diag(FD->getLocation(), diag::err_bad_kernel_param_type)
+ << FieldTy;
+ IsInvalid = true;
+ break;
+ }
+ }
+ return isValid();
+ }
+
+ bool handleOtherType(FieldDecl *FD, QualType FieldTy) final {
+ SemaSYCLRef.Diag(FD->getLocation(), diag::err_bad_kernel_param_type)
+ << FieldTy;
+ IsInvalid = true;
+ return isValid();
+ }
+
+ bool handleStructType(FieldDecl *, QualType FieldTy) final {
+ return isValid();
+ }
+
+ bool handleArrayType(FieldDecl *FD, QualType FieldTy) final {
+ IsInvalid |= checkArrayType(FieldTy, FD->getLocation());
+ return isValid();
+ }
+};
+} // namespace
+
void SemaSYCL::CheckSYCLEntryPointFunctionDecl(FunctionDecl *FD) {
// Ensure that all attributes present on the declaration are consistent
// and warn about any redundant ones.
@@ -665,6 +927,20 @@ OutlinedFunctionDecl *BuildSYCLKernelEntryPointOutline(Sema &SemaRef,
return OFD;
}
+bool verifyKernelArguments(FunctionDecl *FD, SemaSYCL &SemaSYCLRef) {
+ SyclKernelFieldChecker FieldChecker(SemaSYCLRef);
+ SubobjectVisitor Visitor{SemaSYCLRef.getASTContext()};
+ for (auto Param : FD->parameters()) {
+ if (Param->getType()->isRecordType()) {
+ const CXXRecordDecl *ObjRecord = Param->getType()->getAsCXXRecordDecl();
+ assert(ObjRecord && "object is expected");
+ Visitor.visitRecordBases(ObjRecord, FieldChecker);
+ Visitor.visitRecordFields(ObjRecord, FieldChecker);
+ }
+ }
+ return !FieldChecker.isValid();
+}
+
} // unnamed namespace
StmtResult SemaSYCL::BuildSYCLKernelCallStmt(FunctionDecl *FD,
@@ -690,6 +966,8 @@ StmtResult SemaSYCL::BuildSYCLKernelCallStmt(FunctionDecl *FD,
getASTContext().getSYCLKernelInfo(SKEPAttr->getKernelName());
assert(declaresSameEntity(SKI.getKernelEntryPointDecl(), FD) &&
"SYCL kernel name conflict");
+ if (verifyKernelArguments(FD, *this))
+ return StmtError();
// Build the outline of the synthesized device entry point function.
OutlinedFunctionDecl *OFD =
diff --git a/clang/test/SemaSYCL/sycl-kernel-arg-restrictions.cpp b/clang/test/SemaSYCL/sycl-kernel-arg-restrictions.cpp
new file mode 100644
index 0000000000000..6025f89ff12c4
--- /dev/null
+++ b/clang/test/SemaSYCL/sycl-kernel-arg-restrictions.cpp
@@ -0,0 +1,85 @@
+// RUN: %clang_cc1 -triple x86_64-linux-gnu -std=c++17 -fsyntax-only -Wno-vla-cxx-extension -fsycl-is-host -verify %s
+// RUN: %clang_cc1 -triple spirv64 -std=c++17 -fsyntax-only -Wno-vla-cxx-extension -fsycl-is-device -verify %s
+
+// A unique kernel name type is required for each declared kernel entry point.
+template<int, int = 0> struct KN;
+
+// A generic kernel launch function.
+template<typename KNT, typename... Ts>
+void sycl_kernel_launch(const char *, Ts...) {}
+
+// Kernel entry point template definition.
+template<typename KNT, typename T>
+[[clang::sycl_kernel_entry_point(KNT)]]
+void kernel_single_task(T) {}
+
+struct Empty {};
+
+struct S {
+ int a;
+ int &b; //expected-error 2{{'int &' cannot be used as the type of a kernel parameter}}
+};
+
+void fooarr(int (&arr)[5]) {
+}
+
+template <typename T> class Callable {
+ T data; // expected-error {{'int &' cannot be used as the type of a kernel parameter}}
+public:
+ Callable(T d) : data(d) {}
+ void operator()() {
+ }
+};
+
+void refCases(int AS) {
+ int p = 0;
+ double q = 0;
+ float s = 0;
+ kernel_single_task<class KN<1>>( // expected-note {{requested here}}
+ [
+ // expected-error at +1 {{'int &' cannot be used as the type of a kernel parameter}}
+ &p, q,
+ // expected-error at +1 {{'float &' cannot be used as the type of a kernel parameter}}
+ &s] {
+ (void)q;
+ (void)p;
+ (void)s;
+ });
+
+ auto L = [&]() { (void)p;}; // expected-error {{'int &' cannot be used as the type of a kernel parameter}}
+ S Str {p, p};
+ kernel_single_task<class KN<2>>( // expected-note {{requested here}}
+ [=] {
+ (void)L;
+ (void)Str; // no error because fail for L already
+ });
+
+ kernel_single_task<class KN<3>>( // expected-note {{requested here}}
+ [=] {
+ (void)Str;
+ });
+
+ S arr[2] = {Str, Str};
+ kernel_single_task<class KN<4>>( // expected-note {{requested here}}
+ [=] {
+ (void)arr;
+ });
+ int arr1[AS];
+ kernel_single_task<class KN<5>>( // expected-note {{requested here}}
+ [&] {
+ (void)arr1; // expected-error {{'int (&)[AS]' cannot be used as the type of a kernel parameter}}
+ });
+ auto a = &arr1;
+ kernel_single_task<class KN<6>>( // expected-note {{requested here}}
+ [=] {
+ (void)a; // expected-error {{'int[AS]' cannot be used as the type of a kernel parameter}}
+ });
+ int arrayints[5] = {0};
+ kernel_single_task<class KN<7>>( // expected-note {{requested here}}
+ [&] {
+ fooarr(arrayints); // expected-error {{'int (&)[5]' cannot be used as the type of a kernel parameter}}
+ });
+ kernel_single_task<class KN<8>>(Callable<int&>{p}); // expected-note {{requested here}}
+ kernel_single_task<class KN<9>>(Callable<int>{p});
+}
+
More information about the cfe-commits
mailing list