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// Note: returning true in this case results in CheckBuiltinFunctionCall
4124// returning an ExprError
4125bool SemaHLSL::CheckBuiltinFunctionCall(unsigned BuiltinID, CallExpr *TheCall) {
4126 switch (BuiltinID) {
4127 case Builtin::BI__builtin_hlsl_adduint64: {
4128 if (SemaRef.checkArgCount(TheCall, 2))
4129 return true;
4130
4131 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
4133 return true;
4134
4135 // ensure arg integers are 32-bits
4136 if (CheckExpectedBitWidth(&SemaRef, TheCall, 0, 32))
4137 return true;
4138
4139 // ensure both args are vectors of total bit size of a multiple of 64
4140 auto *VTy = TheCall->getArg(0)->getType()->getAs<VectorType>();
4141 int NumElementsArg = VTy->getNumElements();
4142 if (NumElementsArg != 2 && NumElementsArg != 4) {
4143 SemaRef.Diag(TheCall->getBeginLoc(), diag::err_vector_incorrect_bit_count)
4144 << 1 /*a multiple of*/ << 64 << NumElementsArg * 32;
4145 return true;
4146 }
4147
4148 // ensure first arg and second arg have the same type
4149 if (CheckAllArgsHaveSameType(&SemaRef, TheCall))
4150 return true;
4151
4152 ExprResult A = TheCall->getArg(0);
4153 QualType ArgTyA = A.get()->getType();
4154 // return type is the same as the input type
4155 TheCall->setType(ArgTyA);
4156 break;
4157 }
4158 case Builtin::BI__builtin_hlsl_resource_getpointer: {
4159 if (SemaRef.checkArgCountRange(TheCall, 1, 2) ||
4160 CheckResourceHandle(&SemaRef, TheCall, 0) ||
4161 (TheCall->getNumArgs() == 2 && CheckIndexType(&SemaRef, TheCall, 1)))
4162 return true;
4163
4164 auto *ResourceTy =
4165 TheCall->getArg(0)->getType()->castAs<HLSLAttributedResourceType>();
4166 QualType ContainedTy = ResourceTy->getContainedType();
4167 auto ReturnType = SemaRef.Context.getAddrSpaceQualType(
4168 ContainedTy,
4169 getLangASFromResourceClass(ResourceTy->getAttrs().ResourceClass));
4170 ReturnType = SemaRef.Context.getPointerType(ReturnType);
4171 TheCall->setType(ReturnType);
4172
4173 break;
4174 }
4175 case Builtin::BI__builtin_hlsl_resource_getpointer_typed: {
4176 if (SemaRef.checkArgCount(TheCall, 3) ||
4177 CheckResourceHandle(&SemaRef, TheCall, 0) ||
4178 CheckIndexType(&SemaRef, TheCall, 1))
4179 return true;
4180
4181 QualType ElementTy = TheCall->getArg(2)->getType();
4182 assert(ElementTy->isPointerType() &&
4183 "expected pointer type for second argument");
4184 ElementTy = ElementTy->getPointeeType();
4185
4186 // Reject array types
4187 if (ElementTy->isArrayType())
4188 return SemaRef.Diag(
4189 cast<FunctionDecl>(SemaRef.CurContext)->getPointOfInstantiation(),
4190 diag::err_invalid_use_of_array_type);
4191
4192 auto *ResourceTy =
4193 TheCall->getArg(0)->getType()->castAs<HLSLAttributedResourceType>();
4194 auto ReturnType = SemaRef.Context.getAddrSpaceQualType(
4195 ElementTy,
4196 getLangASFromResourceClass(ResourceTy->getAttrs().ResourceClass));
4197 ReturnType = SemaRef.Context.getPointerType(ReturnType);
4198 TheCall->setType(ReturnType);
4199
4200 break;
4201 }
4202 case Builtin::BI__builtin_hlsl_transpose_if_memory_is_row_major: {
4203 if (SemaRef.checkArgCount(TheCall, 2) ||
4204 CheckArgTypeMatches(&SemaRef, TheCall->getArg(1),
4205 SemaRef.getASTContext().IntTy))
4206 return true;
4207
4208 TheCall->setType(TheCall->getArg(0)->getType());
4209
4210 break;
4211 }
4212 case Builtin::BI__builtin_hlsl_resource_load_with_status: {
4213 if (SemaRef.checkArgCount(TheCall, 3) ||
4214 CheckResourceHandle(&SemaRef, TheCall, 0) ||
4215 CheckArgTypeMatches(&SemaRef, TheCall->getArg(1),
4216 SemaRef.getASTContext().UnsignedIntTy) ||
4217 CheckArgTypeMatches(&SemaRef, TheCall->getArg(2),
4218 SemaRef.getASTContext().UnsignedIntTy) ||
4219 CheckModifiableLValue(&SemaRef, TheCall, 2))
4220 return true;
4221
4222 auto *ResourceTy =
4223 TheCall->getArg(0)->getType()->castAs<HLSLAttributedResourceType>();
4224 QualType ReturnType = ResourceTy->getContainedType();
4225 TheCall->setType(ReturnType);
4226
4227 break;
4228 }
4229 case Builtin::BI__builtin_hlsl_resource_load_with_status_typed: {
4230 if (SemaRef.checkArgCount(TheCall, 4) ||
4231 CheckResourceHandle(&SemaRef, TheCall, 0) ||
4232 CheckArgTypeMatches(&SemaRef, TheCall->getArg(1),
4233 SemaRef.getASTContext().UnsignedIntTy) ||
4234 CheckArgTypeMatches(&SemaRef, TheCall->getArg(2),
4235 SemaRef.getASTContext().UnsignedIntTy) ||
4236 CheckModifiableLValue(&SemaRef, TheCall, 2))
4237 return true;
4238
4239 QualType ReturnType = TheCall->getArg(3)->getType();
4240 assert(ReturnType->isPointerType() &&
4241 "expected pointer type for second argument");
4242 ReturnType = ReturnType->getPointeeType();
4243
4244 // Reject array types
4245 if (ReturnType->isArrayType())
4246 return SemaRef.Diag(
4247 cast<FunctionDecl>(SemaRef.CurContext)->getPointOfInstantiation(),
4248 diag::err_invalid_use_of_array_type);
4249
4250 TheCall->setType(ReturnType);
4251
4252 break;
4253 }
4254 case Builtin::BI__builtin_hlsl_resource_load_level:
4255 return CheckLoadLevelBuiltin(SemaRef, TheCall);
4256 case Builtin::BI__builtin_hlsl_resource_load_ms:
4257 return CheckLoadMSBuiltin(SemaRef, TheCall);
4258 case Builtin::BI__builtin_hlsl_resource_sample:
4260 case Builtin::BI__builtin_hlsl_resource_sample_bias:
4262 case Builtin::BI__builtin_hlsl_resource_sample_grad:
4264 case Builtin::BI__builtin_hlsl_resource_sample_level:
4266 case Builtin::BI__builtin_hlsl_resource_sample_cmp:
4268 case Builtin::BI__builtin_hlsl_resource_sample_cmp_level_zero:
4270 case Builtin::BI__builtin_hlsl_resource_calculate_lod:
4271 case Builtin::BI__builtin_hlsl_resource_calculate_lod_unclamped:
4272 return CheckCalculateLodBuiltin(SemaRef, TheCall);
4273 case Builtin::BI__builtin_hlsl_resource_gather:
4274 return CheckGatherBuiltin(SemaRef, TheCall, /*IsCmp=*/false);
4275 case Builtin::BI__builtin_hlsl_resource_gather_cmp:
4276 return CheckGatherBuiltin(SemaRef, TheCall, /*IsCmp=*/true);
4277 case Builtin::BI__builtin_hlsl_resource_uninitializedhandle: {
4278 assert(TheCall->getNumArgs() == 1 && "expected 1 arg");
4279 // Update return type to be the attributed resource type from arg0.
4280 QualType ResourceTy = TheCall->getArg(0)->getType();
4281 TheCall->setType(ResourceTy);
4282 break;
4283 }
4284 case Builtin::BI__builtin_hlsl_resource_handlefrombinding: {
4285 assert(TheCall->getNumArgs() == 6 && "expected 6 args");
4286 // Update return type to be the attributed resource type from arg0.
4287 QualType ResourceTy = TheCall->getArg(0)->getType();
4288 TheCall->setType(ResourceTy);
4289 break;
4290 }
4291 case Builtin::BI__builtin_hlsl_resource_handlefromimplicitbinding: {
4292 assert(TheCall->getNumArgs() == 6 && "expected 6 args");
4293 // Update return type to be the attributed resource type from arg0.
4294 QualType ResourceTy = TheCall->getArg(0)->getType();
4295 TheCall->setType(ResourceTy);
4296 break;
4297 }
4298 case Builtin::BI__builtin_hlsl_resource_counterhandlefromimplicitbinding: {
4299 assert(TheCall->getNumArgs() == 3 && "expected 3 args");
4300 // Update return type to be the attributed resource type from arg0
4301 // with added IsCounter flag.
4302 QualType MainHandleTy = TheCall->getArg(0)->getType();
4303 QualType CounterHandleTy =
4304 createCounterHandleType(SemaRef.getASTContext(), MainHandleTy);
4305 TheCall->setType(CounterHandleTy);
4306 break;
4307 }
4308 case Builtin::BI__builtin_hlsl_resource_handlefromheap: {
4309 if (SemaRef.checkArgCount(TheCall, 2) ||
4310 CheckResourceHandle(&SemaRef, TheCall, 0) ||
4311 CheckArgTypeMatches(&SemaRef, TheCall->getArg(1),
4312 SemaRef.getASTContext().UnsignedIntTy))
4313 return true;
4314
4315 // Update return type to be the attributed resource type from arg0.
4316 QualType ResourceTy = TheCall->getArg(0)->getType();
4317 TheCall->setType(ResourceTy);
4318 break;
4319 }
4320 case Builtin::BI__builtin_hlsl_resource_counterhandlefromheap: {
4321 if (SemaRef.checkArgCount(TheCall, 1) ||
4322 CheckResourceHandle(&SemaRef, TheCall, 0))
4323 return true;
4324 // Update return type to be the attributed resource type from arg0
4325 // with added IsCounter flag.
4326 QualType MainHandleTy = TheCall->getArg(0)->getType();
4327 QualType CounterHandleTy =
4328 createCounterHandleType(SemaRef.getASTContext(), MainHandleTy);
4329 TheCall->setType(CounterHandleTy);
4330 break;
4331 }
4332 case Builtin::BI__builtin_hlsl_and:
4333 case Builtin::BI__builtin_hlsl_or: {
4334 if (SemaRef.checkArgCount(TheCall, 2))
4335 return true;
4336 if (CheckScalarOrVectorOrMatrix(&SemaRef, TheCall, getASTContext().BoolTy,
4337 0))
4338 return true;
4339 if (CheckAllArgsHaveSameType(&SemaRef, TheCall))
4340 return true;
4341
4342 ExprResult A = TheCall->getArg(0);
4343 QualType ArgTyA = A.get()->getType();
4344 // return type is the same as the input type
4345 TheCall->setType(ArgTyA);
4346 break;
4347 }
4348 case Builtin::BI__builtin_hlsl_all:
4349 case Builtin::BI__builtin_hlsl_any: {
4350 if (SemaRef.checkArgCount(TheCall, 1))
4351 return true;
4352 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
4353 return true;
4354 break;
4355 }
4356 case Builtin::BI__builtin_hlsl_asdouble: {
4357 if (SemaRef.checkArgCount(TheCall, 2))
4358 return true;
4360 &SemaRef, TheCall,
4361 /*only check for uint*/ SemaRef.Context.UnsignedIntTy,
4362 /* arg index */ 0))
4363 return true;
4365 &SemaRef, TheCall,
4366 /*only check for uint*/ SemaRef.Context.UnsignedIntTy,
4367 /* arg index */ 1))
4368 return true;
4369 if (CheckAllArgsHaveSameType(&SemaRef, TheCall))
4370 return true;
4371
4372 SetElementTypeAsReturnType(&SemaRef, TheCall, getASTContext().DoubleTy);
4373 break;
4374 }
4375 case Builtin::BI__builtin_hlsl_elementwise_clamp: {
4376 if (SemaRef.BuiltinElementwiseTernaryMath(
4377 TheCall, /*ArgTyRestr=*/
4379 return true;
4380 break;
4381 }
4382 case Builtin::BI__builtin_hlsl_dot: {
4383 // arg count is checked by BuiltinVectorToScalarMath
4384 if (SemaRef.BuiltinVectorToScalarMath(TheCall))
4385 return true;
4387 return true;
4388 break;
4389 }
4390 case Builtin::BI__builtin_hlsl_elementwise_firstbithigh:
4391 case Builtin::BI__builtin_hlsl_elementwise_firstbitlow: {
4392 if (SemaRef.PrepareBuiltinElementwiseMathOneArgCall(TheCall))
4393 return true;
4394
4395 const Expr *Arg = TheCall->getArg(0);
4396 QualType ArgTy = Arg->getType();
4397 QualType EltTy = ArgTy;
4398
4399 QualType ResTy = SemaRef.Context.UnsignedIntTy;
4400
4401 if (auto *VecTy = EltTy->getAs<VectorType>()) {
4402 EltTy = VecTy->getElementType();
4403 ResTy = SemaRef.Context.getExtVectorType(ResTy, VecTy->getNumElements());
4404 }
4405
4406 if (!EltTy->isIntegerType()) {
4407 Diag(Arg->getBeginLoc(), diag::err_builtin_invalid_arg_type)
4408 << 1 << /* scalar or vector of */ 5 << /* integer ty */ 1
4409 << /* no fp */ 0 << ArgTy;
4410 return true;
4411 }
4412
4413 TheCall->setType(ResTy);
4414 break;
4415 }
4416 case Builtin::BI__builtin_hlsl_select: {
4417 if (SemaRef.checkArgCount(TheCall, 3))
4418 return true;
4419 if (CheckScalarOrVector(&SemaRef, TheCall, getASTContext().BoolTy, 0))
4420 return true;
4421 QualType ArgTy = TheCall->getArg(0)->getType();
4422 if (ArgTy->isBooleanType() && CheckBoolSelect(&SemaRef, TheCall))
4423 return true;
4424 auto *VTy = ArgTy->getAs<VectorType>();
4425 if (VTy && VTy->getElementType()->isBooleanType() &&
4426 CheckVectorSelect(&SemaRef, TheCall))
4427 return true;
4428 break;
4429 }
4430 case Builtin::BI__builtin_hlsl_elementwise_saturate:
4431 case Builtin::BI__builtin_hlsl_elementwise_rcp: {
4432 if (SemaRef.checkArgCount(TheCall, 1))
4433 return true;
4434 if (!TheCall->getArg(0)
4435 ->getType()
4436 ->hasFloatingRepresentation()) // half or float or double
4437 return SemaRef.Diag(TheCall->getArg(0)->getBeginLoc(),
4438 diag::err_builtin_invalid_arg_type)
4439 << /* ordinal */ 1 << /* scalar or vector */ 5 << /* no int */ 0
4440 << /* fp */ 1 << TheCall->getArg(0)->getType();
4441 if (SemaRef.PrepareBuiltinElementwiseMathOneArgCall(TheCall))
4442 return true;
4443 break;
4444 }
4445 case Builtin::BI__builtin_hlsl_elementwise_rsqrt:
4446 case Builtin::BI__builtin_hlsl_elementwise_frac:
4447 case Builtin::BI__builtin_hlsl_elementwise_ddx_coarse:
4448 case Builtin::BI__builtin_hlsl_elementwise_ddy_coarse:
4449 case Builtin::BI__builtin_hlsl_elementwise_ddx_fine:
4450 case Builtin::BI__builtin_hlsl_elementwise_ddy_fine: {
4451 if (SemaRef.checkArgCount(TheCall, 1))
4452 return true;
4453 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
4455 return true;
4456 if (SemaRef.PrepareBuiltinElementwiseMathOneArgCall(TheCall))
4457 return true;
4458 break;
4459 }
4460 case Builtin::BI__builtin_hlsl_elementwise_isinf:
4461 case Builtin::BI__builtin_hlsl_elementwise_isnan: {
4462 if (SemaRef.checkArgCount(TheCall, 1))
4463 return true;
4464 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
4466 return true;
4467 if (SemaRef.PrepareBuiltinElementwiseMathOneArgCall(TheCall))
4468 return true;
4470 break;
4471 }
4472 case Builtin::BI__builtin_hlsl_mad: {
4473 if (SemaRef.BuiltinElementwiseTernaryMath(
4474 TheCall, /*ArgTyRestr=*/
4476 return true;
4477 break;
4478 }
4479 case Builtin::BI__builtin_hlsl_mul: {
4480 if (SemaRef.checkArgCount(TheCall, 2))
4481 return true;
4482
4483 Expr *Arg0 = TheCall->getArg(0);
4484 Expr *Arg1 = TheCall->getArg(1);
4485 QualType Ty0 = Arg0->getType();
4486 QualType Ty1 = Arg1->getType();
4487
4488 auto getElemType = [](QualType T) -> QualType {
4489 if (const auto *VTy = T->getAs<VectorType>())
4490 return VTy->getElementType();
4491 if (const auto *MTy = T->getAs<ConstantMatrixType>())
4492 return MTy->getElementType();
4493 return T;
4494 };
4495
4496 QualType EltTy0 = getElemType(Ty0);
4497
4498 bool IsVec0 = Ty0->isVectorType();
4499 bool IsMat0 = Ty0->isConstantMatrixType();
4500 bool IsVec1 = Ty1->isVectorType();
4501 bool IsMat1 = Ty1->isConstantMatrixType();
4502
4503 QualType RetTy;
4504
4505 if (IsVec0 && IsMat1) {
4506 auto *MatTy = Ty1->castAs<ConstantMatrixType>();
4507 RetTy = getASTContext().getExtVectorType(EltTy0, MatTy->getNumColumns());
4508 } else if (IsMat0 && IsVec1) {
4509 auto *MatTy = Ty0->castAs<ConstantMatrixType>();
4510 RetTy = getASTContext().getExtVectorType(EltTy0, MatTy->getNumRows());
4511 } else {
4512 assert(IsMat0 && IsMat1);
4513 auto *MatTy0 = Ty0->castAs<ConstantMatrixType>();
4514 auto *MatTy1 = Ty1->castAs<ConstantMatrixType>();
4516 EltTy0, MatTy0->getNumRows(), MatTy1->getNumColumns());
4517 }
4518
4519 TheCall->setType(RetTy);
4520 break;
4521 }
4522 case Builtin::BI__builtin_elementwise_fma: {
4523 if (SemaRef.checkArgCount(TheCall, 3) ||
4524 CheckAllArgsHaveSameType(&SemaRef, TheCall)) {
4525 return true;
4526 }
4527
4528 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
4530 return true;
4531
4532 ExprResult A = TheCall->getArg(0);
4533 QualType ArgTyA = A.get()->getType();
4534 // return type is the same as input type
4535 TheCall->setType(ArgTyA);
4536 break;
4537 }
4538 case Builtin::BI__builtin_hlsl_transpose: {
4539 if (SemaRef.checkArgCount(TheCall, 1))
4540 return true;
4541
4542 Expr *Arg = TheCall->getArg(0);
4543 QualType ArgTy = Arg->getType();
4544
4545 const auto *MatTy = ArgTy->getAs<ConstantMatrixType>();
4546 if (!MatTy) {
4547 SemaRef.Diag(Arg->getBeginLoc(), diag::err_builtin_invalid_arg_type)
4548 << 1 << /* matrix */ 3 << /* no int */ 0 << /* no fp */ 0 << ArgTy;
4549 return true;
4550 }
4551
4553 MatTy->getElementType(), MatTy->getNumColumns(), MatTy->getNumRows());
4554 TheCall->setType(RetTy);
4555 break;
4556 }
4557 case Builtin::BI__builtin_hlsl_elementwise_sign: {
4558 if (SemaRef.PrepareBuiltinElementwiseMathOneArgCall(TheCall))
4559 return true;
4560 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
4562 return true;
4564 break;
4565 }
4566 case Builtin::BI__builtin_hlsl_wave_active_all_equal: {
4567 if (SemaRef.checkArgCount(TheCall, 1))
4568 return true;
4569
4570 // Ensure input expr type is a scalar/vector
4571 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
4572 return true;
4573
4574 QualType InputTy = TheCall->getArg(0)->getType();
4575 ASTContext &Ctx = getASTContext();
4576
4577 QualType RetTy;
4578
4579 // If vector, construct bool vector of same size
4580 if (const auto *VecTy = InputTy->getAs<ExtVectorType>()) {
4581 unsigned NumElts = VecTy->getNumElements();
4582 RetTy = Ctx.getExtVectorType(Ctx.BoolTy, NumElts);
4583 } else {
4584 // Scalar case
4585 RetTy = Ctx.BoolTy;
4586 }
4587
4588 TheCall->setType(RetTy);
4589 break;
4590 }
4591 case Builtin::BI__builtin_hlsl_wave_active_max:
4592 case Builtin::BI__builtin_hlsl_wave_active_min:
4593 case Builtin::BI__builtin_hlsl_wave_active_sum:
4594 case Builtin::BI__builtin_hlsl_wave_active_product: {
4595 if (SemaRef.checkArgCount(TheCall, 1))
4596 return true;
4597
4598 // Ensure input expr type is a scalar/vector and the same as the return type
4599 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
4600 return true;
4601 if (CheckWaveActive(&SemaRef, TheCall))
4602 return true;
4603 ExprResult Expr = TheCall->getArg(0);
4604 QualType ArgTyExpr = Expr.get()->getType();
4605 TheCall->setType(ArgTyExpr);
4606 break;
4607 }
4608 case Builtin::BI__builtin_hlsl_wave_active_bit_or:
4609 case Builtin::BI__builtin_hlsl_wave_active_bit_xor:
4610 case Builtin::BI__builtin_hlsl_wave_active_bit_and: {
4611 if (SemaRef.checkArgCount(TheCall, 1))
4612 return true;
4613
4614 // Ensure input expr type is a scalar/vector
4615 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
4616 return true;
4617
4618 if (CheckWaveActive(&SemaRef, TheCall))
4619 return true;
4620
4621 // Ensure the expr type is interpretable as a uint or vector<uint>
4622 ExprResult Expr = TheCall->getArg(0);
4623 QualType ArgTyExpr = Expr.get()->getType();
4624 auto *VTy = ArgTyExpr->getAs<VectorType>();
4625 if (!(ArgTyExpr->isIntegerType() ||
4626 (VTy && VTy->getElementType()->isIntegerType()))) {
4627 SemaRef.Diag(TheCall->getArg(0)->getBeginLoc(),
4628 diag::err_builtin_invalid_arg_type)
4629 << ArgTyExpr << SemaRef.Context.UnsignedIntTy << 1 << 0 << 0;
4630 return true;
4631 }
4632
4633 // Ensure input expr type is the same as the return type
4634 TheCall->setType(ArgTyExpr);
4635 break;
4636 }
4637 case Builtin::BI__builtin_hlsl_interlocked_add:
4638 case Builtin::BI__builtin_hlsl_interlocked_and:
4639 case Builtin::BI__builtin_hlsl_interlocked_exchange:
4640 case Builtin::BI__builtin_hlsl_interlocked_max:
4641 case Builtin::BI__builtin_hlsl_interlocked_min:
4642 case Builtin::BI__builtin_hlsl_interlocked_or:
4643 case Builtin::BI__builtin_hlsl_interlocked_xor: {
4644 // The builtin's prototype in Builtins.td is `void (...)`, so direct calls
4645 // to `__builtin_hlsl_interlocked_op` bypass argument checking entirely.
4646 // When reached via the synthesized `InterlockedOp` overload set in
4647 // HLSLExternalSemaSource, overload resolution has already enforced the
4648 // argument count, integer-type matching, and the address-space requirement
4649 // on `dest`. The checks below are a safety net for callers that invoke the
4650 // builtin by its mangled name and would otherwise reach CodeGen unchecked.
4651 // InterlockedExchange always reports the previous value, so it requires
4652 // `original_value` instead of accepting it as an optional argument.
4653 if (BuiltinID == Builtin::BI__builtin_hlsl_interlocked_exchange) {
4654 if (SemaRef.checkArgCount(TheCall, 3))
4655 return true;
4656 } else {
4657 if (TheCall->getNumArgs() < 2) {
4658 SemaRef.Diag(TheCall->getEndLoc(),
4659 diag::err_typecheck_call_too_few_args_at_least)
4660 << /*callee_type=*/0 << /*min_arg_count=*/2 << TheCall->getNumArgs()
4661 << /*is_non_object=*/0 << TheCall->getSourceRange();
4662 return true;
4663 }
4664 if (SemaRef.checkArgCountAtMost(TheCall, 3))
4665 return true;
4666 }
4667
4668 QualType DestTy = TheCall->getArg(0)->getType().getUnqualifiedType();
4669 // InterlockedExchange also operates on float. DXIL lowers that as a
4670 // bitwise exchange of the value's bit pattern, and DXC accepts 32-bit
4671 // float only, so half and double are rejected.
4672 const bool AllowsFloat =
4673 BuiltinID == Builtin::BI__builtin_hlsl_interlocked_exchange;
4674 if (!DestTy->isIntegerType() &&
4675 !(AllowsFloat && DestTy->isSpecificBuiltinType(BuiltinType::Float))) {
4676 SemaRef.Diag(TheCall->getArg(0)->getBeginLoc(),
4677 diag::err_builtin_invalid_arg_type)
4678 << /*ordinal=*/1 << /*scalar*/ 1 << /*integer*/ 1
4679 << /*32 bit floating-point*/ (AllowsFloat ? 3 : 0) << DestTy;
4680 return true;
4681 }
4682
4683 // 64-bit interlocked ops require SM 6.6 on DXIL. The synthesized wrapper
4684 // methods (e.g. RWByteAddressBuffer::InterlockedAdd64) are only declared
4685 // on SM 6.6+, so this defensive check only fires for direct builtin
4686 // calls; skip synthetic invocations (invalid source location).
4687 const TargetInfo &TI = SemaRef.Context.getTargetInfo();
4688 if (TheCall->getBeginLoc().isValid() &&
4689 TI.getTriple().getArch() == llvm::Triple::dxil &&
4690 SemaRef.Context.getTypeSize(DestTy) == 64 &&
4691 TI.getPlatformMinVersion() < VersionTuple(6, 6)) {
4692 SemaRef.Diag(TheCall->getBeginLoc(), diag::err_hlsl_builtin_requires_sm)
4693 << TheCall->getDirectCallee() << VersionTuple(6, 6).getAsString();
4694 return true;
4695 }
4696
4697 if (CheckModifiableLValue(&SemaRef, TheCall, 0))
4698 return true;
4699
4700 if (CheckArgAddrSpaceOneOf(&SemaRef, TheCall, 0,
4702 return true;
4703
4704 if (CheckArgTypeMatches(&SemaRef, TheCall->getArg(1), DestTy))
4705 return true;
4706
4707 if (TheCall->getNumArgs() == 3) {
4708 if (CheckArgTypeMatches(&SemaRef, TheCall->getArg(2), DestTy))
4709 return true;
4710 if (CheckModifiableLValue(&SemaRef, TheCall, 2))
4711 return true;
4712 }
4713
4714 TheCall->setType(SemaRef.Context.VoidTy);
4715 break;
4716 }
4717 // Note these are llvm builtins that we want to catch invalid intrinsic
4718 // generation. Normal handling of these builtins will occur elsewhere.
4719 case Builtin::BI__builtin_elementwise_bitreverse: {
4720 // does not include a check for number of arguments
4721 // because that is done previously
4722 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
4724 return true;
4725 break;
4726 }
4727 case Builtin::BI__builtin_hlsl_wave_prefix_count_bits: {
4728 if (SemaRef.checkArgCount(TheCall, 1))
4729 return true;
4730
4731 QualType ArgType = TheCall->getArg(0)->getType();
4732
4733 if (!(ArgType->isScalarType())) {
4734 SemaRef.Diag(TheCall->getArg(0)->getBeginLoc(),
4735 diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
4736 << ArgType << 0;
4737 return true;
4738 }
4739
4740 if (!(ArgType->isBooleanType())) {
4741 SemaRef.Diag(TheCall->getArg(0)->getBeginLoc(),
4742 diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
4743 << ArgType << 0;
4744 return true;
4745 }
4746
4747 break;
4748 }
4749 case Builtin::BI__builtin_hlsl_wave_read_lane_at: {
4750 if (SemaRef.checkArgCount(TheCall, 2))
4751 return true;
4752
4753 // Ensure index parameter type can be interpreted as a uint
4754 ExprResult Index = TheCall->getArg(1);
4755 QualType ArgTyIndex = Index.get()->getType();
4756 if (!ArgTyIndex->isIntegerType()) {
4757 SemaRef.Diag(TheCall->getArg(1)->getBeginLoc(),
4758 diag::err_typecheck_convert_incompatible)
4759 << ArgTyIndex << SemaRef.Context.UnsignedIntTy << 1 << 0 << 0;
4760 return true;
4761 }
4762
4763 // Ensure input expr type is a scalar/vector and the same as the return type
4764 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
4765 return true;
4766
4767 ExprResult Expr = TheCall->getArg(0);
4768 QualType ArgTyExpr = Expr.get()->getType();
4769 TheCall->setType(ArgTyExpr);
4770 break;
4771 }
4772 case Builtin::BI__builtin_hlsl_wave_read_lane_first: {
4773 if (SemaRef.checkArgCount(TheCall, 1))
4774 return true;
4775
4776 if (CheckAnyScalarOrVectorOrMatrix(&SemaRef, TheCall, 0))
4777 return true;
4778
4779 TheCall->setType(TheCall->getArg(0)->getType());
4780 break;
4781 }
4782 case Builtin::BI__builtin_hlsl_wave_get_lane_index: {
4783 if (SemaRef.checkArgCount(TheCall, 0))
4784 return true;
4785 break;
4786 }
4787 case Builtin::BI__builtin_hlsl_wave_prefix_sum:
4788 case Builtin::BI__builtin_hlsl_wave_prefix_product: {
4789 if (SemaRef.checkArgCount(TheCall, 1))
4790 return true;
4791
4792 // Ensure input expr type is a scalar/vector and the same as the return type
4793 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
4794 return true;
4795 if (CheckWavePrefix(&SemaRef, TheCall))
4796 return true;
4797 ExprResult Expr = TheCall->getArg(0);
4798 QualType ArgTyExpr = Expr.get()->getType();
4799 TheCall->setType(ArgTyExpr);
4800 break;
4801 }
4802 case Builtin::BI__builtin_hlsl_quad_read_across_x:
4803 case Builtin::BI__builtin_hlsl_quad_read_across_y:
4804 case Builtin::BI__builtin_hlsl_quad_read_across_diagonal: {
4805 if (SemaRef.checkArgCount(TheCall, 1))
4806 return true;
4807
4808 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
4809 return true;
4810 if (CheckNotBoolScalarOrVector(&SemaRef, TheCall, 0))
4811 return true;
4812 ExprResult Expr = TheCall->getArg(0);
4813 QualType ArgTyExpr = Expr.get()->getType();
4814 TheCall->setType(ArgTyExpr);
4815 break;
4816 }
4817 case Builtin::BI__builtin_hlsl_elementwise_splitdouble: {
4818 if (SemaRef.checkArgCount(TheCall, 3))
4819 return true;
4820
4821 if (CheckScalarOrVectorOrMatrix(&SemaRef, TheCall, SemaRef.Context.DoubleTy,
4822 0) ||
4824 SemaRef.Context.UnsignedIntTy, 1) ||
4826 SemaRef.Context.UnsignedIntTy, 2))
4827 return true;
4828
4829 if (CheckModifiableLValue(&SemaRef, TheCall, 1) ||
4830 CheckModifiableLValue(&SemaRef, TheCall, 2))
4831 return true;
4832 break;
4833 }
4834 case Builtin::BI__builtin_hlsl_elementwise_clip: {
4835 if (SemaRef.checkArgCount(TheCall, 1))
4836 return true;
4837
4838 if (CheckScalarOrVector(&SemaRef, TheCall, SemaRef.Context.FloatTy, 0))
4839 return true;
4840 break;
4841 }
4842 case Builtin::BI__builtin_elementwise_acos:
4843 case Builtin::BI__builtin_elementwise_asin:
4844 case Builtin::BI__builtin_elementwise_atan:
4845 case Builtin::BI__builtin_elementwise_atan2:
4846 case Builtin::BI__builtin_elementwise_ceil:
4847 case Builtin::BI__builtin_elementwise_cos:
4848 case Builtin::BI__builtin_elementwise_cosh:
4849 case Builtin::BI__builtin_elementwise_exp:
4850 case Builtin::BI__builtin_elementwise_exp2:
4851 case Builtin::BI__builtin_elementwise_exp10:
4852 case Builtin::BI__builtin_elementwise_floor:
4853 case Builtin::BI__builtin_elementwise_fmod:
4854 case Builtin::BI__builtin_elementwise_log:
4855 case Builtin::BI__builtin_elementwise_log2:
4856 case Builtin::BI__builtin_elementwise_log10:
4857 case Builtin::BI__builtin_elementwise_pow:
4858 case Builtin::BI__builtin_elementwise_roundeven:
4859 case Builtin::BI__builtin_elementwise_sin:
4860 case Builtin::BI__builtin_elementwise_sinh:
4861 case Builtin::BI__builtin_elementwise_sqrt:
4862 case Builtin::BI__builtin_elementwise_tan:
4863 case Builtin::BI__builtin_elementwise_tanh:
4864 case Builtin::BI__builtin_elementwise_trunc: {
4865 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
4867 return true;
4868 break;
4869 }
4870 case Builtin::BI__builtin_hlsl_buffer_update_counter: {
4871 assert(TheCall->getNumArgs() == 2 && "expected 2 args");
4872 auto checkResTy = [](const HLSLAttributedResourceType *ResTy) -> bool {
4873 return !(ResTy->getAttrs().ResourceClass == ResourceClass::UAV &&
4874 ResTy->getAttrs().RawBuffer && ResTy->hasContainedType());
4875 };
4876 if (CheckResourceHandle(&SemaRef, TheCall, 0, checkResTy))
4877 return true;
4878 Expr *OffsetExpr = TheCall->getArg(1);
4879 std::optional<llvm::APSInt> Offset =
4880 OffsetExpr->getIntegerConstantExpr(SemaRef.getASTContext());
4881 if (!Offset.has_value() || std::abs(Offset->getExtValue()) != 1) {
4882 SemaRef.Diag(TheCall->getArg(1)->getBeginLoc(),
4883 diag::err_hlsl_expect_arg_const_int_one_or_neg_one)
4884 << 1;
4885 return true;
4886 }
4887 break;
4888 }
4889 case Builtin::BI__builtin_hlsl_elementwise_f16tof32: {
4890 if (SemaRef.checkArgCount(TheCall, 1))
4891 return true;
4892 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
4894 return true;
4895 // ensure arg integers are 32 bits
4896 if (CheckExpectedBitWidth(&SemaRef, TheCall, 0, 32))
4897 return true;
4898 // check it wasn't a bool type
4899 QualType ArgTy = TheCall->getArg(0)->getType();
4900 if (auto *VTy = ArgTy->getAs<VectorType>())
4901 ArgTy = VTy->getElementType();
4902 if (ArgTy->isBooleanType()) {
4903 SemaRef.Diag(TheCall->getArg(0)->getBeginLoc(),
4904 diag::err_builtin_invalid_arg_type)
4905 << 1 << /* scalar or vector of */ 5 << /* unsigned int */ 3
4906 << /* no fp */ 0 << TheCall->getArg(0)->getType();
4907 return true;
4908 }
4909
4910 SetElementTypeAsReturnType(&SemaRef, TheCall, getASTContext().FloatTy);
4911 break;
4912 }
4913 case Builtin::BI__builtin_hlsl_elementwise_f32tof16: {
4914 if (SemaRef.checkArgCount(TheCall, 1))
4915 return true;
4917 return true;
4919 getASTContext().UnsignedIntTy);
4920 break;
4921 }
4922 }
4923 return false;
4924}
4925
4929 WorkList.push_back(BaseTy);
4930 while (!WorkList.empty()) {
4931 QualType T = WorkList.pop_back_val();
4932 T = T.getCanonicalType().getUnqualifiedType();
4933 if (const auto *AT = dyn_cast<ConstantArrayType>(T)) {
4934 llvm::SmallVector<QualType, 16> ElementFields;
4935 // Generally I've avoided recursion in this algorithm, but arrays of
4936 // structs could be time-consuming to flatten and churn through on the
4937 // work list. Hopefully nesting arrays of structs containing arrays
4938 // of structs too many levels deep is unlikely.
4939 BuildFlattenedTypeList(AT->getElementType(), ElementFields);
4940 // Repeat the element's field list n times.
4941 for (uint64_t Ct = 0; Ct < AT->getZExtSize(); ++Ct)
4942 llvm::append_range(List, ElementFields);
4943 continue;
4944 }
4945 // Vectors can only have element types that are builtin types, so this can
4946 // add directly to the list instead of to the WorkList.
4947 if (const auto *VT = dyn_cast<VectorType>(T)) {
4948 List.insert(List.end(), VT->getNumElements(), VT->getElementType());
4949 continue;
4950 }
4951 if (const auto *MT = dyn_cast<ConstantMatrixType>(T)) {
4952 List.insert(List.end(), MT->getNumElementsFlattened(),
4953 MT->getElementType());
4954 continue;
4955 }
4956 if (const auto *RD = T->getAsCXXRecordDecl()) {
4957 if (RD->isStandardLayout())
4958 RD = RD->getStandardLayoutBaseWithFields();
4959
4960 // For types that we shouldn't decompose (unions and non-aggregates), just
4961 // add the type itself to the list.
4962 if (RD->isUnion() || !RD->isAggregate()) {
4963 List.push_back(T);
4964 continue;
4965 }
4966
4968 for (const auto *FD : RD->fields())
4969 if (!FD->isUnnamedBitField())
4970 FieldTypes.push_back(FD->getType());
4971 // Reverse the newly added sub-range.
4972 std::reverse(FieldTypes.begin(), FieldTypes.end());
4973 llvm::append_range(WorkList, FieldTypes);
4974
4975 // If this wasn't a standard layout type we may also have some base
4976 // classes to deal with.
4977 if (!RD->isStandardLayout()) {
4978 FieldTypes.clear();
4979 for (const auto &Base : RD->bases())
4980 FieldTypes.push_back(Base.getType());
4981 std::reverse(FieldTypes.begin(), FieldTypes.end());
4982 llvm::append_range(WorkList, FieldTypes);
4983 }
4984 continue;
4985 }
4986 List.push_back(T);
4987 }
4988}
4989
4991 if (QT.isNull())
4992 return false;
4993
4994 // Must be a class/struct.
4995 const auto *RD = QT->getAsCXXRecordDecl();
4996 if (!RD || RD->isUnion())
4997 return false;
4998
4999 // Cannot be a resource type or contain one.
5000 return !QT->isHLSLIntangibleType();
5001}
5002
5004 // null and array types are not allowed.
5005 if (QT.isNull() || QT->isArrayType())
5006 return false;
5007
5008 // UDT types are not allowed
5009 if (QT->isRecordType())
5010 return false;
5011
5012 if (QT->isBooleanType() || QT->isEnumeralType())
5013 return false;
5014
5015 // the only other valid builtin types are scalars or vectors
5016 if (QT->isArithmeticType()) {
5017 if (SemaRef.Context.getTypeSize(QT) / 8 > 16)
5018 return false;
5019 return true;
5020 }
5021
5022 if (const VectorType *VT = QT->getAs<VectorType>()) {
5023 int ArraySize = VT->getNumElements();
5024
5025 if (ArraySize > 4)
5026 return false;
5027
5028 QualType ElTy = VT->getElementType();
5029 if (ElTy->isBooleanType())
5030 return false;
5031
5032 if (SemaRef.Context.getTypeSize(QT) / 8 > 16)
5033 return false;
5034 return true;
5035 }
5036
5037 return false;
5038}
5039
5041 if (T1.isNull() || T2.isNull())
5042 return false;
5043
5046
5047 // If both types are the same canonical type, they're obviously compatible.
5048 if (SemaRef.getASTContext().hasSameType(T1, T2))
5049 return true;
5050
5052 BuildFlattenedTypeList(T1, T1Types);
5054 BuildFlattenedTypeList(T2, T2Types);
5055
5056 // Check the flattened type list
5057 return llvm::equal(T1Types, T2Types,
5058 [this](QualType LHS, QualType RHS) -> bool {
5059 return SemaRef.IsLayoutCompatible(LHS, RHS);
5060 });
5061}
5062
5064 FunctionDecl *Old) {
5065 if (New->getNumParams() != Old->getNumParams())
5066 return true;
5067
5068 bool HadError = false;
5069
5070 for (unsigned i = 0, e = New->getNumParams(); i != e; ++i) {
5071 ParmVarDecl *NewParam = New->getParamDecl(i);
5072 ParmVarDecl *OldParam = Old->getParamDecl(i);
5073
5074 // HLSL parameter declarations for inout and out must match between
5075 // declarations. In HLSL inout and out are ambiguous at the call site,
5076 // but have different calling behavior, so you cannot overload a
5077 // method based on a difference between inout and out annotations.
5078 const auto *NDAttr = NewParam->getAttr<HLSLParamModifierAttr>();
5079 unsigned NSpellingIdx = (NDAttr ? NDAttr->getSpellingListIndex() : 0);
5080 const auto *ODAttr = OldParam->getAttr<HLSLParamModifierAttr>();
5081 unsigned OSpellingIdx = (ODAttr ? ODAttr->getSpellingListIndex() : 0);
5082
5083 if (NSpellingIdx != OSpellingIdx) {
5084 SemaRef.Diag(NewParam->getLocation(),
5085 diag::err_hlsl_param_qualifier_mismatch)
5086 << NDAttr << NewParam;
5087 SemaRef.Diag(OldParam->getLocation(), diag::note_previous_declaration_as)
5088 << ODAttr;
5089 HadError = true;
5090 }
5091 }
5092 return HadError;
5093}
5094
5095// Generally follows PerformScalarCast, with cases reordered for
5096// clarity of what types are supported
5098
5099 if (!SrcTy->isScalarType() || !DestTy->isScalarType())
5100 return false;
5101
5102 if (SemaRef.getASTContext().hasSameUnqualifiedType(SrcTy, DestTy))
5103 return true;
5104
5105 switch (SrcTy->getScalarTypeKind()) {
5106 case Type::STK_Bool: // casting from bool is like casting from an integer
5107 case Type::STK_Integral:
5108 switch (DestTy->getScalarTypeKind()) {
5109 case Type::STK_Bool:
5110 case Type::STK_Integral:
5111 case Type::STK_Floating:
5112 return true;
5113 case Type::STK_CPointer:
5117 llvm_unreachable("HLSL doesn't support pointers.");
5120 llvm_unreachable("HLSL doesn't support complex types.");
5122 llvm_unreachable("HLSL doesn't support fixed point types.");
5123 }
5124 llvm_unreachable("Should have returned before this");
5125
5126 case Type::STK_Floating:
5127 switch (DestTy->getScalarTypeKind()) {
5128 case Type::STK_Floating:
5129 case Type::STK_Bool:
5130 case Type::STK_Integral:
5131 return true;
5134 llvm_unreachable("HLSL doesn't support complex types.");
5136 llvm_unreachable("HLSL doesn't support fixed point types.");
5137 case Type::STK_CPointer:
5141 llvm_unreachable("HLSL doesn't support pointers.");
5142 }
5143 llvm_unreachable("Should have returned before this");
5144
5146 case Type::STK_CPointer:
5149 llvm_unreachable("HLSL doesn't support pointers.");
5150
5152 llvm_unreachable("HLSL doesn't support fixed point types.");
5153
5156 llvm_unreachable("HLSL doesn't support complex types.");
5157 }
5158
5159 llvm_unreachable("Unhandled scalar cast");
5160}
5161
5162// Can perform an HLSL Aggregate splat cast if the Dest is an aggregate and the
5163// Src is a scalar, a vector of length 1, or a 1x1 matrix
5164// Or if Dest is a vector and Src is a vector of length 1 or a 1x1 matrix
5166
5167 QualType SrcTy = Src->getType();
5168 // Not a valid HLSL Aggregate Splat cast if Dest is a scalar or if this is
5169 // going to be a vector splat from a scalar.
5170 if ((SrcTy->isScalarType() && DestTy->isVectorType()) ||
5171 DestTy->isScalarType())
5172 return false;
5173
5174 const VectorType *SrcVecTy = SrcTy->getAs<VectorType>();
5175 const ConstantMatrixType *SrcMatTy = SrcTy->getAs<ConstantMatrixType>();
5176
5177 // Src isn't a scalar, a vector of length 1, or a 1x1 matrix
5178 if (!SrcTy->isScalarType() &&
5179 !(SrcVecTy && SrcVecTy->getNumElements() == 1) &&
5180 !(SrcMatTy && SrcMatTy->getNumElementsFlattened() == 1))
5181 return false;
5182
5183 if (SrcVecTy)
5184 SrcTy = SrcVecTy->getElementType();
5185 else if (SrcMatTy)
5186 SrcTy = SrcMatTy->getElementType();
5187
5189 BuildFlattenedTypeList(DestTy, DestTypes);
5190
5191 for (unsigned I = 0, Size = DestTypes.size(); I < Size; ++I) {
5192 if (DestTypes[I]->isUnionType())
5193 return false;
5194 if (!CanPerformScalarCast(SrcTy, DestTypes[I]))
5195 return false;
5196 }
5197 return true;
5198}
5199
5200// Can we perform an HLSL Elementwise cast?
5202
5203 // Don't handle casts where LHS and RHS are any combination of scalar/vector
5204 // There must be an aggregate somewhere
5205 QualType SrcTy = Src->getType();
5206 if (SrcTy->isScalarType()) // always a splat and this cast doesn't handle that
5207 return false;
5208
5209 if (SrcTy->isVectorType() &&
5210 (DestTy->isScalarType() || DestTy->isVectorType()))
5211 return false;
5212
5213 if (SrcTy->isConstantMatrixType() &&
5214 (DestTy->isScalarType() || DestTy->isConstantMatrixType()))
5215 return false;
5216
5218 BuildFlattenedTypeList(DestTy, DestTypes);
5220 BuildFlattenedTypeList(SrcTy, SrcTypes);
5221
5222 // Usually the size of SrcTypes must be greater than or equal to the size of
5223 // DestTypes.
5224 if (SrcTypes.size() < DestTypes.size())
5225 return false;
5226
5227 unsigned SrcSize = SrcTypes.size();
5228 unsigned DstSize = DestTypes.size();
5229 unsigned I;
5230 for (I = 0; I < DstSize && I < SrcSize; I++) {
5231 if (SrcTypes[I]->isUnionType() || DestTypes[I]->isUnionType())
5232 return false;
5233 if (!CanPerformScalarCast(SrcTypes[I], DestTypes[I])) {
5234 return false;
5235 }
5236 }
5237
5238 // check the rest of the source type for unions.
5239 for (; I < SrcSize; I++) {
5240 if (SrcTypes[I]->isUnionType())
5241 return false;
5242 }
5243 return true;
5244}
5245
5247 assert(Param->hasAttr<HLSLParamModifierAttr>() &&
5248 "We should not get here without a parameter modifier expression");
5249 const auto *Attr = Param->getAttr<HLSLParamModifierAttr>();
5250 if (Attr->getABI() == ParameterABI::Ordinary)
5251 return ExprResult(Arg);
5252
5253 bool IsInOut = Attr->getABI() == ParameterABI::HLSLInOut;
5254 if (!Arg->isLValue()) {
5255 SemaRef.Diag(Arg->getBeginLoc(), diag::error_hlsl_inout_lvalue)
5256 << Arg << (IsInOut ? 1 : 0);
5257 return ExprError();
5258 }
5259
5260 ASTContext &Ctx = SemaRef.getASTContext();
5261
5262 QualType Ty = Param->getType().getNonLValueExprType(Ctx);
5263
5264 // HLSL allows implicit conversions from scalars to vectors, but not the
5265 // inverse, so we need to disallow `inout` with scalar->vector or
5266 // scalar->matrix conversions.
5267 if (Arg->getType()->isScalarType() != Ty->isScalarType()) {
5268 SemaRef.Diag(Arg->getBeginLoc(), diag::error_hlsl_inout_scalar_extension)
5269 << Arg << (IsInOut ? 1 : 0);
5270 return ExprError();
5271 }
5272
5273 auto *ArgOpV = new (Ctx) OpaqueValueExpr(Param->getBeginLoc(), Arg->getType(),
5274 VK_LValue, OK_Ordinary, Arg);
5275
5276 // Parameters are initialized via copy initialization. This allows for
5277 // overload resolution of argument constructors.
5278 InitializedEntity Entity =
5280 ExprResult Res =
5281 SemaRef.PerformCopyInitialization(Entity, Param->getBeginLoc(), ArgOpV);
5282 if (Res.isInvalid())
5283 return ExprError();
5284 Expr *Base = Res.get();
5285 // After the cast, drop the reference type when creating the exprs.
5286 Ty = Ty.getNonLValueExprType(Ctx);
5287 auto *OpV = new (Ctx)
5288 OpaqueValueExpr(Param->getBeginLoc(), Ty, VK_LValue, OK_Ordinary, Base);
5289
5290 // Writebacks are performed with `=` binary operator, which allows for
5291 // overload resolution on writeback result expressions.
5292 Res = SemaRef.ActOnBinOp(SemaRef.getCurScope(), Arg->getBeginLoc(),
5293 tok::equal, ArgOpV, OpV);
5294
5295 if (Res.isInvalid())
5296 return ExprError();
5297 Expr *Writeback = Res.get();
5298 auto *OutExpr =
5299 HLSLOutArgExpr::Create(Ctx, Ty, ArgOpV, OpV, Writeback, IsInOut);
5300
5301 return ExprResult(OutExpr);
5302}
5303
5305 // If HLSL gains support for references, all the cites that use this will need
5306 // to be updated with semantic checking to produce errors for
5307 // pointers/references.
5308 assert(!Ty->isReferenceType() &&
5309 "Pointer and reference types cannot be inout or out parameters");
5310 Ty = SemaRef.getASTContext().getLValueReferenceType(Ty);
5311 Ty.addRestrict();
5312 return Ty;
5313}
5314
5315// Returns true if the type has a non-empty constant buffer layout (if it is
5316// scalar, vector or matrix, or if it contains any of these.
5318 const Type *Ty = QT->getUnqualifiedDesugaredType();
5319 if (Ty->isScalarType() || Ty->isVectorType() || Ty->isMatrixType())
5320 return true;
5321
5323 return false;
5324
5325 if (const auto *RD = Ty->getAsCXXRecordDecl()) {
5326 for (const auto *FD : RD->fields()) {
5328 return true;
5329 }
5330 assert(RD->getNumBases() <= 1 &&
5331 "HLSL doesn't support multiple inheritance");
5332 return RD->getNumBases()
5333 ? hasConstantBufferLayout(RD->bases_begin()->getType())
5334 : false;
5335 }
5336
5337 if (const auto *AT = dyn_cast<ArrayType>(Ty)) {
5338 if (const auto *CAT = dyn_cast<ConstantArrayType>(AT))
5339 if (isZeroSizedArray(CAT))
5340 return false;
5342 }
5343
5344 return false;
5345}
5346
5347static bool IsDefaultBufferConstantDecl(const ASTContext &Ctx, VarDecl *VD) {
5348 bool IsVulkan =
5349 Ctx.getTargetInfo().getTriple().getOS() == llvm::Triple::Vulkan;
5350 bool IsVKPushConstant = IsVulkan && VD->hasAttr<HLSLVkPushConstantAttr>();
5351 QualType QT = VD->getType();
5352 return VD->getDeclContext()->isTranslationUnit() &&
5353 QT.getAddressSpace() == LangAS::Default &&
5354 VD->getStorageClass() != SC_Static &&
5355 !VD->hasAttr<HLSLVkConstantIdAttr>() && !IsVKPushConstant &&
5357}
5358
5360 // The variable already has an address space (groupshared for ex).
5361 if (Decl->getType().hasAddressSpace())
5362 return;
5363
5364 if (Decl->getType()->isDependentType())
5365 return;
5366
5367 QualType Type = Decl->getType();
5368
5369 if (Decl->hasAttr<HLSLVkExtBuiltinInputAttr>()) {
5370 LangAS ImplAS = LangAS::hlsl_input;
5371 Type = SemaRef.getASTContext().getAddrSpaceQualType(Type, ImplAS);
5372 Decl->setType(Type);
5373 return;
5374 }
5375
5376 if (Decl->hasAttr<HLSLVkExtBuiltinOutputAttr>()) {
5377 LangAS ImplAS = LangAS::hlsl_output;
5378 Type = SemaRef.getASTContext().getAddrSpaceQualType(Type, ImplAS);
5379 Decl->setType(Type);
5380
5381 // HLSL uses `static` differently than C++. For BuiltIn output, the static
5382 // does not imply private to the module scope.
5383 // Marking it as external to reflect the semantic this attribute brings.
5384 // See https://github.com/microsoft/hlsl-specs/issues/350
5385 Decl->setStorageClass(SC_Extern);
5386 return;
5387 }
5388
5389 bool IsVulkan = getASTContext().getTargetInfo().getTriple().getOS() ==
5390 llvm::Triple::Vulkan;
5391 if (IsVulkan && Decl->hasAttr<HLSLVkPushConstantAttr>()) {
5392 if (HasDeclaredAPushConstant)
5393 SemaRef.Diag(Decl->getLocation(), diag::err_hlsl_push_constant_unique);
5394
5396 Type = SemaRef.getASTContext().getAddrSpaceQualType(Type, ImplAS);
5397 Decl->setType(Type);
5398 HasDeclaredAPushConstant = true;
5399 return;
5400 }
5401
5402 if (Type->isSamplerT() || Type->isVoidType())
5403 return;
5404
5405 // Resource handles.
5407 return;
5408
5409 // Only static globals belong to the Private address space.
5410 // Non-static globals belongs to the cbuffer.
5411 if (Decl->getStorageClass() != SC_Static && !Decl->isStaticDataMember())
5412 return;
5413
5415 Type = SemaRef.getASTContext().getAddrSpaceQualType(Type, ImplAS);
5416 Decl->setType(Type);
5417}
5418
5419namespace {
5420
5421// Helper class for assigning bindings to resources declared within a struct.
5422// It keeps track of all binding attributes declared on a struct instance, and
5423// the offsets for each register type that have been assigned so far.
5424// Handles both explicit and implicit bindings.
5425class StructBindingContext {
5426 // Bindings and offsets per register type. We only need to support four
5427 // register types - SRV (u), UAV (t), CBuffer (c), and Sampler (s).
5428 HLSLResourceBindingAttr *RegBindingsAttrs[4];
5429 unsigned RegBindingOffset[4];
5430
5431 // Make sure the RegisterType values are what we expect
5432 static_assert(static_cast<unsigned>(RegisterType::SRV) == 0 &&
5433 static_cast<unsigned>(RegisterType::UAV) == 1 &&
5434 static_cast<unsigned>(RegisterType::CBuffer) == 2 &&
5435 static_cast<unsigned>(RegisterType::Sampler) == 3,
5436 "unexpected register type values");
5437
5438 // Vulkan binding attribute does not vary by register type.
5439 HLSLVkBindingAttr *VkBindingAttr;
5440 unsigned VkBindingOffset;
5441
5442public:
5443 // Constructor: gather all binding attributes on a struct instance and
5444 // initialize offsets.
5445 StructBindingContext(VarDecl *VD) {
5446 for (unsigned i = 0; i < 4; ++i) {
5447 RegBindingsAttrs[i] = nullptr;
5448 RegBindingOffset[i] = 0;
5449 }
5450 VkBindingAttr = nullptr;
5451 VkBindingOffset = 0;
5452
5453 ASTContext &AST = VD->getASTContext();
5454 bool IsSpirv = AST.getTargetInfo().getTriple().isSPIRV();
5455
5456 for (Attr *A : VD->attrs()) {
5457 if (auto *RBA = dyn_cast<HLSLResourceBindingAttr>(A)) {
5458 RegisterType RegType = RBA->getRegisterType();
5459 unsigned RegTypeIdx = static_cast<unsigned>(RegType);
5460 // Ignore unsupported register annotations, such as 'c' or 'i'.
5461 if (RegTypeIdx < 4)
5462 RegBindingsAttrs[RegTypeIdx] = RBA;
5463 continue;
5464 }
5465 // Gather the Vulkan binding attributes only if the target is SPIR-V.
5466 if (IsSpirv) {
5467 if (auto *VBA = dyn_cast<HLSLVkBindingAttr>(A))
5468 VkBindingAttr = VBA;
5469 }
5470 }
5471 }
5472
5473 // Creates a binding attribute for a resource based on the gathered attributes
5474 // and the required register type and range.
5475 Attr *createBindingAttr(SemaHLSL &S, ASTContext &AST, RegisterType RegType,
5476 unsigned Range, bool HasCounter) {
5477 assert(static_cast<unsigned>(RegType) < 4 && "unexpected register type");
5478
5479 if (VkBindingAttr) {
5480 unsigned Offset = VkBindingOffset;
5481 VkBindingOffset += Range;
5482 return HLSLVkBindingAttr::CreateImplicit(
5483 AST, VkBindingAttr->getBinding() + Offset, VkBindingAttr->getSet(),
5484 VkBindingAttr->getRange());
5485 }
5486
5487 HLSLResourceBindingAttr *RBA =
5488 RegBindingsAttrs[static_cast<unsigned>(RegType)];
5489 HLSLResourceBindingAttr *NewAttr = nullptr;
5490
5491 if (RBA && RBA->hasRegisterSlot()) {
5492 // Explicit binding - create a new attribute with offseted slot number
5493 // based on the required register type.
5494 unsigned Offset = RegBindingOffset[static_cast<unsigned>(RegType)];
5495 RegBindingOffset[static_cast<unsigned>(RegType)] += Range;
5496
5497 unsigned NewSlotNumber = RBA->getSlotNumber() + Offset;
5498 StringRef NewSlotNumberStr =
5499 createRegisterString(AST, RBA->getRegisterType(), NewSlotNumber);
5500 NewAttr = HLSLResourceBindingAttr::CreateImplicit(
5501 AST, NewSlotNumberStr, RBA->getSpace(), RBA->getRange());
5502 NewAttr->setBinding(RegType, NewSlotNumber, RBA->getSpaceNumber());
5503 } else {
5504 // No binding attribute or space-only binding - create a binding
5505 // attribute for implicit binding.
5506 NewAttr = HLSLResourceBindingAttr::CreateImplicit(AST, "", "0", {});
5507 NewAttr->setBinding(RegType, std::nullopt,
5508 RBA ? RBA->getSpaceNumber() : 0);
5509 NewAttr->setImplicitBindingOrderID(S.getNextImplicitBindingOrderID());
5510 }
5511 if (HasCounter)
5512 NewAttr->setImplicitCounterBindingOrderID(
5514 return NewAttr;
5515 }
5516};
5517
5518// Creates a global variable declaration for a resource field embedded in a
5519// struct, assigns it a binding, initializes it, and associates it with the
5520// struct declaration via an HLSLAssociatedResourceDeclAttr.
5521static void createGlobalResourceDeclForStruct(
5522 Sema &S, VarDecl *ParentVD, SourceLocation Loc, IdentifierInfo *Id,
5523 QualType ResTy, StructBindingContext &BindingCtx) {
5524 assert(isResourceRecordTypeOrArrayOf(ResTy) &&
5525 "expected resource type or array of resources");
5526
5527 DeclContext *DC = ParentVD->getNonTransparentDeclContext();
5528 assert(DC->isTranslationUnit() && "expected translation unit decl context");
5529
5530 ASTContext &AST = S.getASTContext();
5531 VarDecl *ResDecl =
5532 VarDecl::Create(AST, DC, Loc, Loc, Id, ResTy, nullptr, SC_None);
5533
5534 unsigned Range = 1;
5535 const Type *SingleResTy = ResTy.getTypePtr()->getUnqualifiedDesugaredType();
5536 while (const auto *AT = dyn_cast<ArrayType>(SingleResTy)) {
5537 const auto *CAT = dyn_cast<ConstantArrayType>(AT);
5538 Range = CAT ? (Range * CAT->getSize().getZExtValue()) : 0;
5539 SingleResTy =
5541 }
5542 const HLSLAttributedResourceType *ResHandleTy =
5543 HLSLAttributedResourceType::findHandleTypeOnResource(SingleResTy);
5544
5545 // Add a binding attribute to the global resource declaration.
5546 bool HasCounter = hasCounterHandle(SingleResTy->getAsCXXRecordDecl());
5547 Attr *BindingAttr = BindingCtx.createBindingAttr(
5548 S.HLSL(), AST, getRegisterType(ResHandleTy), Range, HasCounter);
5549 ResDecl->addAttr(BindingAttr);
5550 ResDecl->addAttr(InternalLinkageAttr::CreateImplicit(AST));
5551 ResDecl->setImplicit();
5552
5553 if (Range == 1)
5554 S.HLSL().initGlobalResourceDecl(ResDecl);
5555 else
5556 S.HLSL().initGlobalResourceArrayDecl(ResDecl);
5557
5558 ParentVD->addAttr(
5559 HLSLAssociatedResourceDeclAttr::CreateImplicit(AST, ResDecl));
5560 DC->addDecl(ResDecl);
5561
5562 DeclGroupRef DG(ResDecl);
5564}
5565
5566static void handleArrayOfStructWithResources(
5567 Sema &S, VarDecl *ParentVD, const ConstantArrayType *CAT,
5568 EmbeddedResourceNameBuilder &NameBuilder, StructBindingContext &BindingCtx);
5569
5570// Scans base and all fields of a struct/class type to find all embedded
5571// resources or resource arrays. Creates a global variable for each resource
5572// found.
5573static void handleStructWithResources(Sema &S, VarDecl *ParentVD,
5574 const CXXRecordDecl *RD,
5575 EmbeddedResourceNameBuilder &NameBuilder,
5576 StructBindingContext &BindingCtx) {
5577
5578 // Scan the base classes.
5579 assert(RD->getNumBases() <= 1 && "HLSL doesn't support multiple inheritance");
5580 const auto *BasesIt = RD->bases_begin();
5581 if (BasesIt != RD->bases_end()) {
5582 QualType QT = BasesIt->getType();
5583 if (QT->isHLSLIntangibleType()) {
5584 CXXRecordDecl *BaseRD = QT->getAsCXXRecordDecl();
5585 NameBuilder.pushBaseName(BaseRD->getName());
5586 handleStructWithResources(S, ParentVD, BaseRD, NameBuilder, BindingCtx);
5587 NameBuilder.pop();
5588 }
5589 }
5590 // Process this class fields.
5591 for (const FieldDecl *FD : RD->fields()) {
5592 QualType FDTy = FD->getType().getCanonicalType();
5593 if (!FDTy->isHLSLIntangibleType())
5594 continue;
5595
5596 NameBuilder.pushName(FD->getName());
5597
5599 IdentifierInfo *II = NameBuilder.getNameAsIdentifier(S.getASTContext());
5600 createGlobalResourceDeclForStruct(S, ParentVD, FD->getLocation(), II,
5601 FDTy, BindingCtx);
5602 } else if (const auto *RD = FDTy->getAsCXXRecordDecl()) {
5603 handleStructWithResources(S, ParentVD, RD, NameBuilder, BindingCtx);
5604
5605 } else if (const auto *ArrayTy = dyn_cast<ConstantArrayType>(FDTy)) {
5606 assert(!FDTy->isHLSLResourceRecordArray() &&
5607 "resource arrays should have been already handled");
5608 handleArrayOfStructWithResources(S, ParentVD, ArrayTy, NameBuilder,
5609 BindingCtx);
5610 }
5611 NameBuilder.pop();
5612 }
5613}
5614
5615// Processes array of structs with resources.
5616static void
5617handleArrayOfStructWithResources(Sema &S, VarDecl *ParentVD,
5618 const ConstantArrayType *CAT,
5619 EmbeddedResourceNameBuilder &NameBuilder,
5620 StructBindingContext &BindingCtx) {
5621
5622 QualType ElementTy = CAT->getElementType().getCanonicalType();
5623 assert(ElementTy->isHLSLIntangibleType() && "Expected HLSL intangible type");
5624
5625 const ConstantArrayType *SubCAT = dyn_cast<ConstantArrayType>(ElementTy);
5626 const CXXRecordDecl *ElementRD = ElementTy->getAsCXXRecordDecl();
5627
5628 if (!SubCAT && !ElementRD)
5629 return;
5630
5631 for (unsigned I = 0, E = CAT->getSize().getZExtValue(); I < E; ++I) {
5632 NameBuilder.pushArrayIndex(I);
5633 if (ElementRD)
5634 handleStructWithResources(S, ParentVD, ElementRD, NameBuilder,
5635 BindingCtx);
5636 else
5637 handleArrayOfStructWithResources(S, ParentVD, SubCAT, NameBuilder,
5638 BindingCtx);
5639 NameBuilder.pop();
5640 }
5641}
5642
5643} // namespace
5644
5645// Scans all fields of a user-defined struct (or array of structs)
5646// to find all embedded resources or resource arrays. For each resource
5647// a global variable of the resource type is created and associated
5648// with the parent declaration (VD) through a HLSLAssociatedResourceDeclAttr
5649// attribute.
5650void SemaHLSL::handleGlobalStructOrArrayOfWithResources(VarDecl *VD) {
5651 EmbeddedResourceNameBuilder NameBuilder(VD->getName());
5652 StructBindingContext BindingCtx(VD);
5653
5654 const Type *VDTy = VD->getType().getTypePtr();
5655 assert(VDTy->isHLSLIntangibleType() && !isResourceRecordTypeOrArrayOf(VD) &&
5656 "Expected non-resource struct or array type");
5657
5658 if (const CXXRecordDecl *RD = VDTy->getAsCXXRecordDecl()) {
5659 handleStructWithResources(SemaRef, VD, RD, NameBuilder, BindingCtx);
5660 return;
5661 }
5662
5663 if (const auto *CAT = dyn_cast<ConstantArrayType>(VDTy)) {
5664 handleArrayOfStructWithResources(SemaRef, VD, CAT, NameBuilder, BindingCtx);
5665 return;
5666 }
5667}
5668
5670 if (VD->hasGlobalStorage()) {
5671 // make sure the declaration has a complete type
5672 if (SemaRef.RequireCompleteType(
5673 VD->getLocation(),
5674 SemaRef.getASTContext().getBaseElementType(VD->getType()),
5675 diag::err_typecheck_decl_incomplete_type)) {
5676 VD->setInvalidDecl();
5678 return;
5679 }
5680
5681 // Global variables outside a cbuffer block that are not a resource, static,
5682 // groupshared, or an empty array or struct belong to the default constant
5683 // buffer $Globals (to be created at the end of the translation unit).
5685 // update address space to hlsl_constant
5688 VD->setType(NewTy);
5689 DefaultCBufferDecls.push_back(VD);
5690 }
5691
5692 // find all resources bindings on decl
5693 if (VD->getType()->isHLSLIntangibleType())
5694 collectResourceBindingsOnVarDecl(VD);
5695
5696 if (VD->hasAttr<HLSLVkConstantIdAttr>())
5698
5700 VD->getStorageClass() != SC_Static) {
5701 // Add internal linkage attribute to non-static resource variables. The
5702 // global externally visible storage is accessed through the handle, which
5703 // is a member. The variable itself is not externally visible.
5704 VD->addAttr(InternalLinkageAttr::CreateImplicit(getASTContext()));
5705 }
5706
5707 // process explicit bindings
5708 processExplicitBindingsOnDecl(VD);
5709
5710 // Add implicit binding attribute to non-static resource arrays.
5711 if (VD->getType()->isHLSLResourceRecordArray() &&
5712 VD->getStorageClass() != SC_Static) {
5713 // If the resource array does not have an explicit binding attribute,
5714 // create an implicit one. It will be used to transfer implicit binding
5715 // order_ID to codegen.
5716 ResourceBindingAttrs Binding(VD);
5717 if (!Binding.isExplicit()) {
5718 uint32_t OrderID = getNextImplicitBindingOrderID();
5719 if (Binding.hasBinding())
5720 Binding.setImplicitOrderID(OrderID);
5721 else {
5724 OrderID);
5725 // Re-create the binding object to pick up the new attribute.
5726 Binding = ResourceBindingAttrs(VD);
5727 }
5728 }
5729
5730 // Get to the base type of a potentially multi-dimensional array.
5732
5733 const CXXRecordDecl *RD = Ty->getAsCXXRecordDecl();
5734 if (hasCounterHandle(RD)) {
5735 if (!Binding.hasCounterImplicitOrderID()) {
5736 uint32_t OrderID = getNextImplicitBindingOrderID();
5737 Binding.setCounterImplicitOrderID(OrderID);
5738 }
5739 }
5740 }
5741
5742 // Process resources in user-defined structs, or arrays of such structs.
5743 const Type *VDTy = VD->getType().getTypePtr();
5744 if (VD->getStorageClass() != SC_Static && VDTy->isHLSLIntangibleType() &&
5746 handleGlobalStructOrArrayOfWithResources(VD);
5747
5748 // Mark groupshared variables as extern so they will have
5749 // external storage and won't be default initialized
5750 if (VD->hasAttr<HLSLGroupSharedAddressSpaceAttr>())
5752 }
5753
5755}
5756
5758 assert(VD->getType()->isHLSLResourceRecord() &&
5759 "expected resource record type");
5760
5761 ASTContext &AST = SemaRef.getASTContext();
5762 uint64_t UIntTySize = AST.getTypeSize(AST.UnsignedIntTy);
5763 uint64_t IntTySize = AST.getTypeSize(AST.IntTy);
5764
5765 // Gather resource binding attributes.
5766 ResourceBindingAttrs Binding(VD);
5767
5768 // Find correct initialization method and create its arguments.
5769 QualType ResourceTy = VD->getType();
5770 CXXRecordDecl *ResourceDecl = ResourceTy->getAsCXXRecordDecl();
5771 CXXMethodDecl *CreateMethod = nullptr;
5773
5774 bool HasCounter = hasCounterHandle(ResourceDecl);
5775 const char *CreateMethodName;
5776 if (Binding.isExplicit())
5777 CreateMethodName = HasCounter ? "__createFromBindingWithImplicitCounter"
5778 : "__createFromBinding";
5779 else
5780 CreateMethodName = HasCounter
5781 ? "__createFromImplicitBindingWithImplicitCounter"
5782 : "__createFromImplicitBinding";
5783
5784 CreateMethod =
5785 lookupMethod(SemaRef, ResourceDecl, CreateMethodName, VD->getLocation());
5786
5787 if (!CreateMethod) {
5788 // This can happen if someone creates a struct that looks like an HLSL
5789 // resource record but does not have the required static create method.
5790 // No binding will be generated for it.
5791 assert(!ResourceDecl->isImplicit() &&
5792 "create method lookup should always succeed for built-in resource "
5793 "records");
5794 return false;
5795 }
5796
5797 if (Binding.isExplicit()) {
5798 IntegerLiteral *RegSlot =
5799 IntegerLiteral::Create(AST, llvm::APInt(UIntTySize, Binding.getSlot()),
5801 Args.push_back(RegSlot);
5802 } else {
5803 uint32_t OrderID = (Binding.hasImplicitOrderID())
5804 ? Binding.getImplicitOrderID()
5806 IntegerLiteral *OrderId =
5807 IntegerLiteral::Create(AST, llvm::APInt(UIntTySize, OrderID),
5809 Args.push_back(OrderId);
5810 }
5811
5812 IntegerLiteral *Space =
5813 IntegerLiteral::Create(AST, llvm::APInt(UIntTySize, Binding.getSpace()),
5815 Args.push_back(Space);
5816
5818 AST, llvm::APInt(IntTySize, 1), AST.IntTy, SourceLocation());
5819 Args.push_back(RangeSize);
5820
5822 AST, llvm::APInt(UIntTySize, 0), AST.UnsignedIntTy, SourceLocation());
5823 Args.push_back(Index);
5824
5825 StringRef VarName = VD->getName();
5827 AST, VarName, StringLiteralKind::Ordinary, false,
5828 AST.getStringLiteralArrayType(AST.CharTy.withConst(), VarName.size()),
5829 SourceLocation());
5831 AST, AST.getPointerType(AST.CharTy.withConst()), CK_ArrayToPointerDecay,
5832 Name, nullptr, VK_PRValue, FPOptionsOverride());
5833 Args.push_back(NameCast);
5834
5835 if (HasCounter) {
5836 // Will this be in the correct order?
5837 uint32_t CounterOrderID = getNextImplicitBindingOrderID();
5838 IntegerLiteral *CounterId =
5839 IntegerLiteral::Create(AST, llvm::APInt(UIntTySize, CounterOrderID),
5841 Args.push_back(CounterId);
5842 }
5843
5844 // Make sure the create method template is instantiated and emitted.
5845 if (!CreateMethod->isDefined() && CreateMethod->isTemplateInstantiation())
5846 SemaRef.InstantiateFunctionDefinition(VD->getLocation(), CreateMethod,
5847 true);
5848
5849 // Create CallExpr with a call to the static method and set it as the decl
5850 // initialization.
5852 AST, NestedNameSpecifierLoc(), SourceLocation(), CreateMethod, false,
5853 CreateMethod->getNameInfo(), CreateMethod->getType(), VK_PRValue);
5854
5855 auto *ImpCast = ImplicitCastExpr::Create(
5856 AST, AST.getPointerType(CreateMethod->getType()),
5857 CK_FunctionToPointerDecay, DRE, nullptr, VK_PRValue, FPOptionsOverride());
5858
5859 CallExpr *InitExpr =
5860 CallExpr::Create(AST, ImpCast, Args, ResourceTy, VK_PRValue,
5862 VD->setInit(InitExpr);
5864 SemaRef.CheckCompleteVariableDeclaration(VD);
5865 return true;
5866}
5867
5869 assert(VD->getType()->isHLSLResourceRecordArray() &&
5870 "expected array of resource records");
5871
5872 // Individual resources in a resource array are not initialized here. They
5873 // are initialized later on during codegen when the individual resources are
5874 // accessed. Codegen will emit a call to the resource initialization method
5875 // with the specified array index. We need to make sure though that the method
5876 // for the specific resource type is instantiated, so codegen can emit a call
5877 // to it when the array element is accessed.
5878
5879 // Find correct initialization method based on the resource binding
5880 // information.
5881 ASTContext &AST = SemaRef.getASTContext();
5882 QualType ResElementTy = AST.getBaseElementType(VD->getType());
5883 CXXRecordDecl *ResourceDecl = ResElementTy->getAsCXXRecordDecl();
5884 CXXMethodDecl *CreateMethod = nullptr;
5885
5886 bool HasCounter = hasCounterHandle(ResourceDecl);
5887 ResourceBindingAttrs ResourceAttrs(VD);
5888 if (ResourceAttrs.isExplicit())
5889 // Resource has explicit binding.
5890 CreateMethod =
5891 lookupMethod(SemaRef, ResourceDecl,
5892 HasCounter ? "__createFromBindingWithImplicitCounter"
5893 : "__createFromBinding",
5894 VD->getLocation());
5895 else
5896 // Resource has implicit binding.
5897 CreateMethod = lookupMethod(
5898 SemaRef, ResourceDecl,
5899 HasCounter ? "__createFromImplicitBindingWithImplicitCounter"
5900 : "__createFromImplicitBinding",
5901 VD->getLocation());
5902
5903 if (!CreateMethod)
5904 return false;
5905
5906 // Make sure the create method template is instantiated and emitted.
5907 if (!CreateMethod->isDefined() && CreateMethod->isTemplateInstantiation())
5908 SemaRef.InstantiateFunctionDefinition(VD->getLocation(), CreateMethod,
5909 true);
5910 return true;
5911}
5912
5913// Returns true if the initialization has been handled.
5914// Returns false to use default initialization.
5916 // Objects in the hlsl_constant address space are initialized
5917 // externally, so don't synthesize an implicit initializer.
5919 return true;
5920
5921 if (VD->hasGlobalStorage() && VD->getStorageClass() != SC_Static) {
5922 const Type *Ty = VD->getType().getTypePtr();
5924 return true;
5926 return true;
5927 }
5928
5929 // User-defined structs/classes do not have constructors.
5930 // When declared at a global scope, they are part of the constant buffer
5931 // and should not be initialized by the compiler.
5932 // When declared at a local scope, they are not initialized.
5933 // Also applies to arrays of user-defined structs/classes.
5934 const Type *Ty = VD->getType()->getUnqualifiedDesugaredType();
5935 while (Ty->isArrayType())
5937 if (CXXRecordDecl *RD = Ty->getAsCXXRecordDecl())
5938 return !RD->isHLSLBuiltinRecord();
5939
5940 return false;
5941}
5942
5943std::optional<const DeclBindingInfo *> SemaHLSL::inferGlobalBinding(Expr *E) {
5944 if (auto *Ternary = dyn_cast<ConditionalOperator>(E)) {
5945 auto TrueInfo = inferGlobalBinding(Ternary->getTrueExpr());
5946 auto FalseInfo = inferGlobalBinding(Ternary->getFalseExpr());
5947 if (!TrueInfo || !FalseInfo)
5948 return std::nullopt;
5949 if (*TrueInfo != *FalseInfo)
5950 return std::nullopt;
5951 return TrueInfo;
5952 }
5953
5954 if (auto *ASE = dyn_cast<ArraySubscriptExpr>(E))
5955 E = ASE->getBase()->IgnoreParenImpCasts();
5956
5957 if (DeclRefExpr *DRE = dyn_cast<DeclRefExpr>(E->IgnoreParens()))
5958 if (VarDecl *VD = dyn_cast<VarDecl>(DRE->getDecl())) {
5959 const Type *Ty = VD->getType()->getUnqualifiedDesugaredType();
5960 if (Ty->isArrayType())
5962
5963 if (const auto *AttrResType =
5964 HLSLAttributedResourceType::findHandleTypeOnResource(Ty)) {
5965 ResourceClass RC = AttrResType->getAttrs().ResourceClass;
5966 return Bindings.getDeclBindingInfo(VD, RC);
5967 }
5968 }
5969
5970 return nullptr;
5971}
5972
5973void SemaHLSL::trackLocalResource(VarDecl *VD, Expr *E) {
5974 std::optional<const DeclBindingInfo *> ExprBinding = inferGlobalBinding(E);
5975 if (!ExprBinding) {
5976 SemaRef.Diag(E->getBeginLoc(),
5977 diag::warn_hlsl_assigning_local_resource_is_not_unique)
5978 << E << VD;
5979 return; // Expr use multiple resources
5980 }
5981
5982 if (*ExprBinding == nullptr)
5983 return; // No binding could be inferred to track, return without error
5984
5985 auto PrevBinding = Assigns.find(VD);
5986 if (PrevBinding == Assigns.end()) {
5987 // No previous binding recorded, simply record the new assignment
5988 Assigns.insert({VD, *ExprBinding});
5989 return;
5990 }
5991
5992 // Otherwise, warn if the assignment implies different resource bindings
5993 if (*ExprBinding != PrevBinding->second) {
5994 SemaRef.Diag(E->getBeginLoc(),
5995 diag::warn_hlsl_assigning_local_resource_is_not_unique)
5996 << E << VD;
5997 SemaRef.Diag(VD->getLocation(), diag::note_var_declared_here) << VD;
5998 return;
5999 }
6000
6001 return;
6002}
6003
6005 Expr *RHSExpr, SourceLocation Loc) {
6006 assert((LHSExpr->getType()->isHLSLResourceRecord() ||
6007 LHSExpr->getType()->isHLSLResourceRecordArray()) &&
6008 "expected LHS to be a resource record or array of resource records");
6009 if (Opc != BO_Assign)
6010 return true;
6011
6012 // If LHS is an array subscript, get the underlying declaration.
6013 Expr *E = LHSExpr;
6014 while (auto *ASE = dyn_cast<ArraySubscriptExpr>(E))
6015 E = ASE->getBase()->IgnoreParenImpCasts();
6016
6017 // Report error if LHS is a non-static resource declared at a global scope.
6018 if (DeclRefExpr *DRE = dyn_cast<DeclRefExpr>(E->IgnoreParens())) {
6019 if (VarDecl *VD = dyn_cast<VarDecl>(DRE->getDecl())) {
6020 if (VD->hasGlobalStorage() && VD->getStorageClass() != SC_Static) {
6021 // assignment to global resource is not allowed
6022 SemaRef.Diag(Loc, diag::err_hlsl_assign_to_global_resource) << VD;
6023 SemaRef.Diag(VD->getLocation(), diag::note_var_declared_here) << VD;
6024 return false;
6025 }
6026
6027 trackLocalResource(VD, RHSExpr);
6028 }
6029 }
6030 return true;
6031}
6032
6033// Returns true if the given type can have an overload of the given
6034// binary operator.
6036 CXXRecordDecl *RD = LHSTy->getAsCXXRecordDecl();
6037 if (!RD)
6038 return true;
6039 return RD->isHLSLBuiltinRecord() || Opc != BO_Assign;
6040}
6041
6042// Walks though the global variable declaration, collects all resource binding
6043// requirements and adds them to Bindings
6044void SemaHLSL::collectResourceBindingsOnVarDecl(VarDecl *VD) {
6045 assert(VD->hasGlobalStorage() && VD->getType()->isHLSLIntangibleType() &&
6046 "expected global variable that contains HLSL resource");
6047
6048 // Cbuffers and Tbuffers are HLSLBufferDecl types
6049 if (const HLSLBufferDecl *CBufferOrTBuffer = dyn_cast<HLSLBufferDecl>(VD)) {
6050 Bindings.addDeclBindingInfo(VD, CBufferOrTBuffer->isCBuffer()
6051 ? ResourceClass::CBuffer
6052 : ResourceClass::SRV);
6053 return;
6054 }
6055
6056 // Unwrap arrays
6057 // FIXME: Calculate array size while unwrapping
6058 const Type *Ty = VD->getType()->getUnqualifiedDesugaredType();
6059 while (Ty->isArrayType()) {
6060 const ArrayType *AT = cast<ArrayType>(Ty);
6062 }
6063
6064 // Resource (or array of resources)
6065 if (const HLSLAttributedResourceType *AttrResType =
6066 HLSLAttributedResourceType::findHandleTypeOnResource(Ty)) {
6067 Bindings.addDeclBindingInfo(VD, AttrResType->getAttrs().ResourceClass);
6068 return;
6069 }
6070
6071 // User defined record type
6072 if (const RecordType *RT = dyn_cast<RecordType>(Ty))
6073 collectResourceBindingsOnUserRecordDecl(VD, RT);
6074}
6075
6076// Walks though the explicit resource binding attributes on the declaration,
6077// and makes sure there is a resource that matched the binding and updates
6078// DeclBindingInfoLists
6079void SemaHLSL::processExplicitBindingsOnDecl(VarDecl *VD) {
6080 assert(VD->hasGlobalStorage() && "expected global variable");
6081
6082 bool HasBinding = false;
6083 for (Attr *A : VD->attrs()) {
6084 if (isa<HLSLVkBindingAttr>(A)) {
6085 HasBinding = true;
6086 if (auto PA = VD->getAttr<HLSLVkPushConstantAttr>())
6087 Diag(PA->getLoc(), diag::err_hlsl_attr_incompatible) << A << PA;
6088 }
6089
6090 HLSLResourceBindingAttr *RBA = dyn_cast<HLSLResourceBindingAttr>(A);
6091 if (!RBA || !RBA->hasRegisterSlot())
6092 continue;
6093 HasBinding = true;
6094
6095 RegisterType RT = RBA->getRegisterType();
6096 assert(RT != RegisterType::I && "invalid or obsolete register type should "
6097 "never have an attribute created");
6098
6099 if (RT == RegisterType::C) {
6100 if (Bindings.hasBindingInfoForDecl(VD))
6101 SemaRef.Diag(VD->getLocation(),
6102 diag::warn_hlsl_user_defined_type_missing_member)
6103 << static_cast<int>(RT);
6104 continue;
6105 }
6106
6107 // Find DeclBindingInfo for this binding and update it, or report error
6108 // if it does not exist (user type does to contain resources with the
6109 // expected resource class).
6111 if (DeclBindingInfo *BI = Bindings.getDeclBindingInfo(VD, RC)) {
6112 // update binding info
6113 BI->setBindingAttribute(RBA, BindingType::Explicit);
6114 } else {
6115 SemaRef.Diag(VD->getLocation(),
6116 diag::warn_hlsl_user_defined_type_missing_member)
6117 << static_cast<int>(RT);
6118 }
6119 }
6120
6121 if (!HasBinding && isResourceRecordTypeOrArrayOf(VD))
6122 SemaRef.Diag(VD->getLocation(), diag::warn_hlsl_implicit_binding);
6123}
6124namespace {
6125class InitListTransformer {
6126 Sema &S;
6127 ASTContext &Ctx;
6128 QualType InitTy;
6129 QualType *DstIt = nullptr;
6130 Expr **ArgIt = nullptr;
6131 // Is wrapping the destination type iterator required? This is only used for
6132 // incomplete array types where we loop over the destination type since we
6133 // don't know the full number of elements from the declaration.
6134 bool Wrap;
6135
6136 bool castInitializer(Expr *E) {
6137 assert(DstIt && "This should always be something!");
6138 if (DstIt == DestTypes.end()) {
6139 if (!Wrap) {
6140 ArgExprs.push_back(E);
6141 // This is odd, but it isn't technically a failure due to conversion, we
6142 // handle mismatched counts of arguments differently.
6143 return true;
6144 }
6145 DstIt = DestTypes.begin();
6146 }
6147 InitializedEntity Entity = InitializedEntity::InitializeParameter(
6148 Ctx, *DstIt, /* Consumed (ObjC) */ false);
6149 ExprResult Res = S.PerformCopyInitialization(Entity, E->getBeginLoc(), E);
6150 if (Res.isInvalid())
6151 return false;
6152 Expr *Init = Res.get();
6153 ArgExprs.push_back(Init);
6154 DstIt++;
6155 return true;
6156 }
6157
6158 bool buildInitializerListImpl(Expr *E) {
6159 // If this is an initialization list, traverse the sub initializers.
6160 if (auto *Init = dyn_cast<InitListExpr>(E)) {
6161 for (auto *SubInit : Init->inits())
6162 if (!buildInitializerListImpl(SubInit))
6163 return false;
6164 return true;
6165 }
6166
6167 // If this is a scalar type, just enqueue the expression.
6168 QualType Ty = E->getType().getDesugaredType(Ctx);
6169
6170 if (Ty->isScalarType() || (Ty->isRecordType() && !Ty->isAggregateType()) ||
6172 return castInitializer(E);
6173
6174 // If this is an aggregate type and a prvalue, create an xvalue temporary
6175 // so the member accesses will be xvalues. Wrap it in OpaqueExpr to make
6176 // sure codegen will not generate duplicate copies.
6177 if (E->isPRValue() && Ty->isAggregateType()) {
6179 if (TmpExpr.isInvalid())
6180 return false;
6181 E = TmpExpr.get();
6182 E = new (Ctx) OpaqueValueExpr(E->getBeginLoc(), E->getType(),
6183 E->getValueKind(), E->getObjectKind(), E);
6184 }
6185
6186 if (auto *VecTy = Ty->getAs<VectorType>()) {
6187 uint64_t Size = VecTy->getNumElements();
6188
6189 QualType SizeTy = Ctx.getSizeType();
6190 uint64_t SizeTySize = Ctx.getTypeSize(SizeTy);
6191 for (uint64_t I = 0; I < Size; ++I) {
6192 auto *Idx = IntegerLiteral::Create(Ctx, llvm::APInt(SizeTySize, I),
6193 SizeTy, SourceLocation());
6194
6196 E, E->getBeginLoc(), Idx, E->getEndLoc());
6197 if (ElExpr.isInvalid())
6198 return false;
6199 if (!castInitializer(ElExpr.get()))
6200 return false;
6201 }
6202 return true;
6203 }
6204 if (auto *MTy = Ty->getAs<ConstantMatrixType>()) {
6205 unsigned Rows = MTy->getNumRows();
6206 unsigned Cols = MTy->getNumColumns();
6207 QualType ElemTy = MTy->getElementType();
6208
6209 for (unsigned R = 0; R < Rows; ++R) {
6210 for (unsigned C = 0; C < Cols; ++C) {
6211 // row index literal
6212 Expr *RowIdx = IntegerLiteral::Create(
6213 Ctx, llvm::APInt(Ctx.getIntWidth(Ctx.IntTy), R), Ctx.IntTy,
6214 E->getBeginLoc());
6215 // column index literal
6216 Expr *ColIdx = IntegerLiteral::Create(
6217 Ctx, llvm::APInt(Ctx.getIntWidth(Ctx.IntTy), C), Ctx.IntTy,
6218 E->getBeginLoc());
6220 E, RowIdx, ColIdx, E->getEndLoc());
6221 if (ElExpr.isInvalid())
6222 return false;
6223 if (!castInitializer(ElExpr.get()))
6224 return false;
6225 ElExpr.get()->setType(ElemTy);
6226 }
6227 }
6228 return true;
6229 }
6230
6231 if (auto *ArrTy = dyn_cast<ConstantArrayType>(Ty.getTypePtr())) {
6232 uint64_t Size = ArrTy->getZExtSize();
6233 QualType SizeTy = Ctx.getSizeType();
6234 uint64_t SizeTySize = Ctx.getTypeSize(SizeTy);
6235 for (uint64_t I = 0; I < Size; ++I) {
6236 auto *Idx = IntegerLiteral::Create(Ctx, llvm::APInt(SizeTySize, I),
6237 SizeTy, SourceLocation());
6239 E, E->getBeginLoc(), Idx, E->getEndLoc());
6240 if (ElExpr.isInvalid())
6241 return false;
6242 if (!buildInitializerListImpl(ElExpr.get()))
6243 return false;
6244 }
6245 return true;
6246 }
6247
6248 if (auto *RD = Ty->getAsCXXRecordDecl()) {
6249 llvm::SmallVector<CXXRecordDecl *> RecordDecls;
6250 RecordDecls.push_back(RD);
6251 while (RecordDecls.back()->getNumBases()) {
6252 CXXRecordDecl *D = RecordDecls.back();
6253 assert(D->getNumBases() == 1 &&
6254 "HLSL doesn't support multiple inheritance");
6255 RecordDecls.push_back(
6257 }
6258 while (!RecordDecls.empty()) {
6259 CXXRecordDecl *RD = RecordDecls.pop_back_val();
6260 for (auto *FD : RD->fields()) {
6261 if (FD->isUnnamedBitField())
6262 continue;
6263 DeclAccessPair Found = DeclAccessPair::make(FD, FD->getAccess());
6264 DeclarationNameInfo NameInfo(FD->getDeclName(), E->getBeginLoc());
6266 E, false, E->getBeginLoc(), CXXScopeSpec(), FD, Found, NameInfo);
6267 if (Res.isInvalid())
6268 return false;
6269 if (!buildInitializerListImpl(Res.get()))
6270 return false;
6271 }
6272 }
6273 }
6274 return true;
6275 }
6276
6277 Expr *generateInitListsImpl(QualType Ty) {
6278 Ty = Ty.getDesugaredType(Ctx);
6279 assert(ArgIt != ArgExprs.end() && "Something is off in iteration!");
6280 if (Ty->isScalarType() || (Ty->isRecordType() && !Ty->isAggregateType()) ||
6282 return *(ArgIt++);
6283
6284 llvm::SmallVector<Expr *> Inits;
6285 if (Ty->isVectorType() || Ty->isConstantArrayType() ||
6286 Ty->isConstantMatrixType()) {
6287 QualType ElTy;
6288 uint64_t Size = 0;
6289 if (auto *ATy = Ty->getAs<VectorType>()) {
6290 ElTy = ATy->getElementType();
6291 Size = ATy->getNumElements();
6292 } else if (auto *CMTy = Ty->getAs<ConstantMatrixType>()) {
6293 ElTy = CMTy->getElementType();
6294 Size = CMTy->getNumElementsFlattened();
6295 } else {
6296 auto *VTy = cast<ConstantArrayType>(Ty.getTypePtr());
6297 ElTy = VTy->getElementType();
6298 Size = VTy->getZExtSize();
6299 }
6300 for (uint64_t I = 0; I < Size; ++I)
6301 Inits.push_back(generateInitListsImpl(ElTy));
6302 }
6303 if (auto *RD = Ty->getAsCXXRecordDecl()) {
6304 llvm::SmallVector<CXXRecordDecl *> RecordDecls;
6305 RecordDecls.push_back(RD);
6306 while (RecordDecls.back()->getNumBases()) {
6307 CXXRecordDecl *D = RecordDecls.back();
6308 assert(D->getNumBases() == 1 &&
6309 "HLSL doesn't support multiple inheritance");
6310 RecordDecls.push_back(
6312 }
6313 while (!RecordDecls.empty()) {
6314 CXXRecordDecl *RD = RecordDecls.pop_back_val();
6315 for (auto *FD : RD->fields())
6316 if (!FD->isUnnamedBitField())
6317 Inits.push_back(generateInitListsImpl(FD->getType()));
6318 }
6319 }
6320 auto *NewInit =
6321 new (Ctx) InitListExpr(Ctx, Inits.front()->getBeginLoc(), Inits,
6322 Inits.back()->getEndLoc(), /*isExplicit=*/false);
6323 NewInit->setType(Ty);
6324 return NewInit;
6325 }
6326
6327public:
6328 llvm::SmallVector<QualType, 16> DestTypes;
6329 llvm::SmallVector<Expr *, 16> ArgExprs;
6330 InitListTransformer(Sema &SemaRef, const InitializedEntity &Entity)
6331 : S(SemaRef), Ctx(SemaRef.getASTContext()),
6332 Wrap(Entity.getType()->isIncompleteArrayType()) {
6333 InitTy = Entity.getType().getNonReferenceType();
6334 // When we're generating initializer lists for incomplete array types we
6335 // need to wrap around both when building the initializers and when
6336 // generating the final initializer lists.
6337 if (Wrap) {
6338 assert(InitTy->isIncompleteArrayType());
6339 const IncompleteArrayType *IAT = Ctx.getAsIncompleteArrayType(InitTy);
6340 InitTy = IAT->getElementType();
6341 }
6342 BuildFlattenedTypeList(InitTy, DestTypes);
6343 DstIt = DestTypes.begin();
6344 }
6345
6346 bool buildInitializerList(Expr *E) { return buildInitializerListImpl(E); }
6347
6348 Expr *generateInitLists() {
6349 assert(!ArgExprs.empty() &&
6350 "Call buildInitializerList to generate argument expressions.");
6351 ArgIt = ArgExprs.begin();
6352 if (!Wrap)
6353 return generateInitListsImpl(InitTy);
6354 llvm::SmallVector<Expr *> Inits;
6355 while (ArgIt != ArgExprs.end())
6356 Inits.push_back(generateInitListsImpl(InitTy));
6357
6358 auto *NewInit =
6359 new (Ctx) InitListExpr(Ctx, Inits.front()->getBeginLoc(), Inits,
6360 Inits.back()->getEndLoc(), /*isExplicit=*/false);
6361 llvm::APInt ArySize(64, Inits.size());
6362 NewInit->setType(Ctx.getConstantArrayType(InitTy, ArySize, nullptr,
6363 ArraySizeModifier::Normal, 0));
6364 return NewInit;
6365 }
6366};
6367} // namespace
6368
6369// Recursively detect any incomplete array anywhere in the type graph,
6370// including arrays, struct fields, and base classes.
6372 Ty = Ty.getCanonicalType();
6373
6374 // Array types
6375 if (const ArrayType *AT = dyn_cast<ArrayType>(Ty)) {
6377 return true;
6379 }
6380
6381 // Record (struct/class) types
6382 if (const auto *RT = Ty->getAs<RecordType>()) {
6383 const RecordDecl *RD = RT->getDecl();
6384
6385 // Walk base classes (for C++ / HLSL structs with inheritance)
6386 if (const auto *CXXRD = dyn_cast<CXXRecordDecl>(RD)) {
6387 for (const CXXBaseSpecifier &Base : CXXRD->bases()) {
6388 if (containsIncompleteArrayType(Base.getType()))
6389 return true;
6390 }
6391 }
6392
6393 // Walk fields
6394 for (const FieldDecl *F : RD->fields()) {
6395 if (containsIncompleteArrayType(F->getType()))
6396 return true;
6397 }
6398 }
6399
6400 return false;
6401}
6402
6404 InitListExpr *Init) {
6405 // If the initializer is a scalar, just return it.
6406 if (Init->getType()->isScalarType())
6407 return true;
6408 ASTContext &Ctx = SemaRef.getASTContext();
6409 InitListTransformer ILT(SemaRef, Entity);
6410
6411 for (unsigned I = 0; I < Init->getNumInits(); ++I) {
6412 Expr *E = Init->getInit(I);
6413 if (E->HasSideEffects(Ctx)) {
6414 QualType Ty = E->getType();
6415 if (Ty->isRecordType())
6416 E = new (Ctx) MaterializeTemporaryExpr(Ty, E, E->isLValue());
6417 E = new (Ctx) OpaqueValueExpr(E->getBeginLoc(), Ty, E->getValueKind(),
6418 E->getObjectKind(), E);
6419 Init->setInit(I, E);
6420 }
6421 if (!ILT.buildInitializerList(E))
6422 return false;
6423 }
6424 size_t ExpectedSize = ILT.DestTypes.size();
6425 size_t ActualSize = ILT.ArgExprs.size();
6426 if (ExpectedSize == 0 && ActualSize == 0)
6427 return true;
6428
6429 // Reject empty initializer if *any* incomplete array exists structurally
6430 if (ActualSize == 0 && containsIncompleteArrayType(Entity.getType())) {
6431 QualType InitTy = Entity.getType().getNonReferenceType();
6432 if (InitTy.hasAddressSpace())
6433 InitTy = SemaRef.getASTContext().removeAddrSpaceQualType(InitTy);
6434
6435 SemaRef.Diag(Init->getBeginLoc(), diag::err_hlsl_incorrect_num_initializers)
6436 << /*TooManyOrFew=*/(int)(ExpectedSize < ActualSize) << InitTy
6437 << /*ExpectedSize=*/ExpectedSize << /*ActualSize=*/ActualSize;
6438 return false;
6439 }
6440
6441 // We infer size after validating legality.
6442 // For incomplete arrays it is completely arbitrary to choose whether we think
6443 // the user intended fewer or more elements. This implementation assumes that
6444 // the user intended more, and errors that there are too few initializers to
6445 // complete the final element.
6446 if (Entity.getType()->isIncompleteArrayType()) {
6447 assert(ExpectedSize > 0 &&
6448 "The expected size of an incomplete array type must be at least 1.");
6449 ExpectedSize =
6450 ((ActualSize + ExpectedSize - 1) / ExpectedSize) * ExpectedSize;
6451 }
6452
6453 // An initializer list might be attempting to initialize a reference or
6454 // rvalue-reference. When checking the initializer we should look through
6455 // the reference.
6456 QualType InitTy = Entity.getType().getNonReferenceType();
6457 if (InitTy.hasAddressSpace())
6458 InitTy = SemaRef.getASTContext().removeAddrSpaceQualType(InitTy);
6459 if (ExpectedSize != ActualSize) {
6460 int TooManyOrFew = ActualSize > ExpectedSize ? 1 : 0;
6461 SemaRef.Diag(Init->getBeginLoc(), diag::err_hlsl_incorrect_num_initializers)
6462 << TooManyOrFew << InitTy << ExpectedSize << ActualSize;
6463 return false;
6464 }
6465
6466 // generateInitListsImpl will always return an InitListExpr here, because the
6467 // scalar case is handled above.
6468 auto *NewInit = cast<InitListExpr>(ILT.generateInitLists());
6469 Init->resizeInits(Ctx, NewInit->getNumInits());
6470 for (unsigned I = 0; I < NewInit->getNumInits(); ++I)
6471 Init->updateInit(Ctx, I, NewInit->getInit(I));
6472 return true;
6473}
6474
6475static QualType ReportMatrixInvalidMember(Sema &S, StringRef Name,
6476 StringRef Expected,
6477 SourceLocation OpLoc,
6478 SourceLocation CompLoc) {
6479 S.Diag(OpLoc, diag::err_builtin_matrix_invalid_member)
6480 << Name << Expected << SourceRange(CompLoc);
6481 return QualType();
6482}
6483
6486 const IdentifierInfo *CompName,
6487 SourceLocation CompLoc) {
6488 const auto *MT = baseType->castAs<ConstantMatrixType>();
6489 StringRef AccessorName = CompName->getName();
6490 assert(!AccessorName.empty() && "Matrix Accessor must have a name");
6491
6492 unsigned Rows = MT->getNumRows();
6493 unsigned Cols = MT->getNumColumns();
6494 bool IsZeroBasedAccessor = false;
6495 unsigned ChunkLen = 0;
6496 if (AccessorName.size() < 2)
6497 return ReportMatrixInvalidMember(S, AccessorName,
6498 "length 4 for zero based: \'_mRC\' or "
6499 "length 3 for one-based: \'_RC\' accessor",
6500 OpLoc, CompLoc);
6501
6502 if (AccessorName[0] == '_') {
6503 if (AccessorName[1] == 'm') {
6504 IsZeroBasedAccessor = true;
6505 ChunkLen = 4; // zero-based: "_mRC"
6506 } else {
6507 ChunkLen = 3; // one-based: "_RC"
6508 }
6509 } else
6511 S, AccessorName, "zero based: \'_mRC\' or one-based: \'_RC\' accessor",
6512 OpLoc, CompLoc);
6513
6514 if (AccessorName.size() % ChunkLen != 0) {
6515 const llvm::StringRef Expected = IsZeroBasedAccessor
6516 ? "zero based: '_mRC' accessor"
6517 : "one-based: '_RC' accessor";
6518
6519 return ReportMatrixInvalidMember(S, AccessorName, Expected, OpLoc, CompLoc);
6520 }
6521
6522 auto isDigit = [](char c) { return c >= '0' && c <= '9'; };
6523 auto isZeroBasedIndex = [](unsigned i) { return i <= 3; };
6524 auto isOneBasedIndex = [](unsigned i) { return i >= 1 && i <= 4; };
6525
6526 bool HasRepeated = false;
6527 SmallVector<bool, 16> Seen(Rows * Cols, false);
6528 unsigned NumComponents = 0;
6529 const char *Begin = AccessorName.data();
6530
6531 for (unsigned I = 0, E = AccessorName.size(); I < E; I += ChunkLen) {
6532 const char *Chunk = Begin + I;
6533 char RowChar = 0, ColChar = 0;
6534 if (IsZeroBasedAccessor) {
6535 // Zero-based: "_mRC"
6536 if (Chunk[0] != '_' || Chunk[1] != 'm') {
6537 char Bad = (Chunk[0] != '_') ? Chunk[0] : Chunk[1];
6539 S, StringRef(&Bad, 1), "\'_m\' prefix",
6540 OpLoc.getLocWithOffset(I + (Bad == Chunk[0] ? 1 : 2)), CompLoc);
6541 }
6542 RowChar = Chunk[2];
6543 ColChar = Chunk[3];
6544 } else {
6545 // One-based: "_RC"
6546 if (Chunk[0] != '_')
6548 S, StringRef(&Chunk[0], 1), "\'_\' prefix",
6549 OpLoc.getLocWithOffset(I + 1), CompLoc);
6550 RowChar = Chunk[1];
6551 ColChar = Chunk[2];
6552 }
6553
6554 // Must be digits.
6555 bool IsDigitsError = false;
6556 if (!isDigit(RowChar)) {
6557 unsigned BadPos = IsZeroBasedAccessor ? 2 : 1;
6558 ReportMatrixInvalidMember(S, StringRef(&RowChar, 1), "row as integer",
6559 OpLoc.getLocWithOffset(I + BadPos + 1),
6560 CompLoc);
6561 IsDigitsError = true;
6562 }
6563
6564 if (!isDigit(ColChar)) {
6565 unsigned BadPos = IsZeroBasedAccessor ? 3 : 2;
6566 ReportMatrixInvalidMember(S, StringRef(&ColChar, 1), "column as integer",
6567 OpLoc.getLocWithOffset(I + BadPos + 1),
6568 CompLoc);
6569 IsDigitsError = true;
6570 }
6571 if (IsDigitsError)
6572 return QualType();
6573
6574 unsigned Row = RowChar - '0';
6575 unsigned Col = ColChar - '0';
6576
6577 bool HasIndexingError = false;
6578 if (IsZeroBasedAccessor) {
6579 // 0-based [0..3]
6580 if (!isZeroBasedIndex(Row)) {
6581 S.Diag(OpLoc, diag::err_hlsl_matrix_element_not_in_bounds)
6582 << /*row*/ 0 << /*zero-based*/ 0 << SourceRange(CompLoc);
6583 HasIndexingError = true;
6584 }
6585 if (!isZeroBasedIndex(Col)) {
6586 S.Diag(OpLoc, diag::err_hlsl_matrix_element_not_in_bounds)
6587 << /*col*/ 1 << /*zero-based*/ 0 << SourceRange(CompLoc);
6588 HasIndexingError = true;
6589 }
6590 } else {
6591 // 1-based [1..4]
6592 if (!isOneBasedIndex(Row)) {
6593 S.Diag(OpLoc, diag::err_hlsl_matrix_element_not_in_bounds)
6594 << /*row*/ 0 << /*one-based*/ 1 << SourceRange(CompLoc);
6595 HasIndexingError = true;
6596 }
6597 if (!isOneBasedIndex(Col)) {
6598 S.Diag(OpLoc, diag::err_hlsl_matrix_element_not_in_bounds)
6599 << /*col*/ 1 << /*one-based*/ 1 << SourceRange(CompLoc);
6600 HasIndexingError = true;
6601 }
6602 // Convert to 0-based after range checking.
6603 --Row;
6604 --Col;
6605 }
6606
6607 if (HasIndexingError)
6608 return QualType();
6609
6610 // Note: matrix swizzle index is hard coded. That means Row and Col can
6611 // potentially be larger than Rows and Cols if matrix size is less than
6612 // the max index size.
6613 bool HasBoundsError = false;
6614 if (Row >= Rows) {
6615 Diag(OpLoc, diag::err_hlsl_matrix_index_out_of_bounds)
6616 << /*Row*/ 0 << Row << Rows << SourceRange(CompLoc);
6617 HasBoundsError = true;
6618 }
6619 if (Col >= Cols) {
6620 Diag(OpLoc, diag::err_hlsl_matrix_index_out_of_bounds)
6621 << /*Col*/ 1 << Col << Cols << SourceRange(CompLoc);
6622 HasBoundsError = true;
6623 }
6624 if (HasBoundsError)
6625 return QualType();
6626
6627 unsigned FlatIndex = Row * Cols + Col;
6628 if (Seen[FlatIndex])
6629 HasRepeated = true;
6630 Seen[FlatIndex] = true;
6631 ++NumComponents;
6632 }
6633 if (NumComponents == 0 || NumComponents > 4) {
6634 S.Diag(OpLoc, diag::err_hlsl_matrix_swizzle_invalid_length)
6635 << NumComponents << SourceRange(CompLoc);
6636 return QualType();
6637 }
6638
6639 QualType ElemTy = MT->getElementType();
6640 if (NumComponents == 1)
6641 return ElemTy;
6642 QualType VT = S.Context.getExtVectorType(ElemTy, NumComponents);
6643 if (HasRepeated)
6644 VK = VK_PRValue;
6645
6646 for (Sema::ExtVectorDeclsType::iterator
6648 E = S.ExtVectorDecls.end();
6649 I != E; ++I) {
6650 if ((*I)->getUnderlyingType() == VT)
6652 /*Qualifier=*/std::nullopt, *I);
6653 }
6654
6655 return VT;
6656}
6657
6659 // If initializing a local resource, track the resource binding it is using
6660 if (VDecl->getType()->isHLSLResourceRecord() && !VDecl->hasGlobalStorage())
6661 trackLocalResource(VDecl, Init);
6662
6663 const HLSLVkConstantIdAttr *ConstIdAttr =
6664 VDecl->getAttr<HLSLVkConstantIdAttr>();
6665 if (!ConstIdAttr)
6666 return true;
6667
6668 ASTContext &Context = SemaRef.getASTContext();
6669
6670 APValue InitValue;
6671 if (!Init->isCXX11ConstantExpr(Context, &InitValue)) {
6672 Diag(VDecl->getLocation(), diag::err_specialization_const);
6673 VDecl->setInvalidDecl();
6674 return false;
6675 }
6676
6677 Builtin::ID BID =
6679
6680 // Argument 1: The ID from the attribute
6681 int ConstantID = ConstIdAttr->getId();
6682 llvm::APInt IDVal(Context.getIntWidth(Context.IntTy), ConstantID);
6683 Expr *IdExpr = IntegerLiteral::Create(Context, IDVal, Context.IntTy,
6684 ConstIdAttr->getLocation());
6685
6686 SmallVector<Expr *, 2> Args = {IdExpr, Init};
6687 Expr *C = SemaRef.BuildBuiltinCallExpr(Init->getExprLoc(), BID, Args);
6688 if (C->getType()->getCanonicalTypeUnqualified() !=
6690 C = SemaRef
6691 .BuildCStyleCastExpr(SourceLocation(),
6692 Context.getTrivialTypeSourceInfo(
6693 Init->getType(), Init->getExprLoc()),
6694 SourceLocation(), C)
6695 .get();
6696 }
6697 Init = C;
6698 return true;
6699}
6700
6702 SourceLocation NameLoc) {
6703 if (!Template)
6704 return QualType();
6705
6706 DeclContext *DC = Template->getDeclContext();
6707 if (!DC->isNamespace() || !cast<NamespaceDecl>(DC)->getIdentifier() ||
6708 cast<NamespaceDecl>(DC)->getName() != "hlsl")
6709 return QualType();
6710
6711 TemplateParameterList *Params = Template->getTemplateParameters();
6712 if (!Params || Params->size() != 1)
6713 return QualType();
6714
6715 if (!Template->isImplicit())
6716 return QualType();
6717
6718 // We manually extract default arguments here instead of letting
6719 // CheckTemplateIdType handle it. This ensures that for resource types that
6720 // lack a default argument (like Buffer), we return a null QualType, which
6721 // triggers the "requires template arguments" error rather than a less
6722 // descriptive "too few template arguments" error.
6723 TemplateArgumentListInfo TemplateArgs(NameLoc, NameLoc);
6724 for (NamedDecl *P : *Params) {
6725 if (auto *TTP = dyn_cast<TemplateTypeParmDecl>(P)) {
6726 if (TTP->hasDefaultArgument()) {
6727 TemplateArgs.addArgument(TTP->getDefaultArgument());
6728 continue;
6729 }
6730 } else if (auto *NTTP = dyn_cast<NonTypeTemplateParmDecl>(P)) {
6731 if (NTTP->hasDefaultArgument()) {
6732 TemplateArgs.addArgument(NTTP->getDefaultArgument());
6733 continue;
6734 }
6735 } else if (auto *TTPD = dyn_cast<TemplateTemplateParmDecl>(P)) {
6736 if (TTPD->hasDefaultArgument()) {
6737 TemplateArgs.addArgument(TTPD->getDefaultArgument());
6738 continue;
6739 }
6740 }
6741 return QualType();
6742 }
6743
6744 return SemaRef.CheckTemplateIdType(
6746 TemplateArgs, nullptr, /*ForNestedNameSpecifier=*/false);
6747}
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 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
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 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
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