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