1//===- SemaHLSL.cpp - Semantic Analysis for HLSL constructs ---------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8// This implements Semantic Analysis for HLSL constructs.
9//===----------------------------------------------------------------------===//
10
11#include "clang/Sema/SemaHLSL.h"
12#include "clang/AST/ASTConsumer.h"
13#include "clang/AST/ASTContext.h"
14#include "clang/AST/Attr.h"
15#include "clang/AST/Decl.h"
16#include "clang/AST/DeclBase.h"
17#include "clang/AST/DeclCXX.h"
18#include "clang/AST/DeclarationName.h"
19#include "clang/AST/DynamicRecursiveASTVisitor.h"
20#include "clang/AST/Expr.h"
21#include "clang/AST/HLSLResource.h"
22#include "clang/AST/Type.h"
23#include "clang/AST/TypeBase.h"
24#include "clang/AST/TypeLoc.h"
25#include "clang/Basic/Builtins.h"
26#include "clang/Basic/DiagnosticSema.h"
27#include "clang/Basic/IdentifierTable.h"
28#include "clang/Basic/LLVM.h"
29#include "clang/Basic/SourceLocation.h"
30#include "clang/Basic/Specifiers.h"
31#include "clang/Basic/TargetInfo.h"
32#include "clang/Sema/Initialization.h"
33#include "clang/Sema/Lookup.h"
34#include "clang/Sema/ParsedAttr.h"
35#include "clang/Sema/Sema.h"
36#include "clang/Sema/Template.h"
37#include "llvm/ADT/ArrayRef.h"
38#include "llvm/ADT/STLExtras.h"
39#include "llvm/ADT/SmallVector.h"
40#include "llvm/ADT/StringExtras.h"
41#include "llvm/ADT/StringRef.h"
42#include "llvm/ADT/Twine.h"
43#include "llvm/Frontend/HLSL/HLSLBinding.h"
44#include "llvm/Frontend/HLSL/RootSignatureValidations.h"
45#include "llvm/Support/Casting.h"
46#include "llvm/Support/DXILABI.h"
47#include "llvm/Support/ErrorHandling.h"
48#include "llvm/Support/FormatVariadic.h"
49#include "llvm/TargetParser/Triple.h"
50#include <algorithm>
51#include <cmath>
52#include <cstddef>
53#include <iterator>
54#include <utility>
55
56using namespace clang;
57using namespace clang::hlsl;
58using llvm::hlsl::InterpolationModifier;
59using llvm::hlsl::IOType;
60using llvm::hlsl::SemanticStageInfo;
61using SemanticKind = llvm::dxbc::PSV::SemanticKind;
62using RegisterType = HLSLResourceBindingAttr::RegisterType;
63
64static CXXRecordDecl *createHostLayoutStruct(Sema &S,
65 CXXRecordDecl *StructDecl);
66
67static QualType getScalarComponentType(QualType T) {
68 if (const auto *VT = T->getAs<VectorType>())
69 return VT->getElementType();
70 if (const auto *MT = T->getAs<MatrixType>())
71 return MT->getElementType();
72 return T;
73}
74
75static RegisterType getRegisterType(ResourceClass RC) {
76 switch (RC) {
77 case ResourceClass::SRV:
78 return RegisterType::SRV;
79 case ResourceClass::UAV:
80 return RegisterType::UAV;
81 case ResourceClass::CBuffer:
82 return RegisterType::CBuffer;
83 case ResourceClass::Sampler:
84 return RegisterType::Sampler;
85 }
86 llvm_unreachable("unexpected ResourceClass value");
87}
88
89static RegisterType getRegisterType(const HLSLAttributedResourceType *ResTy) {
90 return getRegisterType(RC: ResTy->getAttrs().ResourceClass);
91}
92
93static LangAS getLangASFromResourceClass(ResourceClass RC) {
94 switch (RC) {
95 case ResourceClass::SRV:
96 case ResourceClass::UAV:
97 return LangAS::hlsl_device;
98 case ResourceClass::CBuffer:
99 return LangAS::hlsl_constant;
100 case ResourceClass::Sampler:
101 return LangAS::hlsl_device;
102 }
103 llvm_unreachable("unexpected ResourceClass value");
104}
105
106// Converts the first letter of string Slot to RegisterType.
107// Returns false if the letter does not correspond to a valid register type.
108static bool convertToRegisterType(StringRef Slot, RegisterType *RT) {
109 assert(RT != nullptr);
110 switch (Slot[0]) {
111 case 't':
112 case 'T':
113 *RT = RegisterType::SRV;
114 return true;
115 case 'u':
116 case 'U':
117 *RT = RegisterType::UAV;
118 return true;
119 case 'b':
120 case 'B':
121 *RT = RegisterType::CBuffer;
122 return true;
123 case 's':
124 case 'S':
125 *RT = RegisterType::Sampler;
126 return true;
127 case 'c':
128 case 'C':
129 *RT = RegisterType::C;
130 return true;
131 case 'i':
132 case 'I':
133 *RT = RegisterType::I;
134 return true;
135 default:
136 return false;
137 }
138}
139
140static char getRegisterTypeChar(RegisterType RT) {
141 switch (RT) {
142 case RegisterType::SRV:
143 return 't';
144 case RegisterType::UAV:
145 return 'u';
146 case RegisterType::CBuffer:
147 return 'b';
148 case RegisterType::Sampler:
149 return 's';
150 case RegisterType::C:
151 return 'c';
152 case RegisterType::I:
153 return 'i';
154 }
155 llvm_unreachable("unexpected RegisterType value");
156}
157
158static ResourceClass getResourceClass(RegisterType RT) {
159 switch (RT) {
160 case RegisterType::SRV:
161 return ResourceClass::SRV;
162 case RegisterType::UAV:
163 return ResourceClass::UAV;
164 case RegisterType::CBuffer:
165 return ResourceClass::CBuffer;
166 case RegisterType::Sampler:
167 return ResourceClass::Sampler;
168 case RegisterType::C:
169 case RegisterType::I:
170 // Deliberately falling through to the unreachable below.
171 break;
172 }
173 llvm_unreachable("unexpected RegisterType value");
174}
175
176static Builtin::ID getSpecConstBuiltinId(const Type *Type) {
177 const auto *BT = dyn_cast<BuiltinType>(Val: Type);
178 if (!BT) {
179 if (!Type->isEnumeralType())
180 return Builtin::NotBuiltin;
181 return Builtin::BI__builtin_get_spirv_spec_constant_int;
182 }
183
184 switch (BT->getKind()) {
185 case BuiltinType::Bool:
186 return Builtin::BI__builtin_get_spirv_spec_constant_bool;
187 case BuiltinType::Short:
188 return Builtin::BI__builtin_get_spirv_spec_constant_short;
189 case BuiltinType::Int:
190 return Builtin::BI__builtin_get_spirv_spec_constant_int;
191 case BuiltinType::LongLong:
192 return Builtin::BI__builtin_get_spirv_spec_constant_longlong;
193 case BuiltinType::UShort:
194 return Builtin::BI__builtin_get_spirv_spec_constant_ushort;
195 case BuiltinType::UInt:
196 return Builtin::BI__builtin_get_spirv_spec_constant_uint;
197 case BuiltinType::ULongLong:
198 return Builtin::BI__builtin_get_spirv_spec_constant_ulonglong;
199 case BuiltinType::Half:
200 return Builtin::BI__builtin_get_spirv_spec_constant_half;
201 case BuiltinType::Float:
202 return Builtin::BI__builtin_get_spirv_spec_constant_float;
203 case BuiltinType::Double:
204 return Builtin::BI__builtin_get_spirv_spec_constant_double;
205 default:
206 return Builtin::NotBuiltin;
207 }
208}
209
210static StringRef createRegisterString(ASTContext &AST, RegisterType RegType,
211 unsigned N) {
212 llvm::SmallString<16> Buffer;
213 llvm::raw_svector_ostream OS(Buffer);
214 OS << getRegisterTypeChar(RT: RegType);
215 OS << N;
216 return AST.backupStr(S: OS.str());
217}
218
219DeclBindingInfo *ResourceBindings::addDeclBindingInfo(const VarDecl *VD,
220 ResourceClass ResClass) {
221 assert(getDeclBindingInfo(VD, ResClass) == nullptr &&
222 "DeclBindingInfo already added");
223 assert(!hasBindingInfoForDecl(VD) || BindingsList.back().Decl == VD);
224 // VarDecl may have multiple entries for different resource classes.
225 // DeclToBindingListIndex stores the index of the first binding we saw
226 // for this decl. If there are any additional ones then that index
227 // shouldn't be updated.
228 DeclToBindingListIndex.try_emplace(Key: VD, Args: BindingsList.size());
229 return &BindingsList.emplace_back(Args&: VD, Args&: ResClass);
230}
231
232DeclBindingInfo *ResourceBindings::getDeclBindingInfo(const VarDecl *VD,
233 ResourceClass ResClass) {
234 auto Entry = DeclToBindingListIndex.find(Val: VD);
235 if (Entry != DeclToBindingListIndex.end()) {
236 for (unsigned Index = Entry->getSecond();
237 Index < BindingsList.size() && BindingsList[Index].Decl == VD;
238 ++Index) {
239 if (BindingsList[Index].ResClass == ResClass)
240 return &BindingsList[Index];
241 }
242 }
243 return nullptr;
244}
245
246bool ResourceBindings::hasBindingInfoForDecl(const VarDecl *VD) const {
247 return DeclToBindingListIndex.contains(Val: VD);
248}
249
250SemaHLSL::SemaHLSL(Sema &S) : SemaBase(S) {}
251
252Decl *SemaHLSL::ActOnStartBuffer(Scope *BufferScope, bool CBuffer,
253 SourceLocation KwLoc, IdentifierInfo *Ident,
254 SourceLocation IdentLoc,
255 SourceLocation LBrace) {
256 // For anonymous namespace, take the location of the left brace.
257 DeclContext *LexicalParent = SemaRef.getCurLexicalContext();
258 HLSLBufferDecl *Result = HLSLBufferDecl::Create(
259 C&: getASTContext(), LexicalParent, CBuffer, KwLoc, ID: Ident, IDLoc: IdentLoc, LBrace);
260
261 // if CBuffer is false, then it's a TBuffer
262 auto RC = CBuffer ? llvm::hlsl::ResourceClass::CBuffer
263 : llvm::hlsl::ResourceClass::SRV;
264 Result->addAttr(A: HLSLResourceClassAttr::CreateImplicit(Ctx&: getASTContext(), ResourceClass: RC));
265
266 SemaRef.PushOnScopeChains(D: Result, S: BufferScope);
267 SemaRef.PushDeclContext(S: BufferScope, DC: Result);
268
269 return Result;
270}
271
272static unsigned calculateLegacyCbufferFieldAlign(const ASTContext &Context,
273 QualType T) {
274 // Arrays, Matrices, and Structs are always aligned to new buffer rows
275 if (T->isArrayType() || T->isStructureType() || T->isConstantMatrixType())
276 return 16;
277
278 // Vectors are aligned to the type they contain
279 if (const VectorType *VT = T->getAs<VectorType>())
280 return calculateLegacyCbufferFieldAlign(Context, T: VT->getElementType());
281
282 assert(Context.getTypeSize(T) <= 64 &&
283 "Scalar bit widths larger than 64 not supported");
284
285 // Scalar types are aligned to their byte width
286 return Context.getTypeSize(T) / 8;
287}
288
289// Calculate the size of a legacy cbuffer type in bytes based on
290// https://learn.microsoft.com/en-us/windows/win32/direct3dhlsl/dx-graphics-hlsl-packing-rules
291static unsigned calculateLegacyCbufferSize(const ASTContext &Context,
292 QualType T) {
293 constexpr unsigned CBufferAlign = 16;
294 if (const auto *RD = T->getAsRecordDecl()) {
295 unsigned Size = 0;
296 for (const FieldDecl *Field : RD->fields()) {
297 QualType Ty = Field->getType();
298 unsigned FieldSize = calculateLegacyCbufferSize(Context, T: Ty);
299 unsigned FieldAlign = calculateLegacyCbufferFieldAlign(Context, T: Ty);
300
301 // If the field crosses the row boundary after alignment it drops to the
302 // next row
303 unsigned AlignSize = llvm::alignTo(Value: Size, Align: FieldAlign);
304 if ((AlignSize % CBufferAlign) + FieldSize > CBufferAlign) {
305 FieldAlign = CBufferAlign;
306 }
307
308 Size = llvm::alignTo(Value: Size, Align: FieldAlign);
309 Size += FieldSize;
310 }
311 return Size;
312 }
313
314 if (const ConstantArrayType *AT = Context.getAsConstantArrayType(T)) {
315 unsigned ElementCount = AT->getSize().getZExtValue();
316 if (ElementCount == 0)
317 return 0;
318
319 unsigned ElementSize =
320 calculateLegacyCbufferSize(Context, T: AT->getElementType());
321 unsigned AlignedElementSize = llvm::alignTo(Value: ElementSize, Align: CBufferAlign);
322 return AlignedElementSize * (ElementCount - 1) + ElementSize;
323 }
324
325 if (const VectorType *VT = T->getAs<VectorType>()) {
326 unsigned ElementCount = VT->getNumElements();
327 unsigned ElementSize =
328 calculateLegacyCbufferSize(Context, T: VT->getElementType());
329 return ElementSize * ElementCount;
330 }
331
332 return Context.getTypeSize(T) / 8;
333}
334
335// Validate packoffset:
336// - if packoffset it used it must be set on all declarations inside the buffer
337// - packoffset ranges must not overlap
338static void validatePackoffset(Sema &S, HLSLBufferDecl *BufDecl) {
339 llvm::SmallVector<std::pair<VarDecl *, HLSLPackOffsetAttr *>> PackOffsetVec;
340
341 // Make sure the packoffset annotations are either on all declarations
342 // or on none.
343 bool HasPackOffset = false;
344 bool HasNonPackOffset = false;
345 for (auto *Field : BufDecl->buffer_decls()) {
346 VarDecl *Var = dyn_cast<VarDecl>(Val: Field);
347 if (!Var)
348 continue;
349 if (Field->hasAttr<HLSLPackOffsetAttr>()) {
350 PackOffsetVec.emplace_back(Args&: Var, Args: Field->getAttr<HLSLPackOffsetAttr>());
351 HasPackOffset = true;
352 } else {
353 HasNonPackOffset = true;
354 }
355 }
356
357 if (!HasPackOffset)
358 return;
359
360 if (HasNonPackOffset)
361 S.Diag(Loc: BufDecl->getLocation(), DiagID: diag::warn_hlsl_packoffset_mix);
362
363 // Make sure there is no overlap in packoffset - sort PackOffsetVec by offset
364 // and compare adjacent values.
365 bool IsValid = true;
366 ASTContext &Context = S.getASTContext();
367 std::sort(first: PackOffsetVec.begin(), last: PackOffsetVec.end(),
368 comp: [](const std::pair<VarDecl *, HLSLPackOffsetAttr *> &LHS,
369 const std::pair<VarDecl *, HLSLPackOffsetAttr *> &RHS) {
370 return LHS.second->getOffsetInBytes() <
371 RHS.second->getOffsetInBytes();
372 });
373 for (unsigned i = 0; i < PackOffsetVec.size() - 1; i++) {
374 VarDecl *Var = PackOffsetVec[i].first;
375 HLSLPackOffsetAttr *Attr = PackOffsetVec[i].second;
376 unsigned Size = calculateLegacyCbufferSize(Context, T: Var->getType());
377 unsigned Begin = Attr->getOffsetInBytes();
378 unsigned End = Begin + Size;
379 unsigned NextBegin = PackOffsetVec[i + 1].second->getOffsetInBytes();
380 if (End > NextBegin) {
381 VarDecl *NextVar = PackOffsetVec[i + 1].first;
382 S.Diag(Loc: NextVar->getLocation(), DiagID: diag::err_hlsl_packoffset_overlap)
383 << NextVar << Var;
384 IsValid = false;
385 }
386 }
387 BufDecl->setHasValidPackoffset(IsValid);
388}
389
390// Returns true if the array has a zero size = if any of the dimensions is 0
391static bool isZeroSizedArray(const ConstantArrayType *CAT) {
392 while (CAT && !CAT->isZeroSize())
393 CAT = dyn_cast<ConstantArrayType>(
394 Val: CAT->getElementType()->getUnqualifiedDesugaredType());
395 return CAT != nullptr;
396}
397
398static bool isResourceRecordTypeOrArrayOf(QualType Ty) {
399 return Ty->isHLSLResourceRecord() || Ty->isHLSLResourceRecordArray();
400}
401
402static bool isResourceRecordTypeOrArrayOf(VarDecl *VD) {
403 return isResourceRecordTypeOrArrayOf(Ty: VD->getType());
404}
405
406static const HLSLAttributedResourceType *
407getResourceArrayHandleType(QualType QT) {
408 assert(QT->isHLSLResourceRecordArray() &&
409 "expected array of resource records");
410 const Type *Ty = QT->getUnqualifiedDesugaredType();
411 while (const ArrayType *AT = dyn_cast<ArrayType>(Val: Ty))
412 Ty = AT->getArrayElementTypeNoTypeQual()->getUnqualifiedDesugaredType();
413 return HLSLAttributedResourceType::findHandleTypeOnResource(RT: Ty);
414}
415
416static const HLSLAttributedResourceType *
417getResourceArrayHandleType(VarDecl *VD) {
418 return getResourceArrayHandleType(QT: VD->getType());
419}
420
421// Returns true if the type is a leaf element type that is not valid to be
422// included in HLSL Buffer, such as a resource class, empty struct, zero-sized
423// array, or a builtin intangible type. Returns false it is a valid leaf element
424// type or if it is a record type that needs to be inspected further.
425static bool isInvalidConstantBufferLeafElementType(const Type *Ty) {
426 Ty = Ty->getUnqualifiedDesugaredType();
427 if (Ty->isHLSLResourceRecord() || Ty->isHLSLResourceRecordArray())
428 return true;
429 if (const auto *RD = Ty->getAsCXXRecordDecl())
430 return RD->isEmpty();
431 if (Ty->isConstantArrayType() &&
432 isZeroSizedArray(CAT: cast<ConstantArrayType>(Val: Ty)))
433 return true;
434 if (Ty->isHLSLBuiltinIntangibleType() || Ty->isHLSLAttributedResourceType())
435 return true;
436 return false;
437}
438
439// Returns true if the struct contains at least one element that prevents it
440// from being included inside HLSL Buffer as is, such as an intangible type,
441// empty struct, or zero-sized array. If it does, a new implicit layout struct
442// needs to be created for HLSL Buffer use that will exclude these unwanted
443// declarations (see createHostLayoutStruct function).
444static bool requiresImplicitBufferLayoutStructure(const CXXRecordDecl *RD) {
445 if (RD->isHLSLIntangible() || RD->isEmpty())
446 return true;
447 // check fields
448 for (const FieldDecl *Field : RD->fields()) {
449 QualType Ty = Field->getType();
450 if (isInvalidConstantBufferLeafElementType(Ty: Ty.getTypePtr()))
451 return true;
452 if (const auto *RD = Ty->getAsCXXRecordDecl();
453 RD && requiresImplicitBufferLayoutStructure(RD))
454 return true;
455 }
456 // check bases
457 for (const CXXBaseSpecifier &Base : RD->bases())
458 if (requiresImplicitBufferLayoutStructure(
459 RD: Base.getType()->castAsCXXRecordDecl()))
460 return true;
461 return false;
462}
463
464static CXXRecordDecl *findRecordDeclInContext(IdentifierInfo *II,
465 DeclContext *DC) {
466 CXXRecordDecl *RD = nullptr;
467 for (NamedDecl *Decl :
468 DC->getNonTransparentContext()->lookup(Name: DeclarationName(II))) {
469 if (CXXRecordDecl *FoundRD = dyn_cast<CXXRecordDecl>(Val: Decl)) {
470 assert(RD == nullptr &&
471 "there should be at most 1 record by a given name in a scope");
472 RD = FoundRD;
473 }
474 }
475 return RD;
476}
477
478// Creates a name for buffer layout struct using the provide name base.
479// If the name must be unique (not previously defined), a suffix is added
480// until a unique name is found.
481static IdentifierInfo *getHostLayoutStructName(Sema &S, NamedDecl *BaseDecl,
482 bool MustBeUnique) {
483 ASTContext &AST = S.getASTContext();
484
485 IdentifierInfo *NameBaseII = BaseDecl->getIdentifier();
486 llvm::SmallString<64> Name("__cblayout_");
487 if (NameBaseII) {
488 Name.append(RHS: NameBaseII->getName());
489 } else {
490 // anonymous struct
491 Name.append(RHS: "anon");
492 MustBeUnique = true;
493 }
494
495 size_t NameLength = Name.size();
496 IdentifierInfo *II = &AST.Idents.get(Name, TokenCode: tok::TokenKind::identifier);
497 if (!MustBeUnique)
498 return II;
499
500 unsigned suffix = 0;
501 while (true) {
502 if (suffix != 0) {
503 Name.append(RHS: "_");
504 Name.append(RHS: llvm::Twine(suffix).str());
505 II = &AST.Idents.get(Name, TokenCode: tok::TokenKind::identifier);
506 }
507 if (!findRecordDeclInContext(II, DC: BaseDecl->getDeclContext()))
508 return II;
509 // declaration with that name already exists - increment suffix and try
510 // again until unique name is found
511 suffix++;
512 Name.truncate(N: NameLength);
513 };
514}
515
516static const Type *createHostLayoutType(Sema &S, const Type *Ty) {
517 ASTContext &AST = S.getASTContext();
518 if (auto *RD = Ty->getAsCXXRecordDecl()) {
519 if (!requiresImplicitBufferLayoutStructure(RD))
520 return Ty;
521 RD = createHostLayoutStruct(S, StructDecl: RD);
522 if (!RD)
523 return nullptr;
524 return AST.getCanonicalTagType(TD: RD)->getTypePtr();
525 }
526
527 if (const auto *CAT = dyn_cast<ConstantArrayType>(Val: Ty)) {
528 const Type *ElementTy = createHostLayoutType(
529 S, Ty: CAT->getElementType()->getUnqualifiedDesugaredType());
530 if (!ElementTy)
531 return nullptr;
532 return AST
533 .getConstantArrayType(EltTy: QualType(ElementTy, 0), ArySize: CAT->getSize(), SizeExpr: nullptr,
534 ASM: CAT->getSizeModifier(),
535 IndexTypeQuals: CAT->getIndexTypeCVRQualifiers())
536 .getTypePtr();
537 }
538 return Ty;
539}
540
541// Returns the type to use for a host layout struct field. For most types this
542// is the unqualified desugared type. Matrix types, however, retain their sugar
543// so that the row_major/column_major orientation (carried as an AttributedType)
544// is preserved; the orientation determines the in-memory cbuffer layout.
545static const Type *getHostLayoutFieldType(QualType QT) {
546 const Type *Desugared = QT->getUnqualifiedDesugaredType();
547 if (Desugared->isConstantMatrixType())
548 return QT.getTypePtr();
549 return Desugared;
550}
551
552// Creates a field declaration of given name and type for HLSL buffer layout
553// struct. Returns nullptr if the type cannot be use in HLSL Buffer layout.
554static FieldDecl *createFieldForHostLayoutStruct(Sema &S, const Type *Ty,
555 IdentifierInfo *II,
556 CXXRecordDecl *LayoutStruct) {
557 if (isInvalidConstantBufferLeafElementType(Ty))
558 return nullptr;
559
560 Ty = createHostLayoutType(S, Ty);
561 if (!Ty)
562 return nullptr;
563
564 QualType QT = QualType(Ty, 0);
565 ASTContext &AST = S.getASTContext();
566 TypeSourceInfo *TSI = AST.getTrivialTypeSourceInfo(T: QT, Loc: SourceLocation());
567 auto *Field = FieldDecl::Create(C: AST, DC: LayoutStruct, StartLoc: SourceLocation(),
568 IdLoc: SourceLocation(), Id: II, T: QT, TInfo: TSI, BW: nullptr, Mutable: false,
569 InitStyle: InClassInitStyle::ICIS_NoInit);
570 Field->setAccess(AccessSpecifier::AS_public);
571 return Field;
572}
573
574// Creates host layout struct for a struct included in HLSL Buffer.
575// The layout struct will include only fields that are allowed in HLSL buffer.
576// These fields will be filtered out:
577// - resource classes
578// - empty structs
579// - zero-sized arrays
580// Returns nullptr if the resulting layout struct would be empty.
581static CXXRecordDecl *createHostLayoutStruct(Sema &S,
582 CXXRecordDecl *StructDecl) {
583 assert(requiresImplicitBufferLayoutStructure(StructDecl) &&
584 "struct is already HLSL buffer compatible");
585
586 ASTContext &AST = S.getASTContext();
587 DeclContext *DC = StructDecl->getDeclContext();
588 IdentifierInfo *II = getHostLayoutStructName(S, BaseDecl: StructDecl, MustBeUnique: false);
589
590 // reuse existing if the layout struct if it already exists
591 if (CXXRecordDecl *RD = findRecordDeclInContext(II, DC))
592 return RD;
593
594 CXXRecordDecl *LS =
595 CXXRecordDecl::Create(C: AST, TK: TagDecl::TagKind::Struct, DC, StartLoc: SourceLocation(),
596 IdLoc: SourceLocation(), Id: II);
597 LS->setImplicit(true);
598 LS->addAttr(A: PackedAttr::CreateImplicit(Ctx&: AST));
599 LS->startDefinition();
600
601 // copy base struct, create HLSL Buffer compatible version if needed
602 if (unsigned NumBases = StructDecl->getNumBases()) {
603 assert(NumBases == 1 && "HLSL supports only one base type");
604 (void)NumBases;
605 CXXBaseSpecifier Base = *StructDecl->bases_begin();
606 CXXRecordDecl *BaseDecl = Base.getType()->castAsCXXRecordDecl();
607 if (requiresImplicitBufferLayoutStructure(RD: BaseDecl)) {
608 BaseDecl = createHostLayoutStruct(S, StructDecl: BaseDecl);
609 if (BaseDecl) {
610 TypeSourceInfo *TSI =
611 AST.getTrivialTypeSourceInfo(T: AST.getCanonicalTagType(TD: BaseDecl));
612 Base = CXXBaseSpecifier(SourceRange(), false, StructDecl->isClass(),
613 AS_none, TSI, SourceLocation());
614 }
615 }
616 if (BaseDecl) {
617 const CXXBaseSpecifier *BasesArray[1] = {&Base};
618 LS->setBases(Bases: BasesArray, NumBases: 1);
619 }
620 }
621
622 // filter struct fields
623 for (const FieldDecl *FD : StructDecl->fields()) {
624 const Type *Ty = getHostLayoutFieldType(QT: FD->getType());
625 if (FieldDecl *NewFD =
626 createFieldForHostLayoutStruct(S, Ty, II: FD->getIdentifier(), LayoutStruct: LS))
627 LS->addDecl(D: NewFD);
628 }
629 LS->completeDefinition();
630
631 if (LS->field_empty() && LS->getNumBases() == 0)
632 return nullptr;
633
634 DC->addDecl(D: LS);
635 return LS;
636}
637
638// Creates host layout struct for HLSL Buffer. The struct will include only
639// fields of types that are allowed in HLSL buffer and it will filter out:
640// - static or groupshared variable declarations
641// - resource classes
642// - empty structs
643// - zero-sized arrays
644// - non-variable declarations
645// The layout struct will be added to the HLSLBufferDecl declarations.
646static void createHostLayoutStructForBuffer(Sema &S, HLSLBufferDecl *BufDecl) {
647 ASTContext &AST = S.getASTContext();
648 IdentifierInfo *II = getHostLayoutStructName(S, BaseDecl: BufDecl, MustBeUnique: true);
649
650 CXXRecordDecl *LS =
651 CXXRecordDecl::Create(C: AST, TK: TagDecl::TagKind::Struct, DC: BufDecl,
652 StartLoc: SourceLocation(), IdLoc: SourceLocation(), Id: II);
653 LS->addAttr(A: PackedAttr::CreateImplicit(Ctx&: AST));
654 LS->setImplicit(true);
655 LS->startDefinition();
656
657 for (Decl *D : BufDecl->buffer_decls()) {
658 VarDecl *VD = dyn_cast<VarDecl>(Val: D);
659 if (!VD || VD->getStorageClass() == SC_Static ||
660 VD->getType().getAddressSpace() == LangAS::hlsl_groupshared)
661 continue;
662 const Type *Ty = getHostLayoutFieldType(QT: VD->getType());
663
664 FieldDecl *FD =
665 createFieldForHostLayoutStruct(S, Ty, II: VD->getIdentifier(), LayoutStruct: LS);
666 // Declarations collected for the default $Globals constant buffer have
667 // already been checked to have non-empty cbuffer layout, so
668 // createFieldForHostLayoutStruct should always succeed. These declarations
669 // already have their address space set to hlsl_constant.
670 // For declarations in a named cbuffer block
671 // createFieldForHostLayoutStruct can still return nullptr if the type
672 // is empty (does not have a cbuffer layout).
673 assert((FD || VD->getType().getAddressSpace() != LangAS::hlsl_constant) &&
674 "host layout field for $Globals decl failed to be created");
675 if (FD) {
676 // Add the field decl to the layout struct.
677 LS->addDecl(D: FD);
678 if (VD->getType().getAddressSpace() != LangAS::hlsl_constant) {
679 // Update address space of the original decl to hlsl_constant.
680 QualType NewTy =
681 AST.getAddrSpaceQualType(T: VD->getType(), AddressSpace: LangAS::hlsl_constant);
682 VD->setType(NewTy);
683 }
684 }
685 }
686 LS->completeDefinition();
687 BufDecl->addLayoutStruct(LS);
688}
689
690static void addImplicitBindingAttrToDecl(Sema &S, Decl *D, RegisterType RT,
691 uint32_t ImplicitBindingOrderID) {
692 auto *Attr =
693 HLSLResourceBindingAttr::CreateImplicit(Ctx&: S.getASTContext(), Slot: "", Space: "0", Range: {});
694 Attr->setBinding(RT, SlotNum: std::nullopt, SpaceNum: 0);
695 Attr->setImplicitBindingOrderID(ImplicitBindingOrderID);
696 D->addAttr(A: Attr);
697}
698
699// Handle end of cbuffer/tbuffer declaration
700void SemaHLSL::ActOnFinishBuffer(Decl *Dcl, SourceLocation RBrace) {
701 auto *BufDecl = cast<HLSLBufferDecl>(Val: Dcl);
702 BufDecl->setRBraceLoc(RBrace);
703
704 validatePackoffset(S&: SemaRef, BufDecl);
705
706 createHostLayoutStructForBuffer(S&: SemaRef, BufDecl);
707
708 // Handle implicit binding if needed.
709 ResourceBindingAttrs ResourceAttrs(Dcl);
710 if (!ResourceAttrs.isExplicit()) {
711 SemaRef.Diag(Loc: Dcl->getLocation(), DiagID: diag::warn_hlsl_implicit_binding);
712 // Use HLSLResourceBindingAttr to transfer implicit binding order_ID
713 // to codegen. If it does not exist, create an implicit attribute.
714 uint32_t OrderID = getNextImplicitBindingOrderID();
715 if (ResourceAttrs.hasBinding())
716 ResourceAttrs.setImplicitOrderID(OrderID);
717 else
718 addImplicitBindingAttrToDecl(S&: SemaRef, D: BufDecl,
719 RT: BufDecl->isCBuffer() ? RegisterType::CBuffer
720 : RegisterType::SRV,
721 ImplicitBindingOrderID: OrderID);
722 }
723
724 SemaRef.PopDeclContext();
725}
726
727HLSLNumThreadsAttr *SemaHLSL::mergeNumThreadsAttr(Decl *D,
728 const AttributeCommonInfo &AL,
729 int X, int Y, int Z) {
730 if (HLSLNumThreadsAttr *NT = D->getAttr<HLSLNumThreadsAttr>()) {
731 if (NT->getX() != X || NT->getY() != Y || NT->getZ() != Z) {
732 Diag(Loc: NT->getLocation(), DiagID: diag::err_hlsl_attribute_param_mismatch) << AL;
733 Diag(Loc: AL.getLoc(), DiagID: diag::note_conflicting_attribute);
734 }
735 return nullptr;
736 }
737 return ::new (getASTContext())
738 HLSLNumThreadsAttr(getASTContext(), AL, X, Y, Z);
739}
740
741HLSLWaveSizeAttr *SemaHLSL::mergeWaveSizeAttr(Decl *D,
742 const AttributeCommonInfo &AL,
743 int Min, int Max, int Preferred,
744 int SpelledArgsCount) {
745 if (HLSLWaveSizeAttr *WS = D->getAttr<HLSLWaveSizeAttr>()) {
746 if (WS->getMin() != Min || WS->getMax() != Max ||
747 WS->getPreferred() != Preferred ||
748 WS->getSpelledArgsCount() != SpelledArgsCount) {
749 Diag(Loc: WS->getLocation(), DiagID: diag::err_hlsl_attribute_param_mismatch) << AL;
750 Diag(Loc: AL.getLoc(), DiagID: diag::note_conflicting_attribute);
751 }
752 return nullptr;
753 }
754 HLSLWaveSizeAttr *Result = ::new (getASTContext())
755 HLSLWaveSizeAttr(getASTContext(), AL, Min, Max, Preferred);
756 Result->setSpelledArgsCount(SpelledArgsCount);
757 return Result;
758}
759
760HLSLVkConstantIdAttr *
761SemaHLSL::mergeVkConstantIdAttr(Decl *D, const AttributeCommonInfo &AL,
762 int Id) {
763
764 auto &TargetInfo = getASTContext().getTargetInfo();
765 if (TargetInfo.getTriple().getArch() != llvm::Triple::spirv) {
766 Diag(Loc: AL.getLoc(), DiagID: diag::warn_attribute_ignored) << AL;
767 return nullptr;
768 }
769
770 auto *VD = cast<VarDecl>(Val: D);
771
772 if (getSpecConstBuiltinId(Type: VD->getType()->getUnqualifiedDesugaredType()) ==
773 Builtin::NotBuiltin) {
774 Diag(Loc: VD->getLocation(), DiagID: diag::err_specialization_const);
775 return nullptr;
776 }
777
778 if (!VD->getType().isConstQualified()) {
779 Diag(Loc: VD->getLocation(), DiagID: diag::err_specialization_const);
780 return nullptr;
781 }
782
783 if (HLSLVkConstantIdAttr *CI = D->getAttr<HLSLVkConstantIdAttr>()) {
784 if (CI->getId() != Id) {
785 Diag(Loc: CI->getLocation(), DiagID: diag::err_hlsl_attribute_param_mismatch) << AL;
786 Diag(Loc: AL.getLoc(), DiagID: diag::note_conflicting_attribute);
787 }
788 return nullptr;
789 }
790
791 HLSLVkConstantIdAttr *Result =
792 ::new (getASTContext()) HLSLVkConstantIdAttr(getASTContext(), AL, Id);
793 return Result;
794}
795
796HLSLShaderAttr *
797SemaHLSL::mergeShaderAttr(Decl *D, const AttributeCommonInfo &AL,
798 llvm::Triple::EnvironmentType ShaderType) {
799 if (HLSLShaderAttr *NT = D->getAttr<HLSLShaderAttr>()) {
800 if (NT->getType() != ShaderType) {
801 Diag(Loc: NT->getLocation(), DiagID: diag::err_hlsl_attribute_param_mismatch) << AL;
802 Diag(Loc: AL.getLoc(), DiagID: diag::note_conflicting_attribute);
803 }
804 return nullptr;
805 }
806 return HLSLShaderAttr::Create(Ctx&: getASTContext(), Type: ShaderType, CommonInfo: AL);
807}
808
809HLSLParamModifierAttr *
810SemaHLSL::mergeParamModifierAttr(Decl *D, const AttributeCommonInfo &AL,
811 HLSLParamModifierAttr::Spelling Spelling) {
812 // We can only merge an `in` attribute with an `out` attribute. All other
813 // combinations of duplicated attributes are ill-formed.
814 if (HLSLParamModifierAttr *PA = D->getAttr<HLSLParamModifierAttr>()) {
815 if ((PA->isIn() && Spelling == HLSLParamModifierAttr::Keyword_out) ||
816 (PA->isOut() && Spelling == HLSLParamModifierAttr::Keyword_in)) {
817 D->dropAttr<HLSLParamModifierAttr>();
818 SourceRange AdjustedRange = {PA->getLocation(), AL.getRange().getEnd()};
819 return HLSLParamModifierAttr::Create(
820 Ctx&: getASTContext(), /*MergedSpelling=*/true, Range: AdjustedRange,
821 S: HLSLParamModifierAttr::Keyword_inout);
822 }
823 Diag(Loc: AL.getLoc(), DiagID: diag::err_hlsl_duplicate_parameter_modifier) << AL;
824 Diag(Loc: PA->getLocation(), DiagID: diag::note_conflicting_attribute);
825 return nullptr;
826 }
827 return HLSLParamModifierAttr::Create(Ctx&: getASTContext(), CommonInfo: AL);
828}
829
830void SemaHLSL::handleInterpolationModifierAttr(Decl *D, const ParsedAttr &AL) {
831 InterpolationModifier Modifier;
832 switch (static_cast<HLSLInterpolationModifierAttr::Spelling>(
833 AL.getSemanticSpelling())) {
834 case HLSLInterpolationModifierAttr::Keyword_nointerpolation:
835 Modifier = InterpolationModifier::NoInterpolation;
836 break;
837 case HLSLInterpolationModifierAttr::Keyword_linear:
838 Modifier = InterpolationModifier::Linear;
839 break;
840 case HLSLInterpolationModifierAttr::Keyword_centroid:
841 Modifier = InterpolationModifier::Centroid;
842 break;
843 case HLSLInterpolationModifierAttr::Keyword_noperspective:
844 Modifier = InterpolationModifier::NoPerspective;
845 break;
846 case HLSLInterpolationModifierAttr::Keyword_sample:
847 Modifier = InterpolationModifier::Sample;
848 break;
849 case HLSLInterpolationModifierAttr::Keyword_center:
850 Modifier = InterpolationModifier::Center;
851 break;
852 case HLSLInterpolationModifierAttr::SpellingNotCalculated:
853 llvm_unreachable("interpolation modifier spelling was not calculated");
854 }
855
856 InterpolationModifier Modifiers = Modifier;
857 if (auto *Previous = D->getAttr<HLSLInterpolationModifierAttr>()) {
858 auto Old = static_cast<InterpolationModifier>(Previous->getModifiers());
859 Modifiers |= Old;
860 if (any(Val: Old & Modifier)) {
861 Diag(Loc: AL.getLoc(), DiagID: diag::warn_hlsl_duplicate_interpolation) << AL;
862 } else if (llvm::hlsl::getInterpolationMode(Modifiers) ==
863 llvm::dxbc::PSV::InterpolationMode::Invalid &&
864 llvm::hlsl::getInterpolationMode(Modifiers: Old) !=
865 llvm::dxbc::PSV::InterpolationMode::Invalid) {
866 Diag(Loc: AL.getLoc(), DiagID: diag::err_hlsl_interpolation_conflict);
867 Diag(Loc: Previous->getLocation(), DiagID: diag::note_conflicting_attribute);
868 D->setInvalidDecl();
869 } else {
870 InterpolationModifier OldLocation =
871 llvm::hlsl::getInterpolationSamplingLocation(Modifiers: Old);
872 InterpolationModifier NewLocation =
873 llvm::hlsl::getInterpolationSamplingLocation(Modifiers: Modifier);
874 if (any(Val: OldLocation) && any(Val: NewLocation)) {
875 Diag(Loc: AL.getLoc(), DiagID: diag::warn_hlsl_interpolation_override)
876 << (std::max(a: OldLocation, b: NewLocation) ==
877 InterpolationModifier::Sample)
878 << (std::min(a: OldLocation, b: NewLocation) ==
879 InterpolationModifier::Centroid);
880 }
881 }
882 D->dropAttr<HLSLInterpolationModifierAttr>();
883 }
884 D->addAttr(A: HLSLInterpolationModifierAttr::Create(
885 Ctx&: getASTContext(), Modifiers: static_cast<unsigned>(Modifiers), CommonInfo: AL));
886}
887
888bool SemaHLSL::checkInterpolationModifiers(
889 const DeclaratorDecl *D, const HLSLInterpolationModifierAttr *Inherited,
890 const HLSLParsedSemanticAttr *Semantic) {
891 if (D->isInvalidDecl())
892 return false;
893 const auto *A = D->getAttr<HLSLInterpolationModifierAttr>();
894 if (!A)
895 A = Inherited;
896 if (!Semantic)
897 Semantic = D->getAttr<HLSLParsedSemanticAttr>();
898
899 const auto *FD = dyn_cast<FunctionDecl>(Val: D);
900 QualType T = FD ? FD->getReturnType() : D->getType();
901 T = getASTContext().getBaseElementType(QT: T.getNonReferenceType());
902 if (T->isDependentType())
903 return true;
904 if (const auto *RT = T->getAs<RecordType>()) {
905 const RecordDecl *RD = RT->getDecl()->getDefinition();
906 if (!RD)
907 return true;
908 bool Valid = true;
909 for (const FieldDecl *Field : RD->fields())
910 Valid &= checkInterpolationModifiers(D: Field, Inherited: A, Semantic);
911 return Valid;
912 }
913 if (!A)
914 return true;
915
916 auto Modifiers = static_cast<InterpolationModifier>(A->getModifiers());
917 if (Modifiers == InterpolationModifier::NoInterpolation) {
918 bool IsPosition =
919 Semantic && llvm::hlsl::getSemanticKind(SemanticName: Semantic->getSemanticName()) ==
920 SemanticKind::Position;
921 if (!IsPosition)
922 return true;
923 Diag(Loc: A->getLocation(), DiagID: diag::err_hlsl_interpolation_position);
924 Diag(Loc: Semantic->getLocation(), DiagID: diag::note_conflicting_attribute);
925 return false;
926 }
927
928 T = getScalarComponentType(T);
929 if (T->isIntegerType() ||
930 (T->isRealFloatingType() && getASTContext().getTypeSize(T) > 32)) {
931 Diag(Loc: A->getLocation(), DiagID: diag::err_hlsl_interpolation_type) << T;
932 return false;
933 }
934 return true;
935}
936
937void SemaHLSL::ActOnTopLevelFunction(FunctionDecl *FD) {
938 auto &TargetInfo = getASTContext().getTargetInfo();
939
940 if (FD->getName() != TargetInfo.getTargetOpts().HLSLEntry)
941 return;
942
943 // If we have specified a root signature to override the entry function then
944 // attach it now
945 HLSLRootSignatureDecl *SignatureDecl =
946 lookupRootSignatureOverrideDecl(DC: FD->getDeclContext());
947 if (SignatureDecl) {
948 FD->dropAttr<RootSignatureAttr>();
949 // We could look up the SourceRange of the macro here as well
950 AttributeCommonInfo AL(RootSigOverrideIdent, AttributeScopeInfo(),
951 SourceRange(), ParsedAttr::Form::Microsoft());
952 FD->addAttr(A: ::new (getASTContext()) RootSignatureAttr(
953 getASTContext(), AL, RootSigOverrideIdent, SignatureDecl));
954 }
955
956 llvm::Triple::EnvironmentType Env = TargetInfo.getTriple().getEnvironment();
957 if (HLSLShaderAttr::isValidShaderType(ShaderType: Env) && Env != llvm::Triple::Library) {
958 if (const auto *Shader = FD->getAttr<HLSLShaderAttr>()) {
959 // The entry point is already annotated - check that it matches the
960 // triple.
961 if (Shader->getType() != Env) {
962 Diag(Loc: Shader->getLocation(), DiagID: diag::err_hlsl_entry_shader_attr_mismatch)
963 << Shader;
964 FD->setInvalidDecl();
965 }
966 } else {
967 // Implicitly add the shader attribute if the entry function isn't
968 // explicitly annotated.
969 FD->addAttr(A: HLSLShaderAttr::CreateImplicit(Ctx&: getASTContext(), Type: Env,
970 Range: FD->getBeginLoc()));
971 }
972 } else {
973 switch (Env) {
974 case llvm::Triple::UnknownEnvironment:
975 case llvm::Triple::Library:
976 break;
977 case llvm::Triple::RootSignature:
978 llvm_unreachable("rootsig environment has no functions");
979 default:
980 llvm_unreachable("Unhandled environment in triple");
981 }
982 }
983}
984
985static bool isVkPipelineBuiltin(const ASTContext &AstContext, FunctionDecl *FD,
986 HLSLAppliedSemanticAttr *Semantic,
987 bool IsInput) {
988 if (AstContext.getTargetInfo().getTriple().getOS() != llvm::Triple::Vulkan)
989 return false;
990
991 const auto *ShaderAttr = FD->getAttr<HLSLShaderAttr>();
992 assert(ShaderAttr && "Entry point has no shader attribute");
993 llvm::Triple::EnvironmentType ST = ShaderAttr->getType();
994 SemanticKind Kind = llvm::hlsl::getSemanticKind(SemanticName: Semantic->getSemanticName());
995
996 switch (Kind) {
997 case SemanticKind::Position:
998 // The SV_Position semantic is lowered to:
999 // - Position built-in for vertex output.
1000 // - FragCoord built-in for fragment input.
1001 return (ST == llvm::Triple::Vertex && !IsInput) ||
1002 (ST == llvm::Triple::Pixel && IsInput);
1003 case SemanticKind::VertexID:
1004 return true;
1005 case SemanticKind::InstanceID:
1006 return ST == llvm::Triple::Vertex && IsInput;
1007 default:
1008 return false;
1009 }
1010}
1011
1012bool SemaHLSL::determineActiveSemanticOnScalar(FunctionDecl *FD,
1013 DeclaratorDecl *OutputDecl,
1014 DeclaratorDecl *D,
1015 SemanticInfo &ActiveSemantic,
1016 SemaHLSL::SemanticContext &SC) {
1017 if (ActiveSemantic.Semantic == nullptr) {
1018 ActiveSemantic.Semantic = D->getAttr<HLSLParsedSemanticAttr>();
1019 if (ActiveSemantic.Semantic)
1020 ActiveSemantic.Index = ActiveSemantic.Semantic->getSemanticIndex();
1021 }
1022
1023 if (!ActiveSemantic.Semantic) {
1024 Diag(Loc: D->getLocation(), DiagID: diag::err_hlsl_missing_semantic_annotation);
1025 return false;
1026 }
1027
1028 auto *A = ::new (getASTContext())
1029 HLSLAppliedSemanticAttr(getASTContext(), *ActiveSemantic.Semantic,
1030 ActiveSemantic.Semantic->getAttrName()->getName(),
1031 ActiveSemantic.Index.value_or(u: 0));
1032 if (!A)
1033 return false;
1034
1035 // Each array element occupies a separate semantic index.
1036 QualType T = D == FD ? FD->getReturnType() : D->getType();
1037 const ConstantArrayType *AT =
1038 getASTContext().getAsConstantArrayType(T: T.getNonReferenceType());
1039 if (isZeroSizedArray(CAT: AT)) {
1040 Diag(Loc: A->getLoc(), DiagID: diag::err_hlsl_semantic_zero_sized_array)
1041 << A->getAttrName();
1042 return false;
1043 }
1044 unsigned ElementCount = AT ? ASTContext::getConstantArrayElementCount(CA: AT) : 1;
1045
1046 checkSemanticAnnotation(EntryPoint: FD, Param: D, SemanticAttr: A, SC, ElementCount);
1047 OutputDecl->addAttr(A);
1048
1049 unsigned Location = ActiveSemantic.Index.value_or(u: 0);
1050
1051 if (!isVkPipelineBuiltin(AstContext: getASTContext(), FD, Semantic: A,
1052 IsInput: any(Val: SC.CurrentIOType & IOType::In))) {
1053 bool HasVkLocation = false;
1054 if (auto *A = D->getAttr<HLSLVkLocationAttr>()) {
1055 HasVkLocation = true;
1056 Location = A->getLocation();
1057 }
1058
1059 if (SC.UsesExplicitVkLocations.value_or(u&: HasVkLocation) != HasVkLocation) {
1060 Diag(Loc: D->getLocation(), DiagID: diag::err_hlsl_semantic_partial_explicit_indexing);
1061 return false;
1062 }
1063 SC.UsesExplicitVkLocations = HasVkLocation;
1064 }
1065
1066 ActiveSemantic.Index = Location + ElementCount;
1067
1068 StringRef BaseName = ActiveSemantic.Semantic->getAttrName()->getName();
1069 std::string LowerName = BaseName.lower();
1070 for (unsigned I = 0; I < ElementCount; ++I) {
1071 auto [It, Inserted] = SC.ActiveSemantics.try_emplace(
1072 Key: (Twine(LowerName) + Twine(Location + I)).str(), Args: D->getLocation());
1073 if (!Inserted) {
1074 Diag(Loc: D->getLocation(), DiagID: diag::err_hlsl_semantic_index_overlap)
1075 << (BaseName + Twine(Location + I)).str();
1076 Diag(Loc: It->second, DiagID: diag::note_previous_use);
1077 return false;
1078 }
1079 }
1080
1081 return true;
1082}
1083
1084bool SemaHLSL::determineActiveSemantic(FunctionDecl *FD,
1085 DeclaratorDecl *OutputDecl,
1086 DeclaratorDecl *D,
1087 SemanticInfo &ActiveSemantic,
1088 SemaHLSL::SemanticContext &SC) {
1089 if (ActiveSemantic.Semantic == nullptr) {
1090 ActiveSemantic.Semantic = D->getAttr<HLSLParsedSemanticAttr>();
1091 if (ActiveSemantic.Semantic)
1092 ActiveSemantic.Index = ActiveSemantic.Semantic->getSemanticIndex();
1093 }
1094
1095 const Type *T = D == FD ? &*FD->getReturnType() : &*D->getType();
1096 T = T->getUnqualifiedDesugaredType();
1097
1098 const RecordType *RT = dyn_cast<RecordType>(Val: T);
1099 if (!RT)
1100 return determineActiveSemanticOnScalar(FD, OutputDecl, D, ActiveSemantic,
1101 SC);
1102
1103 const RecordDecl *RD = RT->getDecl();
1104 for (FieldDecl *Field : RD->fields()) {
1105 SemanticInfo Info = ActiveSemantic;
1106 if (!determineActiveSemantic(FD, OutputDecl, D: Field, ActiveSemantic&: Info, SC)) {
1107 Diag(Loc: Field->getLocation(), DiagID: diag::note_hlsl_semantic_used_here) << Field;
1108 return false;
1109 }
1110 if (ActiveSemantic.Semantic)
1111 ActiveSemantic = Info;
1112 }
1113
1114 return true;
1115}
1116
1117void SemaHLSL::CheckEntryPoint(FunctionDecl *FD) {
1118 const auto *ShaderAttr = FD->getAttr<HLSLShaderAttr>();
1119 assert(ShaderAttr && "Entry point has no shader attribute");
1120 llvm::Triple::EnvironmentType ST = ShaderAttr->getType();
1121 auto &TargetInfo = getASTContext().getTargetInfo();
1122 VersionTuple Ver = TargetInfo.getTriple().getOSVersion();
1123 switch (ST) {
1124 case llvm::Triple::Pixel:
1125 case llvm::Triple::Vertex:
1126 case llvm::Triple::Geometry:
1127 case llvm::Triple::Hull:
1128 case llvm::Triple::Domain:
1129 case llvm::Triple::RayGeneration:
1130 case llvm::Triple::Intersection:
1131 case llvm::Triple::AnyHit:
1132 case llvm::Triple::ClosestHit:
1133 case llvm::Triple::Miss:
1134 case llvm::Triple::Callable:
1135 if (const auto *NT = FD->getAttr<HLSLNumThreadsAttr>()) {
1136 diagnoseAttrStageMismatch(A: NT, Stage: ST,
1137 AllowedStages: {llvm::Triple::Compute,
1138 llvm::Triple::Amplification,
1139 llvm::Triple::Mesh});
1140 FD->setInvalidDecl();
1141 }
1142 if (const auto *WS = FD->getAttr<HLSLWaveSizeAttr>()) {
1143 diagnoseAttrStageMismatch(A: WS, Stage: ST,
1144 AllowedStages: {llvm::Triple::Compute,
1145 llvm::Triple::Amplification,
1146 llvm::Triple::Mesh});
1147 FD->setInvalidDecl();
1148 }
1149 break;
1150
1151 case llvm::Triple::Compute:
1152 case llvm::Triple::Amplification:
1153 case llvm::Triple::Mesh:
1154 if (!FD->hasAttr<HLSLNumThreadsAttr>()) {
1155 Diag(Loc: FD->getLocation(), DiagID: diag::err_hlsl_missing_numthreads)
1156 << llvm::Triple::getEnvironmentTypeName(Kind: ST);
1157 FD->setInvalidDecl();
1158 }
1159 if (const auto *WS = FD->getAttr<HLSLWaveSizeAttr>()) {
1160 if (TargetInfo.getTriple().isSPIRV()) {
1161 Diag(Loc: WS->getLocation(), DiagID: diag::warn_hlsl_wavesize_unsupported_spirv);
1162 } else if (Ver < VersionTuple(6, 6)) {
1163 Diag(Loc: WS->getLocation(), DiagID: diag::err_hlsl_attribute_in_wrong_shader_model)
1164 << WS << "6.6";
1165 FD->setInvalidDecl();
1166 } else if (WS->getSpelledArgsCount() > 1 && Ver < VersionTuple(6, 8)) {
1167 Diag(
1168 Loc: WS->getLocation(),
1169 DiagID: diag::err_hlsl_attribute_number_arguments_insufficient_shader_model)
1170 << WS << WS->getSpelledArgsCount() << "6.8";
1171 FD->setInvalidDecl();
1172 }
1173 }
1174 break;
1175 case llvm::Triple::RootSignature:
1176 llvm_unreachable("rootsig environment has no function entry point");
1177 default:
1178 llvm_unreachable("Unhandled environment in triple");
1179 }
1180
1181 SemaHLSL::SemanticContext InputSC = {};
1182 InputSC.CurrentIOType = IOType::In;
1183 SemaHLSL::SemanticContext OutputSC = {};
1184 OutputSC.CurrentIOType = IOType::Out;
1185
1186 for (ParmVarDecl *Param : FD->parameters()) {
1187 SemanticInfo ActiveSemantic;
1188 ActiveSemantic.Semantic = Param->getAttr<HLSLParsedSemanticAttr>();
1189 if (ActiveSemantic.Semantic)
1190 ActiveSemantic.Index = ActiveSemantic.Semantic->getSemanticIndex();
1191
1192 // FIXME: An `inout` parameter is part of both signatures, but it is only
1193 // verified against the output one here.
1194 const auto *MA = Param->getAttr<HLSLParamModifierAttr>();
1195 SemanticContext &SC = MA && MA->isAnyOut() ? OutputSC : InputSC;
1196
1197 // Interpolation applies to pixel inputs and vertex outputs, including the
1198 // corresponding side of inout parameters.
1199 if (((ST == llvm::Triple::Pixel && (!MA || MA->isAnyIn())) ||
1200 (ST == llvm::Triple::Vertex && MA && MA->isAnyOut())) &&
1201 !checkInterpolationModifiers(D: Param, Inherited: nullptr, Semantic: nullptr))
1202 FD->setInvalidDecl();
1203
1204 if (!determineActiveSemantic(FD, OutputDecl: Param, D: Param, ActiveSemantic, SC)) {
1205 Diag(Loc: Param->getLocation(), DiagID: diag::note_previous_decl) << Param;
1206 FD->setInvalidDecl();
1207 }
1208 }
1209
1210 SemanticInfo ActiveSemantic;
1211 ActiveSemantic.Semantic = FD->getAttr<HLSLParsedSemanticAttr>();
1212 if (ActiveSemantic.Semantic)
1213 ActiveSemantic.Index = ActiveSemantic.Semantic->getSemanticIndex();
1214 if (!FD->getReturnType()->isVoidType()) {
1215 if (ST == llvm::Triple::Vertex &&
1216 !checkInterpolationModifiers(D: FD, Inherited: nullptr, Semantic: nullptr))
1217 FD->setInvalidDecl();
1218 determineActiveSemantic(FD, OutputDecl: FD, D: FD, ActiveSemantic, SC&: OutputSC);
1219 }
1220}
1221
1222void SemaHLSL::checkSemanticAnnotation(
1223 FunctionDecl *EntryPoint, const Decl *Param,
1224 const HLSLAppliedSemanticAttr *SemanticAttr, const SemanticContext &SC,
1225 unsigned ElementCount) {
1226 auto *ShaderAttr = EntryPoint->getAttr<HLSLShaderAttr>();
1227 assert(ShaderAttr && "Entry point has no shader attribute");
1228 llvm::Triple::EnvironmentType ST = ShaderAttr->getType();
1229
1230 SemanticKind Kind =
1231 llvm::hlsl::getSemanticKind(SemanticName: SemanticAttr->getSemanticName());
1232 llvm::hlsl::SemanticInterpretation Interpretation =
1233 llvm::hlsl::getInterpretationKind(SemanticKind: Kind, ShaderStage: ST, IOTy: SC.CurrentIOType);
1234 if (Interpretation == llvm::hlsl::SemanticInterpretation::Invalid) {
1235 diagnoseSemanticStageMismatch(A: SemanticAttr, Stage: ST, CurrentIOType: SC.CurrentIOType, SemanticKind: Kind);
1236 return;
1237 }
1238
1239 // A system-value name can have an arbitrary interpretation, for example
1240 // SV_Position on a vertex input. Only the general type restrictions apply.
1241 if (Interpretation == llvm::hlsl::SemanticInterpretation::Arbitrary) {
1242 diagnoseSemanticType(D: Param, A: SemanticAttr, SemanticKind: SemanticKind::Arbitrary);
1243 return;
1244 }
1245
1246 diagnoseSystemSemanticIndex(A: SemanticAttr, SemanticKind: Kind, ElementCount);
1247 diagnoseSemanticType(D: Param, A: SemanticAttr, SemanticKind: Kind);
1248}
1249
1250void SemaHLSL::diagnoseSystemSemanticIndex(const HLSLAppliedSemanticAttr *A,
1251 SemanticKind Kind,
1252 unsigned ElementCount) {
1253 assert(Kind != SemanticKind::Invalid && Kind != SemanticKind::Arbitrary &&
1254 "expected a recognized system semantic");
1255 assert(ElementCount > 0 && "a semantic covers at least one element");
1256 // The attribute stores the index in an int. Recover its unsigned value
1257 // before widening the arithmetic to detect overflow of the semantic range.
1258 uint32_t FirstIndex = A->getSemanticIndex();
1259 uint64_t LastIndex = uint64_t(FirstIndex) + ElementCount - 1;
1260 constexpr uint32_t MaxSemanticIndex = std::numeric_limits<uint32_t>::max();
1261 if (LastIndex > MaxSemanticIndex) {
1262 Diag(Loc: A->getLoc(), DiagID: diag::err_hlsl_semantic_index_out_of_range)
1263 << A->getAttrName() << LastIndex << MaxSemanticIndex;
1264 return;
1265 }
1266 if (LastIndex == 0)
1267 return;
1268
1269 switch (Kind) {
1270 // These semantics are limited by signature packing, not semantic indices.
1271 case SemanticKind::ClipDistance:
1272 case SemanticKind::CullDistance:
1273 return;
1274 case SemanticKind::Target: {
1275 constexpr unsigned MaxTargetIndex = 7;
1276 if (LastIndex > MaxTargetIndex)
1277 Diag(Loc: A->getLoc(), DiagID: diag::err_hlsl_semantic_index_out_of_range)
1278 << A->getAttrName() << LastIndex << MaxTargetIndex;
1279 return;
1280 }
1281 default:
1282 Diag(Loc: A->getLoc(), DiagID: diag::err_hlsl_semantic_indexing_not_supported)
1283 << A->getAttrName();
1284 return;
1285 }
1286}
1287
1288static QualType getElementTypeOf(QualType T, bool IncludeMatrix) {
1289 if (const auto *VT = T->getAs<clang::VectorType>())
1290 return VT->getElementType();
1291 if (IncludeMatrix)
1292 if (const auto *MT = T->getAs<clang::MatrixType>())
1293 return MT->getElementType();
1294 return T;
1295}
1296
1297static unsigned getComponentCountOf(QualType T) {
1298 if (const auto *VT = T->getAs<clang::VectorType>())
1299 return VT->getNumElements();
1300 return 1;
1301}
1302
1303static bool isFloatOrHalfElement(QualType Elem) {
1304 return Elem->isHalfType() || Elem->isFloat16Type() || Elem->isFloat32Type();
1305}
1306
1307// System-value integer types exclude bool.
1308static bool isIntElementOfWidth(const ASTContext &Ctx, QualType Elem,
1309 uint64_t Width) {
1310 if (!Elem->isIntegerType() || Elem->isBooleanType())
1311 return false;
1312 return Ctx.getTypeSize(T: Elem) == Width;
1313}
1314
1315static bool isIntUpTo32Element(const ASTContext &Ctx, QualType Elem) {
1316 return isIntElementOfWidth(Ctx, Elem, Width: 16) ||
1317 isIntElementOfWidth(Ctx, Elem, Width: 32);
1318}
1319
1320void SemaHLSL::diagnoseSemanticType(const Decl *D,
1321 const HLSLAppliedSemanticAttr *A,
1322 SemanticKind Kind) {
1323 assert(Kind != SemanticKind::Invalid && "expected a valid semantic");
1324 ASTContext &Ctx = getASTContext();
1325
1326 QualType T;
1327 if (const auto *FD = dyn_cast<FunctionDecl>(Val: D))
1328 T = FD->getReturnType();
1329 else
1330 T = cast<ValueDecl>(Val: D)->getType();
1331
1332 // `out` and `inout` parameters are passed by reference.
1333 T = T.getNonReferenceType();
1334
1335 // Array semantics constrain each element's type.
1336 QualType DeclaredTy = T;
1337 while (const ConstantArrayType *AT = Ctx.getAsConstantArrayType(T))
1338 T = AT->getElementType();
1339
1340 QualType ElemTy = getElementTypeOf(T, /*IncludeMatrix=*/false);
1341 unsigned Components = getComponentCountOf(T);
1342
1343 bool IsSPIRV = getASTContext().getTargetInfo().getTriple().isSPIRV();
1344
1345 switch (Kind) {
1346 case SemanticKind::DispatchThreadID:
1347 case SemanticKind::GroupID:
1348 case SemanticKind::GroupThreadID:
1349 if (!isIntUpTo32Element(Ctx, Elem: ElemTy) || Components > 3)
1350 Diag(Loc: A->getLoc(), DiagID: diag::err_hlsl_semantic_invalid_type)
1351 << A->getAttrName() << /* scalar or vector of up to */ 1 << 3
1352 << /* 16 or 32 bit integer */ 0 << DeclaredTy;
1353 return;
1354 case SemanticKind::GroupIndex:
1355 if (!isIntElementOfWidth(Ctx, Elem: ElemTy, Width: 32) || Components != 1)
1356 Diag(Loc: A->getLoc(), DiagID: diag::err_hlsl_semantic_invalid_type)
1357 << A->getAttrName() << /* scalar */ 0 << 1 << /* 32 bit integer */ 1
1358 << DeclaredTy;
1359 return;
1360 case SemanticKind::VertexID:
1361 if (!isIntUpTo32Element(Ctx, Elem: ElemTy) || Components != 1)
1362 Diag(Loc: A->getLoc(), DiagID: diag::err_hlsl_semantic_invalid_type)
1363 << A->getAttrName() << /* scalar */ 0 << 1
1364 << /* 16 or 32 bit integer */ 0 << DeclaredTy;
1365 return;
1366 case SemanticKind::Position:
1367 case SemanticKind::Target:
1368 if (!isFloatOrHalfElement(Elem: ElemTy) || Components > 4)
1369 Diag(Loc: A->getLoc(), DiagID: diag::err_hlsl_semantic_invalid_type)
1370 << A->getAttrName() << /* scalar or vector of up to */ 1 << 4
1371 << /* 16 or 32 bit floating-point */ 2 << DeclaredTy;
1372 return;
1373 case SemanticKind::InstanceID:
1374 // DXIL permits U32 or U16. SPIR-V requires a 32-bit scalar per
1375 // VUID-InstanceIndex-InstanceIndex-04265.
1376 if (!T->isUnsignedIntegerType() ||
1377 !(isIntElementOfWidth(Ctx, Elem: T, Width: 32) ||
1378 (!IsSPIRV && isIntElementOfWidth(Ctx, Elem: T, Width: 16))))
1379 Diag(Loc: A->getLoc(), DiagID: diag::err_hlsl_semantic_invalid_type)
1380 << A->getAttrName() << /* scalar */ 0 << 1
1381 << /* 16 or 32 bit unsigned integer / 32 bit unsigned integer */
1382 (IsSPIRV ? 4 : 3) << DeclaredTy;
1383 return;
1384 default:
1385 // Other semantics only have the general signature type restrictions.
1386 break;
1387 }
1388
1389 // DXIL signatures cannot carry 64-bit components, even for arbitrary
1390 // semantics. SPIR-V interfaces do support these types.
1391 QualType ScalarTy = getElementTypeOf(T, /*IncludeMatrix=*/true);
1392 if (Ctx.getTargetInfo().getTriple().isDXIL() &&
1393 (ScalarTy->isSpecificBuiltinType(K: BuiltinType::Double) ||
1394 isIntElementOfWidth(Ctx, Elem: ScalarTy, Width: 64)))
1395 Diag(Loc: A->getLoc(), DiagID: diag::err_hlsl_semantic_64bit_type)
1396 << A->getAttrName() << DeclaredTy;
1397}
1398
1399void SemaHLSL::diagnoseAttrStageMismatch(
1400 const Attr *A, llvm::Triple::EnvironmentType Stage,
1401 std::initializer_list<llvm::Triple::EnvironmentType> AllowedStages) {
1402 SmallVector<StringRef, 8> StageStrings;
1403 llvm::transform(Range&: AllowedStages, d_first: std::back_inserter(x&: StageStrings),
1404 F: [](llvm::Triple::EnvironmentType ST) {
1405 return StringRef(
1406 HLSLShaderAttr::ConvertEnvironmentTypeToStr(Val: ST));
1407 });
1408 Diag(Loc: A->getLoc(), DiagID: diag::err_hlsl_attr_unsupported_in_stage)
1409 << A->getAttrName() << llvm::Triple::getEnvironmentTypeName(Kind: Stage)
1410 << (AllowedStages.size() != 1) << join(R&: StageStrings, Separator: ", ");
1411}
1412
1413void SemaHLSL::diagnoseSemanticStageMismatch(
1414 const Attr *A, llvm::Triple::EnvironmentType Stage, IOType CurrentIOType,
1415 SemanticKind Kind) {
1416
1417 ArrayRef<SemanticStageInfo> Allowed = llvm::hlsl::getAvailableStages(SemanticKind: Kind);
1418 auto It = llvm::find_if(Range&: Allowed, P: [&Stage](const SemanticStageInfo &Info) {
1419 return Info.Stage == Stage;
1420 });
1421
1422 StringRef CurrentIOTypeName = "patch constants or primitives";
1423 if (any(Val: CurrentIOType & IOType::In))
1424 CurrentIOTypeName = "inputs";
1425 else if (any(Val: CurrentIOType & IOType::Out))
1426 CurrentIOTypeName = "outputs";
1427
1428 // The semantic is not available in this shader stage at all.
1429 if (It == Allowed.end()) {
1430 Diag(Loc: A->getLoc(), DiagID: diag::err_hlsl_semantic_unsupported_iotype_for_stage)
1431 << A->getAttrName() << llvm::Triple::getEnvironmentTypeName(Kind: Stage)
1432 << CurrentIOTypeName;
1433 return;
1434 }
1435
1436 IOType AllowedIOTypes = It->AllowedIOTypesMask;
1437 if (!(AllowedIOTypes & CurrentIOType)) {
1438 Diag(Loc: A->getLoc(), DiagID: diag::err_hlsl_semantic_unsupported_iotype_for_stage)
1439 << A->getAttrName() << llvm::Triple::getEnvironmentTypeName(Kind: Stage)
1440 << CurrentIOTypeName;
1441 return;
1442 }
1443}
1444
1445template <CastKind Kind>
1446static void castVector(Sema &S, ExprResult &E, QualType &Ty, unsigned Sz) {
1447 if (const auto *VTy = Ty->getAs<VectorType>())
1448 Ty = VTy->getElementType();
1449 Ty = S.getASTContext().getExtVectorType(VectorType: Ty, NumElts: Sz);
1450 E = S.ImpCastExprToType(E: E.get(), Type: Ty, CK: Kind);
1451}
1452
1453template <CastKind Kind>
1454static QualType castElement(Sema &S, ExprResult &E, QualType Ty) {
1455 E = S.ImpCastExprToType(E: E.get(), Type: Ty, CK: Kind);
1456 return Ty;
1457}
1458
1459static QualType handleFloatVectorBinOpConversion(
1460 Sema &SemaRef, ExprResult &LHS, ExprResult &RHS, QualType LHSType,
1461 QualType RHSType, QualType LElTy, QualType RElTy, bool IsCompAssign) {
1462 bool LHSFloat = LElTy->isRealFloatingType();
1463 bool RHSFloat = RElTy->isRealFloatingType();
1464
1465 if (LHSFloat && RHSFloat) {
1466 if (IsCompAssign ||
1467 SemaRef.getASTContext().getFloatingTypeOrder(LHS: LElTy, RHS: RElTy) > 0)
1468 return castElement<CK_FloatingCast>(S&: SemaRef, E&: RHS, Ty: LHSType);
1469
1470 return castElement<CK_FloatingCast>(S&: SemaRef, E&: LHS, Ty: RHSType);
1471 }
1472
1473 if (LHSFloat)
1474 return castElement<CK_IntegralToFloating>(S&: SemaRef, E&: RHS, Ty: LHSType);
1475
1476 assert(RHSFloat);
1477 if (IsCompAssign)
1478 return castElement<clang::CK_FloatingToIntegral>(S&: SemaRef, E&: RHS, Ty: LHSType);
1479
1480 return castElement<CK_IntegralToFloating>(S&: SemaRef, E&: LHS, Ty: RHSType);
1481}
1482
1483static QualType handleIntegerVectorBinOpConversion(
1484 Sema &SemaRef, ExprResult &LHS, ExprResult &RHS, QualType LHSType,
1485 QualType RHSType, QualType LElTy, QualType RElTy, bool IsCompAssign) {
1486
1487 int IntOrder = SemaRef.Context.getIntegerTypeOrder(LHS: LElTy, RHS: RElTy);
1488 bool LHSSigned = LElTy->hasSignedIntegerRepresentation();
1489 bool RHSSigned = RElTy->hasSignedIntegerRepresentation();
1490 auto &Ctx = SemaRef.getASTContext();
1491
1492 // If both types have the same signedness, use the higher ranked type.
1493 if (LHSSigned == RHSSigned) {
1494 if (IsCompAssign || IntOrder >= 0)
1495 return castElement<CK_IntegralCast>(S&: SemaRef, E&: RHS, Ty: LHSType);
1496
1497 return castElement<CK_IntegralCast>(S&: SemaRef, E&: LHS, Ty: RHSType);
1498 }
1499
1500 // If the unsigned type has greater than or equal rank of the signed type, use
1501 // the unsigned type.
1502 if (IntOrder != (LHSSigned ? 1 : -1)) {
1503 if (IsCompAssign || RHSSigned)
1504 return castElement<CK_IntegralCast>(S&: SemaRef, E&: RHS, Ty: LHSType);
1505 return castElement<CK_IntegralCast>(S&: SemaRef, E&: LHS, Ty: RHSType);
1506 }
1507
1508 // At this point the signed type has higher rank than the unsigned type, which
1509 // means it will be the same size or bigger. If the signed type is bigger, it
1510 // can represent all the values of the unsigned type, so select it.
1511 if (Ctx.getIntWidth(T: LElTy) != Ctx.getIntWidth(T: RElTy)) {
1512 if (IsCompAssign || LHSSigned)
1513 return castElement<CK_IntegralCast>(S&: SemaRef, E&: RHS, Ty: LHSType);
1514 return castElement<CK_IntegralCast>(S&: SemaRef, E&: LHS, Ty: RHSType);
1515 }
1516
1517 // This is a bit of an odd duck case in HLSL. It shouldn't happen, but can due
1518 // to C/C++ leaking through. The place this happens today is long vs long
1519 // long. When arguments are vector<unsigned long, N> and vector<long long, N>,
1520 // the long long has higher rank than long even though they are the same size.
1521
1522 // If this is a compound assignment cast the right hand side to the left hand
1523 // side's type.
1524 if (IsCompAssign)
1525 return castElement<CK_IntegralCast>(S&: SemaRef, E&: RHS, Ty: LHSType);
1526
1527 // If this isn't a compound assignment we convert to unsigned long long.
1528 QualType ElTy = Ctx.getCorrespondingUnsignedType(T: LHSSigned ? LElTy : RElTy);
1529 QualType NewTy = Ctx.getExtVectorType(
1530 VectorType: ElTy, NumElts: RHSType->castAs<VectorType>()->getNumElements());
1531 (void)castElement<CK_IntegralCast>(S&: SemaRef, E&: RHS, Ty: NewTy);
1532
1533 return castElement<CK_IntegralCast>(S&: SemaRef, E&: LHS, Ty: NewTy);
1534}
1535
1536static CastKind getScalarCastKind(ASTContext &Ctx, QualType DestTy,
1537 QualType SrcTy) {
1538 if (DestTy->isRealFloatingType() && SrcTy->isRealFloatingType())
1539 return CK_FloatingCast;
1540 if (DestTy->isIntegralType(Ctx) && SrcTy->isIntegralType(Ctx))
1541 return CK_IntegralCast;
1542 if (DestTy->isRealFloatingType())
1543 return CK_IntegralToFloating;
1544 assert(SrcTy->isRealFloatingType() && DestTy->isIntegralType(Ctx));
1545 return CK_FloatingToIntegral;
1546}
1547
1548QualType SemaHLSL::handleVectorBinOpConversion(ExprResult &LHS, ExprResult &RHS,
1549 QualType LHSType,
1550 QualType RHSType,
1551 bool IsCompAssign) {
1552 const auto *LVecTy = LHSType->getAs<VectorType>();
1553 const auto *RVecTy = RHSType->getAs<VectorType>();
1554 auto &Ctx = getASTContext();
1555
1556 // If the LHS is not a vector and this is a compound assignment, we truncate
1557 // the argument to a scalar then convert it to the LHS's type.
1558 if (!LVecTy && IsCompAssign) {
1559 QualType RElTy = RHSType->castAs<VectorType>()->getElementType();
1560 RHS = SemaRef.ImpCastExprToType(E: RHS.get(), Type: RElTy, CK: CK_HLSLVectorTruncation);
1561 RHSType = RHS.get()->getType();
1562 if (Ctx.hasSameUnqualifiedType(T1: LHSType, T2: RHSType))
1563 return LHSType;
1564 RHS = SemaRef.ImpCastExprToType(E: RHS.get(), Type: LHSType,
1565 CK: getScalarCastKind(Ctx, DestTy: LHSType, SrcTy: RHSType));
1566 return LHSType;
1567 }
1568
1569 unsigned EndSz = std::numeric_limits<unsigned>::max();
1570 unsigned LSz = 0;
1571 if (LVecTy)
1572 LSz = EndSz = LVecTy->getNumElements();
1573 if (RVecTy)
1574 EndSz = std::min(a: RVecTy->getNumElements(), b: EndSz);
1575 assert(EndSz != std::numeric_limits<unsigned>::max() &&
1576 "one of the above should have had a value");
1577
1578 // In a compound assignment, the left operand does not change type, the right
1579 // operand is converted to the type of the left operand.
1580 if (IsCompAssign && LSz != EndSz) {
1581 Diag(Loc: LHS.get()->getBeginLoc(),
1582 DiagID: diag::err_hlsl_vector_compound_assignment_truncation)
1583 << LHSType << RHSType;
1584 return QualType();
1585 }
1586
1587 if (RVecTy && RVecTy->getNumElements() > EndSz)
1588 castVector<CK_HLSLVectorTruncation>(S&: SemaRef, E&: RHS, Ty&: RHSType, Sz: EndSz);
1589 if (!IsCompAssign && LVecTy && LVecTy->getNumElements() > EndSz)
1590 castVector<CK_HLSLVectorTruncation>(S&: SemaRef, E&: LHS, Ty&: LHSType, Sz: EndSz);
1591
1592 if (!RVecTy)
1593 castVector<CK_VectorSplat>(S&: SemaRef, E&: RHS, Ty&: RHSType, Sz: EndSz);
1594 if (!IsCompAssign && !LVecTy)
1595 castVector<CK_VectorSplat>(S&: SemaRef, E&: LHS, Ty&: LHSType, Sz: EndSz);
1596
1597 // If we're at the same type after resizing we can stop here.
1598 if (Ctx.hasSameUnqualifiedType(T1: LHSType, T2: RHSType))
1599 return Ctx.getCommonSugaredType(X: LHSType, Y: RHSType);
1600
1601 QualType LElTy = LHSType->castAs<VectorType>()->getElementType();
1602 QualType RElTy = RHSType->castAs<VectorType>()->getElementType();
1603
1604 // Handle conversion for floating point vectors.
1605 if (LElTy->isRealFloatingType() || RElTy->isRealFloatingType())
1606 return handleFloatVectorBinOpConversion(SemaRef, LHS, RHS, LHSType, RHSType,
1607 LElTy, RElTy, IsCompAssign);
1608
1609 assert(LElTy->isIntegralType(Ctx) && RElTy->isIntegralType(Ctx) &&
1610 "HLSL Vectors can only contain integer or floating point types");
1611 return handleIntegerVectorBinOpConversion(SemaRef, LHS, RHS, LHSType, RHSType,
1612 LElTy, RElTy, IsCompAssign);
1613}
1614
1615void SemaHLSL::emitLogicalOperatorFixIt(Expr *LHS, Expr *RHS,
1616 BinaryOperatorKind Opc) {
1617 assert((Opc == BO_LOr || Opc == BO_LAnd) &&
1618 "Called with non-logical operator");
1619 llvm::SmallVector<char, 256> Buff;
1620 llvm::raw_svector_ostream OS(Buff);
1621 PrintingPolicy PP(SemaRef.getLangOpts());
1622 StringRef NewFnName = Opc == BO_LOr ? "or" : "and";
1623 OS << NewFnName << "(";
1624 LHS->printPretty(OS, Helper: nullptr, Policy: PP);
1625 OS << ", ";
1626 RHS->printPretty(OS, Helper: nullptr, Policy: PP);
1627 OS << ")";
1628 SourceRange FullRange = SourceRange(LHS->getBeginLoc(), RHS->getEndLoc());
1629 SemaRef.Diag(Loc: LHS->getBeginLoc(), DiagID: diag::note_function_suggestion)
1630 << NewFnName << FixItHint::CreateReplacement(RemoveRange: FullRange, Code: OS.str());
1631}
1632
1633std::pair<IdentifierInfo *, bool>
1634SemaHLSL::ActOnStartRootSignatureDecl(StringRef Signature) {
1635 llvm::hash_code Hash = llvm::hash_value(S: Signature);
1636 std::string IdStr = "__hlsl_rootsig_decl_" + std::to_string(val: Hash);
1637 IdentifierInfo *DeclIdent = &(getASTContext().Idents.get(Name: IdStr));
1638
1639 // Check if we have already found a decl of the same name.
1640 LookupResult R(SemaRef, DeclIdent, SourceLocation(),
1641 Sema::LookupOrdinaryName);
1642 bool Found = SemaRef.LookupQualifiedName(R, LookupCtx: SemaRef.CurContext);
1643 return {DeclIdent, Found};
1644}
1645
1646void SemaHLSL::ActOnFinishRootSignatureDecl(
1647 SourceLocation Loc, IdentifierInfo *DeclIdent,
1648 ArrayRef<hlsl::RootSignatureElement> RootElements) {
1649
1650 if (handleRootSignatureElements(Elements: RootElements))
1651 return;
1652
1653 SmallVector<llvm::hlsl::rootsig::RootElement> Elements;
1654 for (auto &RootSigElement : RootElements)
1655 Elements.push_back(Elt: RootSigElement.getElement());
1656
1657 auto *SignatureDecl = HLSLRootSignatureDecl::Create(
1658 C&: SemaRef.getASTContext(), /*DeclContext=*/DC: SemaRef.CurContext, Loc,
1659 ID: DeclIdent, Version: SemaRef.getLangOpts().HLSLRootSigVer, RootElements: Elements);
1660
1661 SignatureDecl->setImplicit();
1662 SemaRef.PushOnScopeChains(D: SignatureDecl, S: SemaRef.getCurScope());
1663}
1664
1665HLSLRootSignatureDecl *
1666SemaHLSL::lookupRootSignatureOverrideDecl(DeclContext *DC) const {
1667 if (RootSigOverrideIdent) {
1668 LookupResult R(SemaRef, RootSigOverrideIdent, SourceLocation(),
1669 Sema::LookupOrdinaryName);
1670 if (SemaRef.LookupQualifiedName(R, LookupCtx: DC))
1671 return dyn_cast<HLSLRootSignatureDecl>(Val: R.getFoundDecl());
1672 }
1673
1674 return nullptr;
1675}
1676
1677namespace {
1678
1679struct PerVisibilityBindingChecker {
1680 SemaHLSL *S;
1681 // We need one builder per `llvm::dxbc::ShaderVisibility` value.
1682 std::array<llvm::hlsl::BindingInfoBuilder, 8> Builders;
1683
1684 struct ElemInfo {
1685 const hlsl::RootSignatureElement *Elem;
1686 llvm::dxbc::ShaderVisibility Vis;
1687 bool Diagnosed;
1688 };
1689 llvm::SmallVector<ElemInfo> ElemInfoMap;
1690
1691 PerVisibilityBindingChecker(SemaHLSL *S) : S(S) {}
1692
1693 void trackBinding(llvm::dxbc::ShaderVisibility Visibility,
1694 llvm::dxil::ResourceClass RC, uint32_t Space,
1695 uint32_t LowerBound, uint32_t UpperBound,
1696 const hlsl::RootSignatureElement *Elem) {
1697 uint32_t BuilderIndex = llvm::to_underlying(E: Visibility);
1698 assert(BuilderIndex < Builders.size() &&
1699 "Not enough builders for visibility type");
1700 Builders[BuilderIndex].trackBinding(RC, Space, LowerBound, UpperBound,
1701 Cookie: static_cast<const void *>(Elem));
1702
1703 static_assert(llvm::to_underlying(E: llvm::dxbc::ShaderVisibility::All) == 0,
1704 "'All' visibility must come first");
1705 if (Visibility == llvm::dxbc::ShaderVisibility::All)
1706 for (size_t I = 1, E = Builders.size(); I < E; ++I)
1707 Builders[I].trackBinding(RC, Space, LowerBound, UpperBound,
1708 Cookie: static_cast<const void *>(Elem));
1709
1710 ElemInfoMap.push_back(Elt: {.Elem: Elem, .Vis: Visibility, .Diagnosed: false});
1711 }
1712
1713 ElemInfo &getInfo(const hlsl::RootSignatureElement *Elem) {
1714 auto It = llvm::lower_bound(
1715 Range&: ElemInfoMap, Value&: Elem,
1716 C: [](const auto &LHS, const auto &RHS) { return LHS.Elem < RHS; });
1717 assert(It->Elem == Elem && "Element not in map");
1718 return *It;
1719 }
1720
1721 bool checkOverlap() {
1722 llvm::sort(C&: ElemInfoMap, Comp: [](const auto &LHS, const auto &RHS) {
1723 return LHS.Elem < RHS.Elem;
1724 });
1725
1726 bool HadOverlap = false;
1727
1728 using llvm::hlsl::BindingInfoBuilder;
1729 auto ReportOverlap = [this,
1730 &HadOverlap](const BindingInfoBuilder &Builder,
1731 const llvm::hlsl::Binding &Reported) {
1732 HadOverlap = true;
1733
1734 const auto *Elem =
1735 static_cast<const hlsl::RootSignatureElement *>(Reported.Cookie);
1736 const llvm::hlsl::Binding &Previous = Builder.findOverlapping(ReportedBinding: Reported);
1737 const auto *PrevElem =
1738 static_cast<const hlsl::RootSignatureElement *>(Previous.Cookie);
1739
1740 ElemInfo &Info = getInfo(Elem);
1741 // We will have already diagnosed this binding if there's overlap in the
1742 // "All" visibility as well as any particular visibility.
1743 if (Info.Diagnosed)
1744 return;
1745 Info.Diagnosed = true;
1746
1747 ElemInfo &PrevInfo = getInfo(Elem: PrevElem);
1748 llvm::dxbc::ShaderVisibility CommonVis =
1749 Info.Vis == llvm::dxbc::ShaderVisibility::All ? PrevInfo.Vis
1750 : Info.Vis;
1751
1752 this->S->Diag(Loc: Elem->getLocation(), DiagID: diag::err_hlsl_resource_range_overlap)
1753 << llvm::to_underlying(E: Reported.RC) << Reported.LowerBound
1754 << Reported.isUnbounded() << Reported.UpperBound
1755 << llvm::to_underlying(E: Previous.RC) << Previous.LowerBound
1756 << Previous.isUnbounded() << Previous.UpperBound << Reported.Space
1757 << CommonVis;
1758
1759 this->S->Diag(Loc: PrevElem->getLocation(),
1760 DiagID: diag::note_hlsl_resource_range_here);
1761 };
1762
1763 for (BindingInfoBuilder &Builder : Builders)
1764 Builder.calculateBindingInfo(ReportOverlap);
1765
1766 return HadOverlap;
1767 }
1768};
1769
1770static CXXMethodDecl *lookupMethod(Sema &S, CXXRecordDecl *RecordDecl,
1771 StringRef Name, SourceLocation Loc) {
1772 DeclarationName DeclName(&S.getASTContext().Idents.get(Name));
1773 LookupResult Result(S, DeclName, Loc, Sema::LookupMemberName);
1774 if (!S.LookupQualifiedName(R&: Result, LookupCtx: static_cast<DeclContext *>(RecordDecl)))
1775 return nullptr;
1776 return cast<CXXMethodDecl>(Val: Result.getFoundDecl());
1777}
1778
1779} // end anonymous namespace
1780
1781bool SemaHLSL::handleRootSignatureElements(
1782 ArrayRef<hlsl::RootSignatureElement> Elements) {
1783 // Define some common error handling functions
1784 bool HadError = false;
1785 auto ReportError = [this, &HadError](SourceLocation Loc, uint32_t LowerBound,
1786 uint32_t UpperBound) {
1787 HadError = true;
1788 this->Diag(Loc, DiagID: diag::err_hlsl_invalid_rootsig_value)
1789 << LowerBound << UpperBound;
1790 };
1791
1792 auto ReportFloatError = [this, &HadError](SourceLocation Loc,
1793 float LowerBound,
1794 float UpperBound) {
1795 HadError = true;
1796 this->Diag(Loc, DiagID: diag::err_hlsl_invalid_rootsig_value)
1797 << llvm::formatv(Fmt: "{0:f}", Vals&: LowerBound).sstr<6>()
1798 << llvm::formatv(Fmt: "{0:f}", Vals&: UpperBound).sstr<6>();
1799 };
1800
1801 auto VerifyRegister = [ReportError](SourceLocation Loc, uint32_t Register) {
1802 if (!llvm::hlsl::rootsig::verifyRegisterValue(RegisterValue: Register))
1803 ReportError(Loc, 0, 0xfffffffe);
1804 };
1805
1806 auto VerifySpace = [ReportError](SourceLocation Loc, uint32_t Space) {
1807 if (!llvm::hlsl::rootsig::verifyRegisterSpace(RegisterSpace: Space))
1808 ReportError(Loc, 0, 0xffffffef);
1809 };
1810
1811 const uint32_t Version =
1812 llvm::to_underlying(E: SemaRef.getLangOpts().HLSLRootSigVer);
1813 const uint32_t VersionEnum = Version - 1;
1814 auto ReportFlagError = [this, &HadError, VersionEnum](SourceLocation Loc) {
1815 HadError = true;
1816 this->Diag(Loc, DiagID: diag::err_hlsl_invalid_rootsig_flag)
1817 << /*version minor*/ VersionEnum;
1818 };
1819
1820 // Iterate through the elements and do basic validations
1821 for (const hlsl::RootSignatureElement &RootSigElem : Elements) {
1822 SourceLocation Loc = RootSigElem.getLocation();
1823 const llvm::hlsl::rootsig::RootElement &Elem = RootSigElem.getElement();
1824 if (const auto *Descriptor =
1825 std::get_if<llvm::hlsl::rootsig::RootDescriptor>(ptr: &Elem)) {
1826 VerifyRegister(Loc, Descriptor->Reg.Number);
1827 VerifySpace(Loc, Descriptor->Space);
1828
1829 if (!llvm::hlsl::rootsig::verifyRootDescriptorFlag(Version,
1830 Flags: Descriptor->Flags))
1831 ReportFlagError(Loc);
1832 } else if (const auto *Constants =
1833 std::get_if<llvm::hlsl::rootsig::RootConstants>(ptr: &Elem)) {
1834 VerifyRegister(Loc, Constants->Reg.Number);
1835 VerifySpace(Loc, Constants->Space);
1836 } else if (const auto *Sampler =
1837 std::get_if<llvm::hlsl::rootsig::StaticSampler>(ptr: &Elem)) {
1838 VerifyRegister(Loc, Sampler->Reg.Number);
1839 VerifySpace(Loc, Sampler->Space);
1840
1841 assert(!std::isnan(Sampler->MaxLOD) && !std::isnan(Sampler->MinLOD) &&
1842 "By construction, parseFloatParam can't produce a NaN from a "
1843 "float_literal token");
1844
1845 if (!llvm::hlsl::rootsig::verifyMaxAnisotropy(MaxAnisotropy: Sampler->MaxAnisotropy))
1846 ReportError(Loc, 0, 16);
1847 if (!llvm::hlsl::rootsig::verifyMipLODBias(MipLODBias: Sampler->MipLODBias))
1848 ReportFloatError(Loc, -16.f, 15.99f);
1849 } else if (const auto *Clause =
1850 std::get_if<llvm::hlsl::rootsig::DescriptorTableClause>(
1851 ptr: &Elem)) {
1852 VerifyRegister(Loc, Clause->Reg.Number);
1853 VerifySpace(Loc, Clause->Space);
1854
1855 if (!llvm::hlsl::rootsig::verifyNumDescriptors(NumDescriptors: Clause->NumDescriptors)) {
1856 // NumDescriptor could techincally be ~0u but that is reserved for
1857 // unbounded, so the diagnostic will not report that as a valid int
1858 // value
1859 ReportError(Loc, 1, 0xfffffffe);
1860 }
1861
1862 if (!llvm::hlsl::rootsig::verifyDescriptorRangeFlag(Version, Type: Clause->Type,
1863 Flags: Clause->Flags))
1864 ReportFlagError(Loc);
1865 }
1866 }
1867
1868 PerVisibilityBindingChecker BindingChecker(this);
1869 SmallVector<std::pair<const llvm::hlsl::rootsig::DescriptorTableClause *,
1870 const hlsl::RootSignatureElement *>>
1871 UnboundClauses;
1872
1873 for (const hlsl::RootSignatureElement &RootSigElem : Elements) {
1874 const llvm::hlsl::rootsig::RootElement &Elem = RootSigElem.getElement();
1875 if (const auto *Descriptor =
1876 std::get_if<llvm::hlsl::rootsig::RootDescriptor>(ptr: &Elem)) {
1877 uint32_t LowerBound(Descriptor->Reg.Number);
1878 uint32_t UpperBound(LowerBound); // inclusive range
1879
1880 BindingChecker.trackBinding(
1881 Visibility: Descriptor->Visibility,
1882 RC: static_cast<llvm::dxil::ResourceClass>(Descriptor->Type),
1883 Space: Descriptor->Space, LowerBound, UpperBound, Elem: &RootSigElem);
1884 } else if (const auto *Constants =
1885 std::get_if<llvm::hlsl::rootsig::RootConstants>(ptr: &Elem)) {
1886 uint32_t LowerBound(Constants->Reg.Number);
1887 uint32_t UpperBound(LowerBound); // inclusive range
1888
1889 BindingChecker.trackBinding(
1890 Visibility: Constants->Visibility, RC: llvm::dxil::ResourceClass::CBuffer,
1891 Space: Constants->Space, LowerBound, UpperBound, Elem: &RootSigElem);
1892 } else if (const auto *Sampler =
1893 std::get_if<llvm::hlsl::rootsig::StaticSampler>(ptr: &Elem)) {
1894 uint32_t LowerBound(Sampler->Reg.Number);
1895 uint32_t UpperBound(LowerBound); // inclusive range
1896
1897 BindingChecker.trackBinding(
1898 Visibility: Sampler->Visibility, RC: llvm::dxil::ResourceClass::Sampler,
1899 Space: Sampler->Space, LowerBound, UpperBound, Elem: &RootSigElem);
1900 } else if (const auto *Clause =
1901 std::get_if<llvm::hlsl::rootsig::DescriptorTableClause>(
1902 ptr: &Elem)) {
1903 // We'll process these once we see the table element.
1904 UnboundClauses.emplace_back(Args&: Clause, Args: &RootSigElem);
1905 } else if (const auto *Table =
1906 std::get_if<llvm::hlsl::rootsig::DescriptorTable>(ptr: &Elem)) {
1907 assert(UnboundClauses.size() == Table->NumClauses &&
1908 "Number of unbound elements must match the number of clauses");
1909 bool HasAnySampler = false;
1910 bool HasAnyNonSampler = false;
1911 uint64_t Offset = 0;
1912 bool IsPrevUnbound = false;
1913 for (const auto &[Clause, ClauseElem] : UnboundClauses) {
1914 SourceLocation Loc = ClauseElem->getLocation();
1915 if (Clause->Type == llvm::dxil::ResourceClass::Sampler)
1916 HasAnySampler = true;
1917 else
1918 HasAnyNonSampler = true;
1919
1920 if (HasAnySampler && HasAnyNonSampler)
1921 Diag(Loc, DiagID: diag::err_hlsl_invalid_mixed_resources);
1922
1923 // Relevant error will have already been reported above and needs to be
1924 // fixed before we can conduct further analysis, so shortcut error
1925 // return
1926 if (Clause->NumDescriptors == 0)
1927 return true;
1928
1929 bool IsAppending =
1930 Clause->Offset == llvm::hlsl::rootsig::DescriptorTableOffsetAppend;
1931 if (!IsAppending)
1932 Offset = Clause->Offset;
1933
1934 uint64_t RangeBound = llvm::hlsl::rootsig::computeRangeBound(
1935 Offset, Size: Clause->NumDescriptors);
1936
1937 if (IsPrevUnbound && IsAppending)
1938 Diag(Loc, DiagID: diag::err_hlsl_appending_onto_unbound);
1939 else if (!llvm::hlsl::rootsig::verifyNoOverflowedOffset(Offset: RangeBound))
1940 Diag(Loc, DiagID: diag::err_hlsl_offset_overflow) << Offset << RangeBound;
1941
1942 // Update offset to be 1 past this range's bound
1943 Offset = RangeBound + 1;
1944 IsPrevUnbound = Clause->NumDescriptors ==
1945 llvm::hlsl::rootsig::NumDescriptorsUnbounded;
1946
1947 // Compute the register bounds and track resource binding
1948 uint32_t LowerBound(Clause->Reg.Number);
1949 uint32_t UpperBound = llvm::hlsl::rootsig::computeRangeBound(
1950 Offset: LowerBound, Size: Clause->NumDescriptors);
1951
1952 BindingChecker.trackBinding(
1953 Visibility: Table->Visibility,
1954 RC: static_cast<llvm::dxil::ResourceClass>(Clause->Type), Space: Clause->Space,
1955 LowerBound, UpperBound, Elem: ClauseElem);
1956 }
1957 UnboundClauses.clear();
1958 }
1959 }
1960
1961 return BindingChecker.checkOverlap();
1962}
1963
1964void SemaHLSL::handleRootSignatureAttr(Decl *D, const ParsedAttr &AL) {
1965 if (AL.getNumArgs() != 1) {
1966 Diag(Loc: AL.getLoc(), DiagID: diag::err_attribute_wrong_number_arguments) << AL << 1;
1967 return;
1968 }
1969
1970 IdentifierInfo *Ident = AL.getArgAsIdent(Arg: 0)->getIdentifierInfo();
1971 if (auto *RS = D->getAttr<RootSignatureAttr>()) {
1972 if (RS->getSignatureIdent() != Ident) {
1973 Diag(Loc: AL.getLoc(), DiagID: diag::err_disallowed_duplicate_attribute) << RS;
1974 return;
1975 }
1976
1977 Diag(Loc: AL.getLoc(), DiagID: diag::warn_duplicate_attribute_exact) << RS;
1978 return;
1979 }
1980
1981 LookupResult R(SemaRef, Ident, SourceLocation(), Sema::LookupOrdinaryName);
1982 if (SemaRef.LookupQualifiedName(R, LookupCtx: D->getDeclContext()))
1983 if (auto *SignatureDecl =
1984 dyn_cast<HLSLRootSignatureDecl>(Val: R.getFoundDecl())) {
1985 D->addAttr(A: ::new (getASTContext()) RootSignatureAttr(
1986 getASTContext(), AL, Ident, SignatureDecl));
1987 }
1988}
1989
1990void SemaHLSL::handleNumThreadsAttr(Decl *D, const ParsedAttr &AL) {
1991 llvm::VersionTuple SMVersion =
1992 getASTContext().getTargetInfo().getTriple().getOSVersion();
1993 bool IsDXIL = getASTContext().getTargetInfo().getTriple().getArch() ==
1994 llvm::Triple::dxil;
1995
1996 uint32_t ZMax = 1024;
1997 uint32_t ThreadMax = 1024;
1998 if (IsDXIL && SMVersion.getMajor() <= 4) {
1999 ZMax = 1;
2000 ThreadMax = 768;
2001 } else if (IsDXIL && SMVersion.getMajor() == 5) {
2002 ZMax = 64;
2003 ThreadMax = 1024;
2004 }
2005
2006 uint32_t X;
2007 if (!SemaRef.checkUInt32Argument(AI: AL, Expr: AL.getArgAsExpr(Arg: 0), Val&: X))
2008 return;
2009 if (X > 1024) {
2010 Diag(Loc: AL.getArgAsExpr(Arg: 0)->getExprLoc(),
2011 DiagID: diag::err_hlsl_numthreads_argument_oor)
2012 << 0 << 1024;
2013 return;
2014 }
2015 uint32_t Y;
2016 if (!SemaRef.checkUInt32Argument(AI: AL, Expr: AL.getArgAsExpr(Arg: 1), Val&: Y))
2017 return;
2018 if (Y > 1024) {
2019 Diag(Loc: AL.getArgAsExpr(Arg: 1)->getExprLoc(),
2020 DiagID: diag::err_hlsl_numthreads_argument_oor)
2021 << 1 << 1024;
2022 return;
2023 }
2024 uint32_t Z;
2025 if (!SemaRef.checkUInt32Argument(AI: AL, Expr: AL.getArgAsExpr(Arg: 2), Val&: Z))
2026 return;
2027 if (Z > ZMax) {
2028 SemaRef.Diag(Loc: AL.getArgAsExpr(Arg: 2)->getExprLoc(),
2029 DiagID: diag::err_hlsl_numthreads_argument_oor)
2030 << 2 << ZMax;
2031 return;
2032 }
2033
2034 if (X * Y * Z > ThreadMax) {
2035 Diag(Loc: AL.getLoc(), DiagID: diag::err_hlsl_numthreads_invalid) << ThreadMax;
2036 return;
2037 }
2038
2039 HLSLNumThreadsAttr *NewAttr = mergeNumThreadsAttr(D, AL, X, Y, Z);
2040 if (NewAttr)
2041 D->addAttr(A: NewAttr);
2042}
2043
2044static bool isValidWaveSizeValue(unsigned Value) {
2045 return llvm::isPowerOf2_32(Value) && Value >= 4 && Value <= 128;
2046}
2047
2048void SemaHLSL::handleWaveSizeAttr(Decl *D, const ParsedAttr &AL) {
2049 // validate that the wavesize argument is a power of 2 between 4 and 128
2050 // inclusive
2051 unsigned SpelledArgsCount = AL.getNumArgs();
2052 if (SpelledArgsCount == 0 || SpelledArgsCount > 3)
2053 return;
2054
2055 uint32_t Min;
2056 if (!SemaRef.checkUInt32Argument(AI: AL, Expr: AL.getArgAsExpr(Arg: 0), Val&: Min))
2057 return;
2058
2059 uint32_t Max = 0;
2060 if (SpelledArgsCount > 1 &&
2061 !SemaRef.checkUInt32Argument(AI: AL, Expr: AL.getArgAsExpr(Arg: 1), Val&: Max))
2062 return;
2063
2064 uint32_t Preferred = 0;
2065 if (SpelledArgsCount > 2 &&
2066 !SemaRef.checkUInt32Argument(AI: AL, Expr: AL.getArgAsExpr(Arg: 2), Val&: Preferred))
2067 return;
2068
2069 if (SpelledArgsCount > 2) {
2070 if (!isValidWaveSizeValue(Value: Preferred)) {
2071 Diag(Loc: AL.getArgAsExpr(Arg: 2)->getExprLoc(),
2072 DiagID: diag::err_attribute_power_of_two_in_range)
2073 << AL << llvm::dxil::MinWaveSize << llvm::dxil::MaxWaveSize
2074 << Preferred;
2075 return;
2076 }
2077 // Preferred not in range.
2078 if (Preferred < Min || Preferred > Max) {
2079 Diag(Loc: AL.getArgAsExpr(Arg: 2)->getExprLoc(),
2080 DiagID: diag::err_attribute_power_of_two_in_range)
2081 << AL << Min << Max << Preferred;
2082 return;
2083 }
2084 } else if (SpelledArgsCount > 1) {
2085 if (!isValidWaveSizeValue(Value: Max)) {
2086 Diag(Loc: AL.getArgAsExpr(Arg: 1)->getExprLoc(),
2087 DiagID: diag::err_attribute_power_of_two_in_range)
2088 << AL << llvm::dxil::MinWaveSize << llvm::dxil::MaxWaveSize << Max;
2089 return;
2090 }
2091 if (Max < Min) {
2092 Diag(Loc: AL.getLoc(), DiagID: diag::err_attribute_argument_invalid) << AL << 1;
2093 return;
2094 } else if (Max == Min) {
2095 Diag(Loc: AL.getLoc(), DiagID: diag::warn_attr_min_eq_max) << AL;
2096 }
2097 } else {
2098 if (!isValidWaveSizeValue(Value: Min)) {
2099 Diag(Loc: AL.getArgAsExpr(Arg: 0)->getExprLoc(),
2100 DiagID: diag::err_attribute_power_of_two_in_range)
2101 << AL << llvm::dxil::MinWaveSize << llvm::dxil::MaxWaveSize << Min;
2102 return;
2103 }
2104 }
2105
2106 HLSLWaveSizeAttr *NewAttr =
2107 mergeWaveSizeAttr(D, AL, Min, Max, Preferred, SpelledArgsCount);
2108 if (NewAttr)
2109 D->addAttr(A: NewAttr);
2110}
2111
2112void SemaHLSL::handleVkExtBuiltinInputAttr(Decl *D, const ParsedAttr &AL) {
2113 uint32_t ID;
2114 if (!SemaRef.checkUInt32Argument(AI: AL, Expr: AL.getArgAsExpr(Arg: 0), Val&: ID))
2115 return;
2116 D->addAttr(A: ::new (getASTContext())
2117 HLSLVkExtBuiltinInputAttr(getASTContext(), AL, ID));
2118}
2119
2120void SemaHLSL::handleVkExtBuiltinOutputAttr(Decl *D, const ParsedAttr &AL) {
2121 uint32_t ID;
2122 if (!SemaRef.checkUInt32Argument(AI: AL, Expr: AL.getArgAsExpr(Arg: 0), Val&: ID))
2123 return;
2124 D->addAttr(A: ::new (getASTContext())
2125 HLSLVkExtBuiltinOutputAttr(getASTContext(), AL, ID));
2126}
2127
2128void SemaHLSL::handleVkPushConstantAttr(Decl *D, const ParsedAttr &AL) {
2129 D->addAttr(A: ::new (getASTContext())
2130 HLSLVkPushConstantAttr(getASTContext(), AL));
2131}
2132
2133void SemaHLSL::handleVkConstantIdAttr(Decl *D, const ParsedAttr &AL) {
2134 uint32_t Id;
2135 if (!SemaRef.checkUInt32Argument(AI: AL, Expr: AL.getArgAsExpr(Arg: 0), Val&: Id))
2136 return;
2137 HLSLVkConstantIdAttr *NewAttr = mergeVkConstantIdAttr(D, AL, Id);
2138 if (NewAttr)
2139 D->addAttr(A: NewAttr);
2140}
2141
2142void SemaHLSL::handleVkBindingAttr(Decl *D, const ParsedAttr &AL) {
2143 uint32_t Binding = 0;
2144 if (!SemaRef.checkUInt32Argument(AI: AL, Expr: AL.getArgAsExpr(Arg: 0), Val&: Binding))
2145 return;
2146 uint32_t Set = 0;
2147 if (AL.getNumArgs() > 1 &&
2148 !SemaRef.checkUInt32Argument(AI: AL, Expr: AL.getArgAsExpr(Arg: 1), Val&: Set))
2149 return;
2150
2151 D->addAttr(A: ::new (getASTContext())
2152 HLSLVkBindingAttr(getASTContext(), AL, Binding, Set));
2153}
2154
2155void SemaHLSL::handleVkLocationAttr(Decl *D, const ParsedAttr &AL) {
2156 uint32_t Location;
2157 if (!SemaRef.checkUInt32Argument(AI: AL, Expr: AL.getArgAsExpr(Arg: 0), Val&: Location))
2158 return;
2159
2160 D->addAttr(A: ::new (getASTContext())
2161 HLSLVkLocationAttr(getASTContext(), AL, Location));
2162}
2163
2164void SemaHLSL::handleSemanticAttr(Decl *D, const ParsedAttr &AL) {
2165 uint32_t IndexValue(0), ExplicitIndex(0);
2166 if (!SemaRef.checkUInt32Argument(AI: AL, Expr: AL.getArgAsExpr(Arg: 0), Val&: IndexValue) ||
2167 !SemaRef.checkUInt32Argument(AI: AL, Expr: AL.getArgAsExpr(Arg: 1), Val&: ExplicitIndex)) {
2168 assert(0 && "HLSLUnparsedSemantic is expected to have 2 int arguments.");
2169 }
2170 assert(IndexValue > 0 ? ExplicitIndex : true);
2171
2172 SemanticKind Kind = llvm::hlsl::getSemanticKind(SemanticName: AL.getAttrName()->getName());
2173 if (Kind == SemanticKind::Invalid) {
2174 Diag(Loc: AL.getLoc(), DiagID: diag::err_hlsl_unknown_semantic) << AL;
2175 return;
2176 }
2177
2178 switch (Kind) {
2179 // FIXME: These semantics do not yet have CodeGen support.
2180 case SemanticKind::RenderTargetArrayIndex:
2181 case SemanticKind::ViewPortArrayIndex:
2182 case SemanticKind::ClipDistance:
2183 case SemanticKind::CullDistance:
2184 case SemanticKind::OutputControlPointID:
2185 case SemanticKind::DomainLocation:
2186 case SemanticKind::PrimitiveID:
2187 case SemanticKind::GSInstanceID:
2188 case SemanticKind::SampleIndex:
2189 case SemanticKind::IsFrontFace:
2190 case SemanticKind::Coverage:
2191 case SemanticKind::InnerCoverage:
2192 case SemanticKind::Depth:
2193 case SemanticKind::DepthLessEqual:
2194 case SemanticKind::DepthGreaterEqual:
2195 case SemanticKind::StencilRef:
2196 case SemanticKind::TessFactor:
2197 case SemanticKind::InsideTessFactor:
2198 case SemanticKind::ViewID:
2199 case SemanticKind::Barycentrics:
2200 case SemanticKind::ShadingRate:
2201 case SemanticKind::CullPrimitive:
2202 Diag(Loc: AL.getLoc(), DiagID: diag::err_hlsl_unknown_semantic) << AL;
2203 return;
2204 default:
2205 break;
2206 }
2207
2208 D->addAttr(A: HLSLParsedSemanticAttr::Create(
2209 Ctx&: getASTContext(), SemanticName: AL.getAttrName()->getName(), SemanticIndex: IndexValue, CommonInfo: AL));
2210}
2211
2212void SemaHLSL::handlePackOffsetAttr(Decl *D, const ParsedAttr &AL) {
2213 if (!isa<VarDecl>(Val: D) || !isa<HLSLBufferDecl>(Val: D->getDeclContext())) {
2214 Diag(Loc: AL.getLoc(), DiagID: diag::err_hlsl_attr_invalid_ast_node)
2215 << AL << "shader constant in a constant buffer";
2216 return;
2217 }
2218
2219 uint32_t SubComponent;
2220 if (!SemaRef.checkUInt32Argument(AI: AL, Expr: AL.getArgAsExpr(Arg: 0), Val&: SubComponent))
2221 return;
2222 uint32_t Component;
2223 if (!SemaRef.checkUInt32Argument(AI: AL, Expr: AL.getArgAsExpr(Arg: 1), Val&: Component))
2224 return;
2225
2226 QualType T = cast<VarDecl>(Val: D)->getType().getCanonicalType();
2227 // Check if T is an array or struct type.
2228 // TODO: mark matrix type as aggregate type.
2229 bool IsAggregateTy = (T->isArrayType() || T->isStructureType());
2230
2231 // Check Component is valid for T.
2232 if (Component) {
2233 unsigned Size = getASTContext().getTypeSize(T);
2234 if (IsAggregateTy) {
2235 Diag(Loc: AL.getLoc(), DiagID: diag::err_hlsl_invalid_register_or_packoffset);
2236 return;
2237 } else {
2238 // Make sure Component + sizeof(T) <= 4.
2239 if ((Component * 32 + Size) > 128) {
2240 Diag(Loc: AL.getLoc(), DiagID: diag::err_hlsl_packoffset_cross_reg_boundary);
2241 return;
2242 }
2243 QualType EltTy = T;
2244 if (const auto *VT = T->getAs<VectorType>())
2245 EltTy = VT->getElementType();
2246 unsigned Align = getASTContext().getTypeAlign(T: EltTy);
2247 if (Align > 32 && Component == 1) {
2248 // NOTE: Component 3 will hit err_hlsl_packoffset_cross_reg_boundary.
2249 // So we only need to check Component 1 here.
2250 Diag(Loc: AL.getLoc(), DiagID: diag::err_hlsl_packoffset_alignment_mismatch)
2251 << Align << EltTy;
2252 return;
2253 }
2254 }
2255 }
2256
2257 D->addAttr(A: ::new (getASTContext()) HLSLPackOffsetAttr(
2258 getASTContext(), AL, SubComponent, Component));
2259}
2260
2261void SemaHLSL::handleShaderAttr(Decl *D, const ParsedAttr &AL) {
2262 StringRef Str;
2263 SourceLocation ArgLoc;
2264 if (!SemaRef.checkStringLiteralArgumentAttr(Attr: AL, ArgNum: 0, Str, ArgLocation: &ArgLoc))
2265 return;
2266
2267 llvm::Triple::EnvironmentType ShaderType;
2268 if (!HLSLShaderAttr::ConvertStrToEnvironmentType(Val: Str, Out&: ShaderType)) {
2269 Diag(Loc: AL.getLoc(), DiagID: diag::warn_attribute_type_not_supported)
2270 << AL << Str << ArgLoc;
2271 return;
2272 }
2273
2274 // FIXME: check function match the shader stage.
2275
2276 HLSLShaderAttr *NewAttr = mergeShaderAttr(D, AL, ShaderType);
2277 if (NewAttr)
2278 D->addAttr(A: NewAttr);
2279}
2280
2281bool clang::CreateHLSLAttributedResourceType(
2282 Sema &S, QualType Wrapped, ArrayRef<const Attr *> AttrList,
2283 QualType &ResType, HLSLAttributedResourceLocInfo *LocInfo,
2284 Expr *SampleCountExpr) {
2285 assert(AttrList.size() && "expected list of resource attributes");
2286
2287 QualType ContainedTy = QualType();
2288 TypeSourceInfo *ContainedTyInfo = nullptr;
2289 SourceLocation LocBegin = AttrList[0]->getRange().getBegin();
2290 SourceLocation LocEnd = AttrList[0]->getRange().getEnd();
2291
2292 HLSLAttributedResourceType::Attributes ResAttrs;
2293
2294 bool HasResourceClass = false;
2295 bool HasResourceDimension = false;
2296 for (const Attr *A : AttrList) {
2297 if (!A)
2298 continue;
2299 LocEnd = A->getRange().getEnd();
2300 switch (A->getKind()) {
2301 case attr::HLSLResourceClass: {
2302 ResourceClass RC = cast<HLSLResourceClassAttr>(Val: A)->getResourceClass();
2303 if (HasResourceClass) {
2304 S.Diag(Loc: A->getLocation(), DiagID: ResAttrs.ResourceClass == RC
2305 ? diag::warn_duplicate_attribute_exact
2306 : diag::warn_duplicate_attribute)
2307 << A;
2308 return false;
2309 }
2310 ResAttrs.ResourceClass = RC;
2311 HasResourceClass = true;
2312 break;
2313 }
2314 case attr::HLSLResourceDimension: {
2315 llvm::dxil::ResourceDimension RD =
2316 cast<HLSLResourceDimensionAttr>(Val: A)->getDimension();
2317 if (HasResourceDimension) {
2318 S.Diag(Loc: A->getLocation(), DiagID: ResAttrs.ResourceDimension == RD
2319 ? diag::warn_duplicate_attribute_exact
2320 : diag::warn_duplicate_attribute)
2321 << A;
2322 return false;
2323 }
2324 ResAttrs.ResourceDimension = RD;
2325 HasResourceDimension = true;
2326 break;
2327 }
2328 case attr::HLSLIsROV:
2329 if (ResAttrs.IsROV) {
2330 S.Diag(Loc: A->getLocation(), DiagID: diag::warn_duplicate_attribute_exact) << A;
2331 return false;
2332 }
2333 ResAttrs.IsROV = true;
2334 break;
2335 case attr::HLSLRawBuffer:
2336 if (ResAttrs.RawBuffer) {
2337 S.Diag(Loc: A->getLocation(), DiagID: diag::warn_duplicate_attribute_exact) << A;
2338 return false;
2339 }
2340 ResAttrs.RawBuffer = true;
2341 break;
2342 case attr::HLSLIsArray:
2343 if (ResAttrs.IsArray) {
2344 S.Diag(Loc: A->getLocation(), DiagID: diag::warn_duplicate_attribute_exact) << A;
2345 return false;
2346 }
2347 ResAttrs.IsArray = true;
2348 break;
2349 case attr::HLSLIsMultiSampled:
2350 if (ResAttrs.SampleCountExpr) {
2351 S.Diag(Loc: A->getLocation(), DiagID: diag::warn_duplicate_attribute_exact) << A;
2352 return false;
2353 }
2354 // A bare [[hlsl::is_ms]] carries no count, so default it to 0, the same
2355 // value Texture2DMS<T> gets from its template parameter.
2356 ResAttrs.SampleCountExpr =
2357 SampleCountExpr
2358 ? SampleCountExpr
2359 : IntegerLiteral::Create(C: S.Context, V: llvm::APInt(32, 0),
2360 type: S.Context.IntTy, l: A->getLocation());
2361 break;
2362 case attr::HLSLIsCounter:
2363 if (ResAttrs.IsCounter) {
2364 S.Diag(Loc: A->getLocation(), DiagID: diag::warn_duplicate_attribute_exact) << A;
2365 return false;
2366 }
2367 ResAttrs.IsCounter = true;
2368 break;
2369 case attr::HLSLContainedType: {
2370 const HLSLContainedTypeAttr *CTAttr = cast<HLSLContainedTypeAttr>(Val: A);
2371 QualType Ty = CTAttr->getType();
2372 if (!ContainedTy.isNull()) {
2373 S.Diag(Loc: A->getLocation(), DiagID: ContainedTy == Ty
2374 ? diag::warn_duplicate_attribute_exact
2375 : diag::warn_duplicate_attribute)
2376 << A;
2377 return false;
2378 }
2379 ContainedTy = Ty;
2380 ContainedTyInfo = CTAttr->getTypeLoc();
2381 break;
2382 }
2383 default:
2384 llvm_unreachable("unhandled resource attribute type");
2385 }
2386 }
2387
2388 if (!HasResourceClass) {
2389 S.Diag(Loc: AttrList.back()->getRange().getEnd(),
2390 DiagID: diag::err_hlsl_missing_resource_class);
2391 return false;
2392 }
2393
2394 ResType = S.getASTContext().getHLSLAttributedResourceType(
2395 Wrapped, Contained: ContainedTy, Attrs: ResAttrs);
2396
2397 if (LocInfo && ContainedTyInfo) {
2398 LocInfo->Range = SourceRange(LocBegin, LocEnd);
2399 LocInfo->ContainedTyInfo = ContainedTyInfo;
2400 }
2401 return true;
2402}
2403
2404// Validates and creates an HLSL attribute that is applied as type attribute on
2405// HLSL resource. The attributes are collected in HLSLResourcesTypeAttrs and at
2406// the end of the declaration they are applied to the declaration type by
2407// wrapping it in HLSLAttributedResourceType.
2408bool SemaHLSL::handleResourceTypeAttr(QualType T, const ParsedAttr &AL) {
2409 // only allow resource type attributes on intangible types
2410 if (!T->isHLSLResourceType()) {
2411 Diag(Loc: AL.getLoc(), DiagID: diag::err_hlsl_attribute_needs_intangible_type)
2412 << AL << getASTContext().HLSLResourceTy;
2413 return false;
2414 }
2415
2416 // validate number of arguments
2417 if (!AL.checkExactlyNumArgs(S&: SemaRef, Num: AL.getMinArgs()))
2418 return false;
2419
2420 Attr *A = nullptr;
2421
2422 AttributeCommonInfo ACI(
2423 AL.getLoc(), AttributeScopeInfo(AL.getScopeName(), AL.getScopeLoc()),
2424 AttributeCommonInfo::NoSemaHandlerAttribute,
2425 {
2426 AttributeCommonInfo::AS_CXX11, 0, false /*IsAlignas*/,
2427 false /*IsRegularKeywordAttribute*/
2428 });
2429
2430 switch (AL.getKind()) {
2431 case ParsedAttr::AT_HLSLResourceClass: {
2432 StringRef Identifier;
2433 SourceLocation ArgLoc;
2434 if (!SemaRef.checkStringLiteralArgumentAttr(Attr: AL, ArgNum: 0, Str&: Identifier, ArgLocation: &ArgLoc))
2435 return false;
2436
2437 // Validate resource class value
2438 ResourceClass RC;
2439 if (!HLSLResourceClassAttr::ConvertStrToResourceClass(Val: Identifier, Out&: RC)) {
2440 Diag(Loc: ArgLoc, DiagID: diag::warn_attribute_type_not_supported)
2441 << "ResourceClass" << Identifier;
2442 return false;
2443 }
2444 A = HLSLResourceClassAttr::Create(Ctx&: getASTContext(), ResourceClass: RC, CommonInfo: ACI);
2445 break;
2446 }
2447
2448 case ParsedAttr::AT_HLSLResourceDimension: {
2449 StringRef Identifier;
2450 SourceLocation ArgLoc;
2451 if (!SemaRef.checkStringLiteralArgumentAttr(Attr: AL, ArgNum: 0, Str&: Identifier, ArgLocation: &ArgLoc))
2452 return false;
2453
2454 // Validate resource dimension value
2455 llvm::dxil::ResourceDimension RD;
2456 if (!HLSLResourceDimensionAttr::ConvertStrToResourceDimension(Val: Identifier,
2457 Out&: RD)) {
2458 Diag(Loc: ArgLoc, DiagID: diag::warn_attribute_type_not_supported)
2459 << "ResourceDimension" << Identifier;
2460 return false;
2461 }
2462 A = HLSLResourceDimensionAttr::Create(Ctx&: getASTContext(), Dimension: RD, CommonInfo: ACI);
2463 break;
2464 }
2465
2466 case ParsedAttr::AT_HLSLIsROV:
2467 A = HLSLIsROVAttr::Create(Ctx&: getASTContext(), CommonInfo: ACI);
2468 break;
2469
2470 case ParsedAttr::AT_HLSLRawBuffer:
2471 A = HLSLRawBufferAttr::Create(Ctx&: getASTContext(), CommonInfo: ACI);
2472 break;
2473
2474 case ParsedAttr::AT_HLSLIsCounter:
2475 A = HLSLIsCounterAttr::Create(Ctx&: getASTContext(), CommonInfo: ACI);
2476 break;
2477
2478 case ParsedAttr::AT_HLSLIsArray:
2479 A = HLSLIsArrayAttr::Create(Ctx&: getASTContext(), CommonInfo: ACI);
2480 break;
2481
2482 case ParsedAttr::AT_HLSLIsMultiSampled:
2483 A = HLSLIsMultiSampledAttr::Create(Ctx&: getASTContext(), CommonInfo: ACI);
2484 break;
2485
2486 case ParsedAttr::AT_HLSLContainedType: {
2487 if (AL.getNumArgs() != 1 && !AL.hasParsedType()) {
2488 Diag(Loc: AL.getLoc(), DiagID: diag::err_attribute_wrong_number_arguments) << AL << 1;
2489 return false;
2490 }
2491
2492 TypeSourceInfo *TSI = nullptr;
2493 QualType QT = SemaRef.GetTypeFromParser(Ty: AL.getTypeArg(), TInfo: &TSI);
2494 assert(TSI && "no type source info for attribute argument");
2495 if (SemaRef.RequireCompleteType(Loc: TSI->getTypeLoc().getBeginLoc(), T: QT,
2496 DiagID: diag::err_incomplete_type))
2497 return false;
2498 A = HLSLContainedTypeAttr::Create(Ctx&: getASTContext(), Type: TSI, CommonInfo: ACI);
2499 break;
2500 }
2501
2502 default:
2503 llvm_unreachable("unhandled HLSL attribute");
2504 }
2505
2506 HLSLResourcesTypeAttrs.emplace_back(Args&: A);
2507 return true;
2508}
2509
2510// Combines all resource type attributes and creates HLSLAttributedResourceType.
2511QualType SemaHLSL::ProcessResourceTypeAttributes(QualType CurrentType) {
2512 if (!HLSLResourcesTypeAttrs.size())
2513 return CurrentType;
2514
2515 QualType QT = CurrentType;
2516 HLSLAttributedResourceLocInfo LocInfo;
2517 if (CreateHLSLAttributedResourceType(S&: SemaRef, Wrapped: CurrentType,
2518 AttrList: HLSLResourcesTypeAttrs, ResType&: QT, LocInfo: &LocInfo)) {
2519 const HLSLAttributedResourceType *RT =
2520 cast<HLSLAttributedResourceType>(Val: QT.getTypePtr());
2521
2522 // Temporarily store TypeLoc information for the new type.
2523 // It will be transferred to HLSLAttributesResourceTypeLoc
2524 // shortly after the type is created by TypeSpecLocFiller which
2525 // will call the TakeLocForHLSLAttribute method below.
2526 LocsForHLSLAttributedResources.insert(KV: std::pair(RT, LocInfo));
2527 }
2528 HLSLResourcesTypeAttrs.clear();
2529 return QT;
2530}
2531
2532// Returns source location for the HLSLAttributedResourceType
2533HLSLAttributedResourceLocInfo
2534SemaHLSL::TakeLocForHLSLAttribute(const HLSLAttributedResourceType *RT) {
2535 HLSLAttributedResourceLocInfo LocInfo = {};
2536 auto I = LocsForHLSLAttributedResources.find(Val: RT);
2537 if (I != LocsForHLSLAttributedResources.end()) {
2538 LocInfo = I->second;
2539 LocsForHLSLAttributedResources.erase(I);
2540 return LocInfo;
2541 }
2542 LocInfo.Range = SourceRange();
2543 return LocInfo;
2544}
2545
2546// Walks though the global variable declaration, collects all resource binding
2547// requirements and adds them to Bindings
2548void SemaHLSL::collectResourceBindingsOnUserRecordDecl(const VarDecl *VD,
2549 const RecordType *RT) {
2550 const RecordDecl *RD = RT->getDecl()->getDefinitionOrSelf();
2551 for (FieldDecl *FD : RD->fields()) {
2552 const Type *Ty = FD->getType()->getUnqualifiedDesugaredType();
2553
2554 // Unwrap arrays
2555 // FIXME: Calculate array size while unwrapping
2556 assert(!Ty->isIncompleteArrayType() &&
2557 "incomplete arrays inside user defined types are not supported");
2558 while (Ty->isConstantArrayType()) {
2559 const ConstantArrayType *CAT = cast<ConstantArrayType>(Val: Ty);
2560 Ty = CAT->getElementType()->getUnqualifiedDesugaredType();
2561 }
2562
2563 if (!Ty->isRecordType())
2564 continue;
2565
2566 if (const HLSLAttributedResourceType *AttrResType =
2567 HLSLAttributedResourceType::findHandleTypeOnResource(RT: Ty)) {
2568 // Add a new DeclBindingInfo to Bindings if it does not already exist
2569 ResourceClass RC = AttrResType->getAttrs().ResourceClass;
2570 DeclBindingInfo *DBI = Bindings.getDeclBindingInfo(VD, ResClass: RC);
2571 if (!DBI)
2572 Bindings.addDeclBindingInfo(VD, ResClass: RC);
2573 } else if (const RecordType *RT = dyn_cast<RecordType>(Val: Ty)) {
2574 // Recursively scan embedded struct or class; it would be nice to do this
2575 // without recursion, but tricky to correctly calculate the size of the
2576 // binding, which is something we are probably going to need to do later
2577 // on. Hopefully nesting of structs in structs too many levels is
2578 // unlikely.
2579 collectResourceBindingsOnUserRecordDecl(VD, RT);
2580 }
2581 }
2582}
2583
2584// Diagnose localized register binding errors for a single binding; does not
2585// diagnose resource binding on user record types, that will be done later
2586// in processResourceBindingOnDecl based on the information collected in
2587// collectResourceBindingsOnVarDecl.
2588// Returns false if the register binding is not valid.
2589static bool DiagnoseLocalRegisterBinding(Sema &S, SourceLocation &ArgLoc,
2590 Decl *D, RegisterType RegType,
2591 bool SpecifiedSpace) {
2592 int RegTypeNum = static_cast<int>(RegType);
2593
2594 // check if the decl type is groupshared
2595 if (D->hasAttr<HLSLGroupSharedAddressSpaceAttr>()) {
2596 S.Diag(Loc: ArgLoc, DiagID: diag::err_hlsl_binding_type_mismatch) << RegTypeNum;
2597 return false;
2598 }
2599
2600 // Cbuffers and Tbuffers are HLSLBufferDecl types
2601 if (HLSLBufferDecl *CBufferOrTBuffer = dyn_cast<HLSLBufferDecl>(Val: D)) {
2602 ResourceClass RC = CBufferOrTBuffer->isCBuffer() ? ResourceClass::CBuffer
2603 : ResourceClass::SRV;
2604 if (RegType == getRegisterType(RC))
2605 return true;
2606
2607 S.Diag(Loc: D->getLocation(), DiagID: diag::err_hlsl_binding_type_mismatch)
2608 << RegTypeNum;
2609 return false;
2610 }
2611
2612 // Samplers, UAVs, and SRVs are VarDecl types
2613 assert(isa<VarDecl>(D) && "D is expected to be VarDecl or HLSLBufferDecl");
2614 VarDecl *VD = cast<VarDecl>(Val: D);
2615
2616 // Resource
2617 if (const HLSLAttributedResourceType *AttrResType =
2618 HLSLAttributedResourceType::findHandleTypeOnResource(
2619 RT: VD->getType().getTypePtr())) {
2620 if (RegType == getRegisterType(ResTy: AttrResType))
2621 return true;
2622
2623 S.Diag(Loc: D->getLocation(), DiagID: diag::err_hlsl_binding_type_mismatch)
2624 << RegTypeNum;
2625 return false;
2626 }
2627
2628 const clang::Type *Ty = VD->getType().getTypePtr();
2629 while (Ty->isArrayType())
2630 Ty = Ty->getArrayElementTypeNoTypeQual();
2631
2632 // Basic types
2633 if (Ty->isArithmeticType() || Ty->isVectorType()) {
2634 bool DeclaredInCOrTBuffer = isa<HLSLBufferDecl>(Val: D->getDeclContext());
2635 if (SpecifiedSpace && !DeclaredInCOrTBuffer)
2636 S.Diag(Loc: ArgLoc, DiagID: diag::err_hlsl_space_on_global_constant);
2637
2638 if (!DeclaredInCOrTBuffer && (Ty->isIntegralType(Ctx: S.getASTContext()) ||
2639 Ty->isFloatingType() || Ty->isVectorType())) {
2640 // Register annotation on default constant buffer declaration ($Globals)
2641 if (RegType == RegisterType::CBuffer)
2642 S.Diag(Loc: ArgLoc, DiagID: diag::warn_hlsl_deprecated_register_type_b);
2643 else if (RegType != RegisterType::C)
2644 S.Diag(Loc: ArgLoc, DiagID: diag::err_hlsl_binding_type_mismatch) << RegTypeNum;
2645 else
2646 return true;
2647 } else {
2648 if (RegType == RegisterType::C)
2649 S.Diag(Loc: ArgLoc, DiagID: diag::warn_hlsl_register_type_c_packoffset);
2650 else
2651 S.Diag(Loc: ArgLoc, DiagID: diag::err_hlsl_binding_type_mismatch) << RegTypeNum;
2652 }
2653 return false;
2654 }
2655 if (Ty->isRecordType())
2656 // RecordTypes will be diagnosed in processResourceBindingOnDecl
2657 // that is called from ActOnVariableDeclarator
2658 return true;
2659
2660 // Anything else is an error
2661 S.Diag(Loc: ArgLoc, DiagID: diag::err_hlsl_binding_type_mismatch) << RegTypeNum;
2662 return false;
2663}
2664
2665static bool ValidateMultipleRegisterAnnotations(Sema &S, Decl *TheDecl,
2666 RegisterType regType) {
2667 // make sure that there are no two register annotations
2668 // applied to the decl with the same register type
2669 bool RegisterTypesDetected[5] = {false};
2670 RegisterTypesDetected[static_cast<int>(regType)] = true;
2671
2672 for (auto it = TheDecl->attr_begin(); it != TheDecl->attr_end(); ++it) {
2673 if (HLSLResourceBindingAttr *attr =
2674 dyn_cast<HLSLResourceBindingAttr>(Val: *it)) {
2675
2676 RegisterType otherRegType = attr->getRegisterType();
2677 if (RegisterTypesDetected[static_cast<int>(otherRegType)]) {
2678 int otherRegTypeNum = static_cast<int>(otherRegType);
2679 S.Diag(Loc: TheDecl->getLocation(),
2680 DiagID: diag::err_hlsl_duplicate_register_annotation)
2681 << otherRegTypeNum;
2682 return false;
2683 }
2684 RegisterTypesDetected[static_cast<int>(otherRegType)] = true;
2685 }
2686 }
2687 return true;
2688}
2689
2690static bool DiagnoseHLSLRegisterAttribute(Sema &S, SourceLocation &ArgLoc,
2691 Decl *D, RegisterType RegType,
2692 bool SpecifiedSpace) {
2693
2694 // exactly one of these two types should be set
2695 assert(((isa<VarDecl>(D) && !isa<HLSLBufferDecl>(D)) ||
2696 (!isa<VarDecl>(D) && isa<HLSLBufferDecl>(D))) &&
2697 "expecting VarDecl or HLSLBufferDecl");
2698
2699 // check if the declaration contains resource matching the register type
2700 if (!DiagnoseLocalRegisterBinding(S, ArgLoc, D, RegType, SpecifiedSpace))
2701 return false;
2702
2703 // next, if multiple register annotations exist, check that none conflict.
2704 return ValidateMultipleRegisterAnnotations(S, TheDecl: D, regType: RegType);
2705}
2706
2707// return false if the slot count exceeds the limit, true otherwise
2708static bool AccumulateHLSLResourceSlots(QualType Ty, uint64_t &StartSlot,
2709 const uint64_t &Limit,
2710 const ResourceClass ResClass,
2711 ASTContext &Ctx,
2712 uint64_t ArrayCount = 1) {
2713 Ty = Ty.getCanonicalType();
2714 const Type *T = Ty.getTypePtr();
2715
2716 // Early exit if already overflowed
2717 if (StartSlot > Limit)
2718 return false;
2719
2720 // Case 1: array type
2721 if (const auto *AT = dyn_cast<ArrayType>(Val: T)) {
2722 uint64_t Count = 1;
2723
2724 if (const auto *CAT = dyn_cast<ConstantArrayType>(Val: AT))
2725 Count = CAT->getSize().getZExtValue();
2726
2727 QualType ElemTy = AT->getElementType();
2728 return AccumulateHLSLResourceSlots(Ty: ElemTy, StartSlot, Limit, ResClass, Ctx,
2729 ArrayCount: ArrayCount * Count);
2730 }
2731
2732 // Case 2: resource leaf
2733 if (auto ResTy = dyn_cast<HLSLAttributedResourceType>(Val: T)) {
2734 // First ensure this resource counts towards the corresponding
2735 // register type limit.
2736 if (ResTy->getAttrs().ResourceClass != ResClass)
2737 return true;
2738
2739 // Validate highest slot used
2740 uint64_t EndSlot = StartSlot + ArrayCount - 1;
2741 if (EndSlot > Limit)
2742 return false;
2743
2744 // Advance SlotCount past the consumed range
2745 StartSlot = EndSlot + 1;
2746 return true;
2747 }
2748
2749 // Case 3: struct / record
2750 if (const auto *RT = dyn_cast<RecordType>(Val: T)) {
2751 const RecordDecl *RD = RT->getDecl();
2752
2753 if (const auto *CXXRD = dyn_cast<CXXRecordDecl>(Val: RD)) {
2754 for (const CXXBaseSpecifier &Base : CXXRD->bases()) {
2755 if (!AccumulateHLSLResourceSlots(Ty: Base.getType(), StartSlot, Limit,
2756 ResClass, Ctx, ArrayCount))
2757 return false;
2758 }
2759 }
2760
2761 for (const FieldDecl *Field : RD->fields()) {
2762 if (!AccumulateHLSLResourceSlots(Ty: Field->getType(), StartSlot, Limit,
2763 ResClass, Ctx, ArrayCount))
2764 return false;
2765 }
2766
2767 return true;
2768 }
2769
2770 // Case 4: everything else
2771 return true;
2772}
2773
2774// return true if there is something invalid, false otherwise
2775static bool ValidateRegisterNumber(uint64_t SlotNum, Decl *TheDecl,
2776 ASTContext &Ctx, RegisterType RegTy) {
2777 const uint64_t Limit = UINT32_MAX;
2778 if (SlotNum > Limit)
2779 return true;
2780
2781 // after verifying the number doesn't exceed uint32max, we don't need
2782 // to look further into c or i register types
2783 if (RegTy == RegisterType::C || RegTy == RegisterType::I)
2784 return false;
2785
2786 if (VarDecl *VD = dyn_cast<VarDecl>(Val: TheDecl)) {
2787 uint64_t BaseSlot = SlotNum;
2788
2789 if (!AccumulateHLSLResourceSlots(Ty: VD->getType(), StartSlot&: SlotNum, Limit,
2790 ResClass: getResourceClass(RT: RegTy), Ctx))
2791 return true;
2792
2793 // After AccumulateHLSLResourceSlots runs, SlotNum is now
2794 // the first free slot; last used was SlotNum - 1
2795 return (BaseSlot > Limit);
2796 }
2797 // handle the cbuffer/tbuffer case
2798 if (isa<HLSLBufferDecl>(Val: TheDecl))
2799 // resources cannot be put within a cbuffer, so no need
2800 // to analyze the structure since the register number
2801 // won't be pushed any higher.
2802 return (SlotNum > Limit);
2803
2804 // we don't expect any other decl type, so fail
2805 llvm_unreachable("unexpected decl type");
2806}
2807
2808void SemaHLSL::handleResourceBindingAttr(Decl *TheDecl, const ParsedAttr &AL) {
2809 if (VarDecl *VD = dyn_cast<VarDecl>(Val: TheDecl)) {
2810 QualType Ty = VD->getType();
2811 if (const auto *IAT = dyn_cast<IncompleteArrayType>(Val&: Ty))
2812 Ty = IAT->getElementType();
2813 if (SemaRef.RequireCompleteType(Loc: TheDecl->getBeginLoc(), T: Ty,
2814 DiagID: diag::err_incomplete_type))
2815 return;
2816 }
2817
2818 StringRef Slot = "";
2819 StringRef Space = "";
2820 SourceLocation SlotLoc, SpaceLoc;
2821
2822 if (!AL.isArgIdent(Arg: 0)) {
2823 Diag(Loc: AL.getLoc(), DiagID: diag::err_attribute_argument_type)
2824 << AL << AANT_ArgumentIdentifier;
2825 return;
2826 }
2827 IdentifierLoc *Loc = AL.getArgAsIdent(Arg: 0);
2828
2829 if (AL.getNumArgs() == 2) {
2830 Slot = Loc->getIdentifierInfo()->getName();
2831 SlotLoc = Loc->getLoc();
2832 if (!AL.isArgIdent(Arg: 1)) {
2833 Diag(Loc: AL.getLoc(), DiagID: diag::err_attribute_argument_type)
2834 << AL << AANT_ArgumentIdentifier;
2835 return;
2836 }
2837 Loc = AL.getArgAsIdent(Arg: 1);
2838 Space = Loc->getIdentifierInfo()->getName();
2839 SpaceLoc = Loc->getLoc();
2840 } else {
2841 StringRef Str = Loc->getIdentifierInfo()->getName();
2842 if (Str.starts_with(Prefix: "space")) {
2843 Space = Str;
2844 SpaceLoc = Loc->getLoc();
2845 } else {
2846 Slot = Str;
2847 SlotLoc = Loc->getLoc();
2848 Space = "space0";
2849 }
2850 }
2851
2852 RegisterType RegType = RegisterType::SRV;
2853 std::optional<unsigned> SlotNum;
2854 unsigned SpaceNum = 0;
2855
2856 // Validate slot
2857 if (!Slot.empty()) {
2858 if (!convertToRegisterType(Slot, RT: &RegType)) {
2859 Diag(Loc: SlotLoc, DiagID: diag::err_hlsl_binding_type_invalid) << Slot.substr(Start: 0, N: 1);
2860 return;
2861 }
2862 if (RegType == RegisterType::I) {
2863 Diag(Loc: SlotLoc, DiagID: diag::warn_hlsl_deprecated_register_type_i);
2864 return;
2865 }
2866 const StringRef SlotNumStr = Slot.substr(Start: 1);
2867
2868 uint64_t N;
2869
2870 // validate that the slot number is a non-empty number
2871 if (SlotNumStr.getAsInteger(Radix: 10, Result&: N)) {
2872 Diag(Loc: SlotLoc, DiagID: diag::err_hlsl_unsupported_register_number);
2873 return;
2874 }
2875
2876 // Validate register number. It should not exceed UINT32_MAX,
2877 // including if the resource type is an array that starts
2878 // before UINT32_MAX, but ends afterwards.
2879 if (ValidateRegisterNumber(SlotNum: N, TheDecl, Ctx&: getASTContext(), RegTy: RegType)) {
2880 Diag(Loc: SlotLoc, DiagID: diag::err_hlsl_register_number_too_large);
2881 return;
2882 }
2883
2884 // the slot number has been validated and does not exceed UINT32_MAX
2885 SlotNum = (unsigned)N;
2886 }
2887
2888 // Validate space
2889 if (!Space.starts_with(Prefix: "space")) {
2890 Diag(Loc: SpaceLoc, DiagID: diag::err_hlsl_expected_space) << Space;
2891 return;
2892 }
2893 StringRef SpaceNumStr = Space.substr(Start: 5);
2894 if (SpaceNumStr.getAsInteger(Radix: 10, Result&: SpaceNum)) {
2895 Diag(Loc: SpaceLoc, DiagID: diag::err_hlsl_expected_space) << Space;
2896 return;
2897 }
2898
2899 // If we have slot, diagnose it is the right register type for the decl
2900 if (SlotNum.has_value())
2901 if (!DiagnoseHLSLRegisterAttribute(S&: SemaRef, ArgLoc&: SlotLoc, D: TheDecl, RegType,
2902 SpecifiedSpace: !SpaceLoc.isInvalid()))
2903 return;
2904
2905 HLSLResourceBindingAttr *NewAttr =
2906 HLSLResourceBindingAttr::Create(Ctx&: getASTContext(), Slot, Space, CommonInfo: AL);
2907 if (NewAttr) {
2908 NewAttr->setBinding(RT: RegType, SlotNum, SpaceNum);
2909 TheDecl->addAttr(A: NewAttr);
2910 }
2911}
2912
2913void SemaHLSL::handleParamModifierAttr(Decl *D, const ParsedAttr &AL) {
2914 HLSLParamModifierAttr *NewAttr = mergeParamModifierAttr(
2915 D, AL,
2916 Spelling: static_cast<HLSLParamModifierAttr::Spelling>(AL.getSemanticSpelling()));
2917 if (NewAttr)
2918 D->addAttr(A: NewAttr);
2919}
2920
2921static bool isMatrixType(QualType QT) {
2922 const Type *Ty = QT->getUnqualifiedDesugaredType();
2923 return Ty->isDependentType() || Ty->isConstantMatrixType();
2924}
2925
2926/// Walks the existing AttributedType sugar of \p T looking for a previously
2927/// applied HLSLRowMajor/HLSLColumnMajor marker. If one is found, populates
2928/// \p ExistingKind with its attr::Kind and returns true.
2929static bool findExistingMatrixLayoutMarker(QualType T,
2930 attr::Kind &ExistingKind) {
2931 QualType Cur = T;
2932 while (const auto *AT = Cur->getAs<AttributedType>()) {
2933 attr::Kind K = AT->getAttrKind();
2934 if (K == attr::HLSLRowMajor || K == attr::HLSLColumnMajor) {
2935 ExistingKind = K;
2936 return true;
2937 }
2938 Cur = AT->getModifiedType();
2939 }
2940 return false;
2941}
2942
2943Attr *SemaHLSL::buildMatrixLayoutTypeAttr(QualType T, const ParsedAttr &AL) {
2944 if (T.isNull())
2945 return nullptr;
2946
2947 ASTContext &Ctx = getASTContext();
2948 attr::Kind AttrK = AL.getKind() == ParsedAttr::AT_HLSLRowMajor
2949 ? attr::HLSLRowMajor
2950 : attr::HLSLColumnMajor;
2951
2952 // For non-dependent types, the operand must be a matrix.
2953 if (!T->isDependentType() && !isMatrixType(QT: T)) {
2954 Diag(Loc: AL.getLoc(), DiagID: diag::err_hlsl_matrix_layout_non_matrix)
2955 << AL.getAttrName();
2956 AL.setInvalid();
2957 return nullptr;
2958 }
2959
2960 // Conflict / duplicate detection by walking existing sugar.
2961 attr::Kind ExistingKind;
2962 if (findExistingMatrixLayoutMarker(T, ExistingKind)) {
2963 if (ExistingKind == AttrK) {
2964 Diag(Loc: AL.getLoc(), DiagID: diag::warn_duplicate_attribute_exact)
2965 << AL.getAttrName();
2966 Diag(Loc: AL.getLoc(), DiagID: diag::note_previous_attribute);
2967 return nullptr;
2968 }
2969 IdentifierInfo *ExistingII = &Ctx.Idents.get(
2970 Name: ExistingKind == attr::HLSLRowMajor ? "row_major" : "column_major");
2971 Diag(Loc: AL.getLoc(), DiagID: diag::err_hlsl_matrix_layout_conflict)
2972 << AL.getAttrName() << ExistingII;
2973 Diag(Loc: AL.getLoc(), DiagID: diag::note_conflicting_attribute);
2974 AL.setInvalid();
2975 return nullptr;
2976 }
2977
2978 if (AttrK == attr::HLSLRowMajor)
2979 return ::new (Ctx) HLSLRowMajorAttr(Ctx, AL);
2980 return ::new (Ctx) HLSLColumnMajorAttr(Ctx, AL);
2981}
2982
2983// Re-validates an HLSL `row_major` / `column_major` attribute after template
2984// substitution. The parse-time check in `buildMatrixLayoutTypeAttr` is skipped
2985// for dependent types; `TransformAttributedType` calls this once the type is
2986// concrete. Returns `true` (and emits a diagnostic) if the substituted type is
2987// not a matrix or array of matrices, signaling the caller to abort the
2988// transform.
2989bool SemaHLSL::diagnoseMatrixLayoutInstantiation(attr::Kind K, QualType T,
2990 SourceLocation Loc) {
2991 if (K != attr::HLSLRowMajor && K != attr::HLSLColumnMajor)
2992 return false;
2993 if (T.isNull() || T->isDependentType())
2994 return false;
2995 if (isMatrixType(QT: T))
2996 return false;
2997 IdentifierInfo *II = &getASTContext().Idents.get(
2998 Name: K == attr::HLSLRowMajor ? "row_major" : "column_major");
2999 Diag(Loc, DiagID: diag::err_hlsl_matrix_layout_non_matrix) << II;
3000 return true;
3001}
3002
3003// Transpose and matrix mul need to read the destination layout.
3004// Elementwise builtins reuse the operand layout instead.
3005namespace {
3006
3007using llvm::dxil::BarrierMemoryTypeFlag;
3008using llvm::dxil::BarrierSemanticFlag;
3009
3010template <typename T> constexpr uint64_t barrierFlagValue(T Flag) {
3011 return llvm::to_underlying(Flag);
3012}
3013
3014/// This class implements reachable HLSL diagnostics.
3015///
3016/// It diagnoses unavailable APIs in default and relaxed availability modes.
3017/// It also validates Barrier calls in all availability modes.
3018///
3019/// This is done by traversing the AST of all shader entry point functions
3020/// and of all exported functions, and any functions that are referenced
3021/// from this AST. In other words, any functions that are reachable from
3022/// the entry points.
3023class DiagnoseHLSLAvailability : public DynamicRecursiveASTVisitor {
3024 Sema &SemaRef;
3025 bool DiagnoseAvailability;
3026
3027 // Stack of functions to be scaned
3028 llvm::SmallVector<const FunctionDecl *, 8> DeclsToScan;
3029
3030 // Tracks which environments functions have been scanned in.
3031 //
3032 // Maps FunctionDecl to an unsigned number that represents the set of shader
3033 // environments the function has been scanned for.
3034 // The llvm::Triple::EnvironmentType enum values for shader stages guaranteed
3035 // to be numbered from llvm::Triple::Pixel to llvm::Triple::Amplification
3036 // (verified by static_asserts in Triple.cpp), we can use it to index
3037 // individual bits in the set, as long as we shift the values to start with 0
3038 // by subtracting the value of llvm::Triple::Pixel first.
3039 //
3040 // The N'th bit in the set will be set if the function has been scanned
3041 // in shader environment whose llvm::Triple::EnvironmentType integer value
3042 // equals (llvm::Triple::Pixel + N).
3043 //
3044 // For example, if a function has been scanned in compute and pixel stage
3045 // environment, the value will be 0x21 (100001 binary) because:
3046 //
3047 // (int)(llvm::Triple::Pixel - llvm::Triple::Pixel) == 0
3048 // (int)(llvm::Triple::Compute - llvm::Triple::Pixel) == 5
3049 //
3050 // A FunctionDecl is mapped to 0 (or not included in the map) if it has not
3051 // been scanned in any environment.
3052 llvm::DenseMap<const FunctionDecl *, unsigned> ScannedDecls;
3053
3054 // Do not access these directly, use the get/set methods below to make
3055 // sure the values are in sync
3056 llvm::Triple::EnvironmentType CurrentShaderEnvironment;
3057 unsigned CurrentShaderStageBit;
3058
3059 // True if scanning a function that was already scanned in a different
3060 // shader stage context. Suppress stage-independent diagnostics because
3061 // they were reported during the first scan.
3062 bool ReportOnlyShaderStageIssues;
3063
3064 // Helper methods for dealing with current stage context / environment
3065 void SetShaderStageContext(llvm::Triple::EnvironmentType ShaderType) {
3066 static_assert(sizeof(unsigned) >= 4);
3067 assert(HLSLShaderAttr::isValidShaderType(ShaderType));
3068 assert((unsigned)(ShaderType - llvm::Triple::Pixel) < 31 &&
3069 "ShaderType is too big for this bitmap"); // 31 is reserved for
3070 // "unknown"
3071
3072 unsigned bitmapIndex = ShaderType - llvm::Triple::Pixel;
3073 CurrentShaderEnvironment = ShaderType;
3074 CurrentShaderStageBit = (1 << bitmapIndex);
3075 }
3076
3077 void SetUnknownShaderStageContext() {
3078 CurrentShaderEnvironment = llvm::Triple::UnknownEnvironment;
3079 CurrentShaderStageBit = (1 << 31);
3080 }
3081
3082 llvm::Triple::EnvironmentType GetCurrentShaderEnvironment() const {
3083 return CurrentShaderEnvironment;
3084 }
3085
3086 bool InUnknownShaderStageContext() const {
3087 return CurrentShaderEnvironment == llvm::Triple::UnknownEnvironment;
3088 }
3089
3090 // Helper methods for dealing with shader stage bitmap
3091 void AddToScannedFunctions(const FunctionDecl *FD) {
3092 unsigned &ScannedStages = ScannedDecls[FD];
3093 ScannedStages |= CurrentShaderStageBit;
3094 }
3095
3096 unsigned GetScannedStages(const FunctionDecl *FD) { return ScannedDecls[FD]; }
3097
3098 bool WasAlreadyScannedInCurrentStage(const FunctionDecl *FD) {
3099 return WasAlreadyScannedInCurrentStage(ScannerStages: GetScannedStages(FD));
3100 }
3101
3102 bool WasAlreadyScannedInCurrentStage(unsigned ScannerStages) {
3103 return ScannerStages & CurrentShaderStageBit;
3104 }
3105
3106 static bool NeverBeenScanned(unsigned ScannedStages) {
3107 return ScannedStages == 0;
3108 }
3109
3110 // Scanning methods
3111 void HandleFunctionOrMethodRef(FunctionDecl *FD, Expr *RefExpr);
3112 void CheckDeclAvailability(NamedDecl *D, const AvailabilityAttr *AA,
3113 SourceRange Range);
3114 const AvailabilityAttr *FindAvailabilityAttr(const Decl *D);
3115 bool HasMatchingEnvironmentOrNone(const AvailabilityAttr *AA);
3116 void DiagnoseBarrierCall(CallExpr *CE);
3117 uint64_t DiagnoseBarrierGroupMemory(Expr *MemoryArg, uint64_t MemoryFlags,
3118 bool HasVisibleGroup, bool IsAllMemory);
3119 uint64_t DiagnoseBarrierNodeMemory(Expr *MemoryArg, uint64_t MemoryFlags,
3120 bool HasKnownStage, bool IsAllMemory);
3121 void DiagnoseBarrierGroupSemantic(Expr *SemanticArg, uint64_t SemanticFlags,
3122 bool HasVisibleGroup);
3123 void DiagnoseBarrierScope(Expr *SemanticArg, uint64_t MemoryFlags,
3124 uint64_t SemanticFlags);
3125
3126public:
3127 DiagnoseHLSLAvailability(Sema &SemaRef, bool DiagnoseAvailability)
3128 : SemaRef(SemaRef), DiagnoseAvailability(DiagnoseAvailability),
3129 CurrentShaderEnvironment(llvm::Triple::UnknownEnvironment),
3130 CurrentShaderStageBit(0), ReportOnlyShaderStageIssues(false) {}
3131
3132 // AST traversal methods
3133 void RunOnTranslationUnit(const TranslationUnitDecl *TU);
3134 void RunOnFunction(const FunctionDecl *FD);
3135
3136 bool VisitDeclRefExpr(DeclRefExpr *DRE) override {
3137 FunctionDecl *FD = llvm::dyn_cast<FunctionDecl>(Val: DRE->getDecl());
3138 if (FD)
3139 HandleFunctionOrMethodRef(FD, RefExpr: DRE);
3140 return true;
3141 }
3142
3143 bool VisitMemberExpr(MemberExpr *ME) override {
3144 FunctionDecl *FD = llvm::dyn_cast<FunctionDecl>(Val: ME->getMemberDecl());
3145 if (FD)
3146 HandleFunctionOrMethodRef(FD, RefExpr: ME);
3147 return true;
3148 }
3149
3150 bool VisitCallExpr(CallExpr *CE) override {
3151 DiagnoseBarrierCall(CE);
3152 return true;
3153 }
3154};
3155
3156uint64_t DiagnoseHLSLAvailability::DiagnoseBarrierGroupMemory(
3157 Expr *MemoryArg, uint64_t MemoryFlags, bool HasVisibleGroup,
3158 bool IsAllMemory) {
3159 const uint64_t GroupSharedMemory =
3160 barrierFlagValue(Flag: BarrierMemoryTypeFlag::GroupSharedMemory);
3161 if (HasVisibleGroup || (MemoryFlags & GroupSharedMemory) == 0)
3162 return MemoryFlags;
3163
3164 if (!IsAllMemory) {
3165 SemaRef.Diag(Loc: MemoryArg->getExprLoc(),
3166 DiagID: diag::err_hlsl_barrier_flag_requires_group)
3167 << 0;
3168 return MemoryFlags;
3169 }
3170
3171 return MemoryFlags & ~GroupSharedMemory;
3172}
3173
3174uint64_t DiagnoseHLSLAvailability::DiagnoseBarrierNodeMemory(
3175 Expr *MemoryArg, uint64_t MemoryFlags, bool HasKnownStage,
3176 bool IsAllMemory) {
3177 const uint64_t NodeMemory =
3178 barrierFlagValue(Flag: BarrierMemoryTypeFlag::NodeMemory);
3179 if (!HasKnownStage || (MemoryFlags & NodeMemory) == 0)
3180 return MemoryFlags;
3181
3182 if (!IsAllMemory) {
3183 SemaRef.Diag(Loc: MemoryArg->getExprLoc(),
3184 DiagID: diag::err_hlsl_barrier_node_memory_requires_node);
3185 return MemoryFlags;
3186 }
3187
3188 return MemoryFlags & ~NodeMemory;
3189}
3190
3191void DiagnoseHLSLAvailability::DiagnoseBarrierGroupSemantic(
3192 Expr *SemanticArg, uint64_t SemanticFlags, bool HasVisibleGroup) {
3193 if (HasVisibleGroup ||
3194 (SemanticFlags & barrierFlagValue(Flag: BarrierSemanticFlag::GroupFlags)) == 0)
3195 return;
3196
3197 SemaRef.Diag(Loc: SemanticArg->getExprLoc(),
3198 DiagID: diag::err_hlsl_barrier_flag_requires_group)
3199 << ((SemanticFlags & barrierFlagValue(Flag: BarrierSemanticFlag::GroupSync)) !=
3200 0
3201 ? 1
3202 : 2);
3203}
3204
3205void DiagnoseHLSLAvailability::DiagnoseBarrierScope(Expr *SemanticArg,
3206 uint64_t MemoryFlags,
3207 uint64_t SemanticFlags) {
3208 if (ReportOnlyShaderStageIssues)
3209 return;
3210
3211 const uint64_t DeviceScopeMemory =
3212 barrierFlagValue(Flag: BarrierMemoryTypeFlag::UAVMemory) |
3213 barrierFlagValue(Flag: BarrierMemoryTypeFlag::NodeInputMemory);
3214 if ((SemanticFlags & barrierFlagValue(Flag: BarrierSemanticFlag::DeviceScope)) !=
3215 0 &&
3216 (MemoryFlags & DeviceScopeMemory) == 0)
3217 SemaRef.Diag(Loc: SemanticArg->getExprLoc(),
3218 DiagID: diag::err_hlsl_barrier_scope_requires_memory)
3219 << 1;
3220 if ((SemanticFlags & barrierFlagValue(Flag: BarrierSemanticFlag::GroupScope)) !=
3221 0 &&
3222 MemoryFlags == 0)
3223 SemaRef.Diag(Loc: SemanticArg->getExprLoc(),
3224 DiagID: diag::err_hlsl_barrier_scope_requires_memory)
3225 << 0;
3226}
3227
3228void DiagnoseHLSLAvailability::DiagnoseBarrierCall(CallExpr *CE) {
3229 const FunctionDecl *FD = CE->getDirectCallee();
3230 if (!FD || FD->getBuiltinID() != Builtin::BI__builtin_hlsl_barrier)
3231 return;
3232
3233 const llvm::Triple::EnvironmentType Stage = GetCurrentShaderEnvironment();
3234 const bool HasKnownStage = !InUnknownShaderStageContext();
3235 const bool HasVisibleGroup =
3236 !HasKnownStage || Stage == llvm::Triple::Compute ||
3237 Stage == llvm::Triple::Mesh || Stage == llvm::Triple::Amplification;
3238
3239 uint64_t MemoryFlags = barrierFlagValue(Flag: BarrierMemoryTypeFlag::ValidMask);
3240 Expr *MemoryArg = CE->getArg(Arg: 0);
3241 if (MemoryArg->getType()->isUnsignedIntegerType()) {
3242 std::optional<llvm::APSInt> Value =
3243 MemoryArg->getIntegerConstantExpr(Ctx: SemaRef.Context);
3244 if (!Value)
3245 return;
3246 MemoryFlags = Value->getZExtValue();
3247 const bool IsAllMemory =
3248 MemoryFlags == barrierFlagValue(Flag: BarrierMemoryTypeFlag::ValidMask);
3249
3250 MemoryFlags = DiagnoseBarrierGroupMemory(MemoryArg, MemoryFlags,
3251 HasVisibleGroup, IsAllMemory);
3252 MemoryFlags = DiagnoseBarrierNodeMemory(MemoryArg, MemoryFlags,
3253 HasKnownStage, IsAllMemory);
3254 } else if (!HasVisibleGroup) {
3255 SemaRef.Diag(Loc: MemoryArg->getExprLoc(),
3256 DiagID: diag::err_hlsl_barrier_resource_requires_group);
3257 return;
3258 }
3259
3260 Expr *SemanticArg = CE->getArg(Arg: 1);
3261 std::optional<llvm::APSInt> Value =
3262 SemanticArg->getIntegerConstantExpr(Ctx: SemaRef.Context);
3263 if (!Value)
3264 return;
3265 const uint64_t SemanticFlags = Value->getZExtValue();
3266
3267 DiagnoseBarrierGroupSemantic(SemanticArg, SemanticFlags, HasVisibleGroup);
3268
3269 if (MemoryArg->getType()->isUnsignedIntegerType())
3270 DiagnoseBarrierScope(SemanticArg, MemoryFlags, SemanticFlags);
3271}
3272
3273void DiagnoseHLSLAvailability::HandleFunctionOrMethodRef(FunctionDecl *FD,
3274 Expr *RefExpr) {
3275 assert((isa<DeclRefExpr>(RefExpr) || isa<MemberExpr>(RefExpr)) &&
3276 "expected DeclRefExpr or MemberExpr");
3277
3278 if (DiagnoseAvailability)
3279 if (const AvailabilityAttr *AA = FindAvailabilityAttr(D: FD))
3280 CheckDeclAvailability(
3281 D: FD, AA, Range: SourceRange(RefExpr->getBeginLoc(), RefExpr->getEndLoc()));
3282
3283 // has a definition -> add to stack to be scanned
3284 const FunctionDecl *FDWithBody = nullptr;
3285 if (FD->hasBody(Definition&: FDWithBody) && !WasAlreadyScannedInCurrentStage(FD: FDWithBody))
3286 DeclsToScan.push_back(Elt: FDWithBody);
3287}
3288
3289void DiagnoseHLSLAvailability::RunOnTranslationUnit(
3290 const TranslationUnitDecl *TU) {
3291 const TargetInfo &TargetInfo = SemaRef.getASTContext().getTargetInfo();
3292 std::string &EntryName = TargetInfo.getTargetOpts().HLSLEntry;
3293 bool IsLibraryShader = TargetInfo.getTriple().getEnvironment() ==
3294 llvm::Triple::EnvironmentType::Library;
3295 SourceLocation EntryLoc{};
3296
3297 // Iterate over all shader entry functions and library exports, and for those
3298 // that have a body (definiton), run diag scan on each, setting appropriate
3299 // shader environment context based on whether it is a shader entry function
3300 // or an exported function. Exported functions can be in namespaces and in
3301 // export declarations so we need to scan those declaration contexts as well.
3302 llvm::SmallVector<const DeclContext *, 8> DeclContextsToScan;
3303 DeclContextsToScan.push_back(Elt: TU);
3304
3305 while (!DeclContextsToScan.empty()) {
3306 const DeclContext *DC = DeclContextsToScan.pop_back_val();
3307 for (auto &D : DC->decls()) {
3308 // do not scan implicit declaration generated by the implementation
3309 if (D->isImplicit())
3310 continue;
3311
3312 // for namespace or export declaration add the context to the list to be
3313 // scanned later
3314 if (llvm::dyn_cast<NamespaceDecl>(Val: D) || llvm::dyn_cast<ExportDecl>(Val: D)) {
3315 DeclContextsToScan.push_back(Elt: llvm::dyn_cast<DeclContext>(Val: D));
3316 continue;
3317 }
3318
3319 // skip over other decls or function decls without body
3320 const FunctionDecl *FD = llvm::dyn_cast<FunctionDecl>(Val: D);
3321 if (!FD || !FD->isThisDeclarationADefinition())
3322 continue;
3323
3324 // shader entry point
3325 if (HLSLShaderAttr *ShaderAttr = FD->getAttr<HLSLShaderAttr>()) {
3326 if (!IsLibraryShader && FD->getName() == EntryName) {
3327 if (EntryLoc.isValid()) {
3328 SemaRef.Diag(Loc: FD->getLocation(),
3329 DiagID: diag::err_hlsl_ambiguous_entry_point)
3330 << EntryName;
3331 SemaRef.Diag(Loc: EntryLoc, DiagID: diag::note_previous_declaration_as)
3332 << EntryName;
3333 return;
3334 }
3335 EntryLoc = FD->getLocation();
3336 }
3337 SetShaderStageContext(ShaderAttr->getType());
3338 RunOnFunction(FD);
3339 continue;
3340 }
3341 // exported library function
3342 // FIXME: replace this loop with external linkage check once issue #92071
3343 // is resolved
3344 bool isExport = FD->isInExportDeclContext();
3345 if (!isExport) {
3346 for (const auto *Redecl : FD->redecls()) {
3347 if (Redecl->isInExportDeclContext()) {
3348 isExport = true;
3349 break;
3350 }
3351 }
3352 }
3353 if (isExport) {
3354 SetUnknownShaderStageContext();
3355 RunOnFunction(FD);
3356 continue;
3357 }
3358 }
3359 }
3360
3361 if (!IsLibraryShader && EntryLoc.isInvalid()) {
3362 SemaRef.Diag(Loc: TU->getLocation(), DiagID: diag::err_hlsl_missing_entry_point)
3363 << EntryName;
3364 return;
3365 }
3366}
3367
3368void DiagnoseHLSLAvailability::RunOnFunction(const FunctionDecl *FD) {
3369 assert(DeclsToScan.empty() && "DeclsToScan should be empty");
3370 DeclsToScan.push_back(Elt: FD);
3371
3372 while (!DeclsToScan.empty()) {
3373 // Take one decl from the stack and check it by traversing its AST.
3374 // For any CallExpr found during the traversal add it's callee to the top of
3375 // the stack to be processed next. Functions already processed are stored in
3376 // ScannedDecls.
3377 const FunctionDecl *FD = DeclsToScan.pop_back_val();
3378
3379 // Decl was already scanned
3380 const unsigned ScannedStages = GetScannedStages(FD);
3381 if (WasAlreadyScannedInCurrentStage(ScannerStages: ScannedStages))
3382 continue;
3383
3384 ReportOnlyShaderStageIssues = !NeverBeenScanned(ScannedStages);
3385
3386 AddToScannedFunctions(FD);
3387 TraverseStmt(S: FD->getBody());
3388 }
3389}
3390
3391bool DiagnoseHLSLAvailability::HasMatchingEnvironmentOrNone(
3392 const AvailabilityAttr *AA) {
3393 const IdentifierInfo *IIEnvironment = AA->getEnvironment();
3394 if (!IIEnvironment)
3395 return true;
3396
3397 llvm::Triple::EnvironmentType CurrentEnv = GetCurrentShaderEnvironment();
3398 if (CurrentEnv == llvm::Triple::UnknownEnvironment)
3399 return false;
3400
3401 llvm::Triple::EnvironmentType AttrEnv =
3402 AvailabilityAttr::getEnvironmentType(Environment: IIEnvironment->getName());
3403
3404 return CurrentEnv == AttrEnv;
3405}
3406
3407const AvailabilityAttr *
3408DiagnoseHLSLAvailability::FindAvailabilityAttr(const Decl *D) {
3409 AvailabilityAttr const *PartialMatch = nullptr;
3410 // Check each AvailabilityAttr to find the one for this platform.
3411 // For multiple attributes with the same platform try to find one for this
3412 // environment.
3413 for (const auto *A : D->attrs()) {
3414 if (const auto *Avail = dyn_cast<AvailabilityAttr>(Val: A)) {
3415 const AvailabilityAttr *EffectiveAvail = Avail->getEffectiveAttr();
3416 StringRef AttrPlatform = EffectiveAvail->getPlatform()->getName();
3417 StringRef TargetPlatform =
3418 SemaRef.getASTContext().getTargetInfo().getPlatformName();
3419
3420 // Match the platform name.
3421 if (AttrPlatform == TargetPlatform) {
3422 // Find the best matching attribute for this environment
3423 if (HasMatchingEnvironmentOrNone(AA: EffectiveAvail))
3424 return Avail;
3425 PartialMatch = Avail;
3426 }
3427 }
3428 }
3429 return PartialMatch;
3430}
3431
3432// Check availability against target shader model version and current shader
3433// stage and emit diagnostic
3434void DiagnoseHLSLAvailability::CheckDeclAvailability(NamedDecl *D,
3435 const AvailabilityAttr *AA,
3436 SourceRange Range) {
3437
3438 const IdentifierInfo *IIEnv = AA->getEnvironment();
3439
3440 if (!IIEnv) {
3441 // The availability attribute does not have environment -> it depends only
3442 // on shader model version and not on specific the shader stage.
3443
3444 // Skip emitting the diagnostics if the diagnostic mode is set to
3445 // strict (-fhlsl-strict-availability) because all relevant diagnostics
3446 // were already emitted in the DiagnoseUnguardedAvailability scan
3447 // (SemaAvailability.cpp).
3448 if (SemaRef.getLangOpts().HLSLStrictAvailability)
3449 return;
3450
3451 // Do not report shader-stage-independent issues if scanning a function
3452 // that was already scanned in a different shader stage context (they would
3453 // be duplicate)
3454 if (ReportOnlyShaderStageIssues)
3455 return;
3456
3457 } else {
3458 // The availability attribute has environment -> we need to know
3459 // the current stage context to property diagnose it.
3460 if (InUnknownShaderStageContext())
3461 return;
3462 }
3463
3464 // Check introduced version and if environment matches
3465 bool EnvironmentMatches = HasMatchingEnvironmentOrNone(AA);
3466 VersionTuple Introduced = AA->getIntroduced();
3467 VersionTuple TargetVersion =
3468 SemaRef.Context.getTargetInfo().getPlatformMinVersion();
3469
3470 if (TargetVersion >= Introduced && EnvironmentMatches)
3471 return;
3472
3473 // Emit diagnostic message
3474 const TargetInfo &TI = SemaRef.getASTContext().getTargetInfo();
3475 llvm::StringRef PlatformName(
3476 AvailabilityAttr::getPrettyPlatformName(Platform: TI.getPlatformName()));
3477
3478 llvm::StringRef CurrentEnvStr =
3479 llvm::Triple::getEnvironmentTypeName(Kind: GetCurrentShaderEnvironment());
3480
3481 llvm::StringRef AttrEnvStr =
3482 AA->getEnvironment() ? AA->getEnvironment()->getName() : "";
3483 bool UseEnvironment = !AttrEnvStr.empty();
3484
3485 if (EnvironmentMatches) {
3486 SemaRef.Diag(Loc: Range.getBegin(), DiagID: diag::warn_hlsl_availability)
3487 << Range << D << PlatformName << Introduced.getAsString()
3488 << UseEnvironment << CurrentEnvStr;
3489 } else {
3490 SemaRef.Diag(Loc: Range.getBegin(), DiagID: diag::warn_hlsl_availability_unavailable)
3491 << Range << D;
3492 }
3493
3494 SemaRef.Diag(Loc: D->getLocation(), DiagID: diag::note_partial_availability_specified_here)
3495 << D << PlatformName << Introduced.getAsString()
3496 << SemaRef.Context.getTargetInfo().getPlatformMinVersion().getAsString()
3497 << UseEnvironment << AttrEnvStr << CurrentEnvStr;
3498}
3499
3500} // namespace
3501
3502void SemaHLSL::ActOnEndOfTranslationUnit(TranslationUnitDecl *TU) {
3503 // process default CBuffer - create buffer layout struct and invoke codegenCGH
3504 if (!DefaultCBufferDecls.empty()) {
3505 HLSLBufferDecl *DefaultCBuffer = HLSLBufferDecl::CreateDefaultCBuffer(
3506 C&: SemaRef.getASTContext(), LexicalParent: SemaRef.getCurLexicalContext(),
3507 DefaultCBufferDecls);
3508 addImplicitBindingAttrToDecl(S&: SemaRef, D: DefaultCBuffer, RT: RegisterType::CBuffer,
3509 ImplicitBindingOrderID: getNextImplicitBindingOrderID());
3510 SemaRef.getCurLexicalContext()->addDecl(D: DefaultCBuffer);
3511 createHostLayoutStructForBuffer(S&: SemaRef, BufDecl: DefaultCBuffer);
3512
3513 // Set HasValidPackoffset if any of the decls has a register(c#) annotation;
3514 for (const Decl *VD : DefaultCBufferDecls) {
3515 const HLSLResourceBindingAttr *RBA =
3516 VD->getAttr<HLSLResourceBindingAttr>();
3517 if (RBA && RBA->hasRegisterSlot() &&
3518 RBA->getRegisterType() == HLSLResourceBindingAttr::RegisterType::C) {
3519 DefaultCBuffer->setHasValidPackoffset(true);
3520 break;
3521 }
3522 }
3523
3524 DeclGroupRef DG(DefaultCBuffer);
3525 SemaRef.Consumer.HandleTopLevelDecl(D: DG);
3526 }
3527 diagnoseAvailabilityViolations(TU);
3528}
3529
3530// For resource member access through a global struct array, verify that the
3531// array index selecting the struct element is a constant integer expression.
3532// Returns false if the member expression is invalid.
3533bool SemaHLSL::ActOnResourceMemberAccessExpr(MemberExpr *ME) {
3534 assert((ME->getType()->isHLSLResourceRecord() ||
3535 ME->getType()->isHLSLResourceRecordArray()) &&
3536 "expected member expr to have resource record type or array of them");
3537
3538 // Walk the AST from MemberExpr to the VarDecl of the parent struct instance
3539 // and take note of any non-constant array indexing along the way. If the
3540 // VarDecl we find is a global variable, report error if there was any
3541 // non-constant array index in the resource member access along the way.
3542 const Expr *NonConstIndexExpr = nullptr;
3543 const Expr *E = ME->getBase();
3544 while (E) {
3545 if (const DeclRefExpr *DRE = dyn_cast<DeclRefExpr>(Val: E)) {
3546 if (!NonConstIndexExpr)
3547 return true;
3548
3549 const VarDecl *VD = cast<VarDecl>(Val: DRE->getDecl());
3550 if (!VD->hasGlobalStorage())
3551 return true;
3552
3553 SemaRef.Diag(Loc: NonConstIndexExpr->getExprLoc(),
3554 DiagID: diag::err_hlsl_resource_member_array_access_not_constant);
3555 return false;
3556 }
3557
3558 if (const auto *ASE = dyn_cast<ArraySubscriptExpr>(Val: E)) {
3559 const Expr *IdxExpr = ASE->getIdx();
3560 if (!IdxExpr->isIntegerConstantExpr(Ctx: SemaRef.getASTContext()))
3561 NonConstIndexExpr = IdxExpr;
3562 E = ASE->getBase();
3563 } else if (const auto *SubME = dyn_cast<MemberExpr>(Val: E)) {
3564 E = SubME->getBase();
3565 } else if (const auto *ICE = dyn_cast<ImplicitCastExpr>(Val: E)) {
3566 E = ICE->getSubExpr();
3567 } else {
3568 llvm_unreachable("unexpected expr type in resource member access");
3569 }
3570 }
3571 return true;
3572}
3573
3574NamedDecl *SemaHLSL::getConstantBufferConversionFunction(QualType Type,
3575 CXXRecordDecl *RD) {
3576 QualType AddrSpaceType =
3577 SemaRef.Context.getCanonicalType(T: SemaRef.Context.getAddrSpaceQualType(
3578 T: Type.withConst(), AddressSpace: LangAS::hlsl_constant));
3579 QualType ReturnTy = SemaRef.Context.getCanonicalType(
3580 T: SemaRef.Context.getLValueReferenceType(T: AddrSpaceType));
3581
3582 DeclarationName ConvName =
3583 SemaRef.Context.DeclarationNames.getCXXConversionFunctionName(
3584 Ty: CanQualType::CreateUnsafe(Other: ReturnTy));
3585 LookupResult ConvR(SemaRef, ConvName, SourceLocation(),
3586 Sema::LookupOrdinaryName);
3587 [[maybe_unused]] bool LookupSucceeded =
3588 SemaRef.LookupQualifiedName(R&: ConvR, LookupCtx: RD);
3589 assert(LookupSucceeded);
3590
3591 for (NamedDecl *D : ConvR) {
3592 if (isa<CXXConversionDecl>(Val: D->getUnderlyingDecl()))
3593 return D;
3594 }
3595 return nullptr;
3596}
3597
3598std::optional<ExprResult>
3599SemaHLSL::tryPerformConstantBufferConversion(Expr *BaseExpr) {
3600 QualType BaseType = BaseExpr->getType();
3601 const HLSLAttributedResourceType *ResTy =
3602 HLSLAttributedResourceType::findHandleTypeOnResource(
3603 RT: BaseType.getTypePtr());
3604 if (!ResTy ||
3605 ResTy->getAttrs().ResourceClass != llvm::dxil::ResourceClass::CBuffer)
3606 return std::nullopt;
3607
3608 QualType TemplateType = ResTy->getContainedType();
3609
3610 NamedDecl *NamedConversionDecl = getConstantBufferConversionFunction(
3611 Type: TemplateType, RD: BaseType->getAsCXXRecordDecl());
3612 assert(NamedConversionDecl &&
3613 "Could not find conversion function for ConstantBuffer.");
3614 auto *ConversionDecl =
3615 cast<CXXConversionDecl>(Val: NamedConversionDecl->getUnderlyingDecl());
3616
3617 return SemaRef.BuildCXXMemberCallExpr(Exp: BaseExpr, FoundDecl: NamedConversionDecl,
3618 Method: ConversionDecl,
3619 /*HadMultipleCandidates=*/false);
3620}
3621
3622void SemaHLSL::diagnoseAvailabilityViolations(TranslationUnitDecl *TU) {
3623 // Strict mode diagnoses availability during the
3624 // DiagnoseUnguardedAvailability scan in SemaAvailability.cpp. The reachable
3625 // function scan must still run to validate Barrier calls.
3626 const TargetInfo &TI = SemaRef.getASTContext().getTargetInfo();
3627 const bool DiagnoseAvailability =
3628 !SemaRef.getLangOpts().HLSLStrictAvailability ||
3629 TI.getTriple().getEnvironment() == llvm::Triple::EnvironmentType::Library;
3630 DiagnoseHLSLAvailability(SemaRef, DiagnoseAvailability)
3631 .RunOnTranslationUnit(TU);
3632}
3633
3634static bool CheckAllArgsHaveSameType(Sema *S, CallExpr *TheCall) {
3635 assert(TheCall->getNumArgs() > 1);
3636 QualType ArgTy0 = TheCall->getArg(Arg: 0)->getType();
3637
3638 for (unsigned I = 1, N = TheCall->getNumArgs(); I < N; ++I) {
3639 if (!S->getASTContext().hasSameUnqualifiedType(
3640 T1: ArgTy0, T2: TheCall->getArg(Arg: I)->getType())) {
3641 S->Diag(Loc: TheCall->getBeginLoc(), DiagID: diag::err_vec_builtin_incompatible_vector)
3642 << TheCall->getDirectCallee() << /*useAllTerminology*/ true
3643 << SourceRange(TheCall->getArg(Arg: 0)->getBeginLoc(),
3644 TheCall->getArg(Arg: N - 1)->getEndLoc());
3645 return true;
3646 }
3647 }
3648 return false;
3649}
3650
3651static bool CheckArgTypeMatches(Sema *S, Expr *Arg, QualType ExpectedType) {
3652 QualType ArgType = Arg->getType();
3653 if (!S->getASTContext().hasSameUnqualifiedType(T1: ArgType, T2: ExpectedType)) {
3654 S->Diag(Loc: Arg->getBeginLoc(), DiagID: diag::err_typecheck_convert_incompatible)
3655 << ArgType << ExpectedType << 1 << 0 << 0;
3656 return true;
3657 }
3658 return false;
3659}
3660
3661static bool CheckAllArgTypesAreCorrect(
3662 Sema *S, CallExpr *TheCall,
3663 llvm::function_ref<bool(Sema *S, SourceLocation Loc, int ArgOrdinal,
3664 clang::QualType PassedType)>
3665 Check) {
3666 for (unsigned I = 0; I < TheCall->getNumArgs(); ++I) {
3667 Expr *Arg = TheCall->getArg(Arg: I);
3668 if (Check(S, Arg->getBeginLoc(), I + 1, Arg->getType()))
3669 return true;
3670 }
3671 return false;
3672}
3673
3674static bool CheckFloatRepresentation(Sema *S, SourceLocation Loc,
3675 int ArgOrdinal,
3676 clang::QualType PassedType) {
3677 clang::QualType BaseType =
3678 getElementTypeOf(T: PassedType, /*IncludeMatrix=*/false);
3679 if (!BaseType->isFloat32Type())
3680 return S->Diag(Loc, DiagID: diag::err_builtin_invalid_arg_type)
3681 << ArgOrdinal << /* scalar or vector of */ 5 << /* no int */ 0
3682 << /* float */ 1 << PassedType;
3683 return false;
3684}
3685
3686static bool CheckFloatOrHalfRepresentation(Sema *S, SourceLocation Loc,
3687 int ArgOrdinal,
3688 clang::QualType PassedType) {
3689 QualType BaseType = getScalarComponentType(T: PassedType);
3690
3691 if (!BaseType->isHalfType() && !BaseType->isFloat32Type())
3692 return S->Diag(Loc, DiagID: diag::err_builtin_invalid_arg_type)
3693 << ArgOrdinal << /* scalar or vector of */ 5 << /* no int */ 0
3694 << /* half or float */ 2 << PassedType;
3695 return false;
3696}
3697
3698static bool CheckAnyDoubleRepresentation(Sema *S, SourceLocation Loc,
3699 int ArgOrdinal,
3700 clang::QualType PassedType) {
3701 QualType BaseType = getScalarComponentType(T: PassedType);
3702 if (!BaseType->isDoubleType()) {
3703 // FIXME: adopt standard `err_builtin_invalid_arg_type` instead of using
3704 // this custom error.
3705 return S->Diag(Loc, DiagID: diag::err_builtin_requires_double_type)
3706 << ArgOrdinal << PassedType;
3707 }
3708
3709 return false;
3710}
3711
3712static bool CheckModifiableLValue(Sema *S, CallExpr *TheCall,
3713 unsigned ArgIndex) {
3714 auto *Arg = TheCall->getArg(Arg: ArgIndex);
3715 SourceLocation OrigLoc = Arg->getExprLoc();
3716 if (Arg->IgnoreCasts()->isModifiableLvalue(Ctx&: S->Context, Loc: &OrigLoc) ==
3717 Expr::MLV_Valid)
3718 return false;
3719 S->Diag(Loc: OrigLoc, DiagID: diag::error_hlsl_inout_lvalue) << Arg << 0;
3720 return true;
3721}
3722
3723// Verifies that the argument at `ArgIndex` of `TheCall` refers to memory in
3724// one of `AllowedSpaces`. Intended for HLSL builtins (e.g. atomics).
3725static bool CheckArgAddrSpaceOneOf(Sema *S, CallExpr *TheCall,
3726 unsigned ArgIndex,
3727 ArrayRef<LangAS> AllowedSpaces) {
3728 Expr *Arg = TheCall->getArg(Arg: ArgIndex);
3729 QualType LValueTy = Arg->IgnoreCasts()->getType();
3730 if (llvm::is_contained(Range&: AllowedSpaces, Element: LValueTy.getAddressSpace()))
3731 return false;
3732 S->Diag(Loc: Arg->getBeginLoc(), DiagID: diag::err_hlsl_atomic_arg_addr_space)
3733 << (ArgIndex + 1) << LValueTy;
3734 return true;
3735}
3736
3737static bool CheckNoDoubleVectors(Sema *S, SourceLocation Loc, int ArgOrdinal,
3738 clang::QualType PassedType) {
3739 const auto *VecTy = PassedType->getAs<VectorType>();
3740 if (!VecTy)
3741 return false;
3742
3743 if (VecTy->getElementType()->isDoubleType())
3744 return S->Diag(Loc, DiagID: diag::err_builtin_invalid_arg_type)
3745 << ArgOrdinal << /* scalar */ 1 << /* no int */ 0 << /* fp */ 1
3746 << PassedType;
3747 return false;
3748}
3749
3750static bool CheckFloatingOrIntRepresentation(Sema *S, SourceLocation Loc,
3751 int ArgOrdinal,
3752 clang::QualType PassedType) {
3753 if (!PassedType->hasIntegerRepresentation() &&
3754 !PassedType->hasFloatingRepresentation())
3755 return S->Diag(Loc, DiagID: diag::err_builtin_invalid_arg_type)
3756 << ArgOrdinal << /* scalar or vector of */ 5 << /* integer */ 1
3757 << /* fp */ 1 << PassedType;
3758 return false;
3759}
3760
3761static bool CheckUnsignedIntVecRepresentation(Sema *S, SourceLocation Loc,
3762 int ArgOrdinal,
3763 clang::QualType PassedType) {
3764 if (auto *VecTy = PassedType->getAs<VectorType>())
3765 if (VecTy->getElementType()->isUnsignedIntegerType())
3766 return false;
3767
3768 return S->Diag(Loc, DiagID: diag::err_builtin_invalid_arg_type)
3769 << ArgOrdinal << /* vector of */ 4 << /* uint */ 3 << /* no fp */ 0
3770 << PassedType;
3771}
3772
3773// checks for unsigned ints of all sizes
3774static bool CheckUnsignedIntRepresentation(Sema *S, SourceLocation Loc,
3775 int ArgOrdinal,
3776 clang::QualType PassedType) {
3777 if (!PassedType->hasUnsignedIntegerRepresentation())
3778 return S->Diag(Loc, DiagID: diag::err_builtin_invalid_arg_type)
3779 << ArgOrdinal << /* scalar or vector of */ 5 << /* unsigned int */ 3
3780 << /* no fp */ 0 << PassedType;
3781 return false;
3782}
3783
3784static bool CheckExpectedBitWidth(Sema *S, CallExpr *TheCall,
3785 unsigned ArgOrdinal, unsigned Width) {
3786 QualType ArgTy = TheCall->getArg(Arg: 0)->getType();
3787 if (auto *VTy = ArgTy->getAs<VectorType>())
3788 ArgTy = VTy->getElementType();
3789 // ensure arg type has expected bit width
3790 uint64_t ElementBitCount =
3791 S->getASTContext().getTypeSizeInChars(T: ArgTy).getQuantity() * 8;
3792 if (ElementBitCount != Width) {
3793 S->Diag(Loc: TheCall->getArg(Arg: 0)->getBeginLoc(),
3794 DiagID: diag::err_integer_incorrect_bit_count)
3795 << Width << ElementBitCount;
3796 return true;
3797 }
3798 return false;
3799}
3800
3801static void SetElementTypeAsReturnType(Sema *S, CallExpr *TheCall,
3802 QualType ReturnType) {
3803 if (auto *VecTyA = TheCall->getArg(Arg: 0)->getType()->getAs<VectorType>())
3804 ReturnType =
3805 S->Context.getExtVectorType(VectorType: ReturnType, NumElts: VecTyA->getNumElements());
3806 else if (auto *MatTyA =
3807 TheCall->getArg(Arg: 0)->getType()->getAs<ConstantMatrixType>())
3808 ReturnType = S->Context.getConstantMatrixType(
3809 ElementType: ReturnType, NumRows: MatTyA->getNumRows(), NumColumns: MatTyA->getNumColumns());
3810
3811 TheCall->setType(ReturnType);
3812}
3813
3814static bool CheckScalarOrVector(Sema *S, CallExpr *TheCall, QualType Scalar,
3815 unsigned ArgIndex) {
3816 assert(TheCall->getNumArgs() >= ArgIndex);
3817 QualType ArgType = TheCall->getArg(Arg: ArgIndex)->getType();
3818 auto *VTy = ArgType->getAs<VectorType>();
3819 // not the scalar or vector<scalar>
3820 if (!(S->Context.hasSameUnqualifiedType(T1: ArgType, T2: Scalar) ||
3821 (VTy &&
3822 S->Context.hasSameUnqualifiedType(T1: VTy->getElementType(), T2: Scalar)))) {
3823 S->Diag(Loc: TheCall->getArg(Arg: 0)->getBeginLoc(),
3824 DiagID: diag::err_typecheck_expect_scalar_or_vector)
3825 << ArgType << Scalar;
3826 return true;
3827 }
3828 return false;
3829}
3830
3831static bool CheckScalarOrVectorOrMatrix(Sema *S, CallExpr *TheCall,
3832 QualType Scalar, unsigned ArgIndex) {
3833 assert(TheCall->getNumArgs() > ArgIndex);
3834
3835 Expr *Arg = TheCall->getArg(Arg: ArgIndex);
3836 QualType ArgType = Arg->getType();
3837
3838 // Scalar: T
3839 if (S->Context.hasSameUnqualifiedType(T1: ArgType, T2: Scalar))
3840 return false;
3841
3842 // Vector: vector<T>
3843 if (const auto *VTy = ArgType->getAs<VectorType>()) {
3844 if (S->Context.hasSameUnqualifiedType(T1: VTy->getElementType(), T2: Scalar))
3845 return false;
3846 }
3847
3848 // Matrix: ConstantMatrixType with element type T
3849 if (const auto *MTy = ArgType->getAs<ConstantMatrixType>()) {
3850 if (S->Context.hasSameUnqualifiedType(T1: MTy->getElementType(), T2: Scalar))
3851 return false;
3852 }
3853
3854 // Not a scalar/vector/matrix-of-scalar
3855 S->Diag(Loc: Arg->getBeginLoc(),
3856 DiagID: diag::err_typecheck_expect_scalar_or_vector_or_matrix)
3857 << ArgType << Scalar;
3858 return true;
3859}
3860
3861static bool CheckAnyScalarOrVector(Sema *S, CallExpr *TheCall,
3862 unsigned ArgIndex) {
3863 assert(TheCall->getNumArgs() >= ArgIndex);
3864 QualType ArgType = TheCall->getArg(Arg: ArgIndex)->getType();
3865 auto *VTy = ArgType->getAs<VectorType>();
3866 // not the scalar or vector<scalar>
3867 if (!(ArgType->isScalarType() ||
3868 (VTy && VTy->getElementType()->isScalarType()))) {
3869 S->Diag(Loc: TheCall->getArg(Arg: 0)->getBeginLoc(),
3870 DiagID: diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
3871 << ArgType << 1;
3872 return true;
3873 }
3874 return false;
3875}
3876
3877static bool CheckAnyScalarOrVectorOrMatrix(Sema *S, CallExpr *TheCall,
3878 unsigned ArgIndex) {
3879 assert(TheCall->getNumArgs() > ArgIndex);
3880 QualType ArgType = TheCall->getArg(Arg: ArgIndex)->getType();
3881 if (ArgType->isDependentType())
3882 return false;
3883
3884 QualType ElementType = ArgType;
3885 if (const auto *VectorTy = ArgType->getAs<VectorType>())
3886 ElementType = VectorTy->getElementType();
3887 else if (const auto *MatrixTy = ArgType->getAs<ConstantMatrixType>())
3888 ElementType = MatrixTy->getElementType();
3889
3890 if (ElementType->isBooleanType())
3891 return false;
3892
3893 if (ElementType->isIntegerType() || ElementType->isRealFloatingType()) {
3894 unsigned BitWidth = S->Context.getTypeSize(T: ElementType);
3895 if (BitWidth == 16 || BitWidth == 32 || BitWidth == 64)
3896 return false;
3897 }
3898
3899 S->Diag(Loc: TheCall->getArg(Arg: ArgIndex)->getBeginLoc(),
3900 DiagID: diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
3901 << ArgType << 2;
3902 return true;
3903}
3904
3905// Check that the argument is not a bool or vector<bool>
3906// Returns true on error
3907static bool CheckNotBoolScalarOrVector(Sema *S, CallExpr *TheCall,
3908 unsigned ArgIndex) {
3909 QualType BoolType = S->getASTContext().BoolTy;
3910 assert(ArgIndex < TheCall->getNumArgs());
3911 QualType ArgType = TheCall->getArg(Arg: ArgIndex)->getType();
3912 auto *VTy = ArgType->getAs<VectorType>();
3913 // is the bool or vector<bool>
3914 if (S->Context.hasSameUnqualifiedType(T1: ArgType, T2: BoolType) ||
3915 (VTy &&
3916 S->Context.hasSameUnqualifiedType(T1: VTy->getElementType(), T2: BoolType))) {
3917 S->Diag(Loc: TheCall->getArg(Arg: 0)->getBeginLoc(),
3918 DiagID: diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
3919 << ArgType << 0;
3920 return true;
3921 }
3922 return false;
3923}
3924
3925static bool CheckWaveActive(Sema *S, CallExpr *TheCall) {
3926 if (CheckNotBoolScalarOrVector(S, TheCall, ArgIndex: 0))
3927 return true;
3928 return false;
3929}
3930
3931static bool CheckWavePrefix(Sema *S, CallExpr *TheCall) {
3932 if (CheckNotBoolScalarOrVector(S, TheCall, ArgIndex: 0))
3933 return true;
3934 return false;
3935}
3936
3937static bool CheckBoolSelect(Sema *S, CallExpr *TheCall) {
3938 assert(TheCall->getNumArgs() == 3);
3939 Expr *Arg1 = TheCall->getArg(Arg: 1);
3940 Expr *Arg2 = TheCall->getArg(Arg: 2);
3941 if (!S->Context.hasSameUnqualifiedType(T1: Arg1->getType(), T2: Arg2->getType())) {
3942 S->Diag(Loc: TheCall->getBeginLoc(),
3943 DiagID: diag::err_typecheck_call_different_arg_types)
3944 << Arg1->getType() << Arg2->getType() << Arg1->getSourceRange()
3945 << Arg2->getSourceRange();
3946 return true;
3947 }
3948
3949 TheCall->setType(Arg1->getType());
3950 return false;
3951}
3952
3953static bool CheckVectorSelect(Sema *S, CallExpr *TheCall) {
3954 assert(TheCall->getNumArgs() == 3);
3955 Expr *Arg1 = TheCall->getArg(Arg: 1);
3956 QualType Arg1Ty = Arg1->getType();
3957 Expr *Arg2 = TheCall->getArg(Arg: 2);
3958 QualType Arg2Ty = Arg2->getType();
3959
3960 QualType Arg1ScalarTy = Arg1Ty;
3961 if (auto VTy = Arg1ScalarTy->getAs<VectorType>())
3962 Arg1ScalarTy = VTy->getElementType();
3963
3964 QualType Arg2ScalarTy = Arg2Ty;
3965 if (auto VTy = Arg2ScalarTy->getAs<VectorType>())
3966 Arg2ScalarTy = VTy->getElementType();
3967
3968 if (!S->Context.hasSameUnqualifiedType(T1: Arg1ScalarTy, T2: Arg2ScalarTy))
3969 S->Diag(Loc: Arg1->getBeginLoc(), DiagID: diag::err_hlsl_builtin_scalar_vector_mismatch)
3970 << /* second and third */ 1 << TheCall->getCallee() << Arg1Ty << Arg2Ty;
3971
3972 QualType Arg0Ty = TheCall->getArg(Arg: 0)->getType();
3973 unsigned Arg0Length = Arg0Ty->getAs<VectorType>()->getNumElements();
3974 unsigned Arg1Length = Arg1Ty->isVectorType()
3975 ? Arg1Ty->getAs<VectorType>()->getNumElements()
3976 : 0;
3977 unsigned Arg2Length = Arg2Ty->isVectorType()
3978 ? Arg2Ty->getAs<VectorType>()->getNumElements()
3979 : 0;
3980 if (Arg1Length > 0 && Arg0Length != Arg1Length) {
3981 S->Diag(Loc: TheCall->getBeginLoc(),
3982 DiagID: diag::err_typecheck_vector_lengths_not_equal)
3983 << Arg0Ty << Arg1Ty << TheCall->getArg(Arg: 0)->getSourceRange()
3984 << Arg1->getSourceRange();
3985 return true;
3986 }
3987
3988 if (Arg2Length > 0 && Arg0Length != Arg2Length) {
3989 S->Diag(Loc: TheCall->getBeginLoc(),
3990 DiagID: diag::err_typecheck_vector_lengths_not_equal)
3991 << Arg0Ty << Arg2Ty << TheCall->getArg(Arg: 0)->getSourceRange()
3992 << Arg2->getSourceRange();
3993 return true;
3994 }
3995
3996 TheCall->setType(
3997 S->getASTContext().getExtVectorType(VectorType: Arg1ScalarTy, NumElts: Arg0Length));
3998 return false;
3999}
4000
4001static bool CheckMatrixSelect(Sema *S, CallExpr *TheCall) {
4002 assert(TheCall->getNumArgs() == 3);
4003 Expr *Arg1 = TheCall->getArg(Arg: 1);
4004 QualType Arg1Ty = Arg1->getType();
4005 Expr *Arg2 = TheCall->getArg(Arg: 2);
4006 QualType Arg2Ty = Arg2->getType();
4007
4008 QualType Arg1ScalarTy = Arg1Ty;
4009 if (auto MTy = Arg1ScalarTy->getAs<ConstantMatrixType>())
4010 Arg1ScalarTy = MTy->getElementType();
4011
4012 QualType Arg2ScalarTy = Arg2Ty;
4013 if (auto MTy = Arg2ScalarTy->getAs<ConstantMatrixType>())
4014 Arg2ScalarTy = MTy->getElementType();
4015
4016 if (!S->Context.hasSameUnqualifiedType(T1: Arg1ScalarTy, T2: Arg2ScalarTy))
4017 S->Diag(Loc: Arg1->getBeginLoc(), DiagID: diag::err_hlsl_builtin_scalar_vector_mismatch)
4018 << /* second and third */ 1 << TheCall->getCallee() << Arg1Ty << Arg2Ty;
4019
4020 QualType Arg0Ty = TheCall->getArg(Arg: 0)->getType();
4021 auto *Arg0MatTy = Arg0Ty->getAs<ConstantMatrixType>();
4022 unsigned Arg0Rows = Arg0MatTy->getNumRows();
4023 unsigned Arg0Cols = Arg0MatTy->getNumColumns();
4024
4025 for (Expr *Arg : {Arg1, Arg2}) {
4026 auto *MTy = Arg->getType()->getAs<ConstantMatrixType>();
4027 if (MTy &&
4028 (MTy->getNumRows() != Arg0Rows || MTy->getNumColumns() != Arg0Cols)) {
4029 S->Diag(Loc: TheCall->getBeginLoc(),
4030 DiagID: diag::err_typecheck_vector_lengths_not_equal)
4031 << Arg0Ty << Arg->getType() << TheCall->getArg(Arg: 0)->getSourceRange()
4032 << Arg->getSourceRange();
4033 return true;
4034 }
4035 }
4036
4037 TheCall->setType(
4038 S->Context.getConstantMatrixType(ElementType: Arg1ScalarTy, NumRows: Arg0Rows, NumColumns: Arg0Cols));
4039 return false;
4040}
4041
4042static QualType getVectorOrScalarType(Sema &S, QualType BaseType,
4043 unsigned Count) {
4044 return Count > 1 ? S.Context.getExtVectorType(VectorType: BaseType, NumElts: Count) : BaseType;
4045}
4046
4047static bool CheckScalarFloatOperand(Sema &S, CallExpr *TheCall,
4048 unsigned ArgIndex) {
4049 return CheckArgTypeMatches(S: &S, Arg: TheCall->getArg(Arg: ArgIndex), ExpectedType: S.Context.FloatTy);
4050}
4051
4052static bool CheckIndexType(Sema *S, CallExpr *TheCall, unsigned IndexArgIndex) {
4053 assert(TheCall->getNumArgs() > IndexArgIndex && "Index argument missing");
4054 QualType ArgType = TheCall->getArg(Arg: IndexArgIndex)->getType();
4055 QualType IndexTy = ArgType;
4056 unsigned int ActualDim = 1;
4057 if (const auto *VTy = IndexTy->getAs<VectorType>()) {
4058 ActualDim = VTy->getNumElements();
4059 IndexTy = VTy->getElementType();
4060 }
4061 if (!IndexTy->isIntegerType()) {
4062 S->Diag(Loc: TheCall->getArg(Arg: IndexArgIndex)->getBeginLoc(),
4063 DiagID: diag::err_typecheck_expect_int)
4064 << ArgType;
4065 return true;
4066 }
4067
4068 QualType ResourceArgTy = TheCall->getArg(Arg: 0)->getType();
4069 const HLSLAttributedResourceType *ResTy =
4070 ResourceArgTy.getTypePtr()->getAs<HLSLAttributedResourceType>();
4071 assert(ResTy && "Resource argument must be a resource");
4072 HLSLAttributedResourceType::Attributes ResAttrs = ResTy->getAttrs();
4073
4074 unsigned int ExpectedDim = 1;
4075 if (ResAttrs.ResourceDimension != llvm::dxil::ResourceDimension::Unknown)
4076 ExpectedDim = getResourceDimensions(Dim: ResAttrs.ResourceDimension) +
4077 (ResAttrs.IsArray ? 1 : 0);
4078
4079 if (ActualDim != ExpectedDim) {
4080 S->Diag(Loc: TheCall->getArg(Arg: IndexArgIndex)->getBeginLoc(),
4081 DiagID: diag::err_hlsl_builtin_resource_coordinate_dimension_mismatch)
4082 << cast<NamedDecl>(Val: TheCall->getCalleeDecl()) << ExpectedDim
4083 << ActualDim;
4084 return true;
4085 }
4086
4087 return false;
4088}
4089
4090static bool CheckResourceHandle(
4091 Sema *S, CallExpr *TheCall, unsigned ArgIndex,
4092 llvm::function_ref<bool(const HLSLAttributedResourceType *ResType)> Check =
4093 nullptr) {
4094 assert(TheCall->getNumArgs() >= ArgIndex);
4095 QualType ArgType = TheCall->getArg(Arg: ArgIndex)->getType();
4096 const HLSLAttributedResourceType *ResTy =
4097 ArgType.getTypePtr()->getAs<HLSLAttributedResourceType>();
4098 if (!ResTy) {
4099 S->Diag(Loc: TheCall->getArg(Arg: ArgIndex)->getBeginLoc(),
4100 DiagID: diag::err_typecheck_expect_hlsl_resource)
4101 << ArgType;
4102 return true;
4103 }
4104 if (Check && Check(ResTy)) {
4105 S->Diag(Loc: TheCall->getArg(Arg: ArgIndex)->getExprLoc(),
4106 DiagID: diag::err_invalid_hlsl_resource_type)
4107 << ArgType;
4108 return true;
4109 }
4110 return false;
4111}
4112
4113static QualType createCounterHandleType(ASTContext &AST,
4114 QualType MainHandleTy) {
4115 assert(MainHandleTy->isHLSLAttributedResourceType() &&
4116 "expected resource handle type");
4117 auto *MainResType = MainHandleTy->getAs<HLSLAttributedResourceType>();
4118 auto MainAttrs = MainResType->getAttrs();
4119 assert(!MainAttrs.IsCounter && "cannot create a counter from a counter");
4120 MainAttrs.IsCounter = true;
4121 return AST.getHLSLAttributedResourceType(Wrapped: MainResType->getWrappedType(),
4122 Contained: MainResType->getContainedType(),
4123 Attrs: MainAttrs);
4124}
4125
4126enum class SampleKind { Sample, Bias, Grad, Level, Cmp, CmpLevelZero };
4127
4128static StringRef getSampleMethodName(SampleKind Kind) {
4129 switch (Kind) {
4130 case SampleKind::Sample:
4131 return "Sample";
4132 case SampleKind::Bias:
4133 return "SampleBias";
4134 case SampleKind::Grad:
4135 return "SampleGrad";
4136 case SampleKind::Level:
4137 return "SampleLevel";
4138 case SampleKind::Cmp:
4139 return "SampleCmp";
4140 case SampleKind::CmpLevelZero:
4141 return "SampleCmpLevelZero";
4142 }
4143 llvm_unreachable("Invalid SampleKind");
4144}
4145
4146// Returns the name of the resource method whose body the sampling or gather
4147// builtin is being emitted into, which is the name the user called. This
4148// matters for methods that share a builtin, like 'Gather' and 'GatherRed'.
4149// Falls back to DefaultName if the builtin is used outside of a resource
4150// method.
4151static StringRef getCurrentResourceMethodName(Sema &S, StringRef DefaultName) {
4152 const auto *MD = dyn_cast_if_present<CXXMethodDecl>(Val: S.getCurFunctionDecl());
4153 if (!MD || !MD->getDeclName().isIdentifier())
4154 return DefaultName;
4155
4156 QualType RecordTy = S.Context.getCanonicalTagType(TD: MD->getParent());
4157 if (!RecordTy->isHLSLResourceRecord())
4158 return DefaultName;
4159
4160 return MD->getName();
4161}
4162
4163// Sampling from and gathering on resources with a 'double' element type is not
4164// supported. Such resources are still valid declarations whose contents can be
4165// accessed by other means, like Load or the subscript operator.
4166static bool CheckNoDoubleElementType(Sema &S, CallExpr *TheCall,
4167 QualType ContainedType,
4168 StringRef DefaultName) {
4169 QualType EltTy = getElementTypeOf(T: ContainedType, /*IncludeMatrix=*/false);
4170 if (!EltTy->isSpecificBuiltinType(K: BuiltinType::Double))
4171 return false;
4172
4173 S.Diag(Loc: TheCall->getBeginLoc(), DiagID: diag::err_hlsl_sample_double_element_type)
4174 << getCurrentResourceMethodName(S, DefaultName) << ContainedType;
4175 return true;
4176}
4177
4178// Sampling textures with an integer element type was introduced in SM 6.7 as
4179// part of Advanced Texture Operations. The shader model only applies to DirectX
4180// targets; Vulkan has no such restriction.
4181static bool CheckIntegerElementTypeShaderModel(Sema &S, CallExpr *TheCall,
4182 QualType ContainedType,
4183 SampleKind Kind) {
4184 // Comparison sampling requires a floating point element type at every shader
4185 // model, which the caller diagnoses.
4186 if (Kind == SampleKind::Cmp || Kind == SampleKind::CmpLevelZero)
4187 return false;
4188
4189 // 'bool' is an integer type in HLSL, but sampling bool resources is never
4190 // allowed, so it must not be reported as requiring shader model 6.7.
4191 QualType EltTy = getElementTypeOf(T: ContainedType, /*IncludeMatrix=*/false);
4192 if (!EltTy->isIntegerType() || EltTy->isBooleanType())
4193 return false;
4194
4195 const TargetInfo &TI = S.Context.getTargetInfo();
4196 if (!TI.getTriple().isDXIL())
4197 return false;
4198
4199 VersionTuple SMVersion = TI.getPlatformMinVersion();
4200 if (SMVersion >= VersionTuple(6, 7))
4201 return false;
4202
4203 S.Diag(Loc: TheCall->getBeginLoc(), DiagID: diag::err_hlsl_sample_integer_element_type)
4204 << getCurrentResourceMethodName(S, DefaultName: getSampleMethodName(Kind))
4205 << ContainedType << SMVersion.getAsString();
4206 return true;
4207}
4208
4209static bool CheckTextureSamplerAndLocation(Sema &S, CallExpr *TheCall,
4210 bool IncludeArraySlice = true) {
4211 // Check the texture handle.
4212 if (CheckResourceHandle(S: &S, TheCall, ArgIndex: 0,
4213 Check: [](const HLSLAttributedResourceType *ResType) {
4214 return ResType->getAttrs().ResourceDimension ==
4215 llvm::dxil::ResourceDimension::Unknown;
4216 }))
4217 return true;
4218
4219 // Check the sampler handle.
4220 if (CheckResourceHandle(S: &S, TheCall, ArgIndex: 1,
4221 Check: [](const HLSLAttributedResourceType *ResType) {
4222 return ResType->getAttrs().ResourceClass !=
4223 llvm::hlsl::ResourceClass::Sampler;
4224 }))
4225 return true;
4226
4227 auto *ResourceTy =
4228 TheCall->getArg(Arg: 0)->getType()->castAs<HLSLAttributedResourceType>();
4229
4230 // Check the location.
4231 unsigned ExpectedDim =
4232 getResourceDimensions(Dim: ResourceTy->getAttrs().ResourceDimension) +
4233 (IncludeArraySlice && ResourceTy->getAttrs().IsArray ? 1 : 0);
4234 if (CheckArgTypeMatches(
4235 S: &S, Arg: TheCall->getArg(Arg: 2),
4236 ExpectedType: getVectorOrScalarType(S, BaseType: S.Context.FloatTy, Count: ExpectedDim)))
4237 return true;
4238
4239 return false;
4240}
4241
4242static bool CheckCalculateLodBuiltin(Sema &S, CallExpr *TheCall) {
4243 if (S.checkArgCount(Call: TheCall, DesiredArgCount: 3))
4244 return true;
4245
4246 // CalculateLevelOfDetail location uses resource dimension only (e.g. float2
4247 // for 2D), not an extra array slice component like Sample/Gather.
4248 if (CheckTextureSamplerAndLocation(S, TheCall, /*IncludeArraySlice=*/false))
4249 return true;
4250
4251 TheCall->setType(S.Context.FloatTy);
4252 return false;
4253}
4254
4255static bool CheckGatherBuiltin(Sema &S, CallExpr *TheCall, bool IsCmp) {
4256 if (S.checkArgCountRange(Call: TheCall, MinArgCount: IsCmp ? 5 : 4, MaxArgCount: IsCmp ? 6 : 5))
4257 return true;
4258
4259 if (CheckTextureSamplerAndLocation(S, TheCall))
4260 return true;
4261
4262 unsigned NextIdx = 3;
4263 if (IsCmp) {
4264 // Check the compare value.
4265 if (CheckScalarFloatOperand(S, TheCall, ArgIndex: NextIdx))
4266 return true;
4267 NextIdx++;
4268 }
4269
4270 // Check the component operand.
4271 if (CheckArgTypeMatches(S: &S, Arg: TheCall->getArg(Arg: NextIdx),
4272 ExpectedType: S.Context.UnsignedIntTy))
4273 return true;
4274 Expr *ComponentArg = TheCall->getArg(Arg: NextIdx);
4275
4276 // GatherCmp operations on Vulkan target must use component 0 (Red).
4277 if (IsCmp && S.getASTContext().getTargetInfo().getTriple().isSPIRV()) {
4278 std::optional<llvm::APSInt> ComponentOpt =
4279 ComponentArg->getIntegerConstantExpr(Ctx: S.getASTContext());
4280 if (ComponentOpt) {
4281 int64_t ComponentVal = ComponentOpt->getSExtValue();
4282 if (ComponentVal != 0) {
4283 // Issue an error if the component is not 0 (Red).
4284 // 0 -> Red, 1 -> Green, 2 -> Blue, 3 -> Alpha
4285 assert(ComponentVal >= 0 && ComponentVal <= 3 &&
4286 "The component is not in the expected range.");
4287 S.Diag(Loc: ComponentArg->getBeginLoc(),
4288 DiagID: diag::err_hlsl_gathercmp_invalid_component)
4289 << ComponentVal;
4290 return true;
4291 }
4292 }
4293 }
4294
4295 NextIdx++;
4296
4297 // Check the offset operand.
4298 const HLSLAttributedResourceType *ResourceTy =
4299 TheCall->getArg(Arg: 0)->getType()->castAs<HLSLAttributedResourceType>();
4300 if (TheCall->getNumArgs() > NextIdx) {
4301 unsigned ExpectedDim =
4302 getResourceDimensions(Dim: ResourceTy->getAttrs().ResourceDimension);
4303 if (CheckArgTypeMatches(
4304 S: &S, Arg: TheCall->getArg(Arg: NextIdx),
4305 ExpectedType: getVectorOrScalarType(S, BaseType: S.Context.IntTy, Count: ExpectedDim)))
4306 return true;
4307 NextIdx++;
4308 }
4309
4310 assert(ResourceTy->hasContainedType() &&
4311 "Expecting a contained type for resource with a dimension "
4312 "attribute.");
4313 QualType ReturnType = ResourceTy->getContainedType();
4314
4315 if (CheckNoDoubleElementType(S, TheCall, ContainedType: ReturnType,
4316 DefaultName: IsCmp ? "GatherCmp" : "Gather"))
4317 return true;
4318
4319 if (IsCmp) {
4320 if (!ReturnType->hasFloatingRepresentation()) {
4321 S.Diag(Loc: TheCall->getBeginLoc(), DiagID: diag::err_hlsl_samplecmp_requires_float);
4322 return true;
4323 }
4324 }
4325
4326 if (const auto *VecTy = ReturnType->getAs<VectorType>())
4327 ReturnType = VecTy->getElementType();
4328 ReturnType = S.Context.getExtVectorType(VectorType: ReturnType, NumElts: 4);
4329
4330 TheCall->setType(ReturnType);
4331
4332 return false;
4333}
4334static bool CheckLoadLevelBuiltin(Sema &S, CallExpr *TheCall) {
4335 if (S.checkArgCountRange(Call: TheCall, MinArgCount: 2, MaxArgCount: 3))
4336 return true;
4337
4338 // Check the texture handle.
4339 if (CheckResourceHandle(S: &S, TheCall, ArgIndex: 0,
4340 Check: [](const HLSLAttributedResourceType *ResType) {
4341 return ResType->getAttrs().ResourceDimension ==
4342 llvm::dxil::ResourceDimension::Unknown;
4343 }))
4344 return true;
4345
4346 auto *ResourceTy =
4347 TheCall->getArg(Arg: 0)->getType()->castAs<HLSLAttributedResourceType>();
4348
4349 // A UAV descriptor binds a single mip slice, so a RWTexture location has no
4350 // mip component to select, and TextureLoad on a UAV takes no offset.
4351 bool IsUAV =
4352 ResourceTy->getAttrs().ResourceClass == llvm::dxil::ResourceClass::UAV;
4353 if (IsUAV && S.checkArgCount(Call: TheCall, DesiredArgCount: 2))
4354 return true;
4355
4356 // Check the location: int3 for Texture2D and int4 for Texture2DArray, which
4357 // both carry a trailing mip level; int2 and int3 for the RWTexture forms,
4358 // which do not.
4359 unsigned ResourceDim =
4360 getResourceDimensions(Dim: ResourceTy->getAttrs().ResourceDimension);
4361 unsigned LocationDim = ResourceDim + (ResourceTy->getAttrs().IsArray ? 1 : 0);
4362 if (!IsUAV)
4363 ++LocationDim;
4364 if (CheckArgTypeMatches(
4365 S: &S, Arg: TheCall->getArg(Arg: 1),
4366 ExpectedType: getVectorOrScalarType(S, BaseType: S.Context.IntTy, Count: LocationDim)))
4367 return true;
4368
4369 // Check the offset operand (int2 for 2D textures; no array slice).
4370 if (TheCall->getNumArgs() > 2) {
4371 if (CheckArgTypeMatches(
4372 S: &S, Arg: TheCall->getArg(Arg: 2),
4373 ExpectedType: getVectorOrScalarType(S, BaseType: S.Context.IntTy, Count: ResourceDim)))
4374 return true;
4375 }
4376
4377 TheCall->setType(ResourceTy->getContainedType());
4378 return false;
4379}
4380
4381static bool CheckLoadMSBuiltin(Sema &S, CallExpr *TheCall) {
4382 if (S.checkArgCountRange(Call: TheCall, MinArgCount: 3, MaxArgCount: 4))
4383 return true;
4384
4385 // Check the multisampled texture handle.
4386 if (CheckResourceHandle(S: &S, TheCall, ArgIndex: 0,
4387 Check: [](const HLSLAttributedResourceType *ResType) {
4388 return !ResType->isMultiSampled();
4389 }))
4390 return true;
4391
4392 auto *ResourceTy =
4393 TheCall->getArg(Arg: 0)->getType()->castAs<HLSLAttributedResourceType>();
4394
4395 // Check the location (int2 for Texture2DMS, int3 for Texture2DMSArray).
4396 // Unlike Load on regular textures, there is no mip/LOD component.
4397 unsigned ResourceDim =
4398 getResourceDimensions(Dim: ResourceTy->getAttrs().ResourceDimension);
4399 unsigned LocationDim = ResourceDim + (ResourceTy->getAttrs().IsArray ? 1 : 0);
4400 if (CheckArgTypeMatches(
4401 S: &S, Arg: TheCall->getArg(Arg: 1),
4402 ExpectedType: getVectorOrScalarType(S, BaseType: S.Context.IntTy, Count: LocationDim)))
4403 return true;
4404
4405 // Check the sample index operand (scalar int).
4406 if (CheckArgTypeMatches(S: &S, Arg: TheCall->getArg(Arg: 2), ExpectedType: S.Context.IntTy))
4407 return true;
4408
4409 // Check the offset operand (int2 for 2D textures; no array slice).
4410 if (TheCall->getNumArgs() > 3) {
4411 if (CheckArgTypeMatches(
4412 S: &S, Arg: TheCall->getArg(Arg: 3),
4413 ExpectedType: getVectorOrScalarType(S, BaseType: S.Context.IntTy, Count: ResourceDim)))
4414 return true;
4415 }
4416
4417 TheCall->setType(ResourceTy->getContainedType());
4418 return false;
4419}
4420
4421static bool CheckSamplingBuiltin(Sema &S, CallExpr *TheCall, SampleKind Kind) {
4422 unsigned MinArgs, MaxArgs;
4423 if (Kind == SampleKind::Sample) {
4424 MinArgs = 3;
4425 MaxArgs = 5;
4426 } else if (Kind == SampleKind::Bias) {
4427 MinArgs = 4;
4428 MaxArgs = 6;
4429 } else if (Kind == SampleKind::Grad) {
4430 MinArgs = 5;
4431 MaxArgs = 7;
4432 } else if (Kind == SampleKind::Level) {
4433 MinArgs = 4;
4434 MaxArgs = 5;
4435 } else if (Kind == SampleKind::Cmp) {
4436 MinArgs = 4;
4437 MaxArgs = 6;
4438 } else {
4439 assert(Kind == SampleKind::CmpLevelZero);
4440 MinArgs = 4;
4441 MaxArgs = 5;
4442 }
4443
4444 if (S.checkArgCountRange(Call: TheCall, MinArgCount: MinArgs, MaxArgCount: MaxArgs))
4445 return true;
4446
4447 if (CheckTextureSamplerAndLocation(S, TheCall))
4448 return true;
4449
4450 const HLSLAttributedResourceType *ResourceTy =
4451 TheCall->getArg(Arg: 0)->getType()->castAs<HLSLAttributedResourceType>();
4452 unsigned ExpectedDim =
4453 getResourceDimensions(Dim: ResourceTy->getAttrs().ResourceDimension);
4454
4455 unsigned NextIdx = 3;
4456 if (Kind == SampleKind::Bias || Kind == SampleKind::Level ||
4457 Kind == SampleKind::Cmp || Kind == SampleKind::CmpLevelZero) {
4458 // Check the bias, lod level, or compare value, depending on the kind.
4459 // All of them must be a scalar float value.
4460 if (CheckScalarFloatOperand(S, TheCall, ArgIndex: NextIdx))
4461 return true;
4462 NextIdx++;
4463 } else if (Kind == SampleKind::Grad) {
4464 QualType GradTy = getVectorOrScalarType(S, BaseType: S.Context.FloatTy, Count: ExpectedDim);
4465
4466 // Check the DDX operand.
4467 if (CheckArgTypeMatches(S: &S, Arg: TheCall->getArg(Arg: NextIdx), ExpectedType: GradTy))
4468 return true;
4469
4470 // Check the DDY operand.
4471 if (CheckArgTypeMatches(S: &S, Arg: TheCall->getArg(Arg: NextIdx + 1), ExpectedType: GradTy))
4472 return true;
4473 NextIdx += 2;
4474 }
4475
4476 // Check the offset operand (if applicable).
4477 if (hasResourceOffset(Dim: ResourceTy->getAttrs().ResourceDimension) &&
4478 TheCall->getNumArgs() > NextIdx) {
4479 if (CheckArgTypeMatches(
4480 S: &S, Arg: TheCall->getArg(Arg: NextIdx),
4481 ExpectedType: getVectorOrScalarType(S, BaseType: S.Context.IntTy, Count: ExpectedDim)))
4482 return true;
4483 NextIdx++;
4484 }
4485
4486 // Check the clamp operand.
4487 if (Kind != SampleKind::Level && Kind != SampleKind::CmpLevelZero &&
4488 TheCall->getNumArgs() > NextIdx) {
4489 if (CheckScalarFloatOperand(S, TheCall, ArgIndex: NextIdx))
4490 return true;
4491 }
4492
4493 assert(ResourceTy->hasContainedType() &&
4494 "Expecting a contained type for resource with a dimension "
4495 "attribute.");
4496 QualType ReturnType = ResourceTy->getContainedType();
4497
4498 if (CheckNoDoubleElementType(S, TheCall, ContainedType: ReturnType,
4499 DefaultName: getSampleMethodName(Kind)))
4500 return true;
4501
4502 if (CheckIntegerElementTypeShaderModel(S, TheCall, ContainedType: ReturnType, Kind))
4503 return true;
4504
4505 if (Kind == SampleKind::Cmp || Kind == SampleKind::CmpLevelZero) {
4506 if (!ReturnType->hasFloatingRepresentation()) {
4507 S.Diag(Loc: TheCall->getBeginLoc(), DiagID: diag::err_hlsl_samplecmp_requires_float);
4508 return true;
4509 }
4510 ReturnType = S.Context.FloatTy;
4511 }
4512 TheCall->setType(ReturnType);
4513
4514 return false;
4515}
4516
4517/// The `dest` types an interlocked operation accepts. Float is 32-bit only.
4518enum class InterlockedDest { Int, IntOrFloat, Float };
4519
4520/// Check a call to an HLSL interlocked builtin. The builtins are variadic, so
4521/// this is the only check a direct call gets. Overload resolution checks the
4522/// calls that come through the `InterlockedOp` overload sets.
4523static bool CheckInterlockedBuiltin(Sema &S, CallExpr *TheCall,
4524 unsigned MinArgs, unsigned MaxArgs,
4525 InterlockedDest Dest,
4526 bool ReportsOriginalValue) {
4527 if (MinArgs == MaxArgs) {
4528 if (S.checkArgCount(Call: TheCall, DesiredArgCount: MinArgs))
4529 return true;
4530 } else if (TheCall->getNumArgs() < MinArgs) {
4531 S.Diag(Loc: TheCall->getEndLoc(), DiagID: diag::err_typecheck_call_too_few_args_at_least)
4532 << /*callee_type=*/0 << /*min_arg_count=*/MinArgs
4533 << TheCall->getNumArgs() << /*is_non_object=*/0
4534 << TheCall->getSourceRange();
4535 return true;
4536 } else if (S.checkArgCountAtMost(Call: TheCall, MaxArgCount: MaxArgs)) {
4537 return true;
4538 }
4539
4540 QualType DestTy = TheCall->getArg(Arg: 0)->getType().getUnqualifiedType();
4541 const bool DestIsOK =
4542 DestTy->isSpecificBuiltinType(K: BuiltinType::Float)
4543 ? Dest != InterlockedDest::Int
4544 : Dest != InterlockedDest::Float && DestTy->isIntegerType();
4545 if (!DestIsOK) {
4546 S.Diag(Loc: TheCall->getArg(Arg: 0)->getBeginLoc(),
4547 DiagID: diag::err_builtin_invalid_arg_type)
4548 << /*ordinal=*/1 << /*scalar*/ 1
4549 << /*integer*/ (Dest == InterlockedDest::Float ? 0 : 1)
4550 << /*32 bit floating-point*/ (Dest == InterlockedDest::Int ? 0 : 3)
4551 << DestTy;
4552 return true;
4553 }
4554
4555 // 64-bit interlocked ops require SM 6.6 on DXIL. The synthesized wrapper
4556 // methods (e.g. RWByteAddressBuffer::InterlockedAdd64) are only declared on
4557 // SM 6.6+, so this defensive check only fires for direct builtin calls; skip
4558 // synthetic invocations (invalid source location).
4559 const TargetInfo &TI = S.Context.getTargetInfo();
4560 if (TheCall->getBeginLoc().isValid() &&
4561 TI.getTriple().getArch() == llvm::Triple::dxil &&
4562 S.Context.getTypeSize(T: DestTy) == 64 &&
4563 TI.getPlatformMinVersion() < VersionTuple(6, 6)) {
4564 S.Diag(Loc: TheCall->getBeginLoc(), DiagID: diag::err_hlsl_builtin_requires_sm)
4565 << TheCall->getDirectCallee() << VersionTuple(6, 6).getAsString();
4566 return true;
4567 }
4568
4569 if (CheckModifiableLValue(S: &S, TheCall, ArgIndex: 0))
4570 return true;
4571
4572 if (CheckArgAddrSpaceOneOf(S: &S, TheCall, ArgIndex: 0,
4573 AllowedSpaces: {LangAS::hlsl_groupshared, LangAS::hlsl_device}))
4574 return true;
4575
4576 // Every argument after `dest` has the destination's type.
4577 for (unsigned I = 1, E = TheCall->getNumArgs(); I != E; ++I)
4578 if (CheckArgTypeMatches(S: &S, Arg: TheCall->getArg(Arg: I), ExpectedType: DestTy))
4579 return true;
4580
4581 // Operations that report the previous value write it back through their last
4582 // argument.
4583 const unsigned NumArgs = TheCall->getNumArgs();
4584 if (ReportsOriginalValue && NumArgs == MaxArgs &&
4585 CheckModifiableLValue(S: &S, TheCall, ArgIndex: NumArgs - 1))
4586 return true;
4587
4588 TheCall->setType(S.Context.VoidTy);
4589 return false;
4590}
4591
4592// Note: returning true in this case results in CheckBuiltinFunctionCall
4593// returning an ExprError
4594bool SemaHLSL::CheckBuiltinFunctionCall(unsigned BuiltinID, CallExpr *TheCall) {
4595 switch (BuiltinID) {
4596 case Builtin::BI__builtin_hlsl_barrier: {
4597 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 2))
4598 return true;
4599
4600 if (SemaRef.Context.getTargetInfo().getTriple().getArch() !=
4601 llvm::Triple::dxil) {
4602 SemaRef.Diag(Loc: TheCall->getExprLoc(), DiagID: diag::err_hlsl_dxil_only)
4603 << "Barrier";
4604 return true;
4605 }
4606
4607 Expr *MemoryArg = TheCall->getArg(Arg: 0);
4608 if (MemoryArg->getType()->isUnsignedIntegerType()) {
4609 std::optional<llvm::APSInt> MemoryFlags =
4610 MemoryArg->getIntegerConstantExpr(Ctx: SemaRef.Context);
4611 if (!MemoryFlags) {
4612 SemaRef.Diag(Loc: MemoryArg->getExprLoc(),
4613 DiagID: diag::err_constant_integer_arg_type)
4614 << "Barrier";
4615 return true;
4616 }
4617 if ((MemoryFlags->getZExtValue() &
4618 ~barrierFlagValue(Flag: BarrierMemoryTypeFlag::ValidMask)) != 0) {
4619 SemaRef.Diag(Loc: MemoryArg->getExprLoc(),
4620 DiagID: diag::err_hlsl_invalid_barrier_memory_flags);
4621 return true;
4622 }
4623 } else {
4624 const HLSLAttributedResourceType *ResTy =
4625 HLSLAttributedResourceType::findHandleTypeOnResource(
4626 RT: MemoryArg->getType().getTypePtr());
4627 if (!ResTy) {
4628 SemaRef.Diag(Loc: MemoryArg->getExprLoc(),
4629 DiagID: diag::err_typecheck_expect_hlsl_resource)
4630 << MemoryArg->getType();
4631 return true;
4632 }
4633 if (ResTy->getAttrs().ResourceClass != ResourceClass::UAV) {
4634 SemaRef.Diag(Loc: MemoryArg->getExprLoc(),
4635 DiagID: diag::err_invalid_hlsl_resource_type)
4636 << MemoryArg->getType();
4637 return true;
4638 }
4639 }
4640
4641 Expr *SemanticArg = TheCall->getArg(Arg: 1);
4642 std::optional<llvm::APSInt> SemanticFlags =
4643 SemanticArg->getIntegerConstantExpr(Ctx: SemaRef.Context);
4644 if (!SemanticFlags) {
4645 SemaRef.Diag(Loc: SemanticArg->getExprLoc(),
4646 DiagID: diag::err_constant_integer_arg_type)
4647 << "Barrier";
4648 return true;
4649 }
4650 if ((SemanticFlags->getZExtValue() &
4651 ~barrierFlagValue(Flag: BarrierSemanticFlag::ValidMask)) != 0) {
4652 SemaRef.Diag(Loc: SemanticArg->getExprLoc(),
4653 DiagID: diag::err_hlsl_invalid_barrier_semantic_flags);
4654 return true;
4655 }
4656
4657 TheCall->setType(SemaRef.Context.VoidTy);
4658 break;
4659 }
4660 case Builtin::BI__builtin_hlsl_adduint64: {
4661 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 2))
4662 return true;
4663
4664 if (CheckAllArgTypesAreCorrect(S: &SemaRef, TheCall,
4665 Check: CheckUnsignedIntVecRepresentation))
4666 return true;
4667
4668 // ensure arg integers are 32-bits
4669 if (CheckExpectedBitWidth(S: &SemaRef, TheCall, ArgOrdinal: 0, Width: 32))
4670 return true;
4671
4672 // ensure both args are vectors of total bit size of a multiple of 64
4673 auto *VTy = TheCall->getArg(Arg: 0)->getType()->getAs<VectorType>();
4674 int NumElementsArg = VTy->getNumElements();
4675 if (NumElementsArg != 2 && NumElementsArg != 4) {
4676 SemaRef.Diag(Loc: TheCall->getBeginLoc(), DiagID: diag::err_vector_incorrect_bit_count)
4677 << 1 /*a multiple of*/ << 64 << NumElementsArg * 32;
4678 return true;
4679 }
4680
4681 // ensure first arg and second arg have the same type
4682 if (CheckAllArgsHaveSameType(S: &SemaRef, TheCall))
4683 return true;
4684
4685 ExprResult A = TheCall->getArg(Arg: 0);
4686 QualType ArgTyA = A.get()->getType();
4687 // return type is the same as the input type
4688 TheCall->setType(ArgTyA);
4689 break;
4690 }
4691 case Builtin::BI__builtin_hlsl_resource_getpointer: {
4692 if (SemaRef.checkArgCountRange(Call: TheCall, MinArgCount: 1, MaxArgCount: 2) ||
4693 CheckResourceHandle(S: &SemaRef, TheCall, ArgIndex: 0) ||
4694 (TheCall->getNumArgs() == 2 && CheckIndexType(S: &SemaRef, TheCall, IndexArgIndex: 1)))
4695 return true;
4696
4697 auto *ResourceTy =
4698 TheCall->getArg(Arg: 0)->getType()->castAs<HLSLAttributedResourceType>();
4699 QualType ContainedTy = ResourceTy->getContainedType();
4700 auto ReturnType = SemaRef.Context.getAddrSpaceQualType(
4701 T: ContainedTy,
4702 AddressSpace: getLangASFromResourceClass(RC: ResourceTy->getAttrs().ResourceClass));
4703 ReturnType = SemaRef.Context.getPointerType(T: ReturnType);
4704 TheCall->setType(ReturnType);
4705
4706 break;
4707 }
4708 case Builtin::BI__builtin_hlsl_resource_getpointer_typed: {
4709 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 3) ||
4710 CheckResourceHandle(S: &SemaRef, TheCall, ArgIndex: 0) ||
4711 CheckIndexType(S: &SemaRef, TheCall, IndexArgIndex: 1))
4712 return true;
4713
4714 QualType ElementTy = TheCall->getArg(Arg: 2)->getType();
4715 assert(ElementTy->isPointerType() &&
4716 "expected pointer type for second argument");
4717 ElementTy = ElementTy->getPointeeType();
4718
4719 // Reject array types
4720 if (ElementTy->isArrayType())
4721 return SemaRef.Diag(
4722 Loc: cast<FunctionDecl>(Val: SemaRef.CurContext)->getPointOfInstantiation(),
4723 DiagID: diag::err_invalid_use_of_array_type);
4724
4725 auto *ResourceTy =
4726 TheCall->getArg(Arg: 0)->getType()->castAs<HLSLAttributedResourceType>();
4727 auto ReturnType = SemaRef.Context.getAddrSpaceQualType(
4728 T: ElementTy,
4729 AddressSpace: getLangASFromResourceClass(RC: ResourceTy->getAttrs().ResourceClass));
4730 ReturnType = SemaRef.Context.getPointerType(T: ReturnType);
4731 TheCall->setType(ReturnType);
4732
4733 break;
4734 }
4735 case Builtin::BI__builtin_hlsl_transpose_if_memory_is_row_major: {
4736 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 2) ||
4737 CheckArgTypeMatches(S: &SemaRef, Arg: TheCall->getArg(Arg: 1),
4738 ExpectedType: SemaRef.getASTContext().IntTy))
4739 return true;
4740
4741 TheCall->setType(TheCall->getArg(Arg: 0)->getType());
4742
4743 break;
4744 }
4745 case Builtin::BI__builtin_hlsl_resource_load_with_status: {
4746 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 3) ||
4747 CheckResourceHandle(S: &SemaRef, TheCall, ArgIndex: 0) ||
4748 CheckArgTypeMatches(S: &SemaRef, Arg: TheCall->getArg(Arg: 1),
4749 ExpectedType: SemaRef.getASTContext().UnsignedIntTy) ||
4750 CheckArgTypeMatches(S: &SemaRef, Arg: TheCall->getArg(Arg: 2),
4751 ExpectedType: SemaRef.getASTContext().UnsignedIntTy) ||
4752 CheckModifiableLValue(S: &SemaRef, TheCall, ArgIndex: 2))
4753 return true;
4754
4755 auto *ResourceTy =
4756 TheCall->getArg(Arg: 0)->getType()->castAs<HLSLAttributedResourceType>();
4757 QualType ReturnType = ResourceTy->getContainedType();
4758 TheCall->setType(ReturnType);
4759
4760 break;
4761 }
4762 case Builtin::BI__builtin_hlsl_resource_load_with_status_typed: {
4763 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 4) ||
4764 CheckResourceHandle(S: &SemaRef, TheCall, ArgIndex: 0) ||
4765 CheckArgTypeMatches(S: &SemaRef, Arg: TheCall->getArg(Arg: 1),
4766 ExpectedType: SemaRef.getASTContext().UnsignedIntTy) ||
4767 CheckArgTypeMatches(S: &SemaRef, Arg: TheCall->getArg(Arg: 2),
4768 ExpectedType: SemaRef.getASTContext().UnsignedIntTy) ||
4769 CheckModifiableLValue(S: &SemaRef, TheCall, ArgIndex: 2))
4770 return true;
4771
4772 QualType ReturnType = TheCall->getArg(Arg: 3)->getType();
4773 assert(ReturnType->isPointerType() &&
4774 "expected pointer type for second argument");
4775 ReturnType = ReturnType->getPointeeType();
4776
4777 // Reject array types
4778 if (ReturnType->isArrayType())
4779 return SemaRef.Diag(
4780 Loc: cast<FunctionDecl>(Val: SemaRef.CurContext)->getPointOfInstantiation(),
4781 DiagID: diag::err_invalid_use_of_array_type);
4782
4783 TheCall->setType(ReturnType);
4784
4785 break;
4786 }
4787 case Builtin::BI__builtin_hlsl_resource_load_level:
4788 return CheckLoadLevelBuiltin(S&: SemaRef, TheCall);
4789 case Builtin::BI__builtin_hlsl_resource_load_ms:
4790 return CheckLoadMSBuiltin(S&: SemaRef, TheCall);
4791 case Builtin::BI__builtin_hlsl_resource_sample:
4792 return CheckSamplingBuiltin(S&: SemaRef, TheCall, Kind: SampleKind::Sample);
4793 case Builtin::BI__builtin_hlsl_resource_sample_bias:
4794 return CheckSamplingBuiltin(S&: SemaRef, TheCall, Kind: SampleKind::Bias);
4795 case Builtin::BI__builtin_hlsl_resource_sample_grad:
4796 return CheckSamplingBuiltin(S&: SemaRef, TheCall, Kind: SampleKind::Grad);
4797 case Builtin::BI__builtin_hlsl_resource_sample_level:
4798 return CheckSamplingBuiltin(S&: SemaRef, TheCall, Kind: SampleKind::Level);
4799 case Builtin::BI__builtin_hlsl_resource_sample_cmp:
4800 return CheckSamplingBuiltin(S&: SemaRef, TheCall, Kind: SampleKind::Cmp);
4801 case Builtin::BI__builtin_hlsl_resource_sample_cmp_level_zero:
4802 return CheckSamplingBuiltin(S&: SemaRef, TheCall, Kind: SampleKind::CmpLevelZero);
4803 case Builtin::BI__builtin_hlsl_resource_calculate_lod:
4804 case Builtin::BI__builtin_hlsl_resource_calculate_lod_unclamped:
4805 return CheckCalculateLodBuiltin(S&: SemaRef, TheCall);
4806 case Builtin::BI__builtin_hlsl_resource_gather:
4807 return CheckGatherBuiltin(S&: SemaRef, TheCall, /*IsCmp=*/false);
4808 case Builtin::BI__builtin_hlsl_resource_gather_cmp:
4809 return CheckGatherBuiltin(S&: SemaRef, TheCall, /*IsCmp=*/true);
4810 case Builtin::BI__builtin_hlsl_resource_uninitializedhandle: {
4811 assert(TheCall->getNumArgs() == 1 && "expected 1 arg");
4812 // Update return type to be the attributed resource type from arg0.
4813 QualType ResourceTy = TheCall->getArg(Arg: 0)->getType();
4814 TheCall->setType(ResourceTy);
4815 break;
4816 }
4817 case Builtin::BI__builtin_hlsl_resource_handlefrombinding: {
4818 assert(TheCall->getNumArgs() == 6 && "expected 6 args");
4819 // Update return type to be the attributed resource type from arg0.
4820 QualType ResourceTy = TheCall->getArg(Arg: 0)->getType();
4821 TheCall->setType(ResourceTy);
4822 break;
4823 }
4824 case Builtin::BI__builtin_hlsl_resource_handlefromimplicitbinding: {
4825 assert(TheCall->getNumArgs() == 6 && "expected 6 args");
4826 // Update return type to be the attributed resource type from arg0.
4827 QualType ResourceTy = TheCall->getArg(Arg: 0)->getType();
4828 TheCall->setType(ResourceTy);
4829 break;
4830 }
4831 case Builtin::BI__builtin_hlsl_resource_counterhandlefromimplicitbinding: {
4832 assert(TheCall->getNumArgs() == 3 && "expected 3 args");
4833 // Update return type to be the attributed resource type from arg0
4834 // with added IsCounter flag.
4835 QualType MainHandleTy = TheCall->getArg(Arg: 0)->getType();
4836 QualType CounterHandleTy =
4837 createCounterHandleType(AST&: SemaRef.getASTContext(), MainHandleTy);
4838 TheCall->setType(CounterHandleTy);
4839 break;
4840 }
4841 case Builtin::BI__builtin_hlsl_resource_handlefromheap: {
4842 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 2) ||
4843 CheckResourceHandle(S: &SemaRef, TheCall, ArgIndex: 0) ||
4844 CheckArgTypeMatches(S: &SemaRef, Arg: TheCall->getArg(Arg: 1),
4845 ExpectedType: SemaRef.getASTContext().UnsignedIntTy))
4846 return true;
4847
4848 // Update return type to be the attributed resource type from arg0.
4849 QualType ResourceTy = TheCall->getArg(Arg: 0)->getType();
4850 TheCall->setType(ResourceTy);
4851 break;
4852 }
4853 case Builtin::BI__builtin_hlsl_resource_counterhandlefromheap: {
4854 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 1) ||
4855 CheckResourceHandle(S: &SemaRef, TheCall, ArgIndex: 0))
4856 return true;
4857 // Update return type to be the attributed resource type from arg0
4858 // with added IsCounter flag.
4859 QualType MainHandleTy = TheCall->getArg(Arg: 0)->getType();
4860 QualType CounterHandleTy =
4861 createCounterHandleType(AST&: SemaRef.getASTContext(), MainHandleTy);
4862 TheCall->setType(CounterHandleTy);
4863 break;
4864 }
4865 case Builtin::BI__builtin_hlsl_and:
4866 case Builtin::BI__builtin_hlsl_or: {
4867 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 2))
4868 return true;
4869 if (CheckScalarOrVectorOrMatrix(S: &SemaRef, TheCall, Scalar: getASTContext().BoolTy,
4870 ArgIndex: 0))
4871 return true;
4872 if (CheckAllArgsHaveSameType(S: &SemaRef, TheCall))
4873 return true;
4874
4875 ExprResult A = TheCall->getArg(Arg: 0);
4876 QualType ArgTyA = A.get()->getType();
4877 // return type is the same as the input type
4878 TheCall->setType(ArgTyA);
4879 break;
4880 }
4881 case Builtin::BI__builtin_hlsl_all:
4882 case Builtin::BI__builtin_hlsl_any: {
4883 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 1))
4884 return true;
4885 if (CheckAnyScalarOrVector(S: &SemaRef, TheCall, ArgIndex: 0))
4886 return true;
4887 break;
4888 }
4889 case Builtin::BI__builtin_hlsl_asdouble: {
4890 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 2))
4891 return true;
4892 if (CheckScalarOrVector(
4893 S: &SemaRef, TheCall,
4894 /*only check for uint*/ Scalar: SemaRef.Context.UnsignedIntTy,
4895 /* arg index */ ArgIndex: 0))
4896 return true;
4897 if (CheckScalarOrVector(
4898 S: &SemaRef, TheCall,
4899 /*only check for uint*/ Scalar: SemaRef.Context.UnsignedIntTy,
4900 /* arg index */ ArgIndex: 1))
4901 return true;
4902 if (CheckAllArgsHaveSameType(S: &SemaRef, TheCall))
4903 return true;
4904
4905 SetElementTypeAsReturnType(S: &SemaRef, TheCall, ReturnType: getASTContext().DoubleTy);
4906 break;
4907 }
4908 case Builtin::BI__builtin_hlsl_elementwise_clamp: {
4909 if (SemaRef.BuiltinElementwiseTernaryMath(
4910 TheCall, /*ArgTyRestr=*/
4911 Sema::EltwiseBuiltinArgTyRestriction::None))
4912 return true;
4913 break;
4914 }
4915 case Builtin::BI__builtin_hlsl_dot: {
4916 // arg count is checked by BuiltinVectorToScalarMath
4917 if (SemaRef.BuiltinVectorToScalarMath(TheCall))
4918 return true;
4919 if (CheckAllArgTypesAreCorrect(S: &SemaRef, TheCall, Check: CheckNoDoubleVectors))
4920 return true;
4921 break;
4922 }
4923 case Builtin::BI__builtin_hlsl_elementwise_firstbithigh:
4924 case Builtin::BI__builtin_hlsl_elementwise_firstbitlow: {
4925 if (SemaRef.PrepareBuiltinElementwiseMathOneArgCall(TheCall))
4926 return true;
4927
4928 const Expr *Arg = TheCall->getArg(Arg: 0);
4929 QualType ArgTy = Arg->getType();
4930 QualType EltTy = ArgTy;
4931
4932 QualType ResTy = SemaRef.Context.UnsignedIntTy;
4933
4934 if (auto *VecTy = EltTy->getAs<VectorType>()) {
4935 EltTy = VecTy->getElementType();
4936 ResTy = SemaRef.Context.getExtVectorType(VectorType: ResTy, NumElts: VecTy->getNumElements());
4937 }
4938
4939 if (!EltTy->isIntegerType()) {
4940 Diag(Loc: Arg->getBeginLoc(), DiagID: diag::err_builtin_invalid_arg_type)
4941 << 1 << /* scalar or vector of */ 5 << /* integer ty */ 1
4942 << /* no fp */ 0 << ArgTy;
4943 return true;
4944 }
4945
4946 TheCall->setType(ResTy);
4947 break;
4948 }
4949 case Builtin::BI__builtin_hlsl_select: {
4950 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 3))
4951 return true;
4952 if (CheckScalarOrVectorOrMatrix(S: &SemaRef, TheCall, Scalar: getASTContext().BoolTy,
4953 ArgIndex: 0))
4954 return true;
4955 QualType ArgTy = TheCall->getArg(Arg: 0)->getType();
4956 if (ArgTy->isBooleanType() && CheckBoolSelect(S: &SemaRef, TheCall))
4957 return true;
4958 auto *VTy = ArgTy->getAs<VectorType>();
4959 if (VTy && VTy->getElementType()->isBooleanType() &&
4960 CheckVectorSelect(S: &SemaRef, TheCall))
4961 return true;
4962 auto *MTy = ArgTy->getAs<ConstantMatrixType>();
4963 if (MTy && MTy->getElementType()->isBooleanType() &&
4964 CheckMatrixSelect(S: &SemaRef, TheCall))
4965 return true;
4966 break;
4967 }
4968 case Builtin::BI__builtin_hlsl_elementwise_saturate:
4969 case Builtin::BI__builtin_hlsl_elementwise_rcp: {
4970 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 1))
4971 return true;
4972 if (!TheCall->getArg(Arg: 0)
4973 ->getType()
4974 ->hasFloatingRepresentation()) // half or float or double
4975 return SemaRef.Diag(Loc: TheCall->getArg(Arg: 0)->getBeginLoc(),
4976 DiagID: diag::err_builtin_invalid_arg_type)
4977 << /* ordinal */ 1 << /* scalar or vector */ 5 << /* no int */ 0
4978 << /* fp */ 1 << TheCall->getArg(Arg: 0)->getType();
4979 if (SemaRef.PrepareBuiltinElementwiseMathOneArgCall(TheCall))
4980 return true;
4981 break;
4982 }
4983 case Builtin::BI__builtin_hlsl_elementwise_rsqrt:
4984 case Builtin::BI__builtin_hlsl_elementwise_frac:
4985 case Builtin::BI__builtin_hlsl_elementwise_ddx_coarse:
4986 case Builtin::BI__builtin_hlsl_elementwise_ddy_coarse:
4987 case Builtin::BI__builtin_hlsl_elementwise_ddx_fine:
4988 case Builtin::BI__builtin_hlsl_elementwise_ddy_fine: {
4989 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 1))
4990 return true;
4991 if (CheckAllArgTypesAreCorrect(S: &SemaRef, TheCall,
4992 Check: CheckFloatOrHalfRepresentation))
4993 return true;
4994 if (SemaRef.PrepareBuiltinElementwiseMathOneArgCall(TheCall))
4995 return true;
4996 break;
4997 }
4998 case Builtin::BI__builtin_hlsl_elementwise_isfinite:
4999 case Builtin::BI__builtin_hlsl_elementwise_isinf:
5000 case Builtin::BI__builtin_hlsl_elementwise_isnan: {
5001 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 1))
5002 return true;
5003 if (CheckAllArgTypesAreCorrect(S: &SemaRef, TheCall,
5004 Check: CheckFloatOrHalfRepresentation))
5005 return true;
5006 if (SemaRef.PrepareBuiltinElementwiseMathOneArgCall(TheCall))
5007 return true;
5008 SetElementTypeAsReturnType(S: &SemaRef, TheCall, ReturnType: getASTContext().BoolTy);
5009 break;
5010 }
5011 case Builtin::BI__builtin_hlsl_mad: {
5012 if (SemaRef.BuiltinElementwiseTernaryMath(
5013 TheCall, /*ArgTyRestr=*/
5014 Sema::EltwiseBuiltinArgTyRestriction::None))
5015 return true;
5016 break;
5017 }
5018 case Builtin::BI__builtin_hlsl_mul: {
5019 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 2))
5020 return true;
5021
5022 Expr *Arg0 = TheCall->getArg(Arg: 0);
5023 Expr *Arg1 = TheCall->getArg(Arg: 1);
5024 QualType Ty0 = Arg0->getType();
5025 QualType Ty1 = Arg1->getType();
5026
5027 auto getElemType = [](QualType T) -> QualType {
5028 if (const auto *VTy = T->getAs<VectorType>())
5029 return VTy->getElementType();
5030 if (const auto *MTy = T->getAs<ConstantMatrixType>())
5031 return MTy->getElementType();
5032 return T;
5033 };
5034
5035 QualType EltTy0 = getElemType(Ty0);
5036
5037 bool IsVec0 = Ty0->isVectorType();
5038 bool IsMat0 = Ty0->isConstantMatrixType();
5039 bool IsVec1 = Ty1->isVectorType();
5040 bool IsMat1 = Ty1->isConstantMatrixType();
5041
5042 QualType RetTy;
5043
5044 if (IsVec0 && IsMat1) {
5045 auto *MatTy = Ty1->castAs<ConstantMatrixType>();
5046 RetTy = getASTContext().getExtVectorType(VectorType: EltTy0, NumElts: MatTy->getNumColumns());
5047 } else if (IsMat0 && IsVec1) {
5048 auto *MatTy = Ty0->castAs<ConstantMatrixType>();
5049 RetTy = getASTContext().getExtVectorType(VectorType: EltTy0, NumElts: MatTy->getNumRows());
5050 } else {
5051 assert(IsMat0 && IsMat1);
5052 auto *MatTy0 = Ty0->castAs<ConstantMatrixType>();
5053 auto *MatTy1 = Ty1->castAs<ConstantMatrixType>();
5054 RetTy = getASTContext().getConstantMatrixType(
5055 ElementType: EltTy0, NumRows: MatTy0->getNumRows(), NumColumns: MatTy1->getNumColumns());
5056 }
5057
5058 TheCall->setType(RetTy);
5059 break;
5060 }
5061 case Builtin::BI__builtin_elementwise_fma: {
5062 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 3) ||
5063 CheckAllArgsHaveSameType(S: &SemaRef, TheCall)) {
5064 return true;
5065 }
5066
5067 if (CheckAllArgTypesAreCorrect(S: &SemaRef, TheCall,
5068 Check: CheckAnyDoubleRepresentation))
5069 return true;
5070
5071 ExprResult A = TheCall->getArg(Arg: 0);
5072 QualType ArgTyA = A.get()->getType();
5073 // return type is the same as input type
5074 TheCall->setType(ArgTyA);
5075 break;
5076 }
5077 case Builtin::BI__builtin_hlsl_transpose: {
5078 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 1))
5079 return true;
5080
5081 Expr *Arg = TheCall->getArg(Arg: 0);
5082 QualType ArgTy = Arg->getType();
5083
5084 const auto *MatTy = ArgTy->getAs<ConstantMatrixType>();
5085 if (!MatTy) {
5086 SemaRef.Diag(Loc: Arg->getBeginLoc(), DiagID: diag::err_builtin_invalid_arg_type)
5087 << 1 << /* matrix */ 3 << /* no int */ 0 << /* no fp */ 0 << ArgTy;
5088 return true;
5089 }
5090
5091 QualType RetTy = getASTContext().getConstantMatrixType(
5092 ElementType: MatTy->getElementType(), NumRows: MatTy->getNumColumns(), NumColumns: MatTy->getNumRows());
5093 TheCall->setType(RetTy);
5094 break;
5095 }
5096 case Builtin::BI__builtin_hlsl_elementwise_sign: {
5097 if (SemaRef.PrepareBuiltinElementwiseMathOneArgCall(TheCall))
5098 return true;
5099 if (CheckAllArgTypesAreCorrect(S: &SemaRef, TheCall,
5100 Check: CheckFloatingOrIntRepresentation))
5101 return true;
5102 SetElementTypeAsReturnType(S: &SemaRef, TheCall, ReturnType: getASTContext().IntTy);
5103 break;
5104 }
5105 case Builtin::BI__builtin_hlsl_wave_active_all_equal: {
5106 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 1))
5107 return true;
5108
5109 // Ensure input expr type is a scalar/vector
5110 if (CheckAnyScalarOrVector(S: &SemaRef, TheCall, ArgIndex: 0))
5111 return true;
5112
5113 QualType InputTy = TheCall->getArg(Arg: 0)->getType();
5114 ASTContext &Ctx = getASTContext();
5115
5116 QualType RetTy;
5117
5118 // If vector, construct bool vector of same size
5119 if (const auto *VecTy = InputTy->getAs<ExtVectorType>()) {
5120 unsigned NumElts = VecTy->getNumElements();
5121 RetTy = Ctx.getExtVectorType(VectorType: Ctx.BoolTy, NumElts);
5122 } else {
5123 // Scalar case
5124 RetTy = Ctx.BoolTy;
5125 }
5126
5127 TheCall->setType(RetTy);
5128 break;
5129 }
5130 case Builtin::BI__builtin_hlsl_wave_active_max:
5131 case Builtin::BI__builtin_hlsl_wave_active_min:
5132 case Builtin::BI__builtin_hlsl_wave_active_sum:
5133 case Builtin::BI__builtin_hlsl_wave_active_product: {
5134 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 1))
5135 return true;
5136
5137 // Ensure input expr type is a scalar/vector and the same as the return type
5138 if (CheckAnyScalarOrVector(S: &SemaRef, TheCall, ArgIndex: 0))
5139 return true;
5140 if (CheckWaveActive(S: &SemaRef, TheCall))
5141 return true;
5142 ExprResult Expr = TheCall->getArg(Arg: 0);
5143 QualType ArgTyExpr = Expr.get()->getType();
5144 TheCall->setType(ArgTyExpr);
5145 break;
5146 }
5147 case Builtin::BI__builtin_hlsl_wave_active_bit_or:
5148 case Builtin::BI__builtin_hlsl_wave_active_bit_xor:
5149 case Builtin::BI__builtin_hlsl_wave_active_bit_and: {
5150 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 1))
5151 return true;
5152
5153 // Ensure input expr type is a scalar/vector
5154 if (CheckAnyScalarOrVector(S: &SemaRef, TheCall, ArgIndex: 0))
5155 return true;
5156
5157 if (CheckWaveActive(S: &SemaRef, TheCall))
5158 return true;
5159
5160 // Ensure the expr type is interpretable as a uint or vector<uint>
5161 ExprResult Expr = TheCall->getArg(Arg: 0);
5162 QualType ArgTyExpr = Expr.get()->getType();
5163 auto *VTy = ArgTyExpr->getAs<VectorType>();
5164 if (!(ArgTyExpr->isIntegerType() ||
5165 (VTy && VTy->getElementType()->isIntegerType()))) {
5166 SemaRef.Diag(Loc: TheCall->getArg(Arg: 0)->getBeginLoc(),
5167 DiagID: diag::err_builtin_invalid_arg_type)
5168 << ArgTyExpr << SemaRef.Context.UnsignedIntTy << 1 << 0 << 0;
5169 return true;
5170 }
5171
5172 // Ensure input expr type is the same as the return type
5173 TheCall->setType(ArgTyExpr);
5174 break;
5175 }
5176 case Builtin::BI__builtin_hlsl_interlocked_add:
5177 case Builtin::BI__builtin_hlsl_interlocked_and:
5178 case Builtin::BI__builtin_hlsl_interlocked_max:
5179 case Builtin::BI__builtin_hlsl_interlocked_min:
5180 case Builtin::BI__builtin_hlsl_interlocked_or:
5181 case Builtin::BI__builtin_hlsl_interlocked_xor:
5182 if (CheckInterlockedBuiltin(S&: SemaRef, TheCall, /*MinArgs=*/2, /*MaxArgs=*/3,
5183 Dest: InterlockedDest::Int,
5184 /*ReportsOriginalValue=*/true))
5185 return true;
5186 break;
5187 case Builtin::BI__builtin_hlsl_interlocked_exchange:
5188 if (CheckInterlockedBuiltin(S&: SemaRef, TheCall, /*MinArgs=*/3, /*MaxArgs=*/3,
5189 Dest: InterlockedDest::IntOrFloat,
5190 /*ReportsOriginalValue=*/true))
5191 return true;
5192 break;
5193 case Builtin::BI__builtin_hlsl_interlocked_compare_store:
5194 if (CheckInterlockedBuiltin(S&: SemaRef, TheCall, /*MinArgs=*/3, /*MaxArgs=*/3,
5195 Dest: InterlockedDest::Int,
5196 /*ReportsOriginalValue=*/false))
5197 return true;
5198 break;
5199 case Builtin::BI__builtin_hlsl_interlocked_compare_store_float_bitwise:
5200 if (CheckInterlockedBuiltin(S&: SemaRef, TheCall, /*MinArgs=*/3, /*MaxArgs=*/3,
5201 Dest: InterlockedDest::Float,
5202 /*ReportsOriginalValue=*/false))
5203 return true;
5204 break;
5205 case Builtin::BI__builtin_hlsl_interlocked_compare_exchange:
5206 if (CheckInterlockedBuiltin(S&: SemaRef, TheCall, /*MinArgs=*/4, /*MaxArgs=*/4,
5207 Dest: InterlockedDest::Int,
5208 /*ReportsOriginalValue=*/true))
5209 return true;
5210 break;
5211 case Builtin::BI__builtin_hlsl_interlocked_compare_exchange_float_bitwise:
5212 if (CheckInterlockedBuiltin(S&: SemaRef, TheCall, /*MinArgs=*/4, /*MaxArgs=*/4,
5213 Dest: InterlockedDest::Float,
5214 /*ReportsOriginalValue=*/true))
5215 return true;
5216 break;
5217 // Note these are llvm builtins that we want to catch invalid intrinsic
5218 // generation. Normal handling of these builtins will occur elsewhere.
5219 case Builtin::BI__builtin_elementwise_bitreverse: {
5220 // does not include a check for number of arguments
5221 // because that is done previously
5222 if (CheckAllArgTypesAreCorrect(S: &SemaRef, TheCall,
5223 Check: CheckUnsignedIntRepresentation))
5224 return true;
5225 break;
5226 }
5227 case Builtin::BI__builtin_hlsl_wave_prefix_count_bits: {
5228 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 1))
5229 return true;
5230
5231 QualType ArgType = TheCall->getArg(Arg: 0)->getType();
5232
5233 if (!(ArgType->isScalarType())) {
5234 SemaRef.Diag(Loc: TheCall->getArg(Arg: 0)->getBeginLoc(),
5235 DiagID: diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
5236 << ArgType << 0;
5237 return true;
5238 }
5239
5240 if (!(ArgType->isBooleanType())) {
5241 SemaRef.Diag(Loc: TheCall->getArg(Arg: 0)->getBeginLoc(),
5242 DiagID: diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
5243 << ArgType << 0;
5244 return true;
5245 }
5246
5247 break;
5248 }
5249 case Builtin::BI__builtin_hlsl_wave_read_lane_at: {
5250 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 2))
5251 return true;
5252
5253 // Ensure index parameter type can be interpreted as a uint
5254 ExprResult Index = TheCall->getArg(Arg: 1);
5255 QualType ArgTyIndex = Index.get()->getType();
5256 if (!ArgTyIndex->isIntegerType()) {
5257 SemaRef.Diag(Loc: TheCall->getArg(Arg: 1)->getBeginLoc(),
5258 DiagID: diag::err_typecheck_convert_incompatible)
5259 << ArgTyIndex << SemaRef.Context.UnsignedIntTy << 1 << 0 << 0;
5260 return true;
5261 }
5262
5263 // Ensure input expr type is a scalar/vector and the same as the return type
5264 if (CheckAnyScalarOrVector(S: &SemaRef, TheCall, ArgIndex: 0))
5265 return true;
5266
5267 ExprResult Expr = TheCall->getArg(Arg: 0);
5268 QualType ArgTyExpr = Expr.get()->getType();
5269 TheCall->setType(ArgTyExpr);
5270 break;
5271 }
5272 case Builtin::BI__builtin_hlsl_wave_read_lane_first: {
5273 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 1))
5274 return true;
5275
5276 if (CheckAnyScalarOrVectorOrMatrix(S: &SemaRef, TheCall, ArgIndex: 0))
5277 return true;
5278
5279 TheCall->setType(TheCall->getArg(Arg: 0)->getType());
5280 break;
5281 }
5282 case Builtin::BI__builtin_hlsl_wave_get_lane_index: {
5283 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 0))
5284 return true;
5285 break;
5286 }
5287 case Builtin::BI__builtin_hlsl_wave_prefix_sum:
5288 case Builtin::BI__builtin_hlsl_wave_prefix_product: {
5289 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 1))
5290 return true;
5291
5292 // Ensure input expr type is a scalar/vector and the same as the return type
5293 if (CheckAnyScalarOrVector(S: &SemaRef, TheCall, ArgIndex: 0))
5294 return true;
5295 if (CheckWavePrefix(S: &SemaRef, TheCall))
5296 return true;
5297 ExprResult Expr = TheCall->getArg(Arg: 0);
5298 QualType ArgTyExpr = Expr.get()->getType();
5299 TheCall->setType(ArgTyExpr);
5300 break;
5301 }
5302 case Builtin::BI__builtin_hlsl_quad_read_across_x:
5303 case Builtin::BI__builtin_hlsl_quad_read_across_y:
5304 case Builtin::BI__builtin_hlsl_quad_read_across_diagonal: {
5305 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 1))
5306 return true;
5307
5308 if (CheckAnyScalarOrVector(S: &SemaRef, TheCall, ArgIndex: 0))
5309 return true;
5310 if (CheckNotBoolScalarOrVector(S: &SemaRef, TheCall, ArgIndex: 0))
5311 return true;
5312 ExprResult Expr = TheCall->getArg(Arg: 0);
5313 QualType ArgTyExpr = Expr.get()->getType();
5314 TheCall->setType(ArgTyExpr);
5315 break;
5316 }
5317 case Builtin::BI__builtin_hlsl_elementwise_splitdouble: {
5318 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 3))
5319 return true;
5320
5321 if (CheckScalarOrVectorOrMatrix(S: &SemaRef, TheCall, Scalar: SemaRef.Context.DoubleTy,
5322 ArgIndex: 0) ||
5323 CheckScalarOrVectorOrMatrix(S: &SemaRef, TheCall,
5324 Scalar: SemaRef.Context.UnsignedIntTy, ArgIndex: 1) ||
5325 CheckScalarOrVectorOrMatrix(S: &SemaRef, TheCall,
5326 Scalar: SemaRef.Context.UnsignedIntTy, ArgIndex: 2))
5327 return true;
5328
5329 if (CheckModifiableLValue(S: &SemaRef, TheCall, ArgIndex: 1) ||
5330 CheckModifiableLValue(S: &SemaRef, TheCall, ArgIndex: 2))
5331 return true;
5332 break;
5333 }
5334 case Builtin::BI__builtin_hlsl_elementwise_clip: {
5335 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 1))
5336 return true;
5337
5338 if (CheckScalarOrVector(S: &SemaRef, TheCall, Scalar: SemaRef.Context.FloatTy, ArgIndex: 0))
5339 return true;
5340 break;
5341 }
5342 case Builtin::BI__builtin_elementwise_acos:
5343 case Builtin::BI__builtin_elementwise_asin:
5344 case Builtin::BI__builtin_elementwise_atan:
5345 case Builtin::BI__builtin_elementwise_atan2:
5346 case Builtin::BI__builtin_elementwise_ceil:
5347 case Builtin::BI__builtin_elementwise_cos:
5348 case Builtin::BI__builtin_elementwise_cosh:
5349 case Builtin::BI__builtin_elementwise_exp:
5350 case Builtin::BI__builtin_elementwise_exp2:
5351 case Builtin::BI__builtin_elementwise_exp10:
5352 case Builtin::BI__builtin_elementwise_floor:
5353 case Builtin::BI__builtin_elementwise_fmod:
5354 case Builtin::BI__builtin_elementwise_log:
5355 case Builtin::BI__builtin_elementwise_log2:
5356 case Builtin::BI__builtin_elementwise_log10:
5357 case Builtin::BI__builtin_elementwise_pow:
5358 case Builtin::BI__builtin_elementwise_roundeven:
5359 case Builtin::BI__builtin_elementwise_sin:
5360 case Builtin::BI__builtin_elementwise_sinh:
5361 case Builtin::BI__builtin_elementwise_sqrt:
5362 case Builtin::BI__builtin_elementwise_tan:
5363 case Builtin::BI__builtin_elementwise_tanh:
5364 case Builtin::BI__builtin_elementwise_trunc: {
5365 if (CheckAllArgTypesAreCorrect(S: &SemaRef, TheCall,
5366 Check: CheckFloatOrHalfRepresentation))
5367 return true;
5368 break;
5369 }
5370 case Builtin::BI__builtin_hlsl_buffer_update_counter: {
5371 assert(TheCall->getNumArgs() == 2 && "expected 2 args");
5372 auto checkResTy = [](const HLSLAttributedResourceType *ResTy) -> bool {
5373 return !(ResTy->getAttrs().ResourceClass == ResourceClass::UAV &&
5374 ResTy->getAttrs().RawBuffer && ResTy->hasContainedType());
5375 };
5376 if (CheckResourceHandle(S: &SemaRef, TheCall, ArgIndex: 0, Check: checkResTy))
5377 return true;
5378 Expr *OffsetExpr = TheCall->getArg(Arg: 1);
5379 std::optional<llvm::APSInt> Offset =
5380 OffsetExpr->getIntegerConstantExpr(Ctx: SemaRef.getASTContext());
5381 if (!Offset.has_value() || std::abs(i: Offset->getExtValue()) != 1) {
5382 SemaRef.Diag(Loc: TheCall->getArg(Arg: 1)->getBeginLoc(),
5383 DiagID: diag::err_hlsl_expect_arg_const_int_one_or_neg_one)
5384 << 1;
5385 return true;
5386 }
5387 break;
5388 }
5389 case Builtin::BI__builtin_hlsl_elementwise_f16tof32: {
5390 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 1))
5391 return true;
5392 if (CheckAllArgTypesAreCorrect(S: &SemaRef, TheCall,
5393 Check: CheckUnsignedIntRepresentation))
5394 return true;
5395 // ensure arg integers are 32 bits
5396 if (CheckExpectedBitWidth(S: &SemaRef, TheCall, ArgOrdinal: 0, Width: 32))
5397 return true;
5398 // check it wasn't a bool type
5399 QualType ArgTy = TheCall->getArg(Arg: 0)->getType();
5400 if (auto *VTy = ArgTy->getAs<VectorType>())
5401 ArgTy = VTy->getElementType();
5402 if (ArgTy->isBooleanType()) {
5403 SemaRef.Diag(Loc: TheCall->getArg(Arg: 0)->getBeginLoc(),
5404 DiagID: diag::err_builtin_invalid_arg_type)
5405 << 1 << /* scalar or vector of */ 5 << /* unsigned int */ 3
5406 << /* no fp */ 0 << TheCall->getArg(Arg: 0)->getType();
5407 return true;
5408 }
5409
5410 SetElementTypeAsReturnType(S: &SemaRef, TheCall, ReturnType: getASTContext().FloatTy);
5411 break;
5412 }
5413 case Builtin::BI__builtin_hlsl_elementwise_f32tof16: {
5414 if (SemaRef.checkArgCount(Call: TheCall, DesiredArgCount: 1))
5415 return true;
5416 if (CheckAllArgTypesAreCorrect(S: &SemaRef, TheCall, Check: CheckFloatRepresentation))
5417 return true;
5418 SetElementTypeAsReturnType(S: &SemaRef, TheCall,
5419 ReturnType: getASTContext().UnsignedIntTy);
5420 break;
5421 }
5422 }
5423 return false;
5424}
5425
5426static void BuildFlattenedTypeList(QualType BaseTy,
5427 llvm::SmallVectorImpl<QualType> &List) {
5428 llvm::SmallVector<QualType, 16> WorkList;
5429 WorkList.push_back(Elt: BaseTy);
5430 while (!WorkList.empty()) {
5431 QualType T = WorkList.pop_back_val();
5432 T = T.getCanonicalType().getUnqualifiedType();
5433 if (const auto *AT = dyn_cast<ConstantArrayType>(Val&: T)) {
5434 llvm::SmallVector<QualType, 16> ElementFields;
5435 // Generally I've avoided recursion in this algorithm, but arrays of
5436 // structs could be time-consuming to flatten and churn through on the
5437 // work list. Hopefully nesting arrays of structs containing arrays
5438 // of structs too many levels deep is unlikely.
5439 BuildFlattenedTypeList(BaseTy: AT->getElementType(), List&: ElementFields);
5440 // Repeat the element's field list n times.
5441 for (uint64_t Ct = 0; Ct < AT->getZExtSize(); ++Ct)
5442 llvm::append_range(C&: List, R&: ElementFields);
5443 continue;
5444 }
5445 // Vectors can only have element types that are builtin types, so this can
5446 // add directly to the list instead of to the WorkList.
5447 if (const auto *VT = dyn_cast<VectorType>(Val&: T)) {
5448 List.insert(I: List.end(), NumToInsert: VT->getNumElements(), Elt: VT->getElementType());
5449 continue;
5450 }
5451 if (const auto *MT = dyn_cast<ConstantMatrixType>(Val&: T)) {
5452 List.insert(I: List.end(), NumToInsert: MT->getNumElementsFlattened(),
5453 Elt: MT->getElementType());
5454 continue;
5455 }
5456 if (const auto *RD = T->getAsCXXRecordDecl()) {
5457 if (RD->isStandardLayout())
5458 RD = RD->getStandardLayoutBaseWithFields();
5459
5460 // For types that we shouldn't decompose (unions and non-aggregates), just
5461 // add the type itself to the list.
5462 if (RD->isUnion() || !RD->isAggregate()) {
5463 List.push_back(Elt: T);
5464 continue;
5465 }
5466
5467 llvm::SmallVector<QualType, 16> FieldTypes;
5468 for (const auto *FD : RD->fields())
5469 if (!FD->isUnnamedBitField())
5470 FieldTypes.push_back(Elt: FD->getType());
5471 // Reverse the newly added sub-range.
5472 std::reverse(first: FieldTypes.begin(), last: FieldTypes.end());
5473 llvm::append_range(C&: WorkList, R&: FieldTypes);
5474
5475 // If this wasn't a standard layout type we may also have some base
5476 // classes to deal with.
5477 if (!RD->isStandardLayout()) {
5478 FieldTypes.clear();
5479 for (const auto &Base : RD->bases())
5480 FieldTypes.push_back(Elt: Base.getType());
5481 std::reverse(first: FieldTypes.begin(), last: FieldTypes.end());
5482 llvm::append_range(C&: WorkList, R&: FieldTypes);
5483 }
5484 continue;
5485 }
5486 List.push_back(Elt: T);
5487 }
5488}
5489
5490bool SemaHLSL::IsConstantBufferElementCompatible(clang::QualType QT) {
5491 if (QT.isNull())
5492 return false;
5493
5494 // Must be a class/struct.
5495 const auto *RD = QT->getAsCXXRecordDecl();
5496 if (!RD || RD->isUnion())
5497 return false;
5498
5499 // Cannot be a resource type or contain one.
5500 return !QT->isHLSLIntangibleType();
5501}
5502
5503bool SemaHLSL::IsTypedResourceElementCompatible(clang::QualType QT) {
5504 // null and array types are not allowed.
5505 if (QT.isNull() || QT->isArrayType())
5506 return false;
5507
5508 // UDT types are not allowed
5509 if (QT->isRecordType())
5510 return false;
5511
5512 if (QT->isBooleanType() || QT->isEnumeralType())
5513 return false;
5514
5515 // the only other valid builtin types are scalars or vectors
5516 if (QT->isArithmeticType()) {
5517 if (SemaRef.Context.getTypeSize(T: QT) / 8 > 16)
5518 return false;
5519 return true;
5520 }
5521
5522 if (const VectorType *VT = QT->getAs<VectorType>()) {
5523 int ArraySize = VT->getNumElements();
5524
5525 if (ArraySize > 4)
5526 return false;
5527
5528 QualType ElTy = VT->getElementType();
5529 if (ElTy->isBooleanType())
5530 return false;
5531
5532 if (SemaRef.Context.getTypeSize(T: QT) / 8 > 16)
5533 return false;
5534 return true;
5535 }
5536
5537 return false;
5538}
5539
5540bool SemaHLSL::IsScalarizedLayoutCompatible(QualType T1, QualType T2) const {
5541 if (T1.isNull() || T2.isNull())
5542 return false;
5543
5544 T1 = T1.getCanonicalType().getUnqualifiedType();
5545 T2 = T2.getCanonicalType().getUnqualifiedType();
5546
5547 // If both types are the same canonical type, they're obviously compatible.
5548 if (SemaRef.getASTContext().hasSameType(T1, T2))
5549 return true;
5550
5551 llvm::SmallVector<QualType, 16> T1Types;
5552 BuildFlattenedTypeList(BaseTy: T1, List&: T1Types);
5553 llvm::SmallVector<QualType, 16> T2Types;
5554 BuildFlattenedTypeList(BaseTy: T2, List&: T2Types);
5555
5556 // Check the flattened type list
5557 return llvm::equal(LRange&: T1Types, RRange&: T2Types,
5558 P: [this](QualType LHS, QualType RHS) -> bool {
5559 return SemaRef.IsLayoutCompatible(T1: LHS, T2: RHS);
5560 });
5561}
5562
5563bool SemaHLSL::CheckCompatibleParameterABI(FunctionDecl *New,
5564 FunctionDecl *Old) {
5565 if (New->getNumParams() != Old->getNumParams())
5566 return true;
5567
5568 bool HadError = false;
5569
5570 for (unsigned i = 0, e = New->getNumParams(); i != e; ++i) {
5571 ParmVarDecl *NewParam = New->getParamDecl(i);
5572 ParmVarDecl *OldParam = Old->getParamDecl(i);
5573
5574 // HLSL parameter declarations for inout and out must match between
5575 // declarations. In HLSL inout and out are ambiguous at the call site,
5576 // but have different calling behavior, so you cannot overload a
5577 // method based on a difference between inout and out annotations.
5578 const auto *NDAttr = NewParam->getAttr<HLSLParamModifierAttr>();
5579 unsigned NSpellingIdx = (NDAttr ? NDAttr->getSpellingListIndex() : 0);
5580 const auto *ODAttr = OldParam->getAttr<HLSLParamModifierAttr>();
5581 unsigned OSpellingIdx = (ODAttr ? ODAttr->getSpellingListIndex() : 0);
5582
5583 if (NSpellingIdx != OSpellingIdx) {
5584 SemaRef.Diag(Loc: NewParam->getLocation(),
5585 DiagID: diag::err_hlsl_param_qualifier_mismatch)
5586 << NDAttr << NewParam;
5587 SemaRef.Diag(Loc: OldParam->getLocation(), DiagID: diag::note_previous_declaration_as)
5588 << ODAttr;
5589 HadError = true;
5590 }
5591 }
5592 return HadError;
5593}
5594
5595// Generally follows PerformScalarCast, with cases reordered for
5596// clarity of what types are supported
5597bool SemaHLSL::CanPerformScalarCast(QualType SrcTy, QualType DestTy) {
5598
5599 if (!SrcTy->isScalarType() || !DestTy->isScalarType())
5600 return false;
5601
5602 if (SemaRef.getASTContext().hasSameUnqualifiedType(T1: SrcTy, T2: DestTy))
5603 return true;
5604
5605 switch (SrcTy->getScalarTypeKind()) {
5606 case Type::STK_Bool: // casting from bool is like casting from an integer
5607 case Type::STK_Integral:
5608 switch (DestTy->getScalarTypeKind()) {
5609 case Type::STK_Bool:
5610 case Type::STK_Integral:
5611 case Type::STK_Floating:
5612 return true;
5613 case Type::STK_CPointer:
5614 case Type::STK_ObjCObjectPointer:
5615 case Type::STK_BlockPointer:
5616 case Type::STK_MemberPointer:
5617 llvm_unreachable("HLSL doesn't support pointers.");
5618 case Type::STK_IntegralComplex:
5619 case Type::STK_FloatingComplex:
5620 llvm_unreachable("HLSL doesn't support complex types.");
5621 case Type::STK_FixedPoint:
5622 llvm_unreachable("HLSL doesn't support fixed point types.");
5623 }
5624 llvm_unreachable("Should have returned before this");
5625
5626 case Type::STK_Floating:
5627 switch (DestTy->getScalarTypeKind()) {
5628 case Type::STK_Floating:
5629 case Type::STK_Bool:
5630 case Type::STK_Integral:
5631 return true;
5632 case Type::STK_FloatingComplex:
5633 case Type::STK_IntegralComplex:
5634 llvm_unreachable("HLSL doesn't support complex types.");
5635 case Type::STK_FixedPoint:
5636 llvm_unreachable("HLSL doesn't support fixed point types.");
5637 case Type::STK_CPointer:
5638 case Type::STK_ObjCObjectPointer:
5639 case Type::STK_BlockPointer:
5640 case Type::STK_MemberPointer:
5641 llvm_unreachable("HLSL doesn't support pointers.");
5642 }
5643 llvm_unreachable("Should have returned before this");
5644
5645 case Type::STK_MemberPointer:
5646 case Type::STK_CPointer:
5647 case Type::STK_BlockPointer:
5648 case Type::STK_ObjCObjectPointer:
5649 llvm_unreachable("HLSL doesn't support pointers.");
5650
5651 case Type::STK_FixedPoint:
5652 llvm_unreachable("HLSL doesn't support fixed point types.");
5653
5654 case Type::STK_FloatingComplex:
5655 case Type::STK_IntegralComplex:
5656 llvm_unreachable("HLSL doesn't support complex types.");
5657 }
5658
5659 llvm_unreachable("Unhandled scalar cast");
5660}
5661
5662// Can perform an HLSL Aggregate splat cast if the Dest is an aggregate and the
5663// Src is a scalar, a vector of length 1, or a 1x1 matrix
5664// Or if Dest is a vector and Src is a vector of length 1 or a 1x1 matrix
5665bool SemaHLSL::CanPerformAggregateSplatCast(Expr *Src, QualType DestTy) {
5666
5667 QualType SrcTy = Src->getType();
5668 // Not a valid HLSL Aggregate Splat cast if Dest is a scalar or if this is
5669 // going to be a vector splat from a scalar.
5670 if ((SrcTy->isScalarType() && DestTy->isVectorType()) ||
5671 DestTy->isScalarType())
5672 return false;
5673
5674 const VectorType *SrcVecTy = SrcTy->getAs<VectorType>();
5675 const ConstantMatrixType *SrcMatTy = SrcTy->getAs<ConstantMatrixType>();
5676
5677 // Src isn't a scalar, a vector of length 1, or a 1x1 matrix
5678 if (!SrcTy->isScalarType() &&
5679 !(SrcVecTy && SrcVecTy->getNumElements() == 1) &&
5680 !(SrcMatTy && SrcMatTy->getNumElementsFlattened() == 1))
5681 return false;
5682
5683 if (SrcVecTy)
5684 SrcTy = SrcVecTy->getElementType();
5685 else if (SrcMatTy)
5686 SrcTy = SrcMatTy->getElementType();
5687
5688 llvm::SmallVector<QualType> DestTypes;
5689 BuildFlattenedTypeList(BaseTy: DestTy, List&: DestTypes);
5690
5691 for (unsigned I = 0, Size = DestTypes.size(); I < Size; ++I) {
5692 if (DestTypes[I]->isUnionType())
5693 return false;
5694 if (!CanPerformScalarCast(SrcTy, DestTy: DestTypes[I]))
5695 return false;
5696 }
5697 return true;
5698}
5699
5700// Can we perform an HLSL Elementwise cast?
5701bool SemaHLSL::CanPerformElementwiseCast(Expr *Src, QualType DestTy) {
5702
5703 // Don't handle casts where LHS and RHS are any combination of scalar/vector
5704 // There must be an aggregate somewhere
5705 QualType SrcTy = Src->getType();
5706 if (SrcTy->isScalarType()) // always a splat and this cast doesn't handle that
5707 return false;
5708
5709 if (SrcTy->isVectorType() &&
5710 (DestTy->isScalarType() || DestTy->isVectorType()))
5711 return false;
5712
5713 if (SrcTy->isConstantMatrixType() &&
5714 (DestTy->isScalarType() || DestTy->isConstantMatrixType()))
5715 return false;
5716
5717 llvm::SmallVector<QualType> DestTypes;
5718 BuildFlattenedTypeList(BaseTy: DestTy, List&: DestTypes);
5719 llvm::SmallVector<QualType> SrcTypes;
5720 BuildFlattenedTypeList(BaseTy: SrcTy, List&: SrcTypes);
5721
5722 // Usually the size of SrcTypes must be greater than or equal to the size of
5723 // DestTypes.
5724 if (SrcTypes.size() < DestTypes.size())
5725 return false;
5726
5727 unsigned SrcSize = SrcTypes.size();
5728 unsigned DstSize = DestTypes.size();
5729 unsigned I;
5730 for (I = 0; I < DstSize && I < SrcSize; I++) {
5731 if (SrcTypes[I]->isUnionType() || DestTypes[I]->isUnionType())
5732 return false;
5733 if (!CanPerformScalarCast(SrcTy: SrcTypes[I], DestTy: DestTypes[I])) {
5734 return false;
5735 }
5736 }
5737
5738 // check the rest of the source type for unions.
5739 for (; I < SrcSize; I++) {
5740 if (SrcTypes[I]->isUnionType())
5741 return false;
5742 }
5743 return true;
5744}
5745
5746bool SemaHLSL::CanPerformPackedTypeCast(Expr *Src, QualType DestTy) {
5747 ASTContext &Ctx = SemaRef.getASTContext();
5748 QualType UIntTy = Ctx.UnsignedIntTy;
5749 QualType SrcTy = Src->getType();
5750
5751 return (SrcTy->isHLSLBuiltinPackedType() &&
5752 DestTy->isHLSLBuiltinPackedType()) ||
5753 (SrcTy->isHLSLBuiltinPackedType() &&
5754 Ctx.hasSameUnqualifiedType(T1: DestTy, T2: UIntTy)) ||
5755 (DestTy->isHLSLBuiltinPackedType() &&
5756 Ctx.hasSameUnqualifiedType(T1: SrcTy, T2: UIntTy));
5757}
5758
5759ExprResult SemaHLSL::ActOnOutParamExpr(ParmVarDecl *Param, Expr *Arg) {
5760 assert(Param->hasAttr<HLSLParamModifierAttr>() &&
5761 "We should not get here without a parameter modifier expression");
5762 const auto *Attr = Param->getAttr<HLSLParamModifierAttr>();
5763 if (Attr->getABI() == ParameterABI::Ordinary)
5764 return ExprResult(Arg);
5765
5766 bool IsInOut = Attr->getABI() == ParameterABI::HLSLInOut;
5767 if (!Arg->isLValue()) {
5768 SemaRef.Diag(Loc: Arg->getBeginLoc(), DiagID: diag::error_hlsl_inout_lvalue)
5769 << Arg << (IsInOut ? 1 : 0);
5770 return ExprError();
5771 }
5772
5773 ASTContext &Ctx = SemaRef.getASTContext();
5774
5775 QualType Ty = Param->getType().getNonLValueExprType(Context: Ctx);
5776
5777 // HLSL allows implicit conversions from scalars to vectors, but not the
5778 // inverse, so we need to disallow `inout` with scalar->vector or
5779 // scalar->matrix conversions.
5780 if (Arg->getType()->isScalarType() != Ty->isScalarType()) {
5781 SemaRef.Diag(Loc: Arg->getBeginLoc(), DiagID: diag::error_hlsl_inout_scalar_extension)
5782 << Arg << (IsInOut ? 1 : 0);
5783 return ExprError();
5784 }
5785
5786 auto *ArgOpV = new (Ctx) OpaqueValueExpr(Param->getBeginLoc(), Arg->getType(),
5787 VK_LValue, OK_Ordinary, Arg);
5788
5789 // Parameters are initialized via copy initialization. This allows for
5790 // overload resolution of argument constructors.
5791 InitializedEntity Entity =
5792 InitializedEntity::InitializeParameter(Context&: Ctx, Type: Ty, Consumed: false);
5793 ExprResult Res =
5794 SemaRef.PerformCopyInitialization(Entity, EqualLoc: Param->getBeginLoc(), Init: ArgOpV);
5795 if (Res.isInvalid())
5796 return ExprError();
5797 Expr *Base = Res.get();
5798 // After the cast, drop the reference type when creating the exprs.
5799 Ty = Ty.getNonLValueExprType(Context: Ctx);
5800 auto *OpV = new (Ctx)
5801 OpaqueValueExpr(Param->getBeginLoc(), Ty, VK_LValue, OK_Ordinary, Base);
5802
5803 // Writebacks are performed with `=` binary operator, which allows for
5804 // overload resolution on writeback result expressions.
5805 Res = SemaRef.ActOnBinOp(S: SemaRef.getCurScope(), TokLoc: Arg->getBeginLoc(),
5806 Kind: tok::equal, LHSExpr: ArgOpV, RHSExpr: OpV);
5807
5808 if (Res.isInvalid())
5809 return ExprError();
5810 Expr *Writeback = Res.get();
5811 auto *OutExpr =
5812 HLSLOutArgExpr::Create(C: Ctx, Ty, Base: ArgOpV, OpV, WB: Writeback, IsInOut);
5813
5814 return ExprResult(OutExpr);
5815}
5816
5817QualType SemaHLSL::getInoutParameterType(QualType Ty) {
5818 // If HLSL gains support for references, all the cites that use this will need
5819 // to be updated with semantic checking to produce errors for
5820 // pointers/references.
5821 assert(!Ty->isReferenceType() &&
5822 "Pointer and reference types cannot be inout or out parameters");
5823 Ty = SemaRef.getASTContext().getLValueReferenceType(T: Ty);
5824 Ty.addRestrict();
5825 return Ty;
5826}
5827
5828// Returns true if the type has a non-empty constant buffer layout (if it is
5829// scalar, vector or matrix, or if it contains any of these.
5830static bool hasConstantBufferLayout(QualType QT) {
5831 const Type *Ty = QT->getUnqualifiedDesugaredType();
5832 if (Ty->isScalarType() || Ty->isVectorType() || Ty->isMatrixType())
5833 return true;
5834
5835 if (Ty->isHLSLResourceRecord() || Ty->isHLSLResourceRecordArray())
5836 return false;
5837
5838 if (const auto *RD = Ty->getAsCXXRecordDecl()) {
5839 for (const auto *FD : RD->fields()) {
5840 if (hasConstantBufferLayout(QT: FD->getType()))
5841 return true;
5842 }
5843 assert(RD->getNumBases() <= 1 &&
5844 "HLSL doesn't support multiple inheritance");
5845 return RD->getNumBases()
5846 ? hasConstantBufferLayout(QT: RD->bases_begin()->getType())
5847 : false;
5848 }
5849
5850 if (const auto *AT = dyn_cast<ArrayType>(Val: Ty)) {
5851 if (const auto *CAT = dyn_cast<ConstantArrayType>(Val: AT))
5852 if (isZeroSizedArray(CAT))
5853 return false;
5854 return hasConstantBufferLayout(QT: AT->getElementType());
5855 }
5856
5857 return false;
5858}
5859
5860static bool IsDefaultBufferConstantDecl(const ASTContext &Ctx, VarDecl *VD) {
5861 bool IsVulkan =
5862 Ctx.getTargetInfo().getTriple().getOS() == llvm::Triple::Vulkan;
5863 bool IsVKPushConstant = IsVulkan && VD->hasAttr<HLSLVkPushConstantAttr>();
5864 QualType QT = VD->getType();
5865 return VD->getDeclContext()->isTranslationUnit() &&
5866 QT.getAddressSpace() == LangAS::Default &&
5867 VD->getStorageClass() != SC_Static &&
5868 !VD->hasAttr<HLSLVkConstantIdAttr>() && !IsVKPushConstant &&
5869 hasConstantBufferLayout(QT);
5870}
5871
5872void SemaHLSL::deduceAddressSpace(VarDecl *Decl) {
5873 // The variable already has an address space (groupshared for ex).
5874 if (Decl->getType().hasAddressSpace())
5875 return;
5876
5877 if (Decl->getType()->isDependentType())
5878 return;
5879
5880 QualType Type = Decl->getType();
5881
5882 if (Decl->hasAttr<HLSLVkExtBuiltinInputAttr>()) {
5883 LangAS ImplAS = LangAS::hlsl_input;
5884 Type = SemaRef.getASTContext().getAddrSpaceQualType(T: Type, AddressSpace: ImplAS);
5885 Decl->setType(Type);
5886 return;
5887 }
5888
5889 if (Decl->hasAttr<HLSLVkExtBuiltinOutputAttr>()) {
5890 LangAS ImplAS = LangAS::hlsl_output;
5891 Type = SemaRef.getASTContext().getAddrSpaceQualType(T: Type, AddressSpace: ImplAS);
5892 Decl->setType(Type);
5893
5894 // HLSL uses `static` differently than C++. For BuiltIn output, the static
5895 // does not imply private to the module scope.
5896 // Marking it as external to reflect the semantic this attribute brings.
5897 // See https://github.com/microsoft/hlsl-specs/issues/350
5898 Decl->setStorageClass(SC_Extern);
5899 return;
5900 }
5901
5902 bool IsVulkan = getASTContext().getTargetInfo().getTriple().getOS() ==
5903 llvm::Triple::Vulkan;
5904 if (IsVulkan && Decl->hasAttr<HLSLVkPushConstantAttr>()) {
5905 if (HasDeclaredAPushConstant)
5906 SemaRef.Diag(Loc: Decl->getLocation(), DiagID: diag::err_hlsl_push_constant_unique);
5907
5908 LangAS ImplAS = LangAS::hlsl_push_constant;
5909 Type = SemaRef.getASTContext().getAddrSpaceQualType(T: Type, AddressSpace: ImplAS);
5910 Decl->setType(Type);
5911 HasDeclaredAPushConstant = true;
5912 return;
5913 }
5914
5915 if (Type->isSamplerT() || Type->isVoidType())
5916 return;
5917
5918 // Resource handles.
5919 if (Type->isHLSLResourceRecord() || Type->isHLSLResourceRecordArray())
5920 return;
5921
5922 // Only static globals belong to the Private address space.
5923 // Non-static globals belongs to the cbuffer.
5924 if (Decl->getStorageClass() != SC_Static && !Decl->isStaticDataMember())
5925 return;
5926
5927 LangAS ImplAS = LangAS::hlsl_private;
5928 Type = SemaRef.getASTContext().getAddrSpaceQualType(T: Type, AddressSpace: ImplAS);
5929 Decl->setType(Type);
5930}
5931
5932namespace {
5933
5934// Helper class for assigning bindings to resources declared within a struct.
5935// It keeps track of all binding attributes declared on a struct instance, and
5936// the offsets for each register type that have been assigned so far.
5937// Handles both explicit and implicit bindings.
5938class StructBindingContext {
5939 // Bindings and offsets per register type. We only need to support four
5940 // register types - SRV (u), UAV (t), CBuffer (c), and Sampler (s).
5941 HLSLResourceBindingAttr *RegBindingsAttrs[4];
5942 unsigned RegBindingOffset[4];
5943
5944 // Make sure the RegisterType values are what we expect
5945 static_assert(static_cast<unsigned>(RegisterType::SRV) == 0 &&
5946 static_cast<unsigned>(RegisterType::UAV) == 1 &&
5947 static_cast<unsigned>(RegisterType::CBuffer) == 2 &&
5948 static_cast<unsigned>(RegisterType::Sampler) == 3,
5949 "unexpected register type values");
5950
5951 // Vulkan binding attribute does not vary by register type.
5952 HLSLVkBindingAttr *VkBindingAttr;
5953 unsigned VkBindingOffset;
5954
5955public:
5956 // Constructor: gather all binding attributes on a struct instance and
5957 // initialize offsets.
5958 StructBindingContext(VarDecl *VD) {
5959 for (unsigned i = 0; i < 4; ++i) {
5960 RegBindingsAttrs[i] = nullptr;
5961 RegBindingOffset[i] = 0;
5962 }
5963 VkBindingAttr = nullptr;
5964 VkBindingOffset = 0;
5965
5966 ASTContext &AST = VD->getASTContext();
5967 bool IsSpirv = AST.getTargetInfo().getTriple().isSPIRV();
5968
5969 for (Attr *A : VD->attrs()) {
5970 if (auto *RBA = dyn_cast<HLSLResourceBindingAttr>(Val: A)) {
5971 RegisterType RegType = RBA->getRegisterType();
5972 unsigned RegTypeIdx = static_cast<unsigned>(RegType);
5973 // Ignore unsupported register annotations, such as 'c' or 'i'.
5974 if (RegTypeIdx < 4)
5975 RegBindingsAttrs[RegTypeIdx] = RBA;
5976 continue;
5977 }
5978 // Gather the Vulkan binding attributes only if the target is SPIR-V.
5979 if (IsSpirv) {
5980 if (auto *VBA = dyn_cast<HLSLVkBindingAttr>(Val: A))
5981 VkBindingAttr = VBA;
5982 }
5983 }
5984 }
5985
5986 // Creates a binding attribute for a resource based on the gathered attributes
5987 // and the required register type and range.
5988 Attr *createBindingAttr(SemaHLSL &S, ASTContext &AST, RegisterType RegType,
5989 unsigned Range, bool HasCounter) {
5990 assert(static_cast<unsigned>(RegType) < 4 && "unexpected register type");
5991
5992 if (VkBindingAttr) {
5993 unsigned Offset = VkBindingOffset;
5994 VkBindingOffset += Range;
5995 return HLSLVkBindingAttr::CreateImplicit(
5996 Ctx&: AST, Binding: VkBindingAttr->getBinding() + Offset, Set: VkBindingAttr->getSet(),
5997 Range: VkBindingAttr->getRange());
5998 }
5999
6000 HLSLResourceBindingAttr *RBA =
6001 RegBindingsAttrs[static_cast<unsigned>(RegType)];
6002 HLSLResourceBindingAttr *NewAttr = nullptr;
6003
6004 if (RBA && RBA->hasRegisterSlot()) {
6005 // Explicit binding - create a new attribute with offseted slot number
6006 // based on the required register type.
6007 unsigned Offset = RegBindingOffset[static_cast<unsigned>(RegType)];
6008 RegBindingOffset[static_cast<unsigned>(RegType)] += Range;
6009
6010 unsigned NewSlotNumber = RBA->getSlotNumber() + Offset;
6011 StringRef NewSlotNumberStr =
6012 createRegisterString(AST, RegType: RBA->getRegisterType(), N: NewSlotNumber);
6013 NewAttr = HLSLResourceBindingAttr::CreateImplicit(
6014 Ctx&: AST, Slot: NewSlotNumberStr, Space: RBA->getSpace(), Range: RBA->getRange());
6015 NewAttr->setBinding(RT: RegType, SlotNum: NewSlotNumber, SpaceNum: RBA->getSpaceNumber());
6016 } else {
6017 // No binding attribute or space-only binding - create a binding
6018 // attribute for implicit binding.
6019 NewAttr = HLSLResourceBindingAttr::CreateImplicit(Ctx&: AST, Slot: "", Space: "0", Range: {});
6020 NewAttr->setBinding(RT: RegType, SlotNum: std::nullopt,
6021 SpaceNum: RBA ? RBA->getSpaceNumber() : 0);
6022 NewAttr->setImplicitBindingOrderID(S.getNextImplicitBindingOrderID());
6023 }
6024 if (HasCounter)
6025 NewAttr->setImplicitCounterBindingOrderID(
6026 S.getNextImplicitBindingOrderID());
6027 return NewAttr;
6028 }
6029};
6030
6031// Creates a global variable declaration for a resource field embedded in a
6032// struct, assigns it a binding, initializes it, and associates it with the
6033// struct declaration via an HLSLAssociatedResourceDeclAttr.
6034static void createGlobalResourceDeclForStruct(
6035 Sema &S, VarDecl *ParentVD, SourceLocation Loc, IdentifierInfo *Id,
6036 QualType ResTy, StructBindingContext &BindingCtx) {
6037 assert(isResourceRecordTypeOrArrayOf(ResTy) &&
6038 "expected resource type or array of resources");
6039
6040 DeclContext *DC = ParentVD->getNonTransparentDeclContext();
6041 assert(DC->isTranslationUnit() && "expected translation unit decl context");
6042
6043 ASTContext &AST = S.getASTContext();
6044 VarDecl *ResDecl =
6045 VarDecl::Create(C&: AST, DC, StartLoc: Loc, IdLoc: Loc, Id, T: ResTy, TInfo: nullptr, S: SC_None);
6046
6047 unsigned Range = 1;
6048 const Type *SingleResTy = ResTy.getTypePtr()->getUnqualifiedDesugaredType();
6049 while (const auto *AT = dyn_cast<ArrayType>(Val: SingleResTy)) {
6050 const auto *CAT = dyn_cast<ConstantArrayType>(Val: AT);
6051 Range = CAT ? (Range * CAT->getSize().getZExtValue()) : 0;
6052 SingleResTy =
6053 AT->getArrayElementTypeNoTypeQual()->getUnqualifiedDesugaredType();
6054 }
6055 const HLSLAttributedResourceType *ResHandleTy =
6056 HLSLAttributedResourceType::findHandleTypeOnResource(RT: SingleResTy);
6057
6058 // Add a binding attribute to the global resource declaration.
6059 bool HasCounter = hasCounterHandle(RD: SingleResTy->getAsCXXRecordDecl());
6060 Attr *BindingAttr = BindingCtx.createBindingAttr(
6061 S&: S.HLSL(), AST, RegType: getRegisterType(ResTy: ResHandleTy), Range, HasCounter);
6062 ResDecl->addAttr(A: BindingAttr);
6063 ResDecl->addAttr(A: InternalLinkageAttr::CreateImplicit(Ctx&: AST));
6064 ResDecl->setImplicit();
6065
6066 if (Range == 1)
6067 S.HLSL().initGlobalResourceDecl(VD: ResDecl);
6068 else
6069 S.HLSL().initGlobalResourceArrayDecl(VD: ResDecl);
6070
6071 ParentVD->addAttr(
6072 A: HLSLAssociatedResourceDeclAttr::CreateImplicit(Ctx&: AST, ResDecl));
6073 DC->addDecl(D: ResDecl);
6074
6075 DeclGroupRef DG(ResDecl);
6076 S.Consumer.HandleTopLevelDecl(D: DG);
6077}
6078
6079static void handleArrayOfStructWithResources(
6080 Sema &S, VarDecl *ParentVD, const ConstantArrayType *CAT,
6081 EmbeddedResourceNameBuilder &NameBuilder, StructBindingContext &BindingCtx);
6082
6083// Scans base and all fields of a struct/class type to find all embedded
6084// resources or resource arrays. Creates a global variable for each resource
6085// found.
6086static void handleStructWithResources(Sema &S, VarDecl *ParentVD,
6087 const CXXRecordDecl *RD,
6088 EmbeddedResourceNameBuilder &NameBuilder,
6089 StructBindingContext &BindingCtx) {
6090
6091 // Scan the base classes.
6092 assert(RD->getNumBases() <= 1 && "HLSL doesn't support multiple inheritance");
6093 const auto *BasesIt = RD->bases_begin();
6094 if (BasesIt != RD->bases_end()) {
6095 QualType QT = BasesIt->getType();
6096 if (QT->isHLSLIntangibleType()) {
6097 CXXRecordDecl *BaseRD = QT->getAsCXXRecordDecl();
6098 NameBuilder.pushBaseName(N: BaseRD->getName());
6099 handleStructWithResources(S, ParentVD, RD: BaseRD, NameBuilder, BindingCtx);
6100 NameBuilder.pop();
6101 }
6102 }
6103 // Process this class fields.
6104 for (const FieldDecl *FD : RD->fields()) {
6105 QualType FDTy = FD->getType().getCanonicalType();
6106 if (!FDTy->isHLSLIntangibleType())
6107 continue;
6108
6109 NameBuilder.pushName(N: FD->getName());
6110
6111 if (isResourceRecordTypeOrArrayOf(Ty: FDTy)) {
6112 IdentifierInfo *II = NameBuilder.getNameAsIdentifier(AST&: S.getASTContext());
6113 createGlobalResourceDeclForStruct(S, ParentVD, Loc: FD->getLocation(), Id: II,
6114 ResTy: FDTy, BindingCtx);
6115 } else if (const auto *RD = FDTy->getAsCXXRecordDecl()) {
6116 handleStructWithResources(S, ParentVD, RD, NameBuilder, BindingCtx);
6117
6118 } else if (const auto *ArrayTy = dyn_cast<ConstantArrayType>(Val&: FDTy)) {
6119 assert(!FDTy->isHLSLResourceRecordArray() &&
6120 "resource arrays should have been already handled");
6121 handleArrayOfStructWithResources(S, ParentVD, CAT: ArrayTy, NameBuilder,
6122 BindingCtx);
6123 }
6124 NameBuilder.pop();
6125 }
6126}
6127
6128// Processes array of structs with resources.
6129static void
6130handleArrayOfStructWithResources(Sema &S, VarDecl *ParentVD,
6131 const ConstantArrayType *CAT,
6132 EmbeddedResourceNameBuilder &NameBuilder,
6133 StructBindingContext &BindingCtx) {
6134
6135 QualType ElementTy = CAT->getElementType().getCanonicalType();
6136 assert(ElementTy->isHLSLIntangibleType() && "Expected HLSL intangible type");
6137
6138 const ConstantArrayType *SubCAT = dyn_cast<ConstantArrayType>(Val&: ElementTy);
6139 const CXXRecordDecl *ElementRD = ElementTy->getAsCXXRecordDecl();
6140
6141 if (!SubCAT && !ElementRD)
6142 return;
6143
6144 for (unsigned I = 0, E = CAT->getSize().getZExtValue(); I < E; ++I) {
6145 NameBuilder.pushArrayIndex(Index: I);
6146 if (ElementRD)
6147 handleStructWithResources(S, ParentVD, RD: ElementRD, NameBuilder,
6148 BindingCtx);
6149 else
6150 handleArrayOfStructWithResources(S, ParentVD, CAT: SubCAT, NameBuilder,
6151 BindingCtx);
6152 NameBuilder.pop();
6153 }
6154}
6155
6156} // namespace
6157
6158// Scans all fields of a user-defined struct (or array of structs)
6159// to find all embedded resources or resource arrays. For each resource
6160// a global variable of the resource type is created and associated
6161// with the parent declaration (VD) through a HLSLAssociatedResourceDeclAttr
6162// attribute.
6163void SemaHLSL::handleGlobalStructOrArrayOfWithResources(VarDecl *VD) {
6164 EmbeddedResourceNameBuilder NameBuilder(VD->getName());
6165 StructBindingContext BindingCtx(VD);
6166
6167 const Type *VDTy = VD->getType().getTypePtr();
6168 assert(VDTy->isHLSLIntangibleType() && !isResourceRecordTypeOrArrayOf(VD) &&
6169 "Expected non-resource struct or array type");
6170
6171 if (const CXXRecordDecl *RD = VDTy->getAsCXXRecordDecl()) {
6172 handleStructWithResources(S&: SemaRef, ParentVD: VD, RD, NameBuilder, BindingCtx);
6173 return;
6174 }
6175
6176 if (const auto *CAT = dyn_cast<ConstantArrayType>(Val: VDTy)) {
6177 handleArrayOfStructWithResources(S&: SemaRef, ParentVD: VD, CAT, NameBuilder, BindingCtx);
6178 return;
6179 }
6180}
6181
6182void SemaHLSL::ActOnVariableDeclarator(VarDecl *VD) {
6183 if (VD->hasGlobalStorage()) {
6184 // make sure the declaration has a complete type
6185 if (SemaRef.RequireCompleteType(
6186 Loc: VD->getLocation(),
6187 T: SemaRef.getASTContext().getBaseElementType(QT: VD->getType()),
6188 DiagID: diag::err_typecheck_decl_incomplete_type)) {
6189 VD->setInvalidDecl();
6190 deduceAddressSpace(Decl: VD);
6191 return;
6192 }
6193
6194 // Global variables outside a cbuffer block that are not a resource, static,
6195 // groupshared, or an empty array or struct belong to the default constant
6196 // buffer $Globals (to be created at the end of the translation unit).
6197 if (IsDefaultBufferConstantDecl(Ctx: getASTContext(), VD)) {
6198 // update address space to hlsl_constant
6199 QualType NewTy = getASTContext().getAddrSpaceQualType(
6200 T: VD->getType(), AddressSpace: LangAS::hlsl_constant);
6201 VD->setType(NewTy);
6202 DefaultCBufferDecls.push_back(Elt: VD);
6203 }
6204
6205 // find all resources bindings on decl
6206 if (VD->getType()->isHLSLIntangibleType())
6207 collectResourceBindingsOnVarDecl(D: VD);
6208
6209 if (VD->hasAttr<HLSLVkConstantIdAttr>())
6210 VD->setStorageClass(StorageClass::SC_Static);
6211
6212 if (isResourceRecordTypeOrArrayOf(VD) &&
6213 VD->getStorageClass() != SC_Static) {
6214 // Add internal linkage attribute to non-static resource variables. The
6215 // global externally visible storage is accessed through the handle, which
6216 // is a member. The variable itself is not externally visible.
6217 VD->addAttr(A: InternalLinkageAttr::CreateImplicit(Ctx&: getASTContext()));
6218 }
6219
6220 // process explicit bindings
6221 processExplicitBindingsOnDecl(D: VD);
6222
6223 // Add implicit binding attribute to non-static resource arrays.
6224 if (VD->getType()->isHLSLResourceRecordArray() &&
6225 VD->getStorageClass() != SC_Static) {
6226 // If the resource array does not have an explicit binding attribute,
6227 // create an implicit one. It will be used to transfer implicit binding
6228 // order_ID to codegen.
6229 ResourceBindingAttrs Binding(VD);
6230 if (!Binding.isExplicit()) {
6231 uint32_t OrderID = getNextImplicitBindingOrderID();
6232 if (Binding.hasBinding())
6233 Binding.setImplicitOrderID(OrderID);
6234 else {
6235 addImplicitBindingAttrToDecl(
6236 S&: SemaRef, D: VD, RT: getRegisterType(ResTy: getResourceArrayHandleType(VD)),
6237 ImplicitBindingOrderID: OrderID);
6238 // Re-create the binding object to pick up the new attribute.
6239 Binding = ResourceBindingAttrs(VD);
6240 }
6241 }
6242
6243 // Get to the base type of a potentially multi-dimensional array.
6244 QualType Ty = getASTContext().getBaseElementType(QT: VD->getType());
6245
6246 const CXXRecordDecl *RD = Ty->getAsCXXRecordDecl();
6247 if (hasCounterHandle(RD)) {
6248 if (!Binding.hasCounterImplicitOrderID()) {
6249 uint32_t OrderID = getNextImplicitBindingOrderID();
6250 Binding.setCounterImplicitOrderID(OrderID);
6251 }
6252 }
6253 }
6254
6255 // Process resources in user-defined structs, or arrays of such structs.
6256 const Type *VDTy = VD->getType().getTypePtr();
6257 if (VD->getStorageClass() != SC_Static && VDTy->isHLSLIntangibleType() &&
6258 !isResourceRecordTypeOrArrayOf(VD))
6259 handleGlobalStructOrArrayOfWithResources(VD);
6260
6261 // Mark groupshared variables as extern so they will have
6262 // external storage and won't be default initialized
6263 if (VD->hasAttr<HLSLGroupSharedAddressSpaceAttr>())
6264 VD->setStorageClass(StorageClass::SC_Extern);
6265 }
6266
6267 deduceAddressSpace(Decl: VD);
6268}
6269
6270bool SemaHLSL::initGlobalResourceDecl(VarDecl *VD) {
6271 assert(VD->getType()->isHLSLResourceRecord() &&
6272 "expected resource record type");
6273
6274 ASTContext &AST = SemaRef.getASTContext();
6275 uint64_t UIntTySize = AST.getTypeSize(T: AST.UnsignedIntTy);
6276 uint64_t IntTySize = AST.getTypeSize(T: AST.IntTy);
6277
6278 // Gather resource binding attributes.
6279 ResourceBindingAttrs Binding(VD);
6280
6281 // Find correct initialization method and create its arguments.
6282 QualType ResourceTy = VD->getType();
6283 CXXRecordDecl *ResourceDecl = ResourceTy->getAsCXXRecordDecl();
6284 CXXMethodDecl *CreateMethod = nullptr;
6285 llvm::SmallVector<Expr *> Args;
6286
6287 bool HasCounter = hasCounterHandle(RD: ResourceDecl);
6288 const char *CreateMethodName;
6289 if (Binding.isExplicit())
6290 CreateMethodName = HasCounter ? "__createFromBindingWithImplicitCounter"
6291 : "__createFromBinding";
6292 else
6293 CreateMethodName = HasCounter
6294 ? "__createFromImplicitBindingWithImplicitCounter"
6295 : "__createFromImplicitBinding";
6296
6297 CreateMethod =
6298 lookupMethod(S&: SemaRef, RecordDecl: ResourceDecl, Name: CreateMethodName, Loc: VD->getLocation());
6299
6300 if (!CreateMethod) {
6301 // This can happen if someone creates a struct that looks like an HLSL
6302 // resource record but does not have the required static create method.
6303 // No binding will be generated for it.
6304 assert(!ResourceDecl->isImplicit() &&
6305 "create method lookup should always succeed for built-in resource "
6306 "records");
6307 return false;
6308 }
6309
6310 if (Binding.isExplicit()) {
6311 IntegerLiteral *RegSlot =
6312 IntegerLiteral::Create(C: AST, V: llvm::APInt(UIntTySize, Binding.getSlot()),
6313 type: AST.UnsignedIntTy, l: SourceLocation());
6314 Args.push_back(Elt: RegSlot);
6315 } else {
6316 uint32_t OrderID = (Binding.hasImplicitOrderID())
6317 ? Binding.getImplicitOrderID()
6318 : getNextImplicitBindingOrderID();
6319 IntegerLiteral *OrderId =
6320 IntegerLiteral::Create(C: AST, V: llvm::APInt(UIntTySize, OrderID),
6321 type: AST.UnsignedIntTy, l: SourceLocation());
6322 Args.push_back(Elt: OrderId);
6323 }
6324
6325 IntegerLiteral *Space =
6326 IntegerLiteral::Create(C: AST, V: llvm::APInt(UIntTySize, Binding.getSpace()),
6327 type: AST.UnsignedIntTy, l: SourceLocation());
6328 Args.push_back(Elt: Space);
6329
6330 IntegerLiteral *RangeSize = IntegerLiteral::Create(
6331 C: AST, V: llvm::APInt(IntTySize, 1), type: AST.IntTy, l: SourceLocation());
6332 Args.push_back(Elt: RangeSize);
6333
6334 IntegerLiteral *Index = IntegerLiteral::Create(
6335 C: AST, V: llvm::APInt(UIntTySize, 0), type: AST.UnsignedIntTy, l: SourceLocation());
6336 Args.push_back(Elt: Index);
6337
6338 StringRef VarName = VD->getName();
6339 StringLiteral *Name = StringLiteral::Create(
6340 Ctx: AST, Str: VarName, Kind: StringLiteralKind::Ordinary, Pascal: false,
6341 Ty: AST.getStringLiteralArrayType(EltTy: AST.CharTy.withConst(), Length: VarName.size()),
6342 Locs: SourceLocation());
6343 ImplicitCastExpr *NameCast = ImplicitCastExpr::Create(
6344 Context: AST, T: AST.getPointerType(T: AST.CharTy.withConst()), Kind: CK_ArrayToPointerDecay,
6345 Operand: Name, BasePath: nullptr, Cat: VK_PRValue, FPO: FPOptionsOverride());
6346 Args.push_back(Elt: NameCast);
6347
6348 if (HasCounter) {
6349 // Will this be in the correct order?
6350 uint32_t CounterOrderID = getNextImplicitBindingOrderID();
6351 IntegerLiteral *CounterId =
6352 IntegerLiteral::Create(C: AST, V: llvm::APInt(UIntTySize, CounterOrderID),
6353 type: AST.UnsignedIntTy, l: SourceLocation());
6354 Args.push_back(Elt: CounterId);
6355 }
6356
6357 // Make sure the create method template is instantiated and emitted.
6358 if (!CreateMethod->isDefined() && CreateMethod->isTemplateInstantiation())
6359 SemaRef.InstantiateFunctionDefinition(PointOfInstantiation: VD->getLocation(), Function: CreateMethod,
6360 Recursive: true);
6361
6362 // Create CallExpr with a call to the static method and set it as the decl
6363 // initialization.
6364 DeclRefExpr *DRE = DeclRefExpr::Create(
6365 Context: AST, QualifierLoc: NestedNameSpecifierLoc(), TemplateKWLoc: SourceLocation(), D: CreateMethod, RefersToEnclosingVariableOrCapture: false,
6366 NameInfo: CreateMethod->getNameInfo(), T: CreateMethod->getType(), VK: VK_PRValue);
6367
6368 auto *ImpCast = ImplicitCastExpr::Create(
6369 Context: AST, T: AST.getPointerType(T: CreateMethod->getType()),
6370 Kind: CK_FunctionToPointerDecay, Operand: DRE, BasePath: nullptr, Cat: VK_PRValue, FPO: FPOptionsOverride());
6371
6372 CallExpr *InitExpr =
6373 CallExpr::Create(Ctx: AST, Fn: ImpCast, Args, Ty: ResourceTy, VK: VK_PRValue,
6374 RParenLoc: SourceLocation(), FPFeatures: FPOptionsOverride());
6375 VD->setInit(InitExpr);
6376 VD->setInitStyle(VarDecl::CallInit);
6377 SemaRef.CheckCompleteVariableDeclaration(VD);
6378 return true;
6379}
6380
6381bool SemaHLSL::initGlobalResourceArrayDecl(VarDecl *VD) {
6382 assert(VD->getType()->isHLSLResourceRecordArray() &&
6383 "expected array of resource records");
6384
6385 // Individual resources in a resource array are not initialized here. They
6386 // are initialized later on during codegen when the individual resources are
6387 // accessed. Codegen will emit a call to the resource initialization method
6388 // with the specified array index. We need to make sure though that the method
6389 // for the specific resource type is instantiated, so codegen can emit a call
6390 // to it when the array element is accessed.
6391
6392 // Find correct initialization method based on the resource binding
6393 // information.
6394 ASTContext &AST = SemaRef.getASTContext();
6395 QualType ResElementTy = AST.getBaseElementType(QT: VD->getType());
6396 CXXRecordDecl *ResourceDecl = ResElementTy->getAsCXXRecordDecl();
6397 CXXMethodDecl *CreateMethod = nullptr;
6398
6399 bool HasCounter = hasCounterHandle(RD: ResourceDecl);
6400 ResourceBindingAttrs ResourceAttrs(VD);
6401 if (ResourceAttrs.isExplicit())
6402 // Resource has explicit binding.
6403 CreateMethod =
6404 lookupMethod(S&: SemaRef, RecordDecl: ResourceDecl,
6405 Name: HasCounter ? "__createFromBindingWithImplicitCounter"
6406 : "__createFromBinding",
6407 Loc: VD->getLocation());
6408 else
6409 // Resource has implicit binding.
6410 CreateMethod = lookupMethod(
6411 S&: SemaRef, RecordDecl: ResourceDecl,
6412 Name: HasCounter ? "__createFromImplicitBindingWithImplicitCounter"
6413 : "__createFromImplicitBinding",
6414 Loc: VD->getLocation());
6415
6416 if (!CreateMethod)
6417 return false;
6418
6419 // Make sure the create method template is instantiated and emitted.
6420 if (!CreateMethod->isDefined() && CreateMethod->isTemplateInstantiation())
6421 SemaRef.InstantiateFunctionDefinition(PointOfInstantiation: VD->getLocation(), Function: CreateMethod,
6422 Recursive: true);
6423 return true;
6424}
6425
6426// Returns true if the initialization has been handled.
6427// Returns false to use default initialization.
6428bool SemaHLSL::ActOnUninitializedVarDecl(VarDecl *VD) {
6429 // Objects in the hlsl_constant address space are initialized
6430 // externally, so don't synthesize an implicit initializer.
6431 if (VD->getType().getAddressSpace() == LangAS::hlsl_constant)
6432 return true;
6433
6434 if (VD->hasGlobalStorage() && VD->getStorageClass() != SC_Static) {
6435 const Type *Ty = VD->getType().getTypePtr();
6436 if (Ty->isHLSLResourceRecord() && initGlobalResourceDecl(VD))
6437 return true;
6438 if (Ty->isHLSLResourceRecordArray() && initGlobalResourceArrayDecl(VD))
6439 return true;
6440 }
6441
6442 // User-defined structs/classes do not have constructors.
6443 // When declared at a global scope, they are part of the constant buffer
6444 // and should not be initialized by the compiler.
6445 // When declared at a local scope, they are not initialized.
6446 // Also applies to arrays of user-defined structs/classes.
6447 const Type *Ty = VD->getType()->getUnqualifiedDesugaredType();
6448 while (Ty->isArrayType())
6449 Ty = Ty->getArrayElementTypeNoTypeQual()->getUnqualifiedDesugaredType();
6450 if (CXXRecordDecl *RD = Ty->getAsCXXRecordDecl())
6451 return !RD->isHLSLBuiltinRecord();
6452
6453 return false;
6454}
6455
6456std::optional<const DeclBindingInfo *> SemaHLSL::inferGlobalBinding(Expr *E) {
6457 if (auto *Ternary = dyn_cast<ConditionalOperator>(Val: E)) {
6458 auto TrueInfo = inferGlobalBinding(E: Ternary->getTrueExpr());
6459 auto FalseInfo = inferGlobalBinding(E: Ternary->getFalseExpr());
6460 if (!TrueInfo || !FalseInfo)
6461 return std::nullopt;
6462 if (*TrueInfo != *FalseInfo)
6463 return std::nullopt;
6464 return TrueInfo;
6465 }
6466
6467 if (auto *ASE = dyn_cast<ArraySubscriptExpr>(Val: E))
6468 E = ASE->getBase()->IgnoreParenImpCasts();
6469
6470 if (DeclRefExpr *DRE = dyn_cast<DeclRefExpr>(Val: E->IgnoreParens()))
6471 if (VarDecl *VD = dyn_cast<VarDecl>(Val: DRE->getDecl())) {
6472 const Type *Ty = VD->getType()->getUnqualifiedDesugaredType();
6473 if (Ty->isArrayType())
6474 Ty = Ty->getArrayElementTypeNoTypeQual();
6475
6476 if (const auto *AttrResType =
6477 HLSLAttributedResourceType::findHandleTypeOnResource(RT: Ty)) {
6478 ResourceClass RC = AttrResType->getAttrs().ResourceClass;
6479 return Bindings.getDeclBindingInfo(VD, ResClass: RC);
6480 }
6481 }
6482
6483 return nullptr;
6484}
6485
6486void SemaHLSL::trackLocalResource(VarDecl *VD, Expr *E) {
6487 std::optional<const DeclBindingInfo *> ExprBinding = inferGlobalBinding(E);
6488 if (!ExprBinding) {
6489 SemaRef.Diag(Loc: E->getBeginLoc(),
6490 DiagID: diag::warn_hlsl_assigning_local_resource_is_not_unique)
6491 << E << VD;
6492 return; // Expr use multiple resources
6493 }
6494
6495 if (*ExprBinding == nullptr)
6496 return; // No binding could be inferred to track, return without error
6497
6498 auto PrevBinding = Assigns.find(Val: VD);
6499 if (PrevBinding == Assigns.end()) {
6500 // No previous binding recorded, simply record the new assignment
6501 Assigns.insert(KV: {VD, *ExprBinding});
6502 return;
6503 }
6504
6505 // Otherwise, warn if the assignment implies different resource bindings
6506 if (*ExprBinding != PrevBinding->second) {
6507 SemaRef.Diag(Loc: E->getBeginLoc(),
6508 DiagID: diag::warn_hlsl_assigning_local_resource_is_not_unique)
6509 << E << VD;
6510 SemaRef.Diag(Loc: VD->getLocation(), DiagID: diag::note_var_declared_here) << VD;
6511 return;
6512 }
6513
6514 return;
6515}
6516
6517bool SemaHLSL::CheckResourceBinOp(BinaryOperatorKind Opc, Expr *LHSExpr,
6518 Expr *RHSExpr, SourceLocation Loc) {
6519 assert((LHSExpr->getType()->isHLSLResourceRecord() ||
6520 LHSExpr->getType()->isHLSLResourceRecordArray()) &&
6521 "expected LHS to be a resource record or array of resource records");
6522 if (Opc != BO_Assign)
6523 return true;
6524
6525 // If LHS is an array subscript, get the underlying declaration.
6526 Expr *E = LHSExpr;
6527 while (auto *ASE = dyn_cast<ArraySubscriptExpr>(Val: E))
6528 E = ASE->getBase()->IgnoreParenImpCasts();
6529
6530 // Report error if LHS is a non-static resource declared at a global scope.
6531 if (DeclRefExpr *DRE = dyn_cast<DeclRefExpr>(Val: E->IgnoreParens())) {
6532 if (VarDecl *VD = dyn_cast<VarDecl>(Val: DRE->getDecl())) {
6533 if (VD->hasGlobalStorage() && VD->getStorageClass() != SC_Static) {
6534 // assignment to global resource is not allowed
6535 SemaRef.Diag(Loc, DiagID: diag::err_hlsl_assign_to_global_resource) << VD;
6536 SemaRef.Diag(Loc: VD->getLocation(), DiagID: diag::note_var_declared_here) << VD;
6537 return false;
6538 }
6539
6540 trackLocalResource(VD, E: RHSExpr);
6541 }
6542 }
6543 return true;
6544}
6545
6546// Returns true if the given type can have an overload of the given
6547// binary operator.
6548bool SemaHLSL::canHaveOverloadedBinOp(QualType LHSTy, BinaryOperatorKind Opc) {
6549 CXXRecordDecl *RD = LHSTy->getAsCXXRecordDecl();
6550 if (!RD)
6551 return true;
6552 return RD->isHLSLBuiltinRecord() || Opc != BO_Assign;
6553}
6554
6555// Walks though the global variable declaration, collects all resource binding
6556// requirements and adds them to Bindings
6557void SemaHLSL::collectResourceBindingsOnVarDecl(VarDecl *VD) {
6558 assert(VD->hasGlobalStorage() && VD->getType()->isHLSLIntangibleType() &&
6559 "expected global variable that contains HLSL resource");
6560
6561 // Cbuffers and Tbuffers are HLSLBufferDecl types
6562 if (const HLSLBufferDecl *CBufferOrTBuffer = dyn_cast<HLSLBufferDecl>(Val: VD)) {
6563 Bindings.addDeclBindingInfo(VD, ResClass: CBufferOrTBuffer->isCBuffer()
6564 ? ResourceClass::CBuffer
6565 : ResourceClass::SRV);
6566 return;
6567 }
6568
6569 // Unwrap arrays
6570 // FIXME: Calculate array size while unwrapping
6571 const Type *Ty = VD->getType()->getUnqualifiedDesugaredType();
6572 while (Ty->isArrayType()) {
6573 const ArrayType *AT = cast<ArrayType>(Val: Ty);
6574 Ty = AT->getElementType()->getUnqualifiedDesugaredType();
6575 }
6576
6577 // Resource (or array of resources)
6578 if (const HLSLAttributedResourceType *AttrResType =
6579 HLSLAttributedResourceType::findHandleTypeOnResource(RT: Ty)) {
6580 Bindings.addDeclBindingInfo(VD, ResClass: AttrResType->getAttrs().ResourceClass);
6581 return;
6582 }
6583
6584 // User defined record type
6585 if (const RecordType *RT = dyn_cast<RecordType>(Val: Ty))
6586 collectResourceBindingsOnUserRecordDecl(VD, RT);
6587}
6588
6589// Walks though the explicit resource binding attributes on the declaration,
6590// and makes sure there is a resource that matched the binding and updates
6591// DeclBindingInfoLists
6592void SemaHLSL::processExplicitBindingsOnDecl(VarDecl *VD) {
6593 assert(VD->hasGlobalStorage() && "expected global variable");
6594
6595 bool HasBinding = false;
6596 for (Attr *A : VD->attrs()) {
6597 if (isa<HLSLVkBindingAttr>(Val: A)) {
6598 HasBinding = true;
6599 if (auto PA = VD->getAttr<HLSLVkPushConstantAttr>())
6600 Diag(Loc: PA->getLoc(), DiagID: diag::err_hlsl_attr_incompatible) << A << PA;
6601 }
6602
6603 HLSLResourceBindingAttr *RBA = dyn_cast<HLSLResourceBindingAttr>(Val: A);
6604 if (!RBA || !RBA->hasRegisterSlot())
6605 continue;
6606 HasBinding = true;
6607
6608 RegisterType RT = RBA->getRegisterType();
6609 assert(RT != RegisterType::I && "invalid or obsolete register type should "
6610 "never have an attribute created");
6611
6612 if (RT == RegisterType::C) {
6613 if (Bindings.hasBindingInfoForDecl(VD))
6614 SemaRef.Diag(Loc: VD->getLocation(),
6615 DiagID: diag::warn_hlsl_user_defined_type_missing_member)
6616 << static_cast<int>(RT);
6617 continue;
6618 }
6619
6620 // Find DeclBindingInfo for this binding and update it, or report error
6621 // if it does not exist (user type does to contain resources with the
6622 // expected resource class).
6623 ResourceClass RC = getResourceClass(RT);
6624 if (DeclBindingInfo *BI = Bindings.getDeclBindingInfo(VD, ResClass: RC)) {
6625 // update binding info
6626 BI->setBindingAttribute(A: RBA, BT: BindingType::Explicit);
6627 } else {
6628 SemaRef.Diag(Loc: VD->getLocation(),
6629 DiagID: diag::warn_hlsl_user_defined_type_missing_member)
6630 << static_cast<int>(RT);
6631 }
6632 }
6633
6634 if (!HasBinding && isResourceRecordTypeOrArrayOf(VD))
6635 SemaRef.Diag(Loc: VD->getLocation(), DiagID: diag::warn_hlsl_implicit_binding);
6636}
6637namespace {
6638class InitListTransformer {
6639 Sema &S;
6640 ASTContext &Ctx;
6641 QualType InitTy;
6642 QualType *DstIt = nullptr;
6643 Expr **ArgIt = nullptr;
6644 // Is wrapping the destination type iterator required? This is only used for
6645 // incomplete array types where we loop over the destination type since we
6646 // don't know the full number of elements from the declaration.
6647 bool Wrap;
6648
6649 bool castInitializer(Expr *E) {
6650 assert(DstIt && "This should always be something!");
6651 if (DstIt == DestTypes.end()) {
6652 if (!Wrap) {
6653 ArgExprs.push_back(Elt: E);
6654 // This is odd, but it isn't technically a failure due to conversion, we
6655 // handle mismatched counts of arguments differently.
6656 return true;
6657 }
6658 DstIt = DestTypes.begin();
6659 }
6660 InitializedEntity Entity = InitializedEntity::InitializeParameter(
6661 Context&: Ctx, Type: *DstIt, /* Consumed (ObjC) */ Consumed: false);
6662 ExprResult Res = S.PerformCopyInitialization(Entity, EqualLoc: E->getBeginLoc(), Init: E);
6663 if (Res.isInvalid())
6664 return false;
6665 Expr *Init = Res.get();
6666 ArgExprs.push_back(Elt: Init);
6667 DstIt++;
6668 return true;
6669 }
6670
6671 bool buildInitializerListImpl(Expr *E) {
6672 // If this is an initialization list, traverse the sub initializers.
6673 if (auto *Init = dyn_cast<InitListExpr>(Val: E)) {
6674 for (auto *SubInit : Init->inits())
6675 if (!buildInitializerListImpl(E: SubInit))
6676 return false;
6677 return true;
6678 }
6679
6680 // If this is a scalar type, just enqueue the expression.
6681 QualType Ty = E->getType().getDesugaredType(Context: Ctx);
6682
6683 if (Ty->isScalarType() || (Ty->isRecordType() && !Ty->isAggregateType()) ||
6684 Ty->isHLSLAttributedResourceType())
6685 return castInitializer(E);
6686
6687 // If this is an aggregate type and a prvalue, create an xvalue temporary
6688 // so the member accesses will be xvalues. Wrap it in OpaqueExpr to make
6689 // sure codegen will not generate duplicate copies.
6690 if (E->isPRValue() && Ty->isAggregateType()) {
6691 ExprResult TmpExpr = S.TemporaryMaterializationConversion(E);
6692 if (TmpExpr.isInvalid())
6693 return false;
6694 E = TmpExpr.get();
6695 E = new (Ctx) OpaqueValueExpr(E->getBeginLoc(), E->getType(),
6696 E->getValueKind(), E->getObjectKind(), E);
6697 }
6698
6699 if (auto *VecTy = Ty->getAs<VectorType>()) {
6700 uint64_t Size = VecTy->getNumElements();
6701
6702 QualType SizeTy = Ctx.getSizeType();
6703 uint64_t SizeTySize = Ctx.getTypeSize(T: SizeTy);
6704 for (uint64_t I = 0; I < Size; ++I) {
6705 auto *Idx = IntegerLiteral::Create(C: Ctx, V: llvm::APInt(SizeTySize, I),
6706 type: SizeTy, l: SourceLocation());
6707
6708 ExprResult ElExpr = S.CreateBuiltinArraySubscriptExpr(
6709 Base: E, LLoc: E->getBeginLoc(), Idx, RLoc: E->getEndLoc());
6710 if (ElExpr.isInvalid())
6711 return false;
6712 if (!castInitializer(E: ElExpr.get()))
6713 return false;
6714 }
6715 return true;
6716 }
6717 if (auto *MTy = Ty->getAs<ConstantMatrixType>()) {
6718 unsigned Rows = MTy->getNumRows();
6719 unsigned Cols = MTy->getNumColumns();
6720 QualType ElemTy = MTy->getElementType();
6721
6722 for (unsigned R = 0; R < Rows; ++R) {
6723 for (unsigned C = 0; C < Cols; ++C) {
6724 // row index literal
6725 Expr *RowIdx = IntegerLiteral::Create(
6726 C: Ctx, V: llvm::APInt(Ctx.getIntWidth(T: Ctx.IntTy), R), type: Ctx.IntTy,
6727 l: E->getBeginLoc());
6728 // column index literal
6729 Expr *ColIdx = IntegerLiteral::Create(
6730 C: Ctx, V: llvm::APInt(Ctx.getIntWidth(T: Ctx.IntTy), C), type: Ctx.IntTy,
6731 l: E->getBeginLoc());
6732 ExprResult ElExpr = S.CreateBuiltinMatrixSubscriptExpr(
6733 Base: E, RowIdx, ColumnIdx: ColIdx, RBLoc: E->getEndLoc());
6734 if (ElExpr.isInvalid())
6735 return false;
6736 if (!castInitializer(E: ElExpr.get()))
6737 return false;
6738 ElExpr.get()->setType(ElemTy);
6739 }
6740 }
6741 return true;
6742 }
6743
6744 if (auto *ArrTy = dyn_cast<ConstantArrayType>(Val: Ty.getTypePtr())) {
6745 uint64_t Size = ArrTy->getZExtSize();
6746 QualType SizeTy = Ctx.getSizeType();
6747 uint64_t SizeTySize = Ctx.getTypeSize(T: SizeTy);
6748 for (uint64_t I = 0; I < Size; ++I) {
6749 auto *Idx = IntegerLiteral::Create(C: Ctx, V: llvm::APInt(SizeTySize, I),
6750 type: SizeTy, l: SourceLocation());
6751 ExprResult ElExpr = S.CreateBuiltinArraySubscriptExpr(
6752 Base: E, LLoc: E->getBeginLoc(), Idx, RLoc: E->getEndLoc());
6753 if (ElExpr.isInvalid())
6754 return false;
6755 if (!buildInitializerListImpl(E: ElExpr.get()))
6756 return false;
6757 }
6758 return true;
6759 }
6760
6761 if (auto *RD = Ty->getAsCXXRecordDecl()) {
6762 llvm::SmallVector<CXXRecordDecl *> RecordDecls;
6763 RecordDecls.push_back(Elt: RD);
6764 while (RecordDecls.back()->getNumBases()) {
6765 CXXRecordDecl *D = RecordDecls.back();
6766 assert(D->getNumBases() == 1 &&
6767 "HLSL doesn't support multiple inheritance");
6768 RecordDecls.push_back(
6769 Elt: D->bases_begin()->getType()->castAsCXXRecordDecl());
6770 }
6771 while (!RecordDecls.empty()) {
6772 CXXRecordDecl *RD = RecordDecls.pop_back_val();
6773 for (auto *FD : RD->fields()) {
6774 if (FD->isUnnamedBitField())
6775 continue;
6776 DeclAccessPair Found = DeclAccessPair::make(D: FD, AS: FD->getAccess());
6777 DeclarationNameInfo NameInfo(FD->getDeclName(), E->getBeginLoc());
6778 ExprResult Res = S.BuildFieldReferenceExpr(
6779 BaseExpr: E, IsArrow: false, OpLoc: E->getBeginLoc(), SS: CXXScopeSpec(), Field: FD, FoundDecl: Found, MemberNameInfo: NameInfo);
6780 if (Res.isInvalid())
6781 return false;
6782 if (!buildInitializerListImpl(E: Res.get()))
6783 return false;
6784 }
6785 }
6786 }
6787 return true;
6788 }
6789
6790 Expr *generateInitListsImpl(QualType Ty) {
6791 Ty = Ty.getDesugaredType(Context: Ctx);
6792 assert(ArgIt != ArgExprs.end() && "Something is off in iteration!");
6793 if (Ty->isScalarType() || (Ty->isRecordType() && !Ty->isAggregateType()) ||
6794 Ty->isHLSLAttributedResourceType())
6795 return *(ArgIt++);
6796
6797 llvm::SmallVector<Expr *> Inits;
6798 if (Ty->isVectorType() || Ty->isConstantArrayType() ||
6799 Ty->isConstantMatrixType()) {
6800 QualType ElTy;
6801 uint64_t Size = 0;
6802 if (auto *ATy = Ty->getAs<VectorType>()) {
6803 ElTy = ATy->getElementType();
6804 Size = ATy->getNumElements();
6805 } else if (auto *CMTy = Ty->getAs<ConstantMatrixType>()) {
6806 ElTy = CMTy->getElementType();
6807 Size = CMTy->getNumElementsFlattened();
6808 } else {
6809 auto *VTy = cast<ConstantArrayType>(Val: Ty.getTypePtr());
6810 ElTy = VTy->getElementType();
6811 Size = VTy->getZExtSize();
6812 }
6813 for (uint64_t I = 0; I < Size; ++I)
6814 Inits.push_back(Elt: generateInitListsImpl(Ty: ElTy));
6815 }
6816 if (auto *RD = Ty->getAsCXXRecordDecl()) {
6817 llvm::SmallVector<CXXRecordDecl *> RecordDecls;
6818 RecordDecls.push_back(Elt: RD);
6819 while (RecordDecls.back()->getNumBases()) {
6820 CXXRecordDecl *D = RecordDecls.back();
6821 assert(D->getNumBases() == 1 &&
6822 "HLSL doesn't support multiple inheritance");
6823 RecordDecls.push_back(
6824 Elt: D->bases_begin()->getType()->castAsCXXRecordDecl());
6825 }
6826 while (!RecordDecls.empty()) {
6827 CXXRecordDecl *RD = RecordDecls.pop_back_val();
6828 for (auto *FD : RD->fields())
6829 if (!FD->isUnnamedBitField())
6830 Inits.push_back(Elt: generateInitListsImpl(Ty: FD->getType()));
6831 }
6832 }
6833 auto *NewInit =
6834 new (Ctx) InitListExpr(Ctx, Inits.front()->getBeginLoc(), Inits,
6835 Inits.back()->getEndLoc(), /*isExplicit=*/false);
6836 NewInit->setType(Ty);
6837 return NewInit;
6838 }
6839
6840public:
6841 llvm::SmallVector<QualType, 16> DestTypes;
6842 llvm::SmallVector<Expr *, 16> ArgExprs;
6843 InitListTransformer(Sema &SemaRef, const InitializedEntity &Entity)
6844 : S(SemaRef), Ctx(SemaRef.getASTContext()),
6845 Wrap(Entity.getType()->isIncompleteArrayType()) {
6846 InitTy = Entity.getType().getNonReferenceType();
6847 // When we're generating initializer lists for incomplete array types we
6848 // need to wrap around both when building the initializers and when
6849 // generating the final initializer lists.
6850 if (Wrap) {
6851 assert(InitTy->isIncompleteArrayType());
6852 const IncompleteArrayType *IAT = Ctx.getAsIncompleteArrayType(T: InitTy);
6853 InitTy = IAT->getElementType();
6854 }
6855 BuildFlattenedTypeList(BaseTy: InitTy, List&: DestTypes);
6856 DstIt = DestTypes.begin();
6857 }
6858
6859 bool buildInitializerList(Expr *E) { return buildInitializerListImpl(E); }
6860
6861 Expr *generateInitLists() {
6862 assert(!ArgExprs.empty() &&
6863 "Call buildInitializerList to generate argument expressions.");
6864 ArgIt = ArgExprs.begin();
6865 if (!Wrap)
6866 return generateInitListsImpl(Ty: InitTy);
6867 llvm::SmallVector<Expr *> Inits;
6868 while (ArgIt != ArgExprs.end())
6869 Inits.push_back(Elt: generateInitListsImpl(Ty: InitTy));
6870
6871 auto *NewInit =
6872 new (Ctx) InitListExpr(Ctx, Inits.front()->getBeginLoc(), Inits,
6873 Inits.back()->getEndLoc(), /*isExplicit=*/false);
6874 llvm::APInt ArySize(64, Inits.size());
6875 NewInit->setType(Ctx.getConstantArrayType(EltTy: InitTy, ArySize, SizeExpr: nullptr,
6876 ASM: ArraySizeModifier::Normal, IndexTypeQuals: 0));
6877 return NewInit;
6878 }
6879};
6880} // namespace
6881
6882// Recursively detect any incomplete array anywhere in the type graph,
6883// including arrays, struct fields, and base classes.
6884static bool containsIncompleteArrayType(QualType Ty) {
6885 Ty = Ty.getCanonicalType();
6886
6887 // Array types
6888 if (const ArrayType *AT = dyn_cast<ArrayType>(Val&: Ty)) {
6889 if (isa<IncompleteArrayType>(Val: AT))
6890 return true;
6891 return containsIncompleteArrayType(Ty: AT->getElementType());
6892 }
6893
6894 // Record (struct/class) types
6895 if (const auto *RT = Ty->getAs<RecordType>()) {
6896 const RecordDecl *RD = RT->getDecl();
6897
6898 // Walk base classes (for C++ / HLSL structs with inheritance)
6899 if (const auto *CXXRD = dyn_cast<CXXRecordDecl>(Val: RD)) {
6900 for (const CXXBaseSpecifier &Base : CXXRD->bases()) {
6901 if (containsIncompleteArrayType(Ty: Base.getType()))
6902 return true;
6903 }
6904 }
6905
6906 // Walk fields
6907 for (const FieldDecl *F : RD->fields()) {
6908 if (containsIncompleteArrayType(Ty: F->getType()))
6909 return true;
6910 }
6911 }
6912
6913 return false;
6914}
6915
6916bool SemaHLSL::transformInitList(const InitializedEntity &Entity,
6917 InitListExpr *Init) {
6918 // If the initializer is a scalar, just return it.
6919 if (Init->getType()->isScalarType())
6920 return true;
6921 ASTContext &Ctx = SemaRef.getASTContext();
6922 InitListTransformer ILT(SemaRef, Entity);
6923
6924 for (unsigned I = 0; I < Init->getNumInits(); ++I) {
6925 Expr *E = Init->getInit(Init: I);
6926 if (E->HasSideEffects(Ctx)) {
6927 QualType Ty = E->getType();
6928 if (Ty->isRecordType())
6929 E = new (Ctx) MaterializeTemporaryExpr(Ty, E, E->isLValue());
6930 E = new (Ctx) OpaqueValueExpr(E->getBeginLoc(), Ty, E->getValueKind(),
6931 E->getObjectKind(), E);
6932 Init->setInit(Init: I, expr: E);
6933 }
6934 if (!ILT.buildInitializerList(E))
6935 return false;
6936 }
6937 size_t ExpectedSize = ILT.DestTypes.size();
6938 size_t ActualSize = ILT.ArgExprs.size();
6939 if (ExpectedSize == 0 && ActualSize == 0)
6940 return true;
6941
6942 // Reject empty initializer if *any* incomplete array exists structurally
6943 if (ActualSize == 0 && containsIncompleteArrayType(Ty: Entity.getType())) {
6944 QualType InitTy = Entity.getType().getNonReferenceType();
6945 if (InitTy.hasAddressSpace())
6946 InitTy = SemaRef.getASTContext().removeAddrSpaceQualType(T: InitTy);
6947
6948 SemaRef.Diag(Loc: Init->getBeginLoc(), DiagID: diag::err_hlsl_incorrect_num_initializers)
6949 << /*TooManyOrFew=*/(int)(ExpectedSize < ActualSize) << InitTy
6950 << /*ExpectedSize=*/ExpectedSize << /*ActualSize=*/ActualSize;
6951 return false;
6952 }
6953
6954 // We infer size after validating legality.
6955 // For incomplete arrays it is completely arbitrary to choose whether we think
6956 // the user intended fewer or more elements. This implementation assumes that
6957 // the user intended more, and errors that there are too few initializers to
6958 // complete the final element.
6959 if (Entity.getType()->isIncompleteArrayType()) {
6960 assert(ExpectedSize > 0 &&
6961 "The expected size of an incomplete array type must be at least 1.");
6962 ExpectedSize =
6963 ((ActualSize + ExpectedSize - 1) / ExpectedSize) * ExpectedSize;
6964 }
6965
6966 // An initializer list might be attempting to initialize a reference or
6967 // rvalue-reference. When checking the initializer we should look through
6968 // the reference.
6969 QualType InitTy = Entity.getType().getNonReferenceType();
6970 if (InitTy.hasAddressSpace())
6971 InitTy = SemaRef.getASTContext().removeAddrSpaceQualType(T: InitTy);
6972 if (ExpectedSize != ActualSize) {
6973 int TooManyOrFew = ActualSize > ExpectedSize ? 1 : 0;
6974 SemaRef.Diag(Loc: Init->getBeginLoc(), DiagID: diag::err_hlsl_incorrect_num_initializers)
6975 << TooManyOrFew << InitTy << ExpectedSize << ActualSize;
6976 return false;
6977 }
6978
6979 // generateInitListsImpl will always return an InitListExpr here, because the
6980 // scalar case is handled above.
6981 auto *NewInit = cast<InitListExpr>(Val: ILT.generateInitLists());
6982 Init->resizeInits(Context: Ctx, NumInits: NewInit->getNumInits());
6983 for (unsigned I = 0; I < NewInit->getNumInits(); ++I)
6984 Init->updateInit(C: Ctx, Init: I, expr: NewInit->getInit(Init: I));
6985 return true;
6986}
6987
6988static QualType ReportMatrixInvalidMember(Sema &S, StringRef Name,
6989 StringRef Expected,
6990 SourceLocation OpLoc,
6991 SourceLocation CompLoc) {
6992 S.Diag(Loc: OpLoc, DiagID: diag::err_builtin_matrix_invalid_member)
6993 << Name << Expected << SourceRange(CompLoc);
6994 return QualType();
6995}
6996
6997QualType SemaHLSL::checkMatrixComponent(Sema &S, QualType baseType,
6998 ExprValueKind &VK, SourceLocation OpLoc,
6999 const IdentifierInfo *CompName,
7000 SourceLocation CompLoc) {
7001 const auto *MT = baseType->castAs<ConstantMatrixType>();
7002 StringRef AccessorName = CompName->getName();
7003 assert(!AccessorName.empty() && "Matrix Accessor must have a name");
7004
7005 unsigned Rows = MT->getNumRows();
7006 unsigned Cols = MT->getNumColumns();
7007 bool IsZeroBasedAccessor = false;
7008 unsigned ChunkLen = 0;
7009 if (AccessorName.size() < 2)
7010 return ReportMatrixInvalidMember(S, Name: AccessorName,
7011 Expected: "length 4 for zero based: \'_mRC\' or "
7012 "length 3 for one-based: \'_RC\' accessor",
7013 OpLoc, CompLoc);
7014
7015 if (AccessorName[0] == '_') {
7016 if (AccessorName[1] == 'm') {
7017 IsZeroBasedAccessor = true;
7018 ChunkLen = 4; // zero-based: "_mRC"
7019 } else {
7020 ChunkLen = 3; // one-based: "_RC"
7021 }
7022 } else
7023 return ReportMatrixInvalidMember(
7024 S, Name: AccessorName, Expected: "zero based: \'_mRC\' or one-based: \'_RC\' accessor",
7025 OpLoc, CompLoc);
7026
7027 if (AccessorName.size() % ChunkLen != 0) {
7028 const llvm::StringRef Expected = IsZeroBasedAccessor
7029 ? "zero based: '_mRC' accessor"
7030 : "one-based: '_RC' accessor";
7031
7032 return ReportMatrixInvalidMember(S, Name: AccessorName, Expected, OpLoc, CompLoc);
7033 }
7034
7035 auto isDigit = [](char c) { return c >= '0' && c <= '9'; };
7036 auto isZeroBasedIndex = [](unsigned i) { return i <= 3; };
7037 auto isOneBasedIndex = [](unsigned i) { return i >= 1 && i <= 4; };
7038
7039 bool HasRepeated = false;
7040 SmallVector<bool, 16> Seen(Rows * Cols, false);
7041 unsigned NumComponents = 0;
7042 const char *Begin = AccessorName.data();
7043
7044 for (unsigned I = 0, E = AccessorName.size(); I < E; I += ChunkLen) {
7045 const char *Chunk = Begin + I;
7046 char RowChar = 0, ColChar = 0;
7047 if (IsZeroBasedAccessor) {
7048 // Zero-based: "_mRC"
7049 if (Chunk[0] != '_' || Chunk[1] != 'm') {
7050 char Bad = (Chunk[0] != '_') ? Chunk[0] : Chunk[1];
7051 return ReportMatrixInvalidMember(
7052 S, Name: StringRef(&Bad, 1), Expected: "\'_m\' prefix",
7053 OpLoc: OpLoc.getLocWithOffset(Offset: I + (Bad == Chunk[0] ? 1 : 2)), CompLoc);
7054 }
7055 RowChar = Chunk[2];
7056 ColChar = Chunk[3];
7057 } else {
7058 // One-based: "_RC"
7059 if (Chunk[0] != '_')
7060 return ReportMatrixInvalidMember(
7061 S, Name: StringRef(&Chunk[0], 1), Expected: "\'_\' prefix",
7062 OpLoc: OpLoc.getLocWithOffset(Offset: I + 1), CompLoc);
7063 RowChar = Chunk[1];
7064 ColChar = Chunk[2];
7065 }
7066
7067 // Must be digits.
7068 bool IsDigitsError = false;
7069 if (!isDigit(RowChar)) {
7070 unsigned BadPos = IsZeroBasedAccessor ? 2 : 1;
7071 ReportMatrixInvalidMember(S, Name: StringRef(&RowChar, 1), Expected: "row as integer",
7072 OpLoc: OpLoc.getLocWithOffset(Offset: I + BadPos + 1),
7073 CompLoc);
7074 IsDigitsError = true;
7075 }
7076
7077 if (!isDigit(ColChar)) {
7078 unsigned BadPos = IsZeroBasedAccessor ? 3 : 2;
7079 ReportMatrixInvalidMember(S, Name: StringRef(&ColChar, 1), Expected: "column as integer",
7080 OpLoc: OpLoc.getLocWithOffset(Offset: I + BadPos + 1),
7081 CompLoc);
7082 IsDigitsError = true;
7083 }
7084 if (IsDigitsError)
7085 return QualType();
7086
7087 unsigned Row = RowChar - '0';
7088 unsigned Col = ColChar - '0';
7089
7090 bool HasIndexingError = false;
7091 if (IsZeroBasedAccessor) {
7092 // 0-based [0..3]
7093 if (!isZeroBasedIndex(Row)) {
7094 S.Diag(Loc: OpLoc, DiagID: diag::err_hlsl_matrix_element_not_in_bounds)
7095 << /*row*/ 0 << /*zero-based*/ 0 << SourceRange(CompLoc);
7096 HasIndexingError = true;
7097 }
7098 if (!isZeroBasedIndex(Col)) {
7099 S.Diag(Loc: OpLoc, DiagID: diag::err_hlsl_matrix_element_not_in_bounds)
7100 << /*col*/ 1 << /*zero-based*/ 0 << SourceRange(CompLoc);
7101 HasIndexingError = true;
7102 }
7103 } else {
7104 // 1-based [1..4]
7105 if (!isOneBasedIndex(Row)) {
7106 S.Diag(Loc: OpLoc, DiagID: diag::err_hlsl_matrix_element_not_in_bounds)
7107 << /*row*/ 0 << /*one-based*/ 1 << SourceRange(CompLoc);
7108 HasIndexingError = true;
7109 }
7110 if (!isOneBasedIndex(Col)) {
7111 S.Diag(Loc: OpLoc, DiagID: diag::err_hlsl_matrix_element_not_in_bounds)
7112 << /*col*/ 1 << /*one-based*/ 1 << SourceRange(CompLoc);
7113 HasIndexingError = true;
7114 }
7115 // Convert to 0-based after range checking.
7116 --Row;
7117 --Col;
7118 }
7119
7120 if (HasIndexingError)
7121 return QualType();
7122
7123 // Note: matrix swizzle index is hard coded. That means Row and Col can
7124 // potentially be larger than Rows and Cols if matrix size is less than
7125 // the max index size.
7126 bool HasBoundsError = false;
7127 if (Row >= Rows) {
7128 Diag(Loc: OpLoc, DiagID: diag::err_hlsl_matrix_index_out_of_bounds)
7129 << /*Row*/ 0 << Row << Rows << SourceRange(CompLoc);
7130 HasBoundsError = true;
7131 }
7132 if (Col >= Cols) {
7133 Diag(Loc: OpLoc, DiagID: diag::err_hlsl_matrix_index_out_of_bounds)
7134 << /*Col*/ 1 << Col << Cols << SourceRange(CompLoc);
7135 HasBoundsError = true;
7136 }
7137 if (HasBoundsError)
7138 return QualType();
7139
7140 unsigned FlatIndex = Row * Cols + Col;
7141 if (Seen[FlatIndex])
7142 HasRepeated = true;
7143 Seen[FlatIndex] = true;
7144 ++NumComponents;
7145 }
7146 if (NumComponents == 0 || NumComponents > 4) {
7147 S.Diag(Loc: OpLoc, DiagID: diag::err_hlsl_matrix_swizzle_invalid_length)
7148 << NumComponents << SourceRange(CompLoc);
7149 return QualType();
7150 }
7151
7152 QualType ElemTy = MT->getElementType();
7153 if (NumComponents == 1)
7154 return ElemTy;
7155 QualType VT = S.Context.getExtVectorType(VectorType: ElemTy, NumElts: NumComponents);
7156 if (HasRepeated)
7157 VK = VK_PRValue;
7158
7159 for (Sema::ExtVectorDeclsType::iterator
7160 I = S.ExtVectorDecls.begin(source: S.getExternalSource()),
7161 E = S.ExtVectorDecls.end();
7162 I != E; ++I) {
7163 if ((*I)->getUnderlyingType() == VT)
7164 return S.Context.getTypedefType(Keyword: ElaboratedTypeKeyword::None,
7165 /*Qualifier=*/std::nullopt, Decl: *I);
7166 }
7167
7168 return VT;
7169}
7170
7171bool SemaHLSL::handleInitialization(VarDecl *VDecl, Expr *&Init) {
7172 // If initializing a local resource, track the resource binding it is using
7173 if (VDecl->getType()->isHLSLResourceRecord() && !VDecl->hasGlobalStorage())
7174 trackLocalResource(VD: VDecl, E: Init);
7175
7176 const HLSLVkConstantIdAttr *ConstIdAttr =
7177 VDecl->getAttr<HLSLVkConstantIdAttr>();
7178 if (!ConstIdAttr)
7179 return true;
7180
7181 ASTContext &Context = SemaRef.getASTContext();
7182
7183 APValue InitValue;
7184 if (!Init->isCXX11ConstantExpr(Ctx: Context, Result&: InitValue)) {
7185 Diag(Loc: VDecl->getLocation(), DiagID: diag::err_specialization_const);
7186 VDecl->setInvalidDecl();
7187 return false;
7188 }
7189
7190 Builtin::ID BID =
7191 getSpecConstBuiltinId(Type: VDecl->getType()->getUnqualifiedDesugaredType());
7192
7193 // Argument 1: The ID from the attribute
7194 int ConstantID = ConstIdAttr->getId();
7195 llvm::APInt IDVal(Context.getIntWidth(T: Context.IntTy), ConstantID);
7196 Expr *IdExpr = IntegerLiteral::Create(C: Context, V: IDVal, type: Context.IntTy,
7197 l: ConstIdAttr->getLocation());
7198
7199 SmallVector<Expr *, 2> Args = {IdExpr, Init};
7200 Expr *C = SemaRef.BuildBuiltinCallExpr(Loc: Init->getExprLoc(), Id: BID, CallArgs: Args);
7201 if (C->getType()->getCanonicalTypeUnqualified() !=
7202 VDecl->getType()->getCanonicalTypeUnqualified()) {
7203 C = SemaRef
7204 .BuildCStyleCastExpr(LParenLoc: SourceLocation(),
7205 Ty: Context.getTrivialTypeSourceInfo(
7206 T: Init->getType(), Loc: Init->getExprLoc()),
7207 RParenLoc: SourceLocation(), Op: C)
7208 .get();
7209 }
7210 Init = C;
7211 return true;
7212}
7213
7214QualType SemaHLSL::ActOnTemplateShorthand(TemplateDecl *Template,
7215 SourceLocation NameLoc) {
7216 if (!Template)
7217 return QualType();
7218
7219 DeclContext *DC = Template->getDeclContext();
7220 if (!DC->isNamespace() || !cast<NamespaceDecl>(Val: DC)->getIdentifier() ||
7221 cast<NamespaceDecl>(Val: DC)->getName() != "hlsl")
7222 return QualType();
7223
7224 TemplateParameterList *Params = Template->getTemplateParameters();
7225 if (!Params || Params->size() != 1)
7226 return QualType();
7227
7228 if (!Template->isImplicit())
7229 return QualType();
7230
7231 // We manually extract default arguments here instead of letting
7232 // CheckTemplateIdType handle it. This ensures that for resource types that
7233 // lack a default argument (like Buffer), we return a null QualType, which
7234 // triggers the "requires template arguments" error rather than a less
7235 // descriptive "too few template arguments" error.
7236 TemplateArgumentListInfo TemplateArgs(NameLoc, NameLoc);
7237 for (NamedDecl *P : *Params) {
7238 if (auto *TTP = dyn_cast<TemplateTypeParmDecl>(Val: P)) {
7239 if (TTP->hasDefaultArgument()) {
7240 TemplateArgs.addArgument(Loc: TTP->getDefaultArgument());
7241 continue;
7242 }
7243 } else if (auto *NTTP = dyn_cast<NonTypeTemplateParmDecl>(Val: P)) {
7244 if (NTTP->hasDefaultArgument()) {
7245 TemplateArgs.addArgument(Loc: NTTP->getDefaultArgument());
7246 continue;
7247 }
7248 } else if (auto *TTPD = dyn_cast<TemplateTemplateParmDecl>(Val: P)) {
7249 if (TTPD->hasDefaultArgument()) {
7250 TemplateArgs.addArgument(Loc: TTPD->getDefaultArgument());
7251 continue;
7252 }
7253 }
7254 return QualType();
7255 }
7256
7257 return SemaRef.CheckTemplateIdType(
7258 Keyword: ElaboratedTypeKeyword::None, Template: TemplateName(Template), TemplateLoc: NameLoc,
7259 TemplateArgs, Scope: nullptr, /*ForNestedNameSpecifier=*/false);
7260}
7261