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 isMatrixType(QualType QT) {
2682 const Type *Ty = QT->getUnqualifiedDesugaredType();
2683 return Ty->isDependentType() || Ty->isConstantMatrixType();
2684}
2685
2686/// Walks the existing AttributedType sugar of \p T looking for a previously
2687/// applied HLSLRowMajor/HLSLColumnMajor marker. If one is found, populates
2688/// \p ExistingKind with its attr::Kind and returns true.
2690 attr::Kind &ExistingKind) {
2691 QualType Cur = T;
2692 while (const auto *AT = Cur->getAs<AttributedType>()) {
2693 attr::Kind K = AT->getAttrKind();
2694 if (K == attr::HLSLRowMajor || K == attr::HLSLColumnMajor) {
2695 ExistingKind = K;
2696 return true;
2697 }
2698 Cur = AT->getModifiedType();
2699 }
2700 return false;
2701}
2702
2704 if (T.isNull())
2705 return nullptr;
2706
2707 ASTContext &Ctx = getASTContext();
2708 attr::Kind AttrK = AL.getKind() == ParsedAttr::AT_HLSLRowMajor
2709 ? attr::HLSLRowMajor
2710 : attr::HLSLColumnMajor;
2711
2712 // For non-dependent types, the operand must be a matrix.
2713 if (!T->isDependentType() && !isMatrixType(T)) {
2714 Diag(AL.getLoc(), diag::err_hlsl_matrix_layout_non_matrix)
2715 << AL.getAttrName();
2716 AL.setInvalid();
2717 return nullptr;
2718 }
2719
2720 // Conflict / duplicate detection by walking existing sugar.
2721 attr::Kind ExistingKind;
2722 if (findExistingMatrixLayoutMarker(T, ExistingKind)) {
2723 if (ExistingKind == AttrK) {
2724 Diag(AL.getLoc(), diag::warn_duplicate_attribute_exact)
2725 << AL.getAttrName();
2726 Diag(AL.getLoc(), diag::note_previous_attribute);
2727 return nullptr;
2728 }
2729 IdentifierInfo *ExistingII = &Ctx.Idents.get(
2730 ExistingKind == attr::HLSLRowMajor ? "row_major" : "column_major");
2731 Diag(AL.getLoc(), diag::err_hlsl_matrix_layout_conflict)
2732 << AL.getAttrName() << ExistingII;
2733 Diag(AL.getLoc(), diag::note_conflicting_attribute);
2734 AL.setInvalid();
2735 return nullptr;
2736 }
2737
2738 if (AttrK == attr::HLSLRowMajor)
2739 return ::new (Ctx) HLSLRowMajorAttr(Ctx, AL);
2740 return ::new (Ctx) HLSLColumnMajorAttr(Ctx, AL);
2741}
2742
2743// Re-validates an HLSL `row_major` / `column_major` attribute after template
2744// substitution. The parse-time check in `buildMatrixLayoutTypeAttr` is skipped
2745// for dependent types; `TransformAttributedType` calls this once the type is
2746// concrete. Returns `true` (and emits a diagnostic) if the substituted type is
2747// not a matrix or array of matrices, signaling the caller to abort the
2748// transform.
2750 SourceLocation Loc) {
2751 if (K != attr::HLSLRowMajor && K != attr::HLSLColumnMajor)
2752 return false;
2753 if (T.isNull() || T->isDependentType())
2754 return false;
2755 if (isMatrixType(T))
2756 return false;
2758 K == attr::HLSLRowMajor ? "row_major" : "column_major");
2759 Diag(Loc, diag::err_hlsl_matrix_layout_non_matrix) << II;
2760 return true;
2761}
2762
2763// Transpose and matrix mul need to read the destination layout.
2764// Elementwise builtins reuse the operand layout instead.
2765namespace {
2766
2767/// This class implements HLSL availability diagnostics for default
2768/// and relaxed mode
2769///
2770/// The goal of this diagnostic is to emit an error or warning when an
2771/// unavailable API is found in code that is reachable from the shader
2772/// entry function or from an exported function (when compiling a shader
2773/// library).
2774///
2775/// This is done by traversing the AST of all shader entry point functions
2776/// and of all exported functions, and any functions that are referenced
2777/// from this AST. In other words, any functions that are reachable from
2778/// the entry points.
2779class DiagnoseHLSLAvailability : public DynamicRecursiveASTVisitor {
2780 Sema &SemaRef;
2781
2782 // Stack of functions to be scaned
2784
2785 // Tracks which environments functions have been scanned in.
2786 //
2787 // Maps FunctionDecl to an unsigned number that represents the set of shader
2788 // environments the function has been scanned for.
2789 // The llvm::Triple::EnvironmentType enum values for shader stages guaranteed
2790 // to be numbered from llvm::Triple::Pixel to llvm::Triple::Amplification
2791 // (verified by static_asserts in Triple.cpp), we can use it to index
2792 // individual bits in the set, as long as we shift the values to start with 0
2793 // by subtracting the value of llvm::Triple::Pixel first.
2794 //
2795 // The N'th bit in the set will be set if the function has been scanned
2796 // in shader environment whose llvm::Triple::EnvironmentType integer value
2797 // equals (llvm::Triple::Pixel + N).
2798 //
2799 // For example, if a function has been scanned in compute and pixel stage
2800 // environment, the value will be 0x21 (100001 binary) because:
2801 //
2802 // (int)(llvm::Triple::Pixel - llvm::Triple::Pixel) == 0
2803 // (int)(llvm::Triple::Compute - llvm::Triple::Pixel) == 5
2804 //
2805 // A FunctionDecl is mapped to 0 (or not included in the map) if it has not
2806 // been scanned in any environment.
2807 llvm::DenseMap<const FunctionDecl *, unsigned> ScannedDecls;
2808
2809 // Do not access these directly, use the get/set methods below to make
2810 // sure the values are in sync
2811 llvm::Triple::EnvironmentType CurrentShaderEnvironment;
2812 unsigned CurrentShaderStageBit;
2813
2814 // True if scanning a function that was already scanned in a different
2815 // shader stage context, and therefore we should not report issues that
2816 // depend only on shader model version because they would be duplicate.
2817 bool ReportOnlyShaderStageIssues;
2818
2819 // Helper methods for dealing with current stage context / environment
2820 void SetShaderStageContext(llvm::Triple::EnvironmentType ShaderType) {
2821 static_assert(sizeof(unsigned) >= 4);
2822 assert(HLSLShaderAttr::isValidShaderType(ShaderType));
2823 assert((unsigned)(ShaderType - llvm::Triple::Pixel) < 31 &&
2824 "ShaderType is too big for this bitmap"); // 31 is reserved for
2825 // "unknown"
2826
2827 unsigned bitmapIndex = ShaderType - llvm::Triple::Pixel;
2828 CurrentShaderEnvironment = ShaderType;
2829 CurrentShaderStageBit = (1 << bitmapIndex);
2830 }
2831
2832 void SetUnknownShaderStageContext() {
2833 CurrentShaderEnvironment = llvm::Triple::UnknownEnvironment;
2834 CurrentShaderStageBit = (1 << 31);
2835 }
2836
2837 llvm::Triple::EnvironmentType GetCurrentShaderEnvironment() const {
2838 return CurrentShaderEnvironment;
2839 }
2840
2841 bool InUnknownShaderStageContext() const {
2842 return CurrentShaderEnvironment == llvm::Triple::UnknownEnvironment;
2843 }
2844
2845 // Helper methods for dealing with shader stage bitmap
2846 void AddToScannedFunctions(const FunctionDecl *FD) {
2847 unsigned &ScannedStages = ScannedDecls[FD];
2848 ScannedStages |= CurrentShaderStageBit;
2849 }
2850
2851 unsigned GetScannedStages(const FunctionDecl *FD) { return ScannedDecls[FD]; }
2852
2853 bool WasAlreadyScannedInCurrentStage(const FunctionDecl *FD) {
2854 return WasAlreadyScannedInCurrentStage(GetScannedStages(FD));
2855 }
2856
2857 bool WasAlreadyScannedInCurrentStage(unsigned ScannerStages) {
2858 return ScannerStages & CurrentShaderStageBit;
2859 }
2860
2861 static bool NeverBeenScanned(unsigned ScannedStages) {
2862 return ScannedStages == 0;
2863 }
2864
2865 // Scanning methods
2866 void HandleFunctionOrMethodRef(FunctionDecl *FD, Expr *RefExpr);
2867 void CheckDeclAvailability(NamedDecl *D, const AvailabilityAttr *AA,
2868 SourceRange Range);
2869 const AvailabilityAttr *FindAvailabilityAttr(const Decl *D);
2870 bool HasMatchingEnvironmentOrNone(const AvailabilityAttr *AA);
2871
2872public:
2873 DiagnoseHLSLAvailability(Sema &SemaRef)
2874 : SemaRef(SemaRef),
2875 CurrentShaderEnvironment(llvm::Triple::UnknownEnvironment),
2876 CurrentShaderStageBit(0), ReportOnlyShaderStageIssues(false) {}
2877
2878 // AST traversal methods
2879 void RunOnTranslationUnit(const TranslationUnitDecl *TU);
2880 void RunOnFunction(const FunctionDecl *FD);
2881
2882 bool VisitDeclRefExpr(DeclRefExpr *DRE) override {
2883 FunctionDecl *FD = llvm::dyn_cast<FunctionDecl>(DRE->getDecl());
2884 if (FD)
2885 HandleFunctionOrMethodRef(FD, DRE);
2886 return true;
2887 }
2888
2889 bool VisitMemberExpr(MemberExpr *ME) override {
2890 FunctionDecl *FD = llvm::dyn_cast<FunctionDecl>(ME->getMemberDecl());
2891 if (FD)
2892 HandleFunctionOrMethodRef(FD, ME);
2893 return true;
2894 }
2895};
2896
2897void DiagnoseHLSLAvailability::HandleFunctionOrMethodRef(FunctionDecl *FD,
2898 Expr *RefExpr) {
2899 assert((isa<DeclRefExpr>(RefExpr) || isa<MemberExpr>(RefExpr)) &&
2900 "expected DeclRefExpr or MemberExpr");
2901
2902 if (const AvailabilityAttr *AA = FindAvailabilityAttr(FD))
2903 CheckDeclAvailability(
2904 FD, AA, SourceRange(RefExpr->getBeginLoc(), RefExpr->getEndLoc()));
2905
2906 // has a definition -> add to stack to be scanned
2907 const FunctionDecl *FDWithBody = nullptr;
2908 if (FD->hasBody(FDWithBody) && !WasAlreadyScannedInCurrentStage(FDWithBody))
2909 DeclsToScan.push_back(FDWithBody);
2910}
2911
2912void DiagnoseHLSLAvailability::RunOnTranslationUnit(
2913 const TranslationUnitDecl *TU) {
2914 const TargetInfo &TargetInfo = SemaRef.getASTContext().getTargetInfo();
2915 std::string &EntryName = TargetInfo.getTargetOpts().HLSLEntry;
2916 bool IsLibraryShader = TargetInfo.getTriple().getEnvironment() ==
2917 llvm::Triple::EnvironmentType::Library;
2918 SourceLocation EntryLoc{};
2919
2920 // Iterate over all shader entry functions and library exports, and for those
2921 // that have a body (definiton), run diag scan on each, setting appropriate
2922 // shader environment context based on whether it is a shader entry function
2923 // or an exported function. Exported functions can be in namespaces and in
2924 // export declarations so we need to scan those declaration contexts as well.
2926 DeclContextsToScan.push_back(TU);
2927
2928 while (!DeclContextsToScan.empty()) {
2929 const DeclContext *DC = DeclContextsToScan.pop_back_val();
2930 for (auto &D : DC->decls()) {
2931 // do not scan implicit declaration generated by the implementation
2932 if (D->isImplicit())
2933 continue;
2934
2935 // for namespace or export declaration add the context to the list to be
2936 // scanned later
2937 if (llvm::dyn_cast<NamespaceDecl>(D) || llvm::dyn_cast<ExportDecl>(D)) {
2938 DeclContextsToScan.push_back(llvm::dyn_cast<DeclContext>(D));
2939 continue;
2940 }
2941
2942 // skip over other decls or function decls without body
2943 const FunctionDecl *FD = llvm::dyn_cast<FunctionDecl>(D);
2944 if (!FD || !FD->isThisDeclarationADefinition())
2945 continue;
2946
2947 // shader entry point
2948 if (HLSLShaderAttr *ShaderAttr = FD->getAttr<HLSLShaderAttr>()) {
2949 if (!IsLibraryShader && FD->getName() == EntryName) {
2950 if (EntryLoc.isValid()) {
2951 SemaRef.Diag(FD->getLocation(),
2952 diag::err_hlsl_ambiguous_entry_point)
2953 << EntryName;
2954 SemaRef.Diag(EntryLoc, diag::note_previous_declaration_as)
2955 << EntryName;
2956 return;
2957 }
2958 EntryLoc = FD->getLocation();
2959 }
2960 SetShaderStageContext(ShaderAttr->getType());
2961 RunOnFunction(FD);
2962 continue;
2963 }
2964 // exported library function
2965 // FIXME: replace this loop with external linkage check once issue #92071
2966 // is resolved
2967 bool isExport = FD->isInExportDeclContext();
2968 if (!isExport) {
2969 for (const auto *Redecl : FD->redecls()) {
2970 if (Redecl->isInExportDeclContext()) {
2971 isExport = true;
2972 break;
2973 }
2974 }
2975 }
2976 if (isExport) {
2977 SetUnknownShaderStageContext();
2978 RunOnFunction(FD);
2979 continue;
2980 }
2981 }
2982 }
2983
2984 if (!IsLibraryShader && EntryLoc.isInvalid()) {
2985 SemaRef.Diag(TU->getLocation(), diag::err_hlsl_missing_entry_point)
2986 << EntryName;
2987 return;
2988 }
2989}
2990
2991void DiagnoseHLSLAvailability::RunOnFunction(const FunctionDecl *FD) {
2992 assert(DeclsToScan.empty() && "DeclsToScan should be empty");
2993 DeclsToScan.push_back(FD);
2994
2995 while (!DeclsToScan.empty()) {
2996 // Take one decl from the stack and check it by traversing its AST.
2997 // For any CallExpr found during the traversal add it's callee to the top of
2998 // the stack to be processed next. Functions already processed are stored in
2999 // ScannedDecls.
3000 const FunctionDecl *FD = DeclsToScan.pop_back_val();
3001
3002 // Decl was already scanned
3003 const unsigned ScannedStages = GetScannedStages(FD);
3004 if (WasAlreadyScannedInCurrentStage(ScannedStages))
3005 continue;
3006
3007 ReportOnlyShaderStageIssues = !NeverBeenScanned(ScannedStages);
3008
3009 AddToScannedFunctions(FD);
3010 TraverseStmt(FD->getBody());
3011 }
3012}
3013
3014bool DiagnoseHLSLAvailability::HasMatchingEnvironmentOrNone(
3015 const AvailabilityAttr *AA) {
3016 const IdentifierInfo *IIEnvironment = AA->getEnvironment();
3017 if (!IIEnvironment)
3018 return true;
3019
3020 llvm::Triple::EnvironmentType CurrentEnv = GetCurrentShaderEnvironment();
3021 if (CurrentEnv == llvm::Triple::UnknownEnvironment)
3022 return false;
3023
3024 llvm::Triple::EnvironmentType AttrEnv =
3025 AvailabilityAttr::getEnvironmentType(IIEnvironment->getName());
3026
3027 return CurrentEnv == AttrEnv;
3028}
3029
3030const AvailabilityAttr *
3031DiagnoseHLSLAvailability::FindAvailabilityAttr(const Decl *D) {
3032 AvailabilityAttr const *PartialMatch = nullptr;
3033 // Check each AvailabilityAttr to find the one for this platform.
3034 // For multiple attributes with the same platform try to find one for this
3035 // environment.
3036 for (const auto *A : D->attrs()) {
3037 if (const auto *Avail = dyn_cast<AvailabilityAttr>(A)) {
3038 const AvailabilityAttr *EffectiveAvail = Avail->getEffectiveAttr();
3039 StringRef AttrPlatform = EffectiveAvail->getPlatform()->getName();
3040 StringRef TargetPlatform =
3042
3043 // Match the platform name.
3044 if (AttrPlatform == TargetPlatform) {
3045 // Find the best matching attribute for this environment
3046 if (HasMatchingEnvironmentOrNone(EffectiveAvail))
3047 return Avail;
3048 PartialMatch = Avail;
3049 }
3050 }
3051 }
3052 return PartialMatch;
3053}
3054
3055// Check availability against target shader model version and current shader
3056// stage and emit diagnostic
3057void DiagnoseHLSLAvailability::CheckDeclAvailability(NamedDecl *D,
3058 const AvailabilityAttr *AA,
3059 SourceRange Range) {
3060
3061 const IdentifierInfo *IIEnv = AA->getEnvironment();
3062
3063 if (!IIEnv) {
3064 // The availability attribute does not have environment -> it depends only
3065 // on shader model version and not on specific the shader stage.
3066
3067 // Skip emitting the diagnostics if the diagnostic mode is set to
3068 // strict (-fhlsl-strict-availability) because all relevant diagnostics
3069 // were already emitted in the DiagnoseUnguardedAvailability scan
3070 // (SemaAvailability.cpp).
3071 if (SemaRef.getLangOpts().HLSLStrictAvailability)
3072 return;
3073
3074 // Do not report shader-stage-independent issues if scanning a function
3075 // that was already scanned in a different shader stage context (they would
3076 // be duplicate)
3077 if (ReportOnlyShaderStageIssues)
3078 return;
3079
3080 } else {
3081 // The availability attribute has environment -> we need to know
3082 // the current stage context to property diagnose it.
3083 if (InUnknownShaderStageContext())
3084 return;
3085 }
3086
3087 // Check introduced version and if environment matches
3088 bool EnvironmentMatches = HasMatchingEnvironmentOrNone(AA);
3089 VersionTuple Introduced = AA->getIntroduced();
3090 VersionTuple TargetVersion =
3092
3093 if (TargetVersion >= Introduced && EnvironmentMatches)
3094 return;
3095
3096 // Emit diagnostic message
3097 const TargetInfo &TI = SemaRef.getASTContext().getTargetInfo();
3098 llvm::StringRef PlatformName(
3099 AvailabilityAttr::getPrettyPlatformName(TI.getPlatformName()));
3100
3101 llvm::StringRef CurrentEnvStr =
3102 llvm::Triple::getEnvironmentTypeName(GetCurrentShaderEnvironment());
3103
3104 llvm::StringRef AttrEnvStr =
3105 AA->getEnvironment() ? AA->getEnvironment()->getName() : "";
3106 bool UseEnvironment = !AttrEnvStr.empty();
3107
3108 if (EnvironmentMatches) {
3109 SemaRef.Diag(Range.getBegin(), diag::warn_hlsl_availability)
3110 << Range << D << PlatformName << Introduced.getAsString()
3111 << UseEnvironment << CurrentEnvStr;
3112 } else {
3113 SemaRef.Diag(Range.getBegin(), diag::warn_hlsl_availability_unavailable)
3114 << Range << D;
3115 }
3116
3117 SemaRef.Diag(D->getLocation(), diag::note_partial_availability_specified_here)
3118 << D << PlatformName << Introduced.getAsString()
3119 << SemaRef.Context.getTargetInfo().getPlatformMinVersion().getAsString()
3120 << UseEnvironment << AttrEnvStr << CurrentEnvStr;
3121}
3122
3123} // namespace
3124
3126 // process default CBuffer - create buffer layout struct and invoke codegenCGH
3127 if (!DefaultCBufferDecls.empty()) {
3129 SemaRef.getASTContext(), SemaRef.getCurLexicalContext(),
3130 DefaultCBufferDecls);
3131 addImplicitBindingAttrToDecl(SemaRef, DefaultCBuffer, RegisterType::CBuffer,
3133 SemaRef.getCurLexicalContext()->addDecl(DefaultCBuffer);
3135
3136 // Set HasValidPackoffset if any of the decls has a register(c#) annotation;
3137 for (const Decl *VD : DefaultCBufferDecls) {
3138 const HLSLResourceBindingAttr *RBA =
3139 VD->getAttr<HLSLResourceBindingAttr>();
3140 if (RBA && RBA->hasRegisterSlot() &&
3141 RBA->getRegisterType() == HLSLResourceBindingAttr::RegisterType::C) {
3142 DefaultCBuffer->setHasValidPackoffset(true);
3143 break;
3144 }
3145 }
3146
3147 DeclGroupRef DG(DefaultCBuffer);
3148 SemaRef.Consumer.HandleTopLevelDecl(DG);
3149 }
3150 diagnoseAvailabilityViolations(TU);
3151}
3152
3153// For resource member access through a global struct array, verify that the
3154// array index selecting the struct element is a constant integer expression.
3155// Returns false if the member expression is invalid.
3157 assert((ME->getType()->isHLSLResourceRecord() ||
3159 "expected member expr to have resource record type or array of them");
3160
3161 // Walk the AST from MemberExpr to the VarDecl of the parent struct instance
3162 // and take note of any non-constant array indexing along the way. If the
3163 // VarDecl we find is a global variable, report error if there was any
3164 // non-constant array index in the resource member access along the way.
3165 const Expr *NonConstIndexExpr = nullptr;
3166 const Expr *E = ME->getBase();
3167 while (E) {
3168 if (const DeclRefExpr *DRE = dyn_cast<DeclRefExpr>(E)) {
3169 if (!NonConstIndexExpr)
3170 return true;
3171
3172 const VarDecl *VD = cast<VarDecl>(DRE->getDecl());
3173 if (!VD->hasGlobalStorage())
3174 return true;
3175
3176 SemaRef.Diag(NonConstIndexExpr->getExprLoc(),
3177 diag::err_hlsl_resource_member_array_access_not_constant);
3178 return false;
3179 }
3180
3181 if (const auto *ASE = dyn_cast<ArraySubscriptExpr>(E)) {
3182 const Expr *IdxExpr = ASE->getIdx();
3183 if (!IdxExpr->isIntegerConstantExpr(SemaRef.getASTContext()))
3184 NonConstIndexExpr = IdxExpr;
3185 E = ASE->getBase();
3186 } else if (const auto *SubME = dyn_cast<MemberExpr>(E)) {
3187 E = SubME->getBase();
3188 } else if (const auto *ICE = dyn_cast<ImplicitCastExpr>(E)) {
3189 E = ICE->getSubExpr();
3190 } else {
3191 llvm_unreachable("unexpected expr type in resource member access");
3192 }
3193 }
3194 return true;
3195}
3196
3198 CXXRecordDecl *RD) {
3199 QualType AddrSpaceType =
3200 SemaRef.Context.getCanonicalType(SemaRef.Context.getAddrSpaceQualType(
3201 Type.withConst(), LangAS::hlsl_constant));
3202 QualType ReturnTy = SemaRef.Context.getCanonicalType(
3203 SemaRef.Context.getLValueReferenceType(AddrSpaceType));
3204
3205 DeclarationName ConvName =
3206 SemaRef.Context.DeclarationNames.getCXXConversionFunctionName(
3207 CanQualType::CreateUnsafe(ReturnTy));
3208 LookupResult ConvR(SemaRef, ConvName, SourceLocation(),
3210 [[maybe_unused]] bool LookupSucceeded =
3211 SemaRef.LookupQualifiedName(ConvR, RD);
3212 assert(LookupSucceeded);
3213
3214 for (NamedDecl *D : ConvR) {
3216 return D;
3217 }
3218 return nullptr;
3219}
3220
3221std::optional<ExprResult>
3223 QualType BaseType = BaseExpr->getType();
3224 const HLSLAttributedResourceType *ResTy =
3225 HLSLAttributedResourceType::findHandleTypeOnResource(
3226 BaseType.getTypePtr());
3227 if (!ResTy ||
3228 ResTy->getAttrs().ResourceClass != llvm::dxil::ResourceClass::CBuffer)
3229 return std::nullopt;
3230
3231 QualType TemplateType = ResTy->getContainedType();
3232
3233 NamedDecl *NamedConversionDecl = getConstantBufferConversionFunction(
3234 TemplateType, BaseType->getAsCXXRecordDecl());
3235 assert(NamedConversionDecl &&
3236 "Could not find conversion function for ConstantBuffer.");
3237 auto *ConversionDecl =
3238 cast<CXXConversionDecl>(NamedConversionDecl->getUnderlyingDecl());
3239
3240 return SemaRef.BuildCXXMemberCallExpr(BaseExpr, NamedConversionDecl,
3241 ConversionDecl,
3242 /*HadMultipleCandidates=*/false);
3243}
3244
3245void SemaHLSL::diagnoseAvailabilityViolations(TranslationUnitDecl *TU) {
3246 // Skip running the diagnostics scan if the diagnostic mode is
3247 // strict (-fhlsl-strict-availability) and the target shader stage is known
3248 // because all relevant diagnostics were already emitted in the
3249 // DiagnoseUnguardedAvailability scan (SemaAvailability.cpp).
3251 if (SemaRef.getLangOpts().HLSLStrictAvailability &&
3252 TI.getTriple().getEnvironment() != llvm::Triple::EnvironmentType::Library)
3253 return;
3254
3255 DiagnoseHLSLAvailability(SemaRef).RunOnTranslationUnit(TU);
3256}
3257
3258static bool CheckAllArgsHaveSameType(Sema *S, CallExpr *TheCall) {
3259 assert(TheCall->getNumArgs() > 1);
3260 QualType ArgTy0 = TheCall->getArg(0)->getType();
3261
3262 for (unsigned I = 1, N = TheCall->getNumArgs(); I < N; ++I) {
3264 ArgTy0, TheCall->getArg(I)->getType())) {
3265 S->Diag(TheCall->getBeginLoc(), diag::err_vec_builtin_incompatible_vector)
3266 << TheCall->getDirectCallee() << /*useAllTerminology*/ true
3267 << SourceRange(TheCall->getArg(0)->getBeginLoc(),
3268 TheCall->getArg(N - 1)->getEndLoc());
3269 return true;
3270 }
3271 }
3272 return false;
3273}
3274
3276 QualType ArgType = Arg->getType();
3278 S->Diag(Arg->getBeginLoc(), diag::err_typecheck_convert_incompatible)
3279 << ArgType << ExpectedType << 1 << 0 << 0;
3280 return true;
3281 }
3282 return false;
3283}
3284
3286 Sema *S, CallExpr *TheCall,
3287 llvm::function_ref<bool(Sema *S, SourceLocation Loc, int ArgOrdinal,
3288 clang::QualType PassedType)>
3289 Check) {
3290 for (unsigned I = 0; I < TheCall->getNumArgs(); ++I) {
3291 Expr *Arg = TheCall->getArg(I);
3292 if (Check(S, Arg->getBeginLoc(), I + 1, Arg->getType()))
3293 return true;
3294 }
3295 return false;
3296}
3297
3299 int ArgOrdinal,
3300 clang::QualType PassedType) {
3301 clang::QualType BaseType =
3302 PassedType->isVectorType()
3303 ? PassedType->castAs<clang::VectorType>()->getElementType()
3304 : PassedType;
3305 if (!BaseType->isFloat32Type())
3306 return S->Diag(Loc, diag::err_builtin_invalid_arg_type)
3307 << ArgOrdinal << /* scalar or vector of */ 5 << /* no int */ 0
3308 << /* float */ 1 << PassedType;
3309 return false;
3310}
3311
3313 int ArgOrdinal,
3314 clang::QualType PassedType) {
3315 clang::QualType BaseType = PassedType;
3316 if (const auto *VT = PassedType->getAs<clang::VectorType>())
3317 BaseType = VT->getElementType();
3318 else if (const auto *MT = PassedType->getAs<clang::MatrixType>())
3319 BaseType = MT->getElementType();
3320
3321 if (!BaseType->isHalfType() && !BaseType->isFloat32Type())
3322 return S->Diag(Loc, diag::err_builtin_invalid_arg_type)
3323 << ArgOrdinal << /* scalar or vector of */ 5 << /* no int */ 0
3324 << /* half or float */ 2 << PassedType;
3325 return false;
3326}
3327
3329 int ArgOrdinal,
3330 clang::QualType PassedType) {
3331 clang::QualType BaseType =
3332 PassedType->isVectorType()
3333 ? PassedType->castAs<clang::VectorType>()->getElementType()
3334 : PassedType->isMatrixType()
3335 ? PassedType->castAs<clang::MatrixType>()->getElementType()
3336 : PassedType;
3337 if (!BaseType->isDoubleType()) {
3338 // FIXME: adopt standard `err_builtin_invalid_arg_type` instead of using
3339 // this custom error.
3340 return S->Diag(Loc, diag::err_builtin_requires_double_type)
3341 << ArgOrdinal << PassedType;
3342 }
3343
3344 return false;
3345}
3346
3347static bool CheckModifiableLValue(Sema *S, CallExpr *TheCall,
3348 unsigned ArgIndex) {
3349 auto *Arg = TheCall->getArg(ArgIndex);
3350 SourceLocation OrigLoc = Arg->getExprLoc();
3351 if (Arg->IgnoreCasts()->isModifiableLvalue(S->Context, &OrigLoc) ==
3353 return false;
3354 S->Diag(OrigLoc, diag::error_hlsl_inout_lvalue) << Arg << 0;
3355 return true;
3356}
3357
3358// Verifies that the argument at `ArgIndex` of `TheCall` refers to memory in
3359// one of `AllowedSpaces`. Intended for HLSL builtins (e.g. atomics).
3360static bool CheckArgAddrSpaceOneOf(Sema *S, CallExpr *TheCall,
3361 unsigned ArgIndex,
3362 ArrayRef<LangAS> AllowedSpaces) {
3363 Expr *Arg = TheCall->getArg(ArgIndex);
3364 QualType LValueTy = Arg->IgnoreCasts()->getType();
3365 if (llvm::is_contained(AllowedSpaces, LValueTy.getAddressSpace()))
3366 return false;
3367 S->Diag(Arg->getBeginLoc(), diag::err_hlsl_atomic_arg_addr_space)
3368 << (ArgIndex + 1) << LValueTy;
3369 return true;
3370}
3371
3372static bool CheckNoDoubleVectors(Sema *S, SourceLocation Loc, int ArgOrdinal,
3373 clang::QualType PassedType) {
3374 const auto *VecTy = PassedType->getAs<VectorType>();
3375 if (!VecTy)
3376 return false;
3377
3378 if (VecTy->getElementType()->isDoubleType())
3379 return S->Diag(Loc, diag::err_builtin_invalid_arg_type)
3380 << ArgOrdinal << /* scalar */ 1 << /* no int */ 0 << /* fp */ 1
3381 << PassedType;
3382 return false;
3383}
3384
3386 int ArgOrdinal,
3387 clang::QualType PassedType) {
3388 if (!PassedType->hasIntegerRepresentation() &&
3389 !PassedType->hasFloatingRepresentation())
3390 return S->Diag(Loc, diag::err_builtin_invalid_arg_type)
3391 << ArgOrdinal << /* scalar or vector of */ 5 << /* integer */ 1
3392 << /* fp */ 1 << PassedType;
3393 return false;
3394}
3395
3397 int ArgOrdinal,
3398 clang::QualType PassedType) {
3399 if (auto *VecTy = PassedType->getAs<VectorType>())
3400 if (VecTy->getElementType()->isUnsignedIntegerType())
3401 return false;
3402
3403 return S->Diag(Loc, diag::err_builtin_invalid_arg_type)
3404 << ArgOrdinal << /* vector of */ 4 << /* uint */ 3 << /* no fp */ 0
3405 << PassedType;
3406}
3407
3408// checks for unsigned ints of all sizes
3410 int ArgOrdinal,
3411 clang::QualType PassedType) {
3412 if (!PassedType->hasUnsignedIntegerRepresentation())
3413 return S->Diag(Loc, diag::err_builtin_invalid_arg_type)
3414 << ArgOrdinal << /* scalar or vector of */ 5 << /* unsigned int */ 3
3415 << /* no fp */ 0 << PassedType;
3416 return false;
3417}
3418
3419static bool CheckExpectedBitWidth(Sema *S, CallExpr *TheCall,
3420 unsigned ArgOrdinal, unsigned Width) {
3421 QualType ArgTy = TheCall->getArg(0)->getType();
3422 if (auto *VTy = ArgTy->getAs<VectorType>())
3423 ArgTy = VTy->getElementType();
3424 // ensure arg type has expected bit width
3425 uint64_t ElementBitCount =
3427 if (ElementBitCount != Width) {
3428 S->Diag(TheCall->getArg(0)->getBeginLoc(),
3429 diag::err_integer_incorrect_bit_count)
3430 << Width << ElementBitCount;
3431 return true;
3432 }
3433 return false;
3434}
3435
3437 QualType ReturnType) {
3438 if (auto *VecTyA = TheCall->getArg(0)->getType()->getAs<VectorType>())
3439 ReturnType =
3440 S->Context.getExtVectorType(ReturnType, VecTyA->getNumElements());
3441 else if (auto *MatTyA =
3442 TheCall->getArg(0)->getType()->getAs<ConstantMatrixType>())
3443 ReturnType = S->Context.getConstantMatrixType(
3444 ReturnType, MatTyA->getNumRows(), MatTyA->getNumColumns());
3445
3446 TheCall->setType(ReturnType);
3447}
3448
3449static bool CheckScalarOrVector(Sema *S, CallExpr *TheCall, QualType Scalar,
3450 unsigned ArgIndex) {
3451 assert(TheCall->getNumArgs() >= ArgIndex);
3452 QualType ArgType = TheCall->getArg(ArgIndex)->getType();
3453 auto *VTy = ArgType->getAs<VectorType>();
3454 // not the scalar or vector<scalar>
3455 if (!(S->Context.hasSameUnqualifiedType(ArgType, Scalar) ||
3456 (VTy &&
3457 S->Context.hasSameUnqualifiedType(VTy->getElementType(), Scalar)))) {
3458 S->Diag(TheCall->getArg(0)->getBeginLoc(),
3459 diag::err_typecheck_expect_scalar_or_vector)
3460 << ArgType << Scalar;
3461 return true;
3462 }
3463 return false;
3464}
3465
3467 QualType Scalar, unsigned ArgIndex) {
3468 assert(TheCall->getNumArgs() > ArgIndex);
3469
3470 Expr *Arg = TheCall->getArg(ArgIndex);
3471 QualType ArgType = Arg->getType();
3472
3473 // Scalar: T
3474 if (S->Context.hasSameUnqualifiedType(ArgType, Scalar))
3475 return false;
3476
3477 // Vector: vector<T>
3478 if (const auto *VTy = ArgType->getAs<VectorType>()) {
3479 if (S->Context.hasSameUnqualifiedType(VTy->getElementType(), Scalar))
3480 return false;
3481 }
3482
3483 // Matrix: ConstantMatrixType with element type T
3484 if (const auto *MTy = ArgType->getAs<ConstantMatrixType>()) {
3485 if (S->Context.hasSameUnqualifiedType(MTy->getElementType(), Scalar))
3486 return false;
3487 }
3488
3489 // Not a scalar/vector/matrix-of-scalar
3490 S->Diag(Arg->getBeginLoc(),
3491 diag::err_typecheck_expect_scalar_or_vector_or_matrix)
3492 << ArgType << Scalar;
3493 return true;
3494}
3495
3496static bool CheckAnyScalarOrVector(Sema *S, CallExpr *TheCall,
3497 unsigned ArgIndex) {
3498 assert(TheCall->getNumArgs() >= ArgIndex);
3499 QualType ArgType = TheCall->getArg(ArgIndex)->getType();
3500 auto *VTy = ArgType->getAs<VectorType>();
3501 // not the scalar or vector<scalar>
3502 if (!(ArgType->isScalarType() ||
3503 (VTy && VTy->getElementType()->isScalarType()))) {
3504 S->Diag(TheCall->getArg(0)->getBeginLoc(),
3505 diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
3506 << ArgType << 1;
3507 return true;
3508 }
3509 return false;
3510}
3511
3513 unsigned ArgIndex) {
3514 assert(TheCall->getNumArgs() > ArgIndex);
3515 QualType ArgType = TheCall->getArg(ArgIndex)->getType();
3516 if (ArgType->isDependentType())
3517 return false;
3518
3519 QualType ElementType = ArgType;
3520 if (const auto *VectorTy = ArgType->getAs<VectorType>())
3521 ElementType = VectorTy->getElementType();
3522 else if (const auto *MatrixTy = ArgType->getAs<ConstantMatrixType>())
3523 ElementType = MatrixTy->getElementType();
3524
3525 if (ElementType->isBooleanType())
3526 return false;
3527
3528 if (ElementType->isIntegerType() || ElementType->isRealFloatingType()) {
3529 unsigned BitWidth = S->Context.getTypeSize(ElementType);
3530 if (BitWidth == 16 || BitWidth == 32 || BitWidth == 64)
3531 return false;
3532 }
3533
3534 S->Diag(TheCall->getArg(ArgIndex)->getBeginLoc(),
3535 diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
3536 << ArgType << 2;
3537 return true;
3538}
3539
3540// Check that the argument is not a bool or vector<bool>
3541// Returns true on error
3543 unsigned ArgIndex) {
3544 QualType BoolType = S->getASTContext().BoolTy;
3545 assert(ArgIndex < TheCall->getNumArgs());
3546 QualType ArgType = TheCall->getArg(ArgIndex)->getType();
3547 auto *VTy = ArgType->getAs<VectorType>();
3548 // is the bool or vector<bool>
3549 if (S->Context.hasSameUnqualifiedType(ArgType, BoolType) ||
3550 (VTy &&
3551 S->Context.hasSameUnqualifiedType(VTy->getElementType(), BoolType))) {
3552 S->Diag(TheCall->getArg(0)->getBeginLoc(),
3553 diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
3554 << ArgType << 0;
3555 return true;
3556 }
3557 return false;
3558}
3559
3560static bool CheckWaveActive(Sema *S, CallExpr *TheCall) {
3561 if (CheckNotBoolScalarOrVector(S, TheCall, 0))
3562 return true;
3563 return false;
3564}
3565
3566static bool CheckWavePrefix(Sema *S, CallExpr *TheCall) {
3567 if (CheckNotBoolScalarOrVector(S, TheCall, 0))
3568 return true;
3569 return false;
3570}
3571
3572static bool CheckBoolSelect(Sema *S, CallExpr *TheCall) {
3573 assert(TheCall->getNumArgs() == 3);
3574 Expr *Arg1 = TheCall->getArg(1);
3575 Expr *Arg2 = TheCall->getArg(2);
3576 if (!S->Context.hasSameUnqualifiedType(Arg1->getType(), Arg2->getType())) {
3577 S->Diag(TheCall->getBeginLoc(),
3578 diag::err_typecheck_call_different_arg_types)
3579 << Arg1->getType() << Arg2->getType() << Arg1->getSourceRange()
3580 << Arg2->getSourceRange();
3581 return true;
3582 }
3583
3584 TheCall->setType(Arg1->getType());
3585 return false;
3586}
3587
3588static bool CheckVectorSelect(Sema *S, CallExpr *TheCall) {
3589 assert(TheCall->getNumArgs() == 3);
3590 Expr *Arg1 = TheCall->getArg(1);
3591 QualType Arg1Ty = Arg1->getType();
3592 Expr *Arg2 = TheCall->getArg(2);
3593 QualType Arg2Ty = Arg2->getType();
3594
3595 QualType Arg1ScalarTy = Arg1Ty;
3596 if (auto VTy = Arg1ScalarTy->getAs<VectorType>())
3597 Arg1ScalarTy = VTy->getElementType();
3598
3599 QualType Arg2ScalarTy = Arg2Ty;
3600 if (auto VTy = Arg2ScalarTy->getAs<VectorType>())
3601 Arg2ScalarTy = VTy->getElementType();
3602
3603 if (!S->Context.hasSameUnqualifiedType(Arg1ScalarTy, Arg2ScalarTy))
3604 S->Diag(Arg1->getBeginLoc(), diag::err_hlsl_builtin_scalar_vector_mismatch)
3605 << /* second and third */ 1 << TheCall->getCallee() << Arg1Ty << Arg2Ty;
3606
3607 QualType Arg0Ty = TheCall->getArg(0)->getType();
3608 unsigned Arg0Length = Arg0Ty->getAs<VectorType>()->getNumElements();
3609 unsigned Arg1Length = Arg1Ty->isVectorType()
3610 ? Arg1Ty->getAs<VectorType>()->getNumElements()
3611 : 0;
3612 unsigned Arg2Length = Arg2Ty->isVectorType()
3613 ? Arg2Ty->getAs<VectorType>()->getNumElements()
3614 : 0;
3615 if (Arg1Length > 0 && Arg0Length != Arg1Length) {
3616 S->Diag(TheCall->getBeginLoc(),
3617 diag::err_typecheck_vector_lengths_not_equal)
3618 << Arg0Ty << Arg1Ty << TheCall->getArg(0)->getSourceRange()
3619 << Arg1->getSourceRange();
3620 return true;
3621 }
3622
3623 if (Arg2Length > 0 && Arg0Length != Arg2Length) {
3624 S->Diag(TheCall->getBeginLoc(),
3625 diag::err_typecheck_vector_lengths_not_equal)
3626 << Arg0Ty << Arg2Ty << TheCall->getArg(0)->getSourceRange()
3627 << Arg2->getSourceRange();
3628 return true;
3629 }
3630
3631 TheCall->setType(
3632 S->getASTContext().getExtVectorType(Arg1ScalarTy, Arg0Length));
3633 return false;
3634}
3635
3636static bool CheckMatrixSelect(Sema *S, CallExpr *TheCall) {
3637 assert(TheCall->getNumArgs() == 3);
3638 Expr *Arg1 = TheCall->getArg(1);
3639 QualType Arg1Ty = Arg1->getType();
3640 Expr *Arg2 = TheCall->getArg(2);
3641 QualType Arg2Ty = Arg2->getType();
3642
3643 QualType Arg1ScalarTy = Arg1Ty;
3644 if (auto MTy = Arg1ScalarTy->getAs<ConstantMatrixType>())
3645 Arg1ScalarTy = MTy->getElementType();
3646
3647 QualType Arg2ScalarTy = Arg2Ty;
3648 if (auto MTy = Arg2ScalarTy->getAs<ConstantMatrixType>())
3649 Arg2ScalarTy = MTy->getElementType();
3650
3651 if (!S->Context.hasSameUnqualifiedType(Arg1ScalarTy, Arg2ScalarTy))
3652 S->Diag(Arg1->getBeginLoc(), diag::err_hlsl_builtin_scalar_vector_mismatch)
3653 << /* second and third */ 1 << TheCall->getCallee() << Arg1Ty << Arg2Ty;
3654
3655 QualType Arg0Ty = TheCall->getArg(0)->getType();
3656 auto *Arg0MatTy = Arg0Ty->getAs<ConstantMatrixType>();
3657 unsigned Arg0Rows = Arg0MatTy->getNumRows();
3658 unsigned Arg0Cols = Arg0MatTy->getNumColumns();
3659
3660 for (Expr *Arg : {Arg1, Arg2}) {
3661 auto *MTy = Arg->getType()->getAs<ConstantMatrixType>();
3662 if (MTy &&
3663 (MTy->getNumRows() != Arg0Rows || MTy->getNumColumns() != Arg0Cols)) {
3664 S->Diag(TheCall->getBeginLoc(),
3665 diag::err_typecheck_vector_lengths_not_equal)
3666 << Arg0Ty << Arg->getType() << TheCall->getArg(0)->getSourceRange()
3667 << Arg->getSourceRange();
3668 return true;
3669 }
3670 }
3671
3672 TheCall->setType(
3673 S->Context.getConstantMatrixType(Arg1ScalarTy, Arg0Rows, Arg0Cols));
3674 return false;
3675}
3676
3678 unsigned Count) {
3679 return Count > 1 ? S.Context.getExtVectorType(BaseType, Count) : BaseType;
3680}
3681
3682static bool CheckScalarFloatOperand(Sema &S, CallExpr *TheCall,
3683 unsigned ArgIndex) {
3684 return CheckArgTypeMatches(&S, TheCall->getArg(ArgIndex), S.Context.FloatTy);
3685}
3686
3687static bool CheckIndexType(Sema *S, CallExpr *TheCall, unsigned IndexArgIndex) {
3688 assert(TheCall->getNumArgs() > IndexArgIndex && "Index argument missing");
3689 QualType ArgType = TheCall->getArg(IndexArgIndex)->getType();
3690 QualType IndexTy = ArgType;
3691 unsigned int ActualDim = 1;
3692 if (const auto *VTy = IndexTy->getAs<VectorType>()) {
3693 ActualDim = VTy->getNumElements();
3694 IndexTy = VTy->getElementType();
3695 }
3696 if (!IndexTy->isIntegerType()) {
3697 S->Diag(TheCall->getArg(IndexArgIndex)->getBeginLoc(),
3698 diag::err_typecheck_expect_int)
3699 << ArgType;
3700 return true;
3701 }
3702
3703 QualType ResourceArgTy = TheCall->getArg(0)->getType();
3704 const HLSLAttributedResourceType *ResTy =
3705 ResourceArgTy.getTypePtr()->getAs<HLSLAttributedResourceType>();
3706 assert(ResTy && "Resource argument must be a resource");
3707 HLSLAttributedResourceType::Attributes ResAttrs = ResTy->getAttrs();
3708
3709 unsigned int ExpectedDim = 1;
3710 if (ResAttrs.ResourceDimension != llvm::dxil::ResourceDimension::Unknown)
3711 ExpectedDim = getResourceDimensions(ResAttrs.ResourceDimension) +
3712 (ResAttrs.IsArray ? 1 : 0);
3713
3714 if (ActualDim != ExpectedDim) {
3715 S->Diag(TheCall->getArg(IndexArgIndex)->getBeginLoc(),
3716 diag::err_hlsl_builtin_resource_coordinate_dimension_mismatch)
3717 << cast<NamedDecl>(TheCall->getCalleeDecl()) << ExpectedDim
3718 << ActualDim;
3719 return true;
3720 }
3721
3722 return false;
3723}
3724
3726 Sema *S, CallExpr *TheCall, unsigned ArgIndex,
3727 llvm::function_ref<bool(const HLSLAttributedResourceType *ResType)> Check =
3728 nullptr) {
3729 assert(TheCall->getNumArgs() >= ArgIndex);
3730 QualType ArgType = TheCall->getArg(ArgIndex)->getType();
3731 const HLSLAttributedResourceType *ResTy =
3732 ArgType.getTypePtr()->getAs<HLSLAttributedResourceType>();
3733 if (!ResTy) {
3734 S->Diag(TheCall->getArg(ArgIndex)->getBeginLoc(),
3735 diag::err_typecheck_expect_hlsl_resource)
3736 << ArgType;
3737 return true;
3738 }
3739 if (Check && Check(ResTy)) {
3740 S->Diag(TheCall->getArg(ArgIndex)->getExprLoc(),
3741 diag::err_invalid_hlsl_resource_type)
3742 << ArgType;
3743 return true;
3744 }
3745 return false;
3746}
3747
3749 QualType MainHandleTy) {
3750 assert(MainHandleTy->isHLSLAttributedResourceType() &&
3751 "expected resource handle type");
3752 auto *MainResType = MainHandleTy->getAs<HLSLAttributedResourceType>();
3753 auto MainAttrs = MainResType->getAttrs();
3754 assert(!MainAttrs.IsCounter && "cannot create a counter from a counter");
3755 MainAttrs.IsCounter = true;
3756 return AST.getHLSLAttributedResourceType(MainResType->getWrappedType(),
3757 MainResType->getContainedType(),
3758 MainAttrs);
3759}
3760
3761enum class SampleKind { Sample, Bias, Grad, Level, Cmp, CmpLevelZero };
3762
3763static StringRef getSampleMethodName(SampleKind Kind) {
3764 switch (Kind) {
3765 case SampleKind::Sample:
3766 return "Sample";
3767 case SampleKind::Bias:
3768 return "SampleBias";
3769 case SampleKind::Grad:
3770 return "SampleGrad";
3771 case SampleKind::Level:
3772 return "SampleLevel";
3773 case SampleKind::Cmp:
3774 return "SampleCmp";
3776 return "SampleCmpLevelZero";
3777 }
3778 llvm_unreachable("Invalid SampleKind");
3779}
3780
3781// Returns the name of the resource method whose body the sampling or gather
3782// builtin is being emitted into, which is the name the user called. This
3783// matters for methods that share a builtin, like 'Gather' and 'GatherRed'.
3784// Falls back to DefaultName if the builtin is used outside of a resource
3785// method.
3786static StringRef getCurrentResourceMethodName(Sema &S, StringRef DefaultName) {
3787 const auto *MD = dyn_cast_if_present<CXXMethodDecl>(S.getCurFunctionDecl());
3788 if (!MD || !MD->getDeclName().isIdentifier())
3789 return DefaultName;
3790
3791 QualType RecordTy = S.Context.getCanonicalTagType(MD->getParent());
3792 if (!RecordTy->isHLSLResourceRecord())
3793 return DefaultName;
3794
3795 return MD->getName();
3796}
3797
3798// Returns the element type of a typed resource's contained type. Typed resource
3799// element types are scalars or vectors of scalars, so anything that is not a
3800// vector is already the element type.
3802 if (const auto *VecTy = ContainedType->getAs<VectorType>())
3803 return VecTy->getElementType();
3804 return ContainedType;
3805}
3806
3807// Sampling from and gathering on resources with a 'double' element type is not
3808// supported. Such resources are still valid declarations whose contents can be
3809// accessed by other means, like Load or the subscript operator.
3810static bool CheckNoDoubleElementType(Sema &S, CallExpr *TheCall,
3811 QualType ContainedType,
3812 StringRef DefaultName) {
3813 QualType EltTy = getTypedResourceElementType(ContainedType);
3814 if (!EltTy->isSpecificBuiltinType(BuiltinType::Double))
3815 return false;
3816
3817 S.Diag(TheCall->getBeginLoc(), diag::err_hlsl_sample_double_element_type)
3818 << getCurrentResourceMethodName(S, DefaultName) << ContainedType;
3819 return true;
3820}
3821
3822// Sampling textures with an integer element type was introduced in SM 6.7 as
3823// part of Advanced Texture Operations. The shader model only applies to DirectX
3824// targets; Vulkan has no such restriction.
3826 QualType ContainedType,
3827 SampleKind Kind) {
3828 // Comparison sampling requires a floating point element type at every shader
3829 // model, which the caller diagnoses.
3830 if (Kind == SampleKind::Cmp || Kind == SampleKind::CmpLevelZero)
3831 return false;
3832
3833 // 'bool' is an integer type in HLSL, but sampling bool resources is never
3834 // allowed, so it must not be reported as requiring shader model 6.7.
3835 QualType EltTy = getTypedResourceElementType(ContainedType);
3836 if (!EltTy->isIntegerType() || EltTy->isBooleanType())
3837 return false;
3838
3839 const TargetInfo &TI = S.Context.getTargetInfo();
3840 if (!TI.getTriple().isDXIL())
3841 return false;
3842
3843 VersionTuple SMVersion = TI.getPlatformMinVersion();
3844 if (SMVersion >= VersionTuple(6, 7))
3845 return false;
3846
3847 S.Diag(TheCall->getBeginLoc(), diag::err_hlsl_sample_integer_element_type)
3849 << ContainedType << SMVersion.getAsString();
3850 return true;
3851}
3852
3854 bool IncludeArraySlice = true) {
3855 // Check the texture handle.
3856 if (CheckResourceHandle(&S, TheCall, 0,
3857 [](const HLSLAttributedResourceType *ResType) {
3858 return ResType->getAttrs().ResourceDimension ==
3859 llvm::dxil::ResourceDimension::Unknown;
3860 }))
3861 return true;
3862
3863 // Check the sampler handle.
3864 if (CheckResourceHandle(&S, TheCall, 1,
3865 [](const HLSLAttributedResourceType *ResType) {
3866 return ResType->getAttrs().ResourceClass !=
3867 llvm::hlsl::ResourceClass::Sampler;
3868 }))
3869 return true;
3870
3871 auto *ResourceTy =
3872 TheCall->getArg(0)->getType()->castAs<HLSLAttributedResourceType>();
3873
3874 // Check the location.
3875 unsigned ExpectedDim =
3876 getResourceDimensions(ResourceTy->getAttrs().ResourceDimension) +
3877 (IncludeArraySlice && ResourceTy->getAttrs().IsArray ? 1 : 0);
3879 &S, TheCall->getArg(2),
3880 getVectorOrScalarType(S, S.Context.FloatTy, ExpectedDim)))
3881 return true;
3882
3883 return false;
3884}
3885
3886static bool CheckCalculateLodBuiltin(Sema &S, CallExpr *TheCall) {
3887 if (S.checkArgCount(TheCall, 3))
3888 return true;
3889
3890 // CalculateLevelOfDetail location uses resource dimension only (e.g. float2
3891 // for 2D), not an extra array slice component like Sample/Gather.
3892 if (CheckTextureSamplerAndLocation(S, TheCall, /*IncludeArraySlice=*/false))
3893 return true;
3894
3895 TheCall->setType(S.Context.FloatTy);
3896 return false;
3897}
3898
3899static bool CheckGatherBuiltin(Sema &S, CallExpr *TheCall, bool IsCmp) {
3900 if (S.checkArgCountRange(TheCall, IsCmp ? 5 : 4, IsCmp ? 6 : 5))
3901 return true;
3902
3903 if (CheckTextureSamplerAndLocation(S, TheCall))
3904 return true;
3905
3906 unsigned NextIdx = 3;
3907 if (IsCmp) {
3908 // Check the compare value.
3909 if (CheckScalarFloatOperand(S, TheCall, NextIdx))
3910 return true;
3911 NextIdx++;
3912 }
3913
3914 // Check the component operand.
3915 if (CheckArgTypeMatches(&S, TheCall->getArg(NextIdx),
3917 return true;
3918 Expr *ComponentArg = TheCall->getArg(NextIdx);
3919
3920 // GatherCmp operations on Vulkan target must use component 0 (Red).
3921 if (IsCmp && S.getASTContext().getTargetInfo().getTriple().isSPIRV()) {
3922 std::optional<llvm::APSInt> ComponentOpt =
3923 ComponentArg->getIntegerConstantExpr(S.getASTContext());
3924 if (ComponentOpt) {
3925 int64_t ComponentVal = ComponentOpt->getSExtValue();
3926 if (ComponentVal != 0) {
3927 // Issue an error if the component is not 0 (Red).
3928 // 0 -> Red, 1 -> Green, 2 -> Blue, 3 -> Alpha
3929 assert(ComponentVal >= 0 && ComponentVal <= 3 &&
3930 "The component is not in the expected range.");
3931 S.Diag(ComponentArg->getBeginLoc(),
3932 diag::err_hlsl_gathercmp_invalid_component)
3933 << ComponentVal;
3934 return true;
3935 }
3936 }
3937 }
3938
3939 NextIdx++;
3940
3941 // Check the offset operand.
3942 const HLSLAttributedResourceType *ResourceTy =
3943 TheCall->getArg(0)->getType()->castAs<HLSLAttributedResourceType>();
3944 if (TheCall->getNumArgs() > NextIdx) {
3945 unsigned ExpectedDim =
3946 getResourceDimensions(ResourceTy->getAttrs().ResourceDimension);
3948 &S, TheCall->getArg(NextIdx),
3949 getVectorOrScalarType(S, S.Context.IntTy, ExpectedDim)))
3950 return true;
3951 NextIdx++;
3952 }
3953
3954 assert(ResourceTy->hasContainedType() &&
3955 "Expecting a contained type for resource with a dimension "
3956 "attribute.");
3957 QualType ReturnType = ResourceTy->getContainedType();
3958
3959 if (CheckNoDoubleElementType(S, TheCall, ReturnType,
3960 IsCmp ? "GatherCmp" : "Gather"))
3961 return true;
3962
3963 if (IsCmp) {
3964 if (!ReturnType->hasFloatingRepresentation()) {
3965 S.Diag(TheCall->getBeginLoc(), diag::err_hlsl_samplecmp_requires_float);
3966 return true;
3967 }
3968 }
3969
3970 if (const auto *VecTy = ReturnType->getAs<VectorType>())
3971 ReturnType = VecTy->getElementType();
3972 ReturnType = S.Context.getExtVectorType(ReturnType, 4);
3973
3974 TheCall->setType(ReturnType);
3975
3976 return false;
3977}
3978static bool CheckLoadLevelBuiltin(Sema &S, CallExpr *TheCall) {
3979 if (S.checkArgCountRange(TheCall, 2, 3))
3980 return true;
3981
3982 // Check the texture handle.
3983 if (CheckResourceHandle(&S, TheCall, 0,
3984 [](const HLSLAttributedResourceType *ResType) {
3985 return ResType->getAttrs().ResourceDimension ==
3986 llvm::dxil::ResourceDimension::Unknown;
3987 }))
3988 return true;
3989
3990 auto *ResourceTy =
3991 TheCall->getArg(0)->getType()->castAs<HLSLAttributedResourceType>();
3992
3993 // A UAV descriptor binds a single mip slice, so a RWTexture location has no
3994 // mip component to select, and TextureLoad on a UAV takes no offset.
3995 bool IsUAV =
3996 ResourceTy->getAttrs().ResourceClass == llvm::dxil::ResourceClass::UAV;
3997 if (IsUAV && S.checkArgCount(TheCall, 2))
3998 return true;
3999
4000 // Check the location: int3 for Texture2D and int4 for Texture2DArray, which
4001 // both carry a trailing mip level; int2 and int3 for the RWTexture forms,
4002 // which do not.
4003 unsigned ResourceDim =
4004 getResourceDimensions(ResourceTy->getAttrs().ResourceDimension);
4005 unsigned LocationDim = ResourceDim + (ResourceTy->getAttrs().IsArray ? 1 : 0);
4006 if (!IsUAV)
4007 ++LocationDim;
4009 &S, TheCall->getArg(1),
4010 getVectorOrScalarType(S, S.Context.IntTy, LocationDim)))
4011 return true;
4012
4013 // Check the offset operand (int2 for 2D textures; no array slice).
4014 if (TheCall->getNumArgs() > 2) {
4016 &S, TheCall->getArg(2),
4017 getVectorOrScalarType(S, S.Context.IntTy, ResourceDim)))
4018 return true;
4019 }
4020
4021 TheCall->setType(ResourceTy->getContainedType());
4022 return false;
4023}
4024
4025static bool CheckLoadMSBuiltin(Sema &S, CallExpr *TheCall) {
4026 if (S.checkArgCountRange(TheCall, 3, 4))
4027 return true;
4028
4029 // Check the multisampled texture handle.
4030 if (CheckResourceHandle(&S, TheCall, 0,
4031 [](const HLSLAttributedResourceType *ResType) {
4032 return !ResType->isMultiSampled();
4033 }))
4034 return true;
4035
4036 auto *ResourceTy =
4037 TheCall->getArg(0)->getType()->castAs<HLSLAttributedResourceType>();
4038
4039 // Check the location (int2 for Texture2DMS, int3 for Texture2DMSArray).
4040 // Unlike Load on regular textures, there is no mip/LOD component.
4041 unsigned ResourceDim =
4042 getResourceDimensions(ResourceTy->getAttrs().ResourceDimension);
4043 unsigned LocationDim = ResourceDim + (ResourceTy->getAttrs().IsArray ? 1 : 0);
4045 &S, TheCall->getArg(1),
4046 getVectorOrScalarType(S, S.Context.IntTy, LocationDim)))
4047 return true;
4048
4049 // Check the sample index operand (scalar int).
4050 if (CheckArgTypeMatches(&S, TheCall->getArg(2), S.Context.IntTy))
4051 return true;
4052
4053 // Check the offset operand (int2 for 2D textures; no array slice).
4054 if (TheCall->getNumArgs() > 3) {
4056 &S, TheCall->getArg(3),
4057 getVectorOrScalarType(S, S.Context.IntTy, ResourceDim)))
4058 return true;
4059 }
4060
4061 TheCall->setType(ResourceTy->getContainedType());
4062 return false;
4063}
4064
4065static bool CheckSamplingBuiltin(Sema &S, CallExpr *TheCall, SampleKind Kind) {
4066 unsigned MinArgs, MaxArgs;
4067 if (Kind == SampleKind::Sample) {
4068 MinArgs = 3;
4069 MaxArgs = 5;
4070 } else if (Kind == SampleKind::Bias) {
4071 MinArgs = 4;
4072 MaxArgs = 6;
4073 } else if (Kind == SampleKind::Grad) {
4074 MinArgs = 5;
4075 MaxArgs = 7;
4076 } else if (Kind == SampleKind::Level) {
4077 MinArgs = 4;
4078 MaxArgs = 5;
4079 } else if (Kind == SampleKind::Cmp) {
4080 MinArgs = 4;
4081 MaxArgs = 6;
4082 } else {
4083 assert(Kind == SampleKind::CmpLevelZero);
4084 MinArgs = 4;
4085 MaxArgs = 5;
4086 }
4087
4088 if (S.checkArgCountRange(TheCall, MinArgs, MaxArgs))
4089 return true;
4090
4091 if (CheckTextureSamplerAndLocation(S, TheCall))
4092 return true;
4093
4094 const HLSLAttributedResourceType *ResourceTy =
4095 TheCall->getArg(0)->getType()->castAs<HLSLAttributedResourceType>();
4096 unsigned ExpectedDim =
4097 getResourceDimensions(ResourceTy->getAttrs().ResourceDimension);
4098
4099 unsigned NextIdx = 3;
4100 if (Kind == SampleKind::Bias || Kind == SampleKind::Level ||
4101 Kind == SampleKind::Cmp || Kind == SampleKind::CmpLevelZero) {
4102 // Check the bias, lod level, or compare value, depending on the kind.
4103 // All of them must be a scalar float value.
4104 if (CheckScalarFloatOperand(S, TheCall, NextIdx))
4105 return true;
4106 NextIdx++;
4107 } else if (Kind == SampleKind::Grad) {
4108 QualType GradTy = getVectorOrScalarType(S, S.Context.FloatTy, ExpectedDim);
4109
4110 // Check the DDX operand.
4111 if (CheckArgTypeMatches(&S, TheCall->getArg(NextIdx), GradTy))
4112 return true;
4113
4114 // Check the DDY operand.
4115 if (CheckArgTypeMatches(&S, TheCall->getArg(NextIdx + 1), GradTy))
4116 return true;
4117 NextIdx += 2;
4118 }
4119
4120 // Check the offset operand (if applicable).
4121 if (hasResourceOffset(ResourceTy->getAttrs().ResourceDimension) &&
4122 TheCall->getNumArgs() > NextIdx) {
4124 &S, TheCall->getArg(NextIdx),
4125 getVectorOrScalarType(S, S.Context.IntTy, ExpectedDim)))
4126 return true;
4127 NextIdx++;
4128 }
4129
4130 // Check the clamp operand.
4131 if (Kind != SampleKind::Level && Kind != SampleKind::CmpLevelZero &&
4132 TheCall->getNumArgs() > NextIdx) {
4133 if (CheckScalarFloatOperand(S, TheCall, NextIdx))
4134 return true;
4135 }
4136
4137 assert(ResourceTy->hasContainedType() &&
4138 "Expecting a contained type for resource with a dimension "
4139 "attribute.");
4140 QualType ReturnType = ResourceTy->getContainedType();
4141
4142 if (CheckNoDoubleElementType(S, TheCall, ReturnType,
4143 getSampleMethodName(Kind)))
4144 return true;
4145
4146 if (CheckIntegerElementTypeShaderModel(S, TheCall, ReturnType, Kind))
4147 return true;
4148
4149 if (Kind == SampleKind::Cmp || Kind == SampleKind::CmpLevelZero) {
4150 if (!ReturnType->hasFloatingRepresentation()) {
4151 S.Diag(TheCall->getBeginLoc(), diag::err_hlsl_samplecmp_requires_float);
4152 return true;
4153 }
4154 ReturnType = S.Context.FloatTy;
4155 }
4156 TheCall->setType(ReturnType);
4157
4158 return false;
4159}
4160
4161/// The `dest` types an interlocked operation accepts. Float is 32-bit only.
4163
4164/// Check a call to an HLSL interlocked builtin. The builtins are variadic, so
4165/// this is the only check a direct call gets. Overload resolution checks the
4166/// calls that come through the `InterlockedOp` overload sets.
4167static bool CheckInterlockedBuiltin(Sema &S, CallExpr *TheCall,
4168 unsigned MinArgs, unsigned MaxArgs,
4169 InterlockedDest Dest,
4170 bool ReportsOriginalValue) {
4171 if (MinArgs == MaxArgs) {
4172 if (S.checkArgCount(TheCall, MinArgs))
4173 return true;
4174 } else if (TheCall->getNumArgs() < MinArgs) {
4175 S.Diag(TheCall->getEndLoc(), diag::err_typecheck_call_too_few_args_at_least)
4176 << /*callee_type=*/0 << /*min_arg_count=*/MinArgs
4177 << TheCall->getNumArgs() << /*is_non_object=*/0
4178 << TheCall->getSourceRange();
4179 return true;
4180 } else if (S.checkArgCountAtMost(TheCall, MaxArgs)) {
4181 return true;
4182 }
4183
4184 QualType DestTy = TheCall->getArg(0)->getType().getUnqualifiedType();
4185 const bool DestIsOK =
4186 DestTy->isSpecificBuiltinType(BuiltinType::Float)
4187 ? Dest != InterlockedDest::Int
4188 : Dest != InterlockedDest::Float && DestTy->isIntegerType();
4189 if (!DestIsOK) {
4190 S.Diag(TheCall->getArg(0)->getBeginLoc(),
4191 diag::err_builtin_invalid_arg_type)
4192 << /*ordinal=*/1 << /*scalar*/ 1
4193 << /*integer*/ (Dest == InterlockedDest::Float ? 0 : 1)
4194 << /*32 bit floating-point*/ (Dest == InterlockedDest::Int ? 0 : 3)
4195 << DestTy;
4196 return true;
4197 }
4198
4199 // 64-bit interlocked ops require SM 6.6 on DXIL. The synthesized wrapper
4200 // methods (e.g. RWByteAddressBuffer::InterlockedAdd64) are only declared on
4201 // SM 6.6+, so this defensive check only fires for direct builtin calls; skip
4202 // synthetic invocations (invalid source location).
4203 const TargetInfo &TI = S.Context.getTargetInfo();
4204 if (TheCall->getBeginLoc().isValid() &&
4205 TI.getTriple().getArch() == llvm::Triple::dxil &&
4206 S.Context.getTypeSize(DestTy) == 64 &&
4207 TI.getPlatformMinVersion() < VersionTuple(6, 6)) {
4208 S.Diag(TheCall->getBeginLoc(), diag::err_hlsl_builtin_requires_sm)
4209 << TheCall->getDirectCallee() << VersionTuple(6, 6).getAsString();
4210 return true;
4211 }
4212
4213 if (CheckModifiableLValue(&S, TheCall, 0))
4214 return true;
4215
4216 if (CheckArgAddrSpaceOneOf(&S, TheCall, 0,
4218 return true;
4219
4220 // Every argument after `dest` has the destination's type.
4221 for (unsigned I = 1, E = TheCall->getNumArgs(); I != E; ++I)
4222 if (CheckArgTypeMatches(&S, TheCall->getArg(I), DestTy))
4223 return true;
4224
4225 // Operations that report the previous value write it back through their last
4226 // argument.
4227 const unsigned NumArgs = TheCall->getNumArgs();
4228 if (ReportsOriginalValue && NumArgs == MaxArgs &&
4229 CheckModifiableLValue(&S, TheCall, NumArgs - 1))
4230 return true;
4231
4232 TheCall->setType(S.Context.VoidTy);
4233 return false;
4234}
4235
4236// Note: returning true in this case results in CheckBuiltinFunctionCall
4237// returning an ExprError
4238bool SemaHLSL::CheckBuiltinFunctionCall(unsigned BuiltinID, CallExpr *TheCall) {
4239 switch (BuiltinID) {
4240 case Builtin::BI__builtin_hlsl_adduint64: {
4241 if (SemaRef.checkArgCount(TheCall, 2))
4242 return true;
4243
4244 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
4246 return true;
4247
4248 // ensure arg integers are 32-bits
4249 if (CheckExpectedBitWidth(&SemaRef, TheCall, 0, 32))
4250 return true;
4251
4252 // ensure both args are vectors of total bit size of a multiple of 64
4253 auto *VTy = TheCall->getArg(0)->getType()->getAs<VectorType>();
4254 int NumElementsArg = VTy->getNumElements();
4255 if (NumElementsArg != 2 && NumElementsArg != 4) {
4256 SemaRef.Diag(TheCall->getBeginLoc(), diag::err_vector_incorrect_bit_count)
4257 << 1 /*a multiple of*/ << 64 << NumElementsArg * 32;
4258 return true;
4259 }
4260
4261 // ensure first arg and second arg have the same type
4262 if (CheckAllArgsHaveSameType(&SemaRef, TheCall))
4263 return true;
4264
4265 ExprResult A = TheCall->getArg(0);
4266 QualType ArgTyA = A.get()->getType();
4267 // return type is the same as the input type
4268 TheCall->setType(ArgTyA);
4269 break;
4270 }
4271 case Builtin::BI__builtin_hlsl_resource_getpointer: {
4272 if (SemaRef.checkArgCountRange(TheCall, 1, 2) ||
4273 CheckResourceHandle(&SemaRef, TheCall, 0) ||
4274 (TheCall->getNumArgs() == 2 && CheckIndexType(&SemaRef, TheCall, 1)))
4275 return true;
4276
4277 auto *ResourceTy =
4278 TheCall->getArg(0)->getType()->castAs<HLSLAttributedResourceType>();
4279 QualType ContainedTy = ResourceTy->getContainedType();
4280 auto ReturnType = SemaRef.Context.getAddrSpaceQualType(
4281 ContainedTy,
4282 getLangASFromResourceClass(ResourceTy->getAttrs().ResourceClass));
4283 ReturnType = SemaRef.Context.getPointerType(ReturnType);
4284 TheCall->setType(ReturnType);
4285
4286 break;
4287 }
4288 case Builtin::BI__builtin_hlsl_resource_getpointer_typed: {
4289 if (SemaRef.checkArgCount(TheCall, 3) ||
4290 CheckResourceHandle(&SemaRef, TheCall, 0) ||
4291 CheckIndexType(&SemaRef, TheCall, 1))
4292 return true;
4293
4294 QualType ElementTy = TheCall->getArg(2)->getType();
4295 assert(ElementTy->isPointerType() &&
4296 "expected pointer type for second argument");
4297 ElementTy = ElementTy->getPointeeType();
4298
4299 // Reject array types
4300 if (ElementTy->isArrayType())
4301 return SemaRef.Diag(
4302 cast<FunctionDecl>(SemaRef.CurContext)->getPointOfInstantiation(),
4303 diag::err_invalid_use_of_array_type);
4304
4305 auto *ResourceTy =
4306 TheCall->getArg(0)->getType()->castAs<HLSLAttributedResourceType>();
4307 auto ReturnType = SemaRef.Context.getAddrSpaceQualType(
4308 ElementTy,
4309 getLangASFromResourceClass(ResourceTy->getAttrs().ResourceClass));
4310 ReturnType = SemaRef.Context.getPointerType(ReturnType);
4311 TheCall->setType(ReturnType);
4312
4313 break;
4314 }
4315 case Builtin::BI__builtin_hlsl_transpose_if_memory_is_row_major: {
4316 if (SemaRef.checkArgCount(TheCall, 2) ||
4317 CheckArgTypeMatches(&SemaRef, TheCall->getArg(1),
4318 SemaRef.getASTContext().IntTy))
4319 return true;
4320
4321 TheCall->setType(TheCall->getArg(0)->getType());
4322
4323 break;
4324 }
4325 case Builtin::BI__builtin_hlsl_resource_load_with_status: {
4326 if (SemaRef.checkArgCount(TheCall, 3) ||
4327 CheckResourceHandle(&SemaRef, TheCall, 0) ||
4328 CheckArgTypeMatches(&SemaRef, TheCall->getArg(1),
4329 SemaRef.getASTContext().UnsignedIntTy) ||
4330 CheckArgTypeMatches(&SemaRef, TheCall->getArg(2),
4331 SemaRef.getASTContext().UnsignedIntTy) ||
4332 CheckModifiableLValue(&SemaRef, TheCall, 2))
4333 return true;
4334
4335 auto *ResourceTy =
4336 TheCall->getArg(0)->getType()->castAs<HLSLAttributedResourceType>();
4337 QualType ReturnType = ResourceTy->getContainedType();
4338 TheCall->setType(ReturnType);
4339
4340 break;
4341 }
4342 case Builtin::BI__builtin_hlsl_resource_load_with_status_typed: {
4343 if (SemaRef.checkArgCount(TheCall, 4) ||
4344 CheckResourceHandle(&SemaRef, TheCall, 0) ||
4345 CheckArgTypeMatches(&SemaRef, TheCall->getArg(1),
4346 SemaRef.getASTContext().UnsignedIntTy) ||
4347 CheckArgTypeMatches(&SemaRef, TheCall->getArg(2),
4348 SemaRef.getASTContext().UnsignedIntTy) ||
4349 CheckModifiableLValue(&SemaRef, TheCall, 2))
4350 return true;
4351
4352 QualType ReturnType = TheCall->getArg(3)->getType();
4353 assert(ReturnType->isPointerType() &&
4354 "expected pointer type for second argument");
4355 ReturnType = ReturnType->getPointeeType();
4356
4357 // Reject array types
4358 if (ReturnType->isArrayType())
4359 return SemaRef.Diag(
4360 cast<FunctionDecl>(SemaRef.CurContext)->getPointOfInstantiation(),
4361 diag::err_invalid_use_of_array_type);
4362
4363 TheCall->setType(ReturnType);
4364
4365 break;
4366 }
4367 case Builtin::BI__builtin_hlsl_resource_load_level:
4368 return CheckLoadLevelBuiltin(SemaRef, TheCall);
4369 case Builtin::BI__builtin_hlsl_resource_load_ms:
4370 return CheckLoadMSBuiltin(SemaRef, TheCall);
4371 case Builtin::BI__builtin_hlsl_resource_sample:
4373 case Builtin::BI__builtin_hlsl_resource_sample_bias:
4375 case Builtin::BI__builtin_hlsl_resource_sample_grad:
4377 case Builtin::BI__builtin_hlsl_resource_sample_level:
4379 case Builtin::BI__builtin_hlsl_resource_sample_cmp:
4381 case Builtin::BI__builtin_hlsl_resource_sample_cmp_level_zero:
4383 case Builtin::BI__builtin_hlsl_resource_calculate_lod:
4384 case Builtin::BI__builtin_hlsl_resource_calculate_lod_unclamped:
4385 return CheckCalculateLodBuiltin(SemaRef, TheCall);
4386 case Builtin::BI__builtin_hlsl_resource_gather:
4387 return CheckGatherBuiltin(SemaRef, TheCall, /*IsCmp=*/false);
4388 case Builtin::BI__builtin_hlsl_resource_gather_cmp:
4389 return CheckGatherBuiltin(SemaRef, TheCall, /*IsCmp=*/true);
4390 case Builtin::BI__builtin_hlsl_resource_uninitializedhandle: {
4391 assert(TheCall->getNumArgs() == 1 && "expected 1 arg");
4392 // Update return type to be the attributed resource type from arg0.
4393 QualType ResourceTy = TheCall->getArg(0)->getType();
4394 TheCall->setType(ResourceTy);
4395 break;
4396 }
4397 case Builtin::BI__builtin_hlsl_resource_handlefrombinding: {
4398 assert(TheCall->getNumArgs() == 6 && "expected 6 args");
4399 // Update return type to be the attributed resource type from arg0.
4400 QualType ResourceTy = TheCall->getArg(0)->getType();
4401 TheCall->setType(ResourceTy);
4402 break;
4403 }
4404 case Builtin::BI__builtin_hlsl_resource_handlefromimplicitbinding: {
4405 assert(TheCall->getNumArgs() == 6 && "expected 6 args");
4406 // Update return type to be the attributed resource type from arg0.
4407 QualType ResourceTy = TheCall->getArg(0)->getType();
4408 TheCall->setType(ResourceTy);
4409 break;
4410 }
4411 case Builtin::BI__builtin_hlsl_resource_counterhandlefromimplicitbinding: {
4412 assert(TheCall->getNumArgs() == 3 && "expected 3 args");
4413 // Update return type to be the attributed resource type from arg0
4414 // with added IsCounter flag.
4415 QualType MainHandleTy = TheCall->getArg(0)->getType();
4416 QualType CounterHandleTy =
4417 createCounterHandleType(SemaRef.getASTContext(), MainHandleTy);
4418 TheCall->setType(CounterHandleTy);
4419 break;
4420 }
4421 case Builtin::BI__builtin_hlsl_resource_handlefromheap: {
4422 if (SemaRef.checkArgCount(TheCall, 2) ||
4423 CheckResourceHandle(&SemaRef, TheCall, 0) ||
4424 CheckArgTypeMatches(&SemaRef, TheCall->getArg(1),
4425 SemaRef.getASTContext().UnsignedIntTy))
4426 return true;
4427
4428 // Update return type to be the attributed resource type from arg0.
4429 QualType ResourceTy = TheCall->getArg(0)->getType();
4430 TheCall->setType(ResourceTy);
4431 break;
4432 }
4433 case Builtin::BI__builtin_hlsl_resource_counterhandlefromheap: {
4434 if (SemaRef.checkArgCount(TheCall, 1) ||
4435 CheckResourceHandle(&SemaRef, TheCall, 0))
4436 return true;
4437 // Update return type to be the attributed resource type from arg0
4438 // with added IsCounter flag.
4439 QualType MainHandleTy = TheCall->getArg(0)->getType();
4440 QualType CounterHandleTy =
4441 createCounterHandleType(SemaRef.getASTContext(), MainHandleTy);
4442 TheCall->setType(CounterHandleTy);
4443 break;
4444 }
4445 case Builtin::BI__builtin_hlsl_and:
4446 case Builtin::BI__builtin_hlsl_or: {
4447 if (SemaRef.checkArgCount(TheCall, 2))
4448 return true;
4449 if (CheckScalarOrVectorOrMatrix(&SemaRef, TheCall, getASTContext().BoolTy,
4450 0))
4451 return true;
4452 if (CheckAllArgsHaveSameType(&SemaRef, TheCall))
4453 return true;
4454
4455 ExprResult A = TheCall->getArg(0);
4456 QualType ArgTyA = A.get()->getType();
4457 // return type is the same as the input type
4458 TheCall->setType(ArgTyA);
4459 break;
4460 }
4461 case Builtin::BI__builtin_hlsl_all:
4462 case Builtin::BI__builtin_hlsl_any: {
4463 if (SemaRef.checkArgCount(TheCall, 1))
4464 return true;
4465 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
4466 return true;
4467 break;
4468 }
4469 case Builtin::BI__builtin_hlsl_asdouble: {
4470 if (SemaRef.checkArgCount(TheCall, 2))
4471 return true;
4473 &SemaRef, TheCall,
4474 /*only check for uint*/ SemaRef.Context.UnsignedIntTy,
4475 /* arg index */ 0))
4476 return true;
4478 &SemaRef, TheCall,
4479 /*only check for uint*/ SemaRef.Context.UnsignedIntTy,
4480 /* arg index */ 1))
4481 return true;
4482 if (CheckAllArgsHaveSameType(&SemaRef, TheCall))
4483 return true;
4484
4485 SetElementTypeAsReturnType(&SemaRef, TheCall, getASTContext().DoubleTy);
4486 break;
4487 }
4488 case Builtin::BI__builtin_hlsl_elementwise_clamp: {
4489 if (SemaRef.BuiltinElementwiseTernaryMath(
4490 TheCall, /*ArgTyRestr=*/
4492 return true;
4493 break;
4494 }
4495 case Builtin::BI__builtin_hlsl_dot: {
4496 // arg count is checked by BuiltinVectorToScalarMath
4497 if (SemaRef.BuiltinVectorToScalarMath(TheCall))
4498 return true;
4500 return true;
4501 break;
4502 }
4503 case Builtin::BI__builtin_hlsl_elementwise_firstbithigh:
4504 case Builtin::BI__builtin_hlsl_elementwise_firstbitlow: {
4505 if (SemaRef.PrepareBuiltinElementwiseMathOneArgCall(TheCall))
4506 return true;
4507
4508 const Expr *Arg = TheCall->getArg(0);
4509 QualType ArgTy = Arg->getType();
4510 QualType EltTy = ArgTy;
4511
4512 QualType ResTy = SemaRef.Context.UnsignedIntTy;
4513
4514 if (auto *VecTy = EltTy->getAs<VectorType>()) {
4515 EltTy = VecTy->getElementType();
4516 ResTy = SemaRef.Context.getExtVectorType(ResTy, VecTy->getNumElements());
4517 }
4518
4519 if (!EltTy->isIntegerType()) {
4520 Diag(Arg->getBeginLoc(), diag::err_builtin_invalid_arg_type)
4521 << 1 << /* scalar or vector of */ 5 << /* integer ty */ 1
4522 << /* no fp */ 0 << ArgTy;
4523 return true;
4524 }
4525
4526 TheCall->setType(ResTy);
4527 break;
4528 }
4529 case Builtin::BI__builtin_hlsl_select: {
4530 if (SemaRef.checkArgCount(TheCall, 3))
4531 return true;
4532 if (CheckScalarOrVectorOrMatrix(&SemaRef, TheCall, getASTContext().BoolTy,
4533 0))
4534 return true;
4535 QualType ArgTy = TheCall->getArg(0)->getType();
4536 if (ArgTy->isBooleanType() && CheckBoolSelect(&SemaRef, TheCall))
4537 return true;
4538 auto *VTy = ArgTy->getAs<VectorType>();
4539 if (VTy && VTy->getElementType()->isBooleanType() &&
4540 CheckVectorSelect(&SemaRef, TheCall))
4541 return true;
4542 auto *MTy = ArgTy->getAs<ConstantMatrixType>();
4543 if (MTy && MTy->getElementType()->isBooleanType() &&
4544 CheckMatrixSelect(&SemaRef, TheCall))
4545 return true;
4546 break;
4547 }
4548 case Builtin::BI__builtin_hlsl_elementwise_saturate:
4549 case Builtin::BI__builtin_hlsl_elementwise_rcp: {
4550 if (SemaRef.checkArgCount(TheCall, 1))
4551 return true;
4552 if (!TheCall->getArg(0)
4553 ->getType()
4554 ->hasFloatingRepresentation()) // half or float or double
4555 return SemaRef.Diag(TheCall->getArg(0)->getBeginLoc(),
4556 diag::err_builtin_invalid_arg_type)
4557 << /* ordinal */ 1 << /* scalar or vector */ 5 << /* no int */ 0
4558 << /* fp */ 1 << TheCall->getArg(0)->getType();
4559 if (SemaRef.PrepareBuiltinElementwiseMathOneArgCall(TheCall))
4560 return true;
4561 break;
4562 }
4563 case Builtin::BI__builtin_hlsl_elementwise_rsqrt:
4564 case Builtin::BI__builtin_hlsl_elementwise_frac:
4565 case Builtin::BI__builtin_hlsl_elementwise_ddx_coarse:
4566 case Builtin::BI__builtin_hlsl_elementwise_ddy_coarse:
4567 case Builtin::BI__builtin_hlsl_elementwise_ddx_fine:
4568 case Builtin::BI__builtin_hlsl_elementwise_ddy_fine: {
4569 if (SemaRef.checkArgCount(TheCall, 1))
4570 return true;
4571 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
4573 return true;
4574 if (SemaRef.PrepareBuiltinElementwiseMathOneArgCall(TheCall))
4575 return true;
4576 break;
4577 }
4578 case Builtin::BI__builtin_hlsl_elementwise_isinf:
4579 case Builtin::BI__builtin_hlsl_elementwise_isnan: {
4580 if (SemaRef.checkArgCount(TheCall, 1))
4581 return true;
4582 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
4584 return true;
4585 if (SemaRef.PrepareBuiltinElementwiseMathOneArgCall(TheCall))
4586 return true;
4588 break;
4589 }
4590 case Builtin::BI__builtin_hlsl_mad: {
4591 if (SemaRef.BuiltinElementwiseTernaryMath(
4592 TheCall, /*ArgTyRestr=*/
4594 return true;
4595 break;
4596 }
4597 case Builtin::BI__builtin_hlsl_mul: {
4598 if (SemaRef.checkArgCount(TheCall, 2))
4599 return true;
4600
4601 Expr *Arg0 = TheCall->getArg(0);
4602 Expr *Arg1 = TheCall->getArg(1);
4603 QualType Ty0 = Arg0->getType();
4604 QualType Ty1 = Arg1->getType();
4605
4606 auto getElemType = [](QualType T) -> QualType {
4607 if (const auto *VTy = T->getAs<VectorType>())
4608 return VTy->getElementType();
4609 if (const auto *MTy = T->getAs<ConstantMatrixType>())
4610 return MTy->getElementType();
4611 return T;
4612 };
4613
4614 QualType EltTy0 = getElemType(Ty0);
4615
4616 bool IsVec0 = Ty0->isVectorType();
4617 bool IsMat0 = Ty0->isConstantMatrixType();
4618 bool IsVec1 = Ty1->isVectorType();
4619 bool IsMat1 = Ty1->isConstantMatrixType();
4620
4621 QualType RetTy;
4622
4623 if (IsVec0 && IsMat1) {
4624 auto *MatTy = Ty1->castAs<ConstantMatrixType>();
4625 RetTy = getASTContext().getExtVectorType(EltTy0, MatTy->getNumColumns());
4626 } else if (IsMat0 && IsVec1) {
4627 auto *MatTy = Ty0->castAs<ConstantMatrixType>();
4628 RetTy = getASTContext().getExtVectorType(EltTy0, MatTy->getNumRows());
4629 } else {
4630 assert(IsMat0 && IsMat1);
4631 auto *MatTy0 = Ty0->castAs<ConstantMatrixType>();
4632 auto *MatTy1 = Ty1->castAs<ConstantMatrixType>();
4634 EltTy0, MatTy0->getNumRows(), MatTy1->getNumColumns());
4635 }
4636
4637 TheCall->setType(RetTy);
4638 break;
4639 }
4640 case Builtin::BI__builtin_elementwise_fma: {
4641 if (SemaRef.checkArgCount(TheCall, 3) ||
4642 CheckAllArgsHaveSameType(&SemaRef, TheCall)) {
4643 return true;
4644 }
4645
4646 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
4648 return true;
4649
4650 ExprResult A = TheCall->getArg(0);
4651 QualType ArgTyA = A.get()->getType();
4652 // return type is the same as input type
4653 TheCall->setType(ArgTyA);
4654 break;
4655 }
4656 case Builtin::BI__builtin_hlsl_transpose: {
4657 if (SemaRef.checkArgCount(TheCall, 1))
4658 return true;
4659
4660 Expr *Arg = TheCall->getArg(0);
4661 QualType ArgTy = Arg->getType();
4662
4663 const auto *MatTy = ArgTy->getAs<ConstantMatrixType>();
4664 if (!MatTy) {
4665 SemaRef.Diag(Arg->getBeginLoc(), diag::err_builtin_invalid_arg_type)
4666 << 1 << /* matrix */ 3 << /* no int */ 0 << /* no fp */ 0 << ArgTy;
4667 return true;
4668 }
4669
4671 MatTy->getElementType(), MatTy->getNumColumns(), MatTy->getNumRows());
4672 TheCall->setType(RetTy);
4673 break;
4674 }
4675 case Builtin::BI__builtin_hlsl_elementwise_sign: {
4676 if (SemaRef.PrepareBuiltinElementwiseMathOneArgCall(TheCall))
4677 return true;
4678 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
4680 return true;
4682 break;
4683 }
4684 case Builtin::BI__builtin_hlsl_wave_active_all_equal: {
4685 if (SemaRef.checkArgCount(TheCall, 1))
4686 return true;
4687
4688 // Ensure input expr type is a scalar/vector
4689 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
4690 return true;
4691
4692 QualType InputTy = TheCall->getArg(0)->getType();
4693 ASTContext &Ctx = getASTContext();
4694
4695 QualType RetTy;
4696
4697 // If vector, construct bool vector of same size
4698 if (const auto *VecTy = InputTy->getAs<ExtVectorType>()) {
4699 unsigned NumElts = VecTy->getNumElements();
4700 RetTy = Ctx.getExtVectorType(Ctx.BoolTy, NumElts);
4701 } else {
4702 // Scalar case
4703 RetTy = Ctx.BoolTy;
4704 }
4705
4706 TheCall->setType(RetTy);
4707 break;
4708 }
4709 case Builtin::BI__builtin_hlsl_wave_active_max:
4710 case Builtin::BI__builtin_hlsl_wave_active_min:
4711 case Builtin::BI__builtin_hlsl_wave_active_sum:
4712 case Builtin::BI__builtin_hlsl_wave_active_product: {
4713 if (SemaRef.checkArgCount(TheCall, 1))
4714 return true;
4715
4716 // Ensure input expr type is a scalar/vector and the same as the return type
4717 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
4718 return true;
4719 if (CheckWaveActive(&SemaRef, TheCall))
4720 return true;
4721 ExprResult Expr = TheCall->getArg(0);
4722 QualType ArgTyExpr = Expr.get()->getType();
4723 TheCall->setType(ArgTyExpr);
4724 break;
4725 }
4726 case Builtin::BI__builtin_hlsl_wave_active_bit_or:
4727 case Builtin::BI__builtin_hlsl_wave_active_bit_xor:
4728 case Builtin::BI__builtin_hlsl_wave_active_bit_and: {
4729 if (SemaRef.checkArgCount(TheCall, 1))
4730 return true;
4731
4732 // Ensure input expr type is a scalar/vector
4733 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
4734 return true;
4735
4736 if (CheckWaveActive(&SemaRef, TheCall))
4737 return true;
4738
4739 // Ensure the expr type is interpretable as a uint or vector<uint>
4740 ExprResult Expr = TheCall->getArg(0);
4741 QualType ArgTyExpr = Expr.get()->getType();
4742 auto *VTy = ArgTyExpr->getAs<VectorType>();
4743 if (!(ArgTyExpr->isIntegerType() ||
4744 (VTy && VTy->getElementType()->isIntegerType()))) {
4745 SemaRef.Diag(TheCall->getArg(0)->getBeginLoc(),
4746 diag::err_builtin_invalid_arg_type)
4747 << ArgTyExpr << SemaRef.Context.UnsignedIntTy << 1 << 0 << 0;
4748 return true;
4749 }
4750
4751 // Ensure input expr type is the same as the return type
4752 TheCall->setType(ArgTyExpr);
4753 break;
4754 }
4755 case Builtin::BI__builtin_hlsl_interlocked_add:
4756 case Builtin::BI__builtin_hlsl_interlocked_and:
4757 case Builtin::BI__builtin_hlsl_interlocked_max:
4758 case Builtin::BI__builtin_hlsl_interlocked_min:
4759 case Builtin::BI__builtin_hlsl_interlocked_or:
4760 case Builtin::BI__builtin_hlsl_interlocked_xor:
4761 if (CheckInterlockedBuiltin(SemaRef, TheCall, /*MinArgs=*/2, /*MaxArgs=*/3,
4763 /*ReportsOriginalValue=*/true))
4764 return true;
4765 break;
4766 case Builtin::BI__builtin_hlsl_interlocked_exchange:
4767 if (CheckInterlockedBuiltin(SemaRef, TheCall, /*MinArgs=*/3, /*MaxArgs=*/3,
4769 /*ReportsOriginalValue=*/true))
4770 return true;
4771 break;
4772 case Builtin::BI__builtin_hlsl_interlocked_compare_store:
4773 if (CheckInterlockedBuiltin(SemaRef, TheCall, /*MinArgs=*/3, /*MaxArgs=*/3,
4775 /*ReportsOriginalValue=*/false))
4776 return true;
4777 break;
4778 case Builtin::BI__builtin_hlsl_interlocked_compare_store_float_bitwise:
4779 if (CheckInterlockedBuiltin(SemaRef, TheCall, /*MinArgs=*/3, /*MaxArgs=*/3,
4781 /*ReportsOriginalValue=*/false))
4782 return true;
4783 break;
4784 case Builtin::BI__builtin_hlsl_interlocked_compare_exchange:
4785 if (CheckInterlockedBuiltin(SemaRef, TheCall, /*MinArgs=*/4, /*MaxArgs=*/4,
4787 /*ReportsOriginalValue=*/true))
4788 return true;
4789 break;
4790 case Builtin::BI__builtin_hlsl_interlocked_compare_exchange_float_bitwise:
4791 if (CheckInterlockedBuiltin(SemaRef, TheCall, /*MinArgs=*/4, /*MaxArgs=*/4,
4793 /*ReportsOriginalValue=*/true))
4794 return true;
4795 break;
4796 // Note these are llvm builtins that we want to catch invalid intrinsic
4797 // generation. Normal handling of these builtins will occur elsewhere.
4798 case Builtin::BI__builtin_elementwise_bitreverse: {
4799 // does not include a check for number of arguments
4800 // because that is done previously
4801 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
4803 return true;
4804 break;
4805 }
4806 case Builtin::BI__builtin_hlsl_wave_prefix_count_bits: {
4807 if (SemaRef.checkArgCount(TheCall, 1))
4808 return true;
4809
4810 QualType ArgType = TheCall->getArg(0)->getType();
4811
4812 if (!(ArgType->isScalarType())) {
4813 SemaRef.Diag(TheCall->getArg(0)->getBeginLoc(),
4814 diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
4815 << ArgType << 0;
4816 return true;
4817 }
4818
4819 if (!(ArgType->isBooleanType())) {
4820 SemaRef.Diag(TheCall->getArg(0)->getBeginLoc(),
4821 diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
4822 << ArgType << 0;
4823 return true;
4824 }
4825
4826 break;
4827 }
4828 case Builtin::BI__builtin_hlsl_wave_read_lane_at: {
4829 if (SemaRef.checkArgCount(TheCall, 2))
4830 return true;
4831
4832 // Ensure index parameter type can be interpreted as a uint
4833 ExprResult Index = TheCall->getArg(1);
4834 QualType ArgTyIndex = Index.get()->getType();
4835 if (!ArgTyIndex->isIntegerType()) {
4836 SemaRef.Diag(TheCall->getArg(1)->getBeginLoc(),
4837 diag::err_typecheck_convert_incompatible)
4838 << ArgTyIndex << SemaRef.Context.UnsignedIntTy << 1 << 0 << 0;
4839 return true;
4840 }
4841
4842 // Ensure input expr type is a scalar/vector and the same as the return type
4843 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
4844 return true;
4845
4846 ExprResult Expr = TheCall->getArg(0);
4847 QualType ArgTyExpr = Expr.get()->getType();
4848 TheCall->setType(ArgTyExpr);
4849 break;
4850 }
4851 case Builtin::BI__builtin_hlsl_wave_read_lane_first: {
4852 if (SemaRef.checkArgCount(TheCall, 1))
4853 return true;
4854
4855 if (CheckAnyScalarOrVectorOrMatrix(&SemaRef, TheCall, 0))
4856 return true;
4857
4858 TheCall->setType(TheCall->getArg(0)->getType());
4859 break;
4860 }
4861 case Builtin::BI__builtin_hlsl_wave_get_lane_index: {
4862 if (SemaRef.checkArgCount(TheCall, 0))
4863 return true;
4864 break;
4865 }
4866 case Builtin::BI__builtin_hlsl_wave_prefix_sum:
4867 case Builtin::BI__builtin_hlsl_wave_prefix_product: {
4868 if (SemaRef.checkArgCount(TheCall, 1))
4869 return true;
4870
4871 // Ensure input expr type is a scalar/vector and the same as the return type
4872 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
4873 return true;
4874 if (CheckWavePrefix(&SemaRef, TheCall))
4875 return true;
4876 ExprResult Expr = TheCall->getArg(0);
4877 QualType ArgTyExpr = Expr.get()->getType();
4878 TheCall->setType(ArgTyExpr);
4879 break;
4880 }
4881 case Builtin::BI__builtin_hlsl_quad_read_across_x:
4882 case Builtin::BI__builtin_hlsl_quad_read_across_y:
4883 case Builtin::BI__builtin_hlsl_quad_read_across_diagonal: {
4884 if (SemaRef.checkArgCount(TheCall, 1))
4885 return true;
4886
4887 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
4888 return true;
4889 if (CheckNotBoolScalarOrVector(&SemaRef, TheCall, 0))
4890 return true;
4891 ExprResult Expr = TheCall->getArg(0);
4892 QualType ArgTyExpr = Expr.get()->getType();
4893 TheCall->setType(ArgTyExpr);
4894 break;
4895 }
4896 case Builtin::BI__builtin_hlsl_elementwise_splitdouble: {
4897 if (SemaRef.checkArgCount(TheCall, 3))
4898 return true;
4899
4900 if (CheckScalarOrVectorOrMatrix(&SemaRef, TheCall, SemaRef.Context.DoubleTy,
4901 0) ||
4903 SemaRef.Context.UnsignedIntTy, 1) ||
4905 SemaRef.Context.UnsignedIntTy, 2))
4906 return true;
4907
4908 if (CheckModifiableLValue(&SemaRef, TheCall, 1) ||
4909 CheckModifiableLValue(&SemaRef, TheCall, 2))
4910 return true;
4911 break;
4912 }
4913 case Builtin::BI__builtin_hlsl_elementwise_clip: {
4914 if (SemaRef.checkArgCount(TheCall, 1))
4915 return true;
4916
4917 if (CheckScalarOrVector(&SemaRef, TheCall, SemaRef.Context.FloatTy, 0))
4918 return true;
4919 break;
4920 }
4921 case Builtin::BI__builtin_elementwise_acos:
4922 case Builtin::BI__builtin_elementwise_asin:
4923 case Builtin::BI__builtin_elementwise_atan:
4924 case Builtin::BI__builtin_elementwise_atan2:
4925 case Builtin::BI__builtin_elementwise_ceil:
4926 case Builtin::BI__builtin_elementwise_cos:
4927 case Builtin::BI__builtin_elementwise_cosh:
4928 case Builtin::BI__builtin_elementwise_exp:
4929 case Builtin::BI__builtin_elementwise_exp2:
4930 case Builtin::BI__builtin_elementwise_exp10:
4931 case Builtin::BI__builtin_elementwise_floor:
4932 case Builtin::BI__builtin_elementwise_fmod:
4933 case Builtin::BI__builtin_elementwise_log:
4934 case Builtin::BI__builtin_elementwise_log2:
4935 case Builtin::BI__builtin_elementwise_log10:
4936 case Builtin::BI__builtin_elementwise_pow:
4937 case Builtin::BI__builtin_elementwise_roundeven:
4938 case Builtin::BI__builtin_elementwise_sin:
4939 case Builtin::BI__builtin_elementwise_sinh:
4940 case Builtin::BI__builtin_elementwise_sqrt:
4941 case Builtin::BI__builtin_elementwise_tan:
4942 case Builtin::BI__builtin_elementwise_tanh:
4943 case Builtin::BI__builtin_elementwise_trunc: {
4944 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
4946 return true;
4947 break;
4948 }
4949 case Builtin::BI__builtin_hlsl_buffer_update_counter: {
4950 assert(TheCall->getNumArgs() == 2 && "expected 2 args");
4951 auto checkResTy = [](const HLSLAttributedResourceType *ResTy) -> bool {
4952 return !(ResTy->getAttrs().ResourceClass == ResourceClass::UAV &&
4953 ResTy->getAttrs().RawBuffer && ResTy->hasContainedType());
4954 };
4955 if (CheckResourceHandle(&SemaRef, TheCall, 0, checkResTy))
4956 return true;
4957 Expr *OffsetExpr = TheCall->getArg(1);
4958 std::optional<llvm::APSInt> Offset =
4959 OffsetExpr->getIntegerConstantExpr(SemaRef.getASTContext());
4960 if (!Offset.has_value() || std::abs(Offset->getExtValue()) != 1) {
4961 SemaRef.Diag(TheCall->getArg(1)->getBeginLoc(),
4962 diag::err_hlsl_expect_arg_const_int_one_or_neg_one)
4963 << 1;
4964 return true;
4965 }
4966 break;
4967 }
4968 case Builtin::BI__builtin_hlsl_elementwise_f16tof32: {
4969 if (SemaRef.checkArgCount(TheCall, 1))
4970 return true;
4971 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
4973 return true;
4974 // ensure arg integers are 32 bits
4975 if (CheckExpectedBitWidth(&SemaRef, TheCall, 0, 32))
4976 return true;
4977 // check it wasn't a bool type
4978 QualType ArgTy = TheCall->getArg(0)->getType();
4979 if (auto *VTy = ArgTy->getAs<VectorType>())
4980 ArgTy = VTy->getElementType();
4981 if (ArgTy->isBooleanType()) {
4982 SemaRef.Diag(TheCall->getArg(0)->getBeginLoc(),
4983 diag::err_builtin_invalid_arg_type)
4984 << 1 << /* scalar or vector of */ 5 << /* unsigned int */ 3
4985 << /* no fp */ 0 << TheCall->getArg(0)->getType();
4986 return true;
4987 }
4988
4989 SetElementTypeAsReturnType(&SemaRef, TheCall, getASTContext().FloatTy);
4990 break;
4991 }
4992 case Builtin::BI__builtin_hlsl_elementwise_f32tof16: {
4993 if (SemaRef.checkArgCount(TheCall, 1))
4994 return true;
4996 return true;
4998 getASTContext().UnsignedIntTy);
4999 break;
5000 }
5001 }
5002 return false;
5003}
5004
5008 WorkList.push_back(BaseTy);
5009 while (!WorkList.empty()) {
5010 QualType T = WorkList.pop_back_val();
5011 T = T.getCanonicalType().getUnqualifiedType();
5012 if (const auto *AT = dyn_cast<ConstantArrayType>(T)) {
5013 llvm::SmallVector<QualType, 16> ElementFields;
5014 // Generally I've avoided recursion in this algorithm, but arrays of
5015 // structs could be time-consuming to flatten and churn through on the
5016 // work list. Hopefully nesting arrays of structs containing arrays
5017 // of structs too many levels deep is unlikely.
5018 BuildFlattenedTypeList(AT->getElementType(), ElementFields);
5019 // Repeat the element's field list n times.
5020 for (uint64_t Ct = 0; Ct < AT->getZExtSize(); ++Ct)
5021 llvm::append_range(List, ElementFields);
5022 continue;
5023 }
5024 // Vectors can only have element types that are builtin types, so this can
5025 // add directly to the list instead of to the WorkList.
5026 if (const auto *VT = dyn_cast<VectorType>(T)) {
5027 List.insert(List.end(), VT->getNumElements(), VT->getElementType());
5028 continue;
5029 }
5030 if (const auto *MT = dyn_cast<ConstantMatrixType>(T)) {
5031 List.insert(List.end(), MT->getNumElementsFlattened(),
5032 MT->getElementType());
5033 continue;
5034 }
5035 if (const auto *RD = T->getAsCXXRecordDecl()) {
5036 if (RD->isStandardLayout())
5037 RD = RD->getStandardLayoutBaseWithFields();
5038
5039 // For types that we shouldn't decompose (unions and non-aggregates), just
5040 // add the type itself to the list.
5041 if (RD->isUnion() || !RD->isAggregate()) {
5042 List.push_back(T);
5043 continue;
5044 }
5045
5047 for (const auto *FD : RD->fields())
5048 if (!FD->isUnnamedBitField())
5049 FieldTypes.push_back(FD->getType());
5050 // Reverse the newly added sub-range.
5051 std::reverse(FieldTypes.begin(), FieldTypes.end());
5052 llvm::append_range(WorkList, FieldTypes);
5053
5054 // If this wasn't a standard layout type we may also have some base
5055 // classes to deal with.
5056 if (!RD->isStandardLayout()) {
5057 FieldTypes.clear();
5058 for (const auto &Base : RD->bases())
5059 FieldTypes.push_back(Base.getType());
5060 std::reverse(FieldTypes.begin(), FieldTypes.end());
5061 llvm::append_range(WorkList, FieldTypes);
5062 }
5063 continue;
5064 }
5065 List.push_back(T);
5066 }
5067}
5068
5070 if (QT.isNull())
5071 return false;
5072
5073 // Must be a class/struct.
5074 const auto *RD = QT->getAsCXXRecordDecl();
5075 if (!RD || RD->isUnion())
5076 return false;
5077
5078 // Cannot be a resource type or contain one.
5079 return !QT->isHLSLIntangibleType();
5080}
5081
5083 // null and array types are not allowed.
5084 if (QT.isNull() || QT->isArrayType())
5085 return false;
5086
5087 // UDT types are not allowed
5088 if (QT->isRecordType())
5089 return false;
5090
5091 if (QT->isBooleanType() || QT->isEnumeralType())
5092 return false;
5093
5094 // the only other valid builtin types are scalars or vectors
5095 if (QT->isArithmeticType()) {
5096 if (SemaRef.Context.getTypeSize(QT) / 8 > 16)
5097 return false;
5098 return true;
5099 }
5100
5101 if (const VectorType *VT = QT->getAs<VectorType>()) {
5102 int ArraySize = VT->getNumElements();
5103
5104 if (ArraySize > 4)
5105 return false;
5106
5107 QualType ElTy = VT->getElementType();
5108 if (ElTy->isBooleanType())
5109 return false;
5110
5111 if (SemaRef.Context.getTypeSize(QT) / 8 > 16)
5112 return false;
5113 return true;
5114 }
5115
5116 return false;
5117}
5118
5120 if (T1.isNull() || T2.isNull())
5121 return false;
5122
5125
5126 // If both types are the same canonical type, they're obviously compatible.
5127 if (SemaRef.getASTContext().hasSameType(T1, T2))
5128 return true;
5129
5131 BuildFlattenedTypeList(T1, T1Types);
5133 BuildFlattenedTypeList(T2, T2Types);
5134
5135 // Check the flattened type list
5136 return llvm::equal(T1Types, T2Types,
5137 [this](QualType LHS, QualType RHS) -> bool {
5138 return SemaRef.IsLayoutCompatible(LHS, RHS);
5139 });
5140}
5141
5143 FunctionDecl *Old) {
5144 if (New->getNumParams() != Old->getNumParams())
5145 return true;
5146
5147 bool HadError = false;
5148
5149 for (unsigned i = 0, e = New->getNumParams(); i != e; ++i) {
5150 ParmVarDecl *NewParam = New->getParamDecl(i);
5151 ParmVarDecl *OldParam = Old->getParamDecl(i);
5152
5153 // HLSL parameter declarations for inout and out must match between
5154 // declarations. In HLSL inout and out are ambiguous at the call site,
5155 // but have different calling behavior, so you cannot overload a
5156 // method based on a difference between inout and out annotations.
5157 const auto *NDAttr = NewParam->getAttr<HLSLParamModifierAttr>();
5158 unsigned NSpellingIdx = (NDAttr ? NDAttr->getSpellingListIndex() : 0);
5159 const auto *ODAttr = OldParam->getAttr<HLSLParamModifierAttr>();
5160 unsigned OSpellingIdx = (ODAttr ? ODAttr->getSpellingListIndex() : 0);
5161
5162 if (NSpellingIdx != OSpellingIdx) {
5163 SemaRef.Diag(NewParam->getLocation(),
5164 diag::err_hlsl_param_qualifier_mismatch)
5165 << NDAttr << NewParam;
5166 SemaRef.Diag(OldParam->getLocation(), diag::note_previous_declaration_as)
5167 << ODAttr;
5168 HadError = true;
5169 }
5170 }
5171 return HadError;
5172}
5173
5174// Generally follows PerformScalarCast, with cases reordered for
5175// clarity of what types are supported
5177
5178 if (!SrcTy->isScalarType() || !DestTy->isScalarType())
5179 return false;
5180
5181 if (SemaRef.getASTContext().hasSameUnqualifiedType(SrcTy, DestTy))
5182 return true;
5183
5184 switch (SrcTy->getScalarTypeKind()) {
5185 case Type::STK_Bool: // casting from bool is like casting from an integer
5186 case Type::STK_Integral:
5187 switch (DestTy->getScalarTypeKind()) {
5188 case Type::STK_Bool:
5189 case Type::STK_Integral:
5190 case Type::STK_Floating:
5191 return true;
5192 case Type::STK_CPointer:
5196 llvm_unreachable("HLSL doesn't support pointers.");
5199 llvm_unreachable("HLSL doesn't support complex types.");
5201 llvm_unreachable("HLSL doesn't support fixed point types.");
5202 }
5203 llvm_unreachable("Should have returned before this");
5204
5205 case Type::STK_Floating:
5206 switch (DestTy->getScalarTypeKind()) {
5207 case Type::STK_Floating:
5208 case Type::STK_Bool:
5209 case Type::STK_Integral:
5210 return true;
5213 llvm_unreachable("HLSL doesn't support complex types.");
5215 llvm_unreachable("HLSL doesn't support fixed point types.");
5216 case Type::STK_CPointer:
5220 llvm_unreachable("HLSL doesn't support pointers.");
5221 }
5222 llvm_unreachable("Should have returned before this");
5223
5225 case Type::STK_CPointer:
5228 llvm_unreachable("HLSL doesn't support pointers.");
5229
5231 llvm_unreachable("HLSL doesn't support fixed point types.");
5232
5235 llvm_unreachable("HLSL doesn't support complex types.");
5236 }
5237
5238 llvm_unreachable("Unhandled scalar cast");
5239}
5240
5241// Can perform an HLSL Aggregate splat cast if the Dest is an aggregate and the
5242// Src is a scalar, a vector of length 1, or a 1x1 matrix
5243// Or if Dest is a vector and Src is a vector of length 1 or a 1x1 matrix
5245
5246 QualType SrcTy = Src->getType();
5247 // Not a valid HLSL Aggregate Splat cast if Dest is a scalar or if this is
5248 // going to be a vector splat from a scalar.
5249 if ((SrcTy->isScalarType() && DestTy->isVectorType()) ||
5250 DestTy->isScalarType())
5251 return false;
5252
5253 const VectorType *SrcVecTy = SrcTy->getAs<VectorType>();
5254 const ConstantMatrixType *SrcMatTy = SrcTy->getAs<ConstantMatrixType>();
5255
5256 // Src isn't a scalar, a vector of length 1, or a 1x1 matrix
5257 if (!SrcTy->isScalarType() &&
5258 !(SrcVecTy && SrcVecTy->getNumElements() == 1) &&
5259 !(SrcMatTy && SrcMatTy->getNumElementsFlattened() == 1))
5260 return false;
5261
5262 if (SrcVecTy)
5263 SrcTy = SrcVecTy->getElementType();
5264 else if (SrcMatTy)
5265 SrcTy = SrcMatTy->getElementType();
5266
5268 BuildFlattenedTypeList(DestTy, DestTypes);
5269
5270 for (unsigned I = 0, Size = DestTypes.size(); I < Size; ++I) {
5271 if (DestTypes[I]->isUnionType())
5272 return false;
5273 if (!CanPerformScalarCast(SrcTy, DestTypes[I]))
5274 return false;
5275 }
5276 return true;
5277}
5278
5279// Can we perform an HLSL Elementwise cast?
5281
5282 // Don't handle casts where LHS and RHS are any combination of scalar/vector
5283 // There must be an aggregate somewhere
5284 QualType SrcTy = Src->getType();
5285 if (SrcTy->isScalarType()) // always a splat and this cast doesn't handle that
5286 return false;
5287
5288 if (SrcTy->isVectorType() &&
5289 (DestTy->isScalarType() || DestTy->isVectorType()))
5290 return false;
5291
5292 if (SrcTy->isConstantMatrixType() &&
5293 (DestTy->isScalarType() || DestTy->isConstantMatrixType()))
5294 return false;
5295
5297 BuildFlattenedTypeList(DestTy, DestTypes);
5299 BuildFlattenedTypeList(SrcTy, SrcTypes);
5300
5301 // Usually the size of SrcTypes must be greater than or equal to the size of
5302 // DestTypes.
5303 if (SrcTypes.size() < DestTypes.size())
5304 return false;
5305
5306 unsigned SrcSize = SrcTypes.size();
5307 unsigned DstSize = DestTypes.size();
5308 unsigned I;
5309 for (I = 0; I < DstSize && I < SrcSize; I++) {
5310 if (SrcTypes[I]->isUnionType() || DestTypes[I]->isUnionType())
5311 return false;
5312 if (!CanPerformScalarCast(SrcTypes[I], DestTypes[I])) {
5313 return false;
5314 }
5315 }
5316
5317 // check the rest of the source type for unions.
5318 for (; I < SrcSize; I++) {
5319 if (SrcTypes[I]->isUnionType())
5320 return false;
5321 }
5322 return true;
5323}
5324
5326 ASTContext &Ctx = SemaRef.getASTContext();
5327 QualType UIntTy = Ctx.UnsignedIntTy;
5328 QualType SrcTy = Src->getType();
5329
5330 return (SrcTy->isHLSLBuiltinPackedType() &&
5331 DestTy->isHLSLBuiltinPackedType()) ||
5332 (SrcTy->isHLSLBuiltinPackedType() &&
5333 Ctx.hasSameUnqualifiedType(DestTy, UIntTy)) ||
5334 (DestTy->isHLSLBuiltinPackedType() &&
5335 Ctx.hasSameUnqualifiedType(SrcTy, UIntTy));
5336}
5337
5339 assert(Param->hasAttr<HLSLParamModifierAttr>() &&
5340 "We should not get here without a parameter modifier expression");
5341 const auto *Attr = Param->getAttr<HLSLParamModifierAttr>();
5342 if (Attr->getABI() == ParameterABI::Ordinary)
5343 return ExprResult(Arg);
5344
5345 bool IsInOut = Attr->getABI() == ParameterABI::HLSLInOut;
5346 if (!Arg->isLValue()) {
5347 SemaRef.Diag(Arg->getBeginLoc(), diag::error_hlsl_inout_lvalue)
5348 << Arg << (IsInOut ? 1 : 0);
5349 return ExprError();
5350 }
5351
5352 ASTContext &Ctx = SemaRef.getASTContext();
5353
5354 QualType Ty = Param->getType().getNonLValueExprType(Ctx);
5355
5356 // HLSL allows implicit conversions from scalars to vectors, but not the
5357 // inverse, so we need to disallow `inout` with scalar->vector or
5358 // scalar->matrix conversions.
5359 if (Arg->getType()->isScalarType() != Ty->isScalarType()) {
5360 SemaRef.Diag(Arg->getBeginLoc(), diag::error_hlsl_inout_scalar_extension)
5361 << Arg << (IsInOut ? 1 : 0);
5362 return ExprError();
5363 }
5364
5365 auto *ArgOpV = new (Ctx) OpaqueValueExpr(Param->getBeginLoc(), Arg->getType(),
5366 VK_LValue, OK_Ordinary, Arg);
5367
5368 // Parameters are initialized via copy initialization. This allows for
5369 // overload resolution of argument constructors.
5370 InitializedEntity Entity =
5372 ExprResult Res =
5373 SemaRef.PerformCopyInitialization(Entity, Param->getBeginLoc(), ArgOpV);
5374 if (Res.isInvalid())
5375 return ExprError();
5376 Expr *Base = Res.get();
5377 // After the cast, drop the reference type when creating the exprs.
5378 Ty = Ty.getNonLValueExprType(Ctx);
5379 auto *OpV = new (Ctx)
5380 OpaqueValueExpr(Param->getBeginLoc(), Ty, VK_LValue, OK_Ordinary, Base);
5381
5382 // Writebacks are performed with `=` binary operator, which allows for
5383 // overload resolution on writeback result expressions.
5384 Res = SemaRef.ActOnBinOp(SemaRef.getCurScope(), Arg->getBeginLoc(),
5385 tok::equal, ArgOpV, OpV);
5386
5387 if (Res.isInvalid())
5388 return ExprError();
5389 Expr *Writeback = Res.get();
5390 auto *OutExpr =
5391 HLSLOutArgExpr::Create(Ctx, Ty, ArgOpV, OpV, Writeback, IsInOut);
5392
5393 return ExprResult(OutExpr);
5394}
5395
5397 // If HLSL gains support for references, all the cites that use this will need
5398 // to be updated with semantic checking to produce errors for
5399 // pointers/references.
5400 assert(!Ty->isReferenceType() &&
5401 "Pointer and reference types cannot be inout or out parameters");
5402 Ty = SemaRef.getASTContext().getLValueReferenceType(Ty);
5403 Ty.addRestrict();
5404 return Ty;
5405}
5406
5407// Returns true if the type has a non-empty constant buffer layout (if it is
5408// scalar, vector or matrix, or if it contains any of these.
5410 const Type *Ty = QT->getUnqualifiedDesugaredType();
5411 if (Ty->isScalarType() || Ty->isVectorType() || Ty->isMatrixType())
5412 return true;
5413
5415 return false;
5416
5417 if (const auto *RD = Ty->getAsCXXRecordDecl()) {
5418 for (const auto *FD : RD->fields()) {
5420 return true;
5421 }
5422 assert(RD->getNumBases() <= 1 &&
5423 "HLSL doesn't support multiple inheritance");
5424 return RD->getNumBases()
5425 ? hasConstantBufferLayout(RD->bases_begin()->getType())
5426 : false;
5427 }
5428
5429 if (const auto *AT = dyn_cast<ArrayType>(Ty)) {
5430 if (const auto *CAT = dyn_cast<ConstantArrayType>(AT))
5431 if (isZeroSizedArray(CAT))
5432 return false;
5434 }
5435
5436 return false;
5437}
5438
5439static bool IsDefaultBufferConstantDecl(const ASTContext &Ctx, VarDecl *VD) {
5440 bool IsVulkan =
5441 Ctx.getTargetInfo().getTriple().getOS() == llvm::Triple::Vulkan;
5442 bool IsVKPushConstant = IsVulkan && VD->hasAttr<HLSLVkPushConstantAttr>();
5443 QualType QT = VD->getType();
5444 return VD->getDeclContext()->isTranslationUnit() &&
5445 QT.getAddressSpace() == LangAS::Default &&
5446 VD->getStorageClass() != SC_Static &&
5447 !VD->hasAttr<HLSLVkConstantIdAttr>() && !IsVKPushConstant &&
5449}
5450
5452 // The variable already has an address space (groupshared for ex).
5453 if (Decl->getType().hasAddressSpace())
5454 return;
5455
5456 if (Decl->getType()->isDependentType())
5457 return;
5458
5459 QualType Type = Decl->getType();
5460
5461 if (Decl->hasAttr<HLSLVkExtBuiltinInputAttr>()) {
5462 LangAS ImplAS = LangAS::hlsl_input;
5463 Type = SemaRef.getASTContext().getAddrSpaceQualType(Type, ImplAS);
5464 Decl->setType(Type);
5465 return;
5466 }
5467
5468 if (Decl->hasAttr<HLSLVkExtBuiltinOutputAttr>()) {
5469 LangAS ImplAS = LangAS::hlsl_output;
5470 Type = SemaRef.getASTContext().getAddrSpaceQualType(Type, ImplAS);
5471 Decl->setType(Type);
5472
5473 // HLSL uses `static` differently than C++. For BuiltIn output, the static
5474 // does not imply private to the module scope.
5475 // Marking it as external to reflect the semantic this attribute brings.
5476 // See https://github.com/microsoft/hlsl-specs/issues/350
5477 Decl->setStorageClass(SC_Extern);
5478 return;
5479 }
5480
5481 bool IsVulkan = getASTContext().getTargetInfo().getTriple().getOS() ==
5482 llvm::Triple::Vulkan;
5483 if (IsVulkan && Decl->hasAttr<HLSLVkPushConstantAttr>()) {
5484 if (HasDeclaredAPushConstant)
5485 SemaRef.Diag(Decl->getLocation(), diag::err_hlsl_push_constant_unique);
5486
5488 Type = SemaRef.getASTContext().getAddrSpaceQualType(Type, ImplAS);
5489 Decl->setType(Type);
5490 HasDeclaredAPushConstant = true;
5491 return;
5492 }
5493
5494 if (Type->isSamplerT() || Type->isVoidType())
5495 return;
5496
5497 // Resource handles.
5499 return;
5500
5501 // Only static globals belong to the Private address space.
5502 // Non-static globals belongs to the cbuffer.
5503 if (Decl->getStorageClass() != SC_Static && !Decl->isStaticDataMember())
5504 return;
5505
5507 Type = SemaRef.getASTContext().getAddrSpaceQualType(Type, ImplAS);
5508 Decl->setType(Type);
5509}
5510
5511namespace {
5512
5513// Helper class for assigning bindings to resources declared within a struct.
5514// It keeps track of all binding attributes declared on a struct instance, and
5515// the offsets for each register type that have been assigned so far.
5516// Handles both explicit and implicit bindings.
5517class StructBindingContext {
5518 // Bindings and offsets per register type. We only need to support four
5519 // register types - SRV (u), UAV (t), CBuffer (c), and Sampler (s).
5520 HLSLResourceBindingAttr *RegBindingsAttrs[4];
5521 unsigned RegBindingOffset[4];
5522
5523 // Make sure the RegisterType values are what we expect
5524 static_assert(static_cast<unsigned>(RegisterType::SRV) == 0 &&
5525 static_cast<unsigned>(RegisterType::UAV) == 1 &&
5526 static_cast<unsigned>(RegisterType::CBuffer) == 2 &&
5527 static_cast<unsigned>(RegisterType::Sampler) == 3,
5528 "unexpected register type values");
5529
5530 // Vulkan binding attribute does not vary by register type.
5531 HLSLVkBindingAttr *VkBindingAttr;
5532 unsigned VkBindingOffset;
5533
5534public:
5535 // Constructor: gather all binding attributes on a struct instance and
5536 // initialize offsets.
5537 StructBindingContext(VarDecl *VD) {
5538 for (unsigned i = 0; i < 4; ++i) {
5539 RegBindingsAttrs[i] = nullptr;
5540 RegBindingOffset[i] = 0;
5541 }
5542 VkBindingAttr = nullptr;
5543 VkBindingOffset = 0;
5544
5545 ASTContext &AST = VD->getASTContext();
5546 bool IsSpirv = AST.getTargetInfo().getTriple().isSPIRV();
5547
5548 for (Attr *A : VD->attrs()) {
5549 if (auto *RBA = dyn_cast<HLSLResourceBindingAttr>(A)) {
5550 RegisterType RegType = RBA->getRegisterType();
5551 unsigned RegTypeIdx = static_cast<unsigned>(RegType);
5552 // Ignore unsupported register annotations, such as 'c' or 'i'.
5553 if (RegTypeIdx < 4)
5554 RegBindingsAttrs[RegTypeIdx] = RBA;
5555 continue;
5556 }
5557 // Gather the Vulkan binding attributes only if the target is SPIR-V.
5558 if (IsSpirv) {
5559 if (auto *VBA = dyn_cast<HLSLVkBindingAttr>(A))
5560 VkBindingAttr = VBA;
5561 }
5562 }
5563 }
5564
5565 // Creates a binding attribute for a resource based on the gathered attributes
5566 // and the required register type and range.
5567 Attr *createBindingAttr(SemaHLSL &S, ASTContext &AST, RegisterType RegType,
5568 unsigned Range, bool HasCounter) {
5569 assert(static_cast<unsigned>(RegType) < 4 && "unexpected register type");
5570
5571 if (VkBindingAttr) {
5572 unsigned Offset = VkBindingOffset;
5573 VkBindingOffset += Range;
5574 return HLSLVkBindingAttr::CreateImplicit(
5575 AST, VkBindingAttr->getBinding() + Offset, VkBindingAttr->getSet(),
5576 VkBindingAttr->getRange());
5577 }
5578
5579 HLSLResourceBindingAttr *RBA =
5580 RegBindingsAttrs[static_cast<unsigned>(RegType)];
5581 HLSLResourceBindingAttr *NewAttr = nullptr;
5582
5583 if (RBA && RBA->hasRegisterSlot()) {
5584 // Explicit binding - create a new attribute with offseted slot number
5585 // based on the required register type.
5586 unsigned Offset = RegBindingOffset[static_cast<unsigned>(RegType)];
5587 RegBindingOffset[static_cast<unsigned>(RegType)] += Range;
5588
5589 unsigned NewSlotNumber = RBA->getSlotNumber() + Offset;
5590 StringRef NewSlotNumberStr =
5591 createRegisterString(AST, RBA->getRegisterType(), NewSlotNumber);
5592 NewAttr = HLSLResourceBindingAttr::CreateImplicit(
5593 AST, NewSlotNumberStr, RBA->getSpace(), RBA->getRange());
5594 NewAttr->setBinding(RegType, NewSlotNumber, RBA->getSpaceNumber());
5595 } else {
5596 // No binding attribute or space-only binding - create a binding
5597 // attribute for implicit binding.
5598 NewAttr = HLSLResourceBindingAttr::CreateImplicit(AST, "", "0", {});
5599 NewAttr->setBinding(RegType, std::nullopt,
5600 RBA ? RBA->getSpaceNumber() : 0);
5601 NewAttr->setImplicitBindingOrderID(S.getNextImplicitBindingOrderID());
5602 }
5603 if (HasCounter)
5604 NewAttr->setImplicitCounterBindingOrderID(
5606 return NewAttr;
5607 }
5608};
5609
5610// Creates a global variable declaration for a resource field embedded in a
5611// struct, assigns it a binding, initializes it, and associates it with the
5612// struct declaration via an HLSLAssociatedResourceDeclAttr.
5613static void createGlobalResourceDeclForStruct(
5614 Sema &S, VarDecl *ParentVD, SourceLocation Loc, IdentifierInfo *Id,
5615 QualType ResTy, StructBindingContext &BindingCtx) {
5616 assert(isResourceRecordTypeOrArrayOf(ResTy) &&
5617 "expected resource type or array of resources");
5618
5619 DeclContext *DC = ParentVD->getNonTransparentDeclContext();
5620 assert(DC->isTranslationUnit() && "expected translation unit decl context");
5621
5622 ASTContext &AST = S.getASTContext();
5623 VarDecl *ResDecl =
5624 VarDecl::Create(AST, DC, Loc, Loc, Id, ResTy, nullptr, SC_None);
5625
5626 unsigned Range = 1;
5627 const Type *SingleResTy = ResTy.getTypePtr()->getUnqualifiedDesugaredType();
5628 while (const auto *AT = dyn_cast<ArrayType>(SingleResTy)) {
5629 const auto *CAT = dyn_cast<ConstantArrayType>(AT);
5630 Range = CAT ? (Range * CAT->getSize().getZExtValue()) : 0;
5631 SingleResTy =
5633 }
5634 const HLSLAttributedResourceType *ResHandleTy =
5635 HLSLAttributedResourceType::findHandleTypeOnResource(SingleResTy);
5636
5637 // Add a binding attribute to the global resource declaration.
5638 bool HasCounter = hasCounterHandle(SingleResTy->getAsCXXRecordDecl());
5639 Attr *BindingAttr = BindingCtx.createBindingAttr(
5640 S.HLSL(), AST, getRegisterType(ResHandleTy), Range, HasCounter);
5641 ResDecl->addAttr(BindingAttr);
5642 ResDecl->addAttr(InternalLinkageAttr::CreateImplicit(AST));
5643 ResDecl->setImplicit();
5644
5645 if (Range == 1)
5646 S.HLSL().initGlobalResourceDecl(ResDecl);
5647 else
5648 S.HLSL().initGlobalResourceArrayDecl(ResDecl);
5649
5650 ParentVD->addAttr(
5651 HLSLAssociatedResourceDeclAttr::CreateImplicit(AST, ResDecl));
5652 DC->addDecl(ResDecl);
5653
5654 DeclGroupRef DG(ResDecl);
5656}
5657
5658static void handleArrayOfStructWithResources(
5659 Sema &S, VarDecl *ParentVD, const ConstantArrayType *CAT,
5660 EmbeddedResourceNameBuilder &NameBuilder, StructBindingContext &BindingCtx);
5661
5662// Scans base and all fields of a struct/class type to find all embedded
5663// resources or resource arrays. Creates a global variable for each resource
5664// found.
5665static void handleStructWithResources(Sema &S, VarDecl *ParentVD,
5666 const CXXRecordDecl *RD,
5667 EmbeddedResourceNameBuilder &NameBuilder,
5668 StructBindingContext &BindingCtx) {
5669
5670 // Scan the base classes.
5671 assert(RD->getNumBases() <= 1 && "HLSL doesn't support multiple inheritance");
5672 const auto *BasesIt = RD->bases_begin();
5673 if (BasesIt != RD->bases_end()) {
5674 QualType QT = BasesIt->getType();
5675 if (QT->isHLSLIntangibleType()) {
5676 CXXRecordDecl *BaseRD = QT->getAsCXXRecordDecl();
5677 NameBuilder.pushBaseName(BaseRD->getName());
5678 handleStructWithResources(S, ParentVD, BaseRD, NameBuilder, BindingCtx);
5679 NameBuilder.pop();
5680 }
5681 }
5682 // Process this class fields.
5683 for (const FieldDecl *FD : RD->fields()) {
5684 QualType FDTy = FD->getType().getCanonicalType();
5685 if (!FDTy->isHLSLIntangibleType())
5686 continue;
5687
5688 NameBuilder.pushName(FD->getName());
5689
5691 IdentifierInfo *II = NameBuilder.getNameAsIdentifier(S.getASTContext());
5692 createGlobalResourceDeclForStruct(S, ParentVD, FD->getLocation(), II,
5693 FDTy, BindingCtx);
5694 } else if (const auto *RD = FDTy->getAsCXXRecordDecl()) {
5695 handleStructWithResources(S, ParentVD, RD, NameBuilder, BindingCtx);
5696
5697 } else if (const auto *ArrayTy = dyn_cast<ConstantArrayType>(FDTy)) {
5698 assert(!FDTy->isHLSLResourceRecordArray() &&
5699 "resource arrays should have been already handled");
5700 handleArrayOfStructWithResources(S, ParentVD, ArrayTy, NameBuilder,
5701 BindingCtx);
5702 }
5703 NameBuilder.pop();
5704 }
5705}
5706
5707// Processes array of structs with resources.
5708static void
5709handleArrayOfStructWithResources(Sema &S, VarDecl *ParentVD,
5710 const ConstantArrayType *CAT,
5711 EmbeddedResourceNameBuilder &NameBuilder,
5712 StructBindingContext &BindingCtx) {
5713
5714 QualType ElementTy = CAT->getElementType().getCanonicalType();
5715 assert(ElementTy->isHLSLIntangibleType() && "Expected HLSL intangible type");
5716
5717 const ConstantArrayType *SubCAT = dyn_cast<ConstantArrayType>(ElementTy);
5718 const CXXRecordDecl *ElementRD = ElementTy->getAsCXXRecordDecl();
5719
5720 if (!SubCAT && !ElementRD)
5721 return;
5722
5723 for (unsigned I = 0, E = CAT->getSize().getZExtValue(); I < E; ++I) {
5724 NameBuilder.pushArrayIndex(I);
5725 if (ElementRD)
5726 handleStructWithResources(S, ParentVD, ElementRD, NameBuilder,
5727 BindingCtx);
5728 else
5729 handleArrayOfStructWithResources(S, ParentVD, SubCAT, NameBuilder,
5730 BindingCtx);
5731 NameBuilder.pop();
5732 }
5733}
5734
5735} // namespace
5736
5737// Scans all fields of a user-defined struct (or array of structs)
5738// to find all embedded resources or resource arrays. For each resource
5739// a global variable of the resource type is created and associated
5740// with the parent declaration (VD) through a HLSLAssociatedResourceDeclAttr
5741// attribute.
5742void SemaHLSL::handleGlobalStructOrArrayOfWithResources(VarDecl *VD) {
5743 EmbeddedResourceNameBuilder NameBuilder(VD->getName());
5744 StructBindingContext BindingCtx(VD);
5745
5746 const Type *VDTy = VD->getType().getTypePtr();
5747 assert(VDTy->isHLSLIntangibleType() && !isResourceRecordTypeOrArrayOf(VD) &&
5748 "Expected non-resource struct or array type");
5749
5750 if (const CXXRecordDecl *RD = VDTy->getAsCXXRecordDecl()) {
5751 handleStructWithResources(SemaRef, VD, RD, NameBuilder, BindingCtx);
5752 return;
5753 }
5754
5755 if (const auto *CAT = dyn_cast<ConstantArrayType>(VDTy)) {
5756 handleArrayOfStructWithResources(SemaRef, VD, CAT, NameBuilder, BindingCtx);
5757 return;
5758 }
5759}
5760
5762 if (VD->hasGlobalStorage()) {
5763 // make sure the declaration has a complete type
5764 if (SemaRef.RequireCompleteType(
5765 VD->getLocation(),
5766 SemaRef.getASTContext().getBaseElementType(VD->getType()),
5767 diag::err_typecheck_decl_incomplete_type)) {
5768 VD->setInvalidDecl();
5770 return;
5771 }
5772
5773 // Global variables outside a cbuffer block that are not a resource, static,
5774 // groupshared, or an empty array or struct belong to the default constant
5775 // buffer $Globals (to be created at the end of the translation unit).
5777 // update address space to hlsl_constant
5780 VD->setType(NewTy);
5781 DefaultCBufferDecls.push_back(VD);
5782 }
5783
5784 // find all resources bindings on decl
5785 if (VD->getType()->isHLSLIntangibleType())
5786 collectResourceBindingsOnVarDecl(VD);
5787
5788 if (VD->hasAttr<HLSLVkConstantIdAttr>())
5790
5792 VD->getStorageClass() != SC_Static) {
5793 // Add internal linkage attribute to non-static resource variables. The
5794 // global externally visible storage is accessed through the handle, which
5795 // is a member. The variable itself is not externally visible.
5796 VD->addAttr(InternalLinkageAttr::CreateImplicit(getASTContext()));
5797 }
5798
5799 // process explicit bindings
5800 processExplicitBindingsOnDecl(VD);
5801
5802 // Add implicit binding attribute to non-static resource arrays.
5803 if (VD->getType()->isHLSLResourceRecordArray() &&
5804 VD->getStorageClass() != SC_Static) {
5805 // If the resource array does not have an explicit binding attribute,
5806 // create an implicit one. It will be used to transfer implicit binding
5807 // order_ID to codegen.
5808 ResourceBindingAttrs Binding(VD);
5809 if (!Binding.isExplicit()) {
5810 uint32_t OrderID = getNextImplicitBindingOrderID();
5811 if (Binding.hasBinding())
5812 Binding.setImplicitOrderID(OrderID);
5813 else {
5816 OrderID);
5817 // Re-create the binding object to pick up the new attribute.
5818 Binding = ResourceBindingAttrs(VD);
5819 }
5820 }
5821
5822 // Get to the base type of a potentially multi-dimensional array.
5824
5825 const CXXRecordDecl *RD = Ty->getAsCXXRecordDecl();
5826 if (hasCounterHandle(RD)) {
5827 if (!Binding.hasCounterImplicitOrderID()) {
5828 uint32_t OrderID = getNextImplicitBindingOrderID();
5829 Binding.setCounterImplicitOrderID(OrderID);
5830 }
5831 }
5832 }
5833
5834 // Process resources in user-defined structs, or arrays of such structs.
5835 const Type *VDTy = VD->getType().getTypePtr();
5836 if (VD->getStorageClass() != SC_Static && VDTy->isHLSLIntangibleType() &&
5838 handleGlobalStructOrArrayOfWithResources(VD);
5839
5840 // Mark groupshared variables as extern so they will have
5841 // external storage and won't be default initialized
5842 if (VD->hasAttr<HLSLGroupSharedAddressSpaceAttr>())
5844 }
5845
5847}
5848
5850 assert(VD->getType()->isHLSLResourceRecord() &&
5851 "expected resource record type");
5852
5853 ASTContext &AST = SemaRef.getASTContext();
5854 uint64_t UIntTySize = AST.getTypeSize(AST.UnsignedIntTy);
5855 uint64_t IntTySize = AST.getTypeSize(AST.IntTy);
5856
5857 // Gather resource binding attributes.
5858 ResourceBindingAttrs Binding(VD);
5859
5860 // Find correct initialization method and create its arguments.
5861 QualType ResourceTy = VD->getType();
5862 CXXRecordDecl *ResourceDecl = ResourceTy->getAsCXXRecordDecl();
5863 CXXMethodDecl *CreateMethod = nullptr;
5865
5866 bool HasCounter = hasCounterHandle(ResourceDecl);
5867 const char *CreateMethodName;
5868 if (Binding.isExplicit())
5869 CreateMethodName = HasCounter ? "__createFromBindingWithImplicitCounter"
5870 : "__createFromBinding";
5871 else
5872 CreateMethodName = HasCounter
5873 ? "__createFromImplicitBindingWithImplicitCounter"
5874 : "__createFromImplicitBinding";
5875
5876 CreateMethod =
5877 lookupMethod(SemaRef, ResourceDecl, CreateMethodName, VD->getLocation());
5878
5879 if (!CreateMethod) {
5880 // This can happen if someone creates a struct that looks like an HLSL
5881 // resource record but does not have the required static create method.
5882 // No binding will be generated for it.
5883 assert(!ResourceDecl->isImplicit() &&
5884 "create method lookup should always succeed for built-in resource "
5885 "records");
5886 return false;
5887 }
5888
5889 if (Binding.isExplicit()) {
5890 IntegerLiteral *RegSlot =
5891 IntegerLiteral::Create(AST, llvm::APInt(UIntTySize, Binding.getSlot()),
5893 Args.push_back(RegSlot);
5894 } else {
5895 uint32_t OrderID = (Binding.hasImplicitOrderID())
5896 ? Binding.getImplicitOrderID()
5898 IntegerLiteral *OrderId =
5899 IntegerLiteral::Create(AST, llvm::APInt(UIntTySize, OrderID),
5901 Args.push_back(OrderId);
5902 }
5903
5904 IntegerLiteral *Space =
5905 IntegerLiteral::Create(AST, llvm::APInt(UIntTySize, Binding.getSpace()),
5907 Args.push_back(Space);
5908
5910 AST, llvm::APInt(IntTySize, 1), AST.IntTy, SourceLocation());
5911 Args.push_back(RangeSize);
5912
5914 AST, llvm::APInt(UIntTySize, 0), AST.UnsignedIntTy, SourceLocation());
5915 Args.push_back(Index);
5916
5917 StringRef VarName = VD->getName();
5919 AST, VarName, StringLiteralKind::Ordinary, false,
5920 AST.getStringLiteralArrayType(AST.CharTy.withConst(), VarName.size()),
5921 SourceLocation());
5923 AST, AST.getPointerType(AST.CharTy.withConst()), CK_ArrayToPointerDecay,
5924 Name, nullptr, VK_PRValue, FPOptionsOverride());
5925 Args.push_back(NameCast);
5926
5927 if (HasCounter) {
5928 // Will this be in the correct order?
5929 uint32_t CounterOrderID = getNextImplicitBindingOrderID();
5930 IntegerLiteral *CounterId =
5931 IntegerLiteral::Create(AST, llvm::APInt(UIntTySize, CounterOrderID),
5933 Args.push_back(CounterId);
5934 }
5935
5936 // Make sure the create method template is instantiated and emitted.
5937 if (!CreateMethod->isDefined() && CreateMethod->isTemplateInstantiation())
5938 SemaRef.InstantiateFunctionDefinition(VD->getLocation(), CreateMethod,
5939 true);
5940
5941 // Create CallExpr with a call to the static method and set it as the decl
5942 // initialization.
5944 AST, NestedNameSpecifierLoc(), SourceLocation(), CreateMethod, false,
5945 CreateMethod->getNameInfo(), CreateMethod->getType(), VK_PRValue);
5946
5947 auto *ImpCast = ImplicitCastExpr::Create(
5948 AST, AST.getPointerType(CreateMethod->getType()),
5949 CK_FunctionToPointerDecay, DRE, nullptr, VK_PRValue, FPOptionsOverride());
5950
5951 CallExpr *InitExpr =
5952 CallExpr::Create(AST, ImpCast, Args, ResourceTy, VK_PRValue,
5954 VD->setInit(InitExpr);
5956 SemaRef.CheckCompleteVariableDeclaration(VD);
5957 return true;
5958}
5959
5961 assert(VD->getType()->isHLSLResourceRecordArray() &&
5962 "expected array of resource records");
5963
5964 // Individual resources in a resource array are not initialized here. They
5965 // are initialized later on during codegen when the individual resources are
5966 // accessed. Codegen will emit a call to the resource initialization method
5967 // with the specified array index. We need to make sure though that the method
5968 // for the specific resource type is instantiated, so codegen can emit a call
5969 // to it when the array element is accessed.
5970
5971 // Find correct initialization method based on the resource binding
5972 // information.
5973 ASTContext &AST = SemaRef.getASTContext();
5974 QualType ResElementTy = AST.getBaseElementType(VD->getType());
5975 CXXRecordDecl *ResourceDecl = ResElementTy->getAsCXXRecordDecl();
5976 CXXMethodDecl *CreateMethod = nullptr;
5977
5978 bool HasCounter = hasCounterHandle(ResourceDecl);
5979 ResourceBindingAttrs ResourceAttrs(VD);
5980 if (ResourceAttrs.isExplicit())
5981 // Resource has explicit binding.
5982 CreateMethod =
5983 lookupMethod(SemaRef, ResourceDecl,
5984 HasCounter ? "__createFromBindingWithImplicitCounter"
5985 : "__createFromBinding",
5986 VD->getLocation());
5987 else
5988 // Resource has implicit binding.
5989 CreateMethod = lookupMethod(
5990 SemaRef, ResourceDecl,
5991 HasCounter ? "__createFromImplicitBindingWithImplicitCounter"
5992 : "__createFromImplicitBinding",
5993 VD->getLocation());
5994
5995 if (!CreateMethod)
5996 return false;
5997
5998 // Make sure the create method template is instantiated and emitted.
5999 if (!CreateMethod->isDefined() && CreateMethod->isTemplateInstantiation())
6000 SemaRef.InstantiateFunctionDefinition(VD->getLocation(), CreateMethod,
6001 true);
6002 return true;
6003}
6004
6005// Returns true if the initialization has been handled.
6006// Returns false to use default initialization.
6008 // Objects in the hlsl_constant address space are initialized
6009 // externally, so don't synthesize an implicit initializer.
6011 return true;
6012
6013 if (VD->hasGlobalStorage() && VD->getStorageClass() != SC_Static) {
6014 const Type *Ty = VD->getType().getTypePtr();
6016 return true;
6018 return true;
6019 }
6020
6021 // User-defined structs/classes do not have constructors.
6022 // When declared at a global scope, they are part of the constant buffer
6023 // and should not be initialized by the compiler.
6024 // When declared at a local scope, they are not initialized.
6025 // Also applies to arrays of user-defined structs/classes.
6026 const Type *Ty = VD->getType()->getUnqualifiedDesugaredType();
6027 while (Ty->isArrayType())
6029 if (CXXRecordDecl *RD = Ty->getAsCXXRecordDecl())
6030 return !RD->isHLSLBuiltinRecord();
6031
6032 return false;
6033}
6034
6035std::optional<const DeclBindingInfo *> SemaHLSL::inferGlobalBinding(Expr *E) {
6036 if (auto *Ternary = dyn_cast<ConditionalOperator>(E)) {
6037 auto TrueInfo = inferGlobalBinding(Ternary->getTrueExpr());
6038 auto FalseInfo = inferGlobalBinding(Ternary->getFalseExpr());
6039 if (!TrueInfo || !FalseInfo)
6040 return std::nullopt;
6041 if (*TrueInfo != *FalseInfo)
6042 return std::nullopt;
6043 return TrueInfo;
6044 }
6045
6046 if (auto *ASE = dyn_cast<ArraySubscriptExpr>(E))
6047 E = ASE->getBase()->IgnoreParenImpCasts();
6048
6049 if (DeclRefExpr *DRE = dyn_cast<DeclRefExpr>(E->IgnoreParens()))
6050 if (VarDecl *VD = dyn_cast<VarDecl>(DRE->getDecl())) {
6051 const Type *Ty = VD->getType()->getUnqualifiedDesugaredType();
6052 if (Ty->isArrayType())
6054
6055 if (const auto *AttrResType =
6056 HLSLAttributedResourceType::findHandleTypeOnResource(Ty)) {
6057 ResourceClass RC = AttrResType->getAttrs().ResourceClass;
6058 return Bindings.getDeclBindingInfo(VD, RC);
6059 }
6060 }
6061
6062 return nullptr;
6063}
6064
6065void SemaHLSL::trackLocalResource(VarDecl *VD, Expr *E) {
6066 std::optional<const DeclBindingInfo *> ExprBinding = inferGlobalBinding(E);
6067 if (!ExprBinding) {
6068 SemaRef.Diag(E->getBeginLoc(),
6069 diag::warn_hlsl_assigning_local_resource_is_not_unique)
6070 << E << VD;
6071 return; // Expr use multiple resources
6072 }
6073
6074 if (*ExprBinding == nullptr)
6075 return; // No binding could be inferred to track, return without error
6076
6077 auto PrevBinding = Assigns.find(VD);
6078 if (PrevBinding == Assigns.end()) {
6079 // No previous binding recorded, simply record the new assignment
6080 Assigns.insert({VD, *ExprBinding});
6081 return;
6082 }
6083
6084 // Otherwise, warn if the assignment implies different resource bindings
6085 if (*ExprBinding != PrevBinding->second) {
6086 SemaRef.Diag(E->getBeginLoc(),
6087 diag::warn_hlsl_assigning_local_resource_is_not_unique)
6088 << E << VD;
6089 SemaRef.Diag(VD->getLocation(), diag::note_var_declared_here) << VD;
6090 return;
6091 }
6092
6093 return;
6094}
6095
6097 Expr *RHSExpr, SourceLocation Loc) {
6098 assert((LHSExpr->getType()->isHLSLResourceRecord() ||
6099 LHSExpr->getType()->isHLSLResourceRecordArray()) &&
6100 "expected LHS to be a resource record or array of resource records");
6101 if (Opc != BO_Assign)
6102 return true;
6103
6104 // If LHS is an array subscript, get the underlying declaration.
6105 Expr *E = LHSExpr;
6106 while (auto *ASE = dyn_cast<ArraySubscriptExpr>(E))
6107 E = ASE->getBase()->IgnoreParenImpCasts();
6108
6109 // Report error if LHS is a non-static resource declared at a global scope.
6110 if (DeclRefExpr *DRE = dyn_cast<DeclRefExpr>(E->IgnoreParens())) {
6111 if (VarDecl *VD = dyn_cast<VarDecl>(DRE->getDecl())) {
6112 if (VD->hasGlobalStorage() && VD->getStorageClass() != SC_Static) {
6113 // assignment to global resource is not allowed
6114 SemaRef.Diag(Loc, diag::err_hlsl_assign_to_global_resource) << VD;
6115 SemaRef.Diag(VD->getLocation(), diag::note_var_declared_here) << VD;
6116 return false;
6117 }
6118
6119 trackLocalResource(VD, RHSExpr);
6120 }
6121 }
6122 return true;
6123}
6124
6125// Returns true if the given type can have an overload of the given
6126// binary operator.
6128 CXXRecordDecl *RD = LHSTy->getAsCXXRecordDecl();
6129 if (!RD)
6130 return true;
6131 return RD->isHLSLBuiltinRecord() || Opc != BO_Assign;
6132}
6133
6134// Walks though the global variable declaration, collects all resource binding
6135// requirements and adds them to Bindings
6136void SemaHLSL::collectResourceBindingsOnVarDecl(VarDecl *VD) {
6137 assert(VD->hasGlobalStorage() && VD->getType()->isHLSLIntangibleType() &&
6138 "expected global variable that contains HLSL resource");
6139
6140 // Cbuffers and Tbuffers are HLSLBufferDecl types
6141 if (const HLSLBufferDecl *CBufferOrTBuffer = dyn_cast<HLSLBufferDecl>(VD)) {
6142 Bindings.addDeclBindingInfo(VD, CBufferOrTBuffer->isCBuffer()
6143 ? ResourceClass::CBuffer
6144 : ResourceClass::SRV);
6145 return;
6146 }
6147
6148 // Unwrap arrays
6149 // FIXME: Calculate array size while unwrapping
6150 const Type *Ty = VD->getType()->getUnqualifiedDesugaredType();
6151 while (Ty->isArrayType()) {
6152 const ArrayType *AT = cast<ArrayType>(Ty);
6154 }
6155
6156 // Resource (or array of resources)
6157 if (const HLSLAttributedResourceType *AttrResType =
6158 HLSLAttributedResourceType::findHandleTypeOnResource(Ty)) {
6159 Bindings.addDeclBindingInfo(VD, AttrResType->getAttrs().ResourceClass);
6160 return;
6161 }
6162
6163 // User defined record type
6164 if (const RecordType *RT = dyn_cast<RecordType>(Ty))
6165 collectResourceBindingsOnUserRecordDecl(VD, RT);
6166}
6167
6168// Walks though the explicit resource binding attributes on the declaration,
6169// and makes sure there is a resource that matched the binding and updates
6170// DeclBindingInfoLists
6171void SemaHLSL::processExplicitBindingsOnDecl(VarDecl *VD) {
6172 assert(VD->hasGlobalStorage() && "expected global variable");
6173
6174 bool HasBinding = false;
6175 for (Attr *A : VD->attrs()) {
6176 if (isa<HLSLVkBindingAttr>(A)) {
6177 HasBinding = true;
6178 if (auto PA = VD->getAttr<HLSLVkPushConstantAttr>())
6179 Diag(PA->getLoc(), diag::err_hlsl_attr_incompatible) << A << PA;
6180 }
6181
6182 HLSLResourceBindingAttr *RBA = dyn_cast<HLSLResourceBindingAttr>(A);
6183 if (!RBA || !RBA->hasRegisterSlot())
6184 continue;
6185 HasBinding = true;
6186
6187 RegisterType RT = RBA->getRegisterType();
6188 assert(RT != RegisterType::I && "invalid or obsolete register type should "
6189 "never have an attribute created");
6190
6191 if (RT == RegisterType::C) {
6192 if (Bindings.hasBindingInfoForDecl(VD))
6193 SemaRef.Diag(VD->getLocation(),
6194 diag::warn_hlsl_user_defined_type_missing_member)
6195 << static_cast<int>(RT);
6196 continue;
6197 }
6198
6199 // Find DeclBindingInfo for this binding and update it, or report error
6200 // if it does not exist (user type does to contain resources with the
6201 // expected resource class).
6203 if (DeclBindingInfo *BI = Bindings.getDeclBindingInfo(VD, RC)) {
6204 // update binding info
6205 BI->setBindingAttribute(RBA, BindingType::Explicit);
6206 } else {
6207 SemaRef.Diag(VD->getLocation(),
6208 diag::warn_hlsl_user_defined_type_missing_member)
6209 << static_cast<int>(RT);
6210 }
6211 }
6212
6213 if (!HasBinding && isResourceRecordTypeOrArrayOf(VD))
6214 SemaRef.Diag(VD->getLocation(), diag::warn_hlsl_implicit_binding);
6215}
6216namespace {
6217class InitListTransformer {
6218 Sema &S;
6219 ASTContext &Ctx;
6220 QualType InitTy;
6221 QualType *DstIt = nullptr;
6222 Expr **ArgIt = nullptr;
6223 // Is wrapping the destination type iterator required? This is only used for
6224 // incomplete array types where we loop over the destination type since we
6225 // don't know the full number of elements from the declaration.
6226 bool Wrap;
6227
6228 bool castInitializer(Expr *E) {
6229 assert(DstIt && "This should always be something!");
6230 if (DstIt == DestTypes.end()) {
6231 if (!Wrap) {
6232 ArgExprs.push_back(E);
6233 // This is odd, but it isn't technically a failure due to conversion, we
6234 // handle mismatched counts of arguments differently.
6235 return true;
6236 }
6237 DstIt = DestTypes.begin();
6238 }
6239 InitializedEntity Entity = InitializedEntity::InitializeParameter(
6240 Ctx, *DstIt, /* Consumed (ObjC) */ false);
6241 ExprResult Res = S.PerformCopyInitialization(Entity, E->getBeginLoc(), E);
6242 if (Res.isInvalid())
6243 return false;
6244 Expr *Init = Res.get();
6245 ArgExprs.push_back(Init);
6246 DstIt++;
6247 return true;
6248 }
6249
6250 bool buildInitializerListImpl(Expr *E) {
6251 // If this is an initialization list, traverse the sub initializers.
6252 if (auto *Init = dyn_cast<InitListExpr>(E)) {
6253 for (auto *SubInit : Init->inits())
6254 if (!buildInitializerListImpl(SubInit))
6255 return false;
6256 return true;
6257 }
6258
6259 // If this is a scalar type, just enqueue the expression.
6260 QualType Ty = E->getType().getDesugaredType(Ctx);
6261
6262 if (Ty->isScalarType() || (Ty->isRecordType() && !Ty->isAggregateType()) ||
6264 return castInitializer(E);
6265
6266 // If this is an aggregate type and a prvalue, create an xvalue temporary
6267 // so the member accesses will be xvalues. Wrap it in OpaqueExpr to make
6268 // sure codegen will not generate duplicate copies.
6269 if (E->isPRValue() && Ty->isAggregateType()) {
6271 if (TmpExpr.isInvalid())
6272 return false;
6273 E = TmpExpr.get();
6274 E = new (Ctx) OpaqueValueExpr(E->getBeginLoc(), E->getType(),
6275 E->getValueKind(), E->getObjectKind(), E);
6276 }
6277
6278 if (auto *VecTy = Ty->getAs<VectorType>()) {
6279 uint64_t Size = VecTy->getNumElements();
6280
6281 QualType SizeTy = Ctx.getSizeType();
6282 uint64_t SizeTySize = Ctx.getTypeSize(SizeTy);
6283 for (uint64_t I = 0; I < Size; ++I) {
6284 auto *Idx = IntegerLiteral::Create(Ctx, llvm::APInt(SizeTySize, I),
6285 SizeTy, SourceLocation());
6286
6288 E, E->getBeginLoc(), Idx, E->getEndLoc());
6289 if (ElExpr.isInvalid())
6290 return false;
6291 if (!castInitializer(ElExpr.get()))
6292 return false;
6293 }
6294 return true;
6295 }
6296 if (auto *MTy = Ty->getAs<ConstantMatrixType>()) {
6297 unsigned Rows = MTy->getNumRows();
6298 unsigned Cols = MTy->getNumColumns();
6299 QualType ElemTy = MTy->getElementType();
6300
6301 for (unsigned R = 0; R < Rows; ++R) {
6302 for (unsigned C = 0; C < Cols; ++C) {
6303 // row index literal
6304 Expr *RowIdx = IntegerLiteral::Create(
6305 Ctx, llvm::APInt(Ctx.getIntWidth(Ctx.IntTy), R), Ctx.IntTy,
6306 E->getBeginLoc());
6307 // column index literal
6308 Expr *ColIdx = IntegerLiteral::Create(
6309 Ctx, llvm::APInt(Ctx.getIntWidth(Ctx.IntTy), C), Ctx.IntTy,
6310 E->getBeginLoc());
6312 E, RowIdx, ColIdx, E->getEndLoc());
6313 if (ElExpr.isInvalid())
6314 return false;
6315 if (!castInitializer(ElExpr.get()))
6316 return false;
6317 ElExpr.get()->setType(ElemTy);
6318 }
6319 }
6320 return true;
6321 }
6322
6323 if (auto *ArrTy = dyn_cast<ConstantArrayType>(Ty.getTypePtr())) {
6324 uint64_t Size = ArrTy->getZExtSize();
6325 QualType SizeTy = Ctx.getSizeType();
6326 uint64_t SizeTySize = Ctx.getTypeSize(SizeTy);
6327 for (uint64_t I = 0; I < Size; ++I) {
6328 auto *Idx = IntegerLiteral::Create(Ctx, llvm::APInt(SizeTySize, I),
6329 SizeTy, SourceLocation());
6331 E, E->getBeginLoc(), Idx, E->getEndLoc());
6332 if (ElExpr.isInvalid())
6333 return false;
6334 if (!buildInitializerListImpl(ElExpr.get()))
6335 return false;
6336 }
6337 return true;
6338 }
6339
6340 if (auto *RD = Ty->getAsCXXRecordDecl()) {
6341 llvm::SmallVector<CXXRecordDecl *> RecordDecls;
6342 RecordDecls.push_back(RD);
6343 while (RecordDecls.back()->getNumBases()) {
6344 CXXRecordDecl *D = RecordDecls.back();
6345 assert(D->getNumBases() == 1 &&
6346 "HLSL doesn't support multiple inheritance");
6347 RecordDecls.push_back(
6349 }
6350 while (!RecordDecls.empty()) {
6351 CXXRecordDecl *RD = RecordDecls.pop_back_val();
6352 for (auto *FD : RD->fields()) {
6353 if (FD->isUnnamedBitField())
6354 continue;
6355 DeclAccessPair Found = DeclAccessPair::make(FD, FD->getAccess());
6356 DeclarationNameInfo NameInfo(FD->getDeclName(), E->getBeginLoc());
6358 E, false, E->getBeginLoc(), CXXScopeSpec(), FD, Found, NameInfo);
6359 if (Res.isInvalid())
6360 return false;
6361 if (!buildInitializerListImpl(Res.get()))
6362 return false;
6363 }
6364 }
6365 }
6366 return true;
6367 }
6368
6369 Expr *generateInitListsImpl(QualType Ty) {
6370 Ty = Ty.getDesugaredType(Ctx);
6371 assert(ArgIt != ArgExprs.end() && "Something is off in iteration!");
6372 if (Ty->isScalarType() || (Ty->isRecordType() && !Ty->isAggregateType()) ||
6374 return *(ArgIt++);
6375
6376 llvm::SmallVector<Expr *> Inits;
6377 if (Ty->isVectorType() || Ty->isConstantArrayType() ||
6378 Ty->isConstantMatrixType()) {
6379 QualType ElTy;
6380 uint64_t Size = 0;
6381 if (auto *ATy = Ty->getAs<VectorType>()) {
6382 ElTy = ATy->getElementType();
6383 Size = ATy->getNumElements();
6384 } else if (auto *CMTy = Ty->getAs<ConstantMatrixType>()) {
6385 ElTy = CMTy->getElementType();
6386 Size = CMTy->getNumElementsFlattened();
6387 } else {
6388 auto *VTy = cast<ConstantArrayType>(Ty.getTypePtr());
6389 ElTy = VTy->getElementType();
6390 Size = VTy->getZExtSize();
6391 }
6392 for (uint64_t I = 0; I < Size; ++I)
6393 Inits.push_back(generateInitListsImpl(ElTy));
6394 }
6395 if (auto *RD = Ty->getAsCXXRecordDecl()) {
6396 llvm::SmallVector<CXXRecordDecl *> RecordDecls;
6397 RecordDecls.push_back(RD);
6398 while (RecordDecls.back()->getNumBases()) {
6399 CXXRecordDecl *D = RecordDecls.back();
6400 assert(D->getNumBases() == 1 &&
6401 "HLSL doesn't support multiple inheritance");
6402 RecordDecls.push_back(
6404 }
6405 while (!RecordDecls.empty()) {
6406 CXXRecordDecl *RD = RecordDecls.pop_back_val();
6407 for (auto *FD : RD->fields())
6408 if (!FD->isUnnamedBitField())
6409 Inits.push_back(generateInitListsImpl(FD->getType()));
6410 }
6411 }
6412 auto *NewInit =
6413 new (Ctx) InitListExpr(Ctx, Inits.front()->getBeginLoc(), Inits,
6414 Inits.back()->getEndLoc(), /*isExplicit=*/false);
6415 NewInit->setType(Ty);
6416 return NewInit;
6417 }
6418
6419public:
6420 llvm::SmallVector<QualType, 16> DestTypes;
6421 llvm::SmallVector<Expr *, 16> ArgExprs;
6422 InitListTransformer(Sema &SemaRef, const InitializedEntity &Entity)
6423 : S(SemaRef), Ctx(SemaRef.getASTContext()),
6424 Wrap(Entity.getType()->isIncompleteArrayType()) {
6425 InitTy = Entity.getType().getNonReferenceType();
6426 // When we're generating initializer lists for incomplete array types we
6427 // need to wrap around both when building the initializers and when
6428 // generating the final initializer lists.
6429 if (Wrap) {
6430 assert(InitTy->isIncompleteArrayType());
6431 const IncompleteArrayType *IAT = Ctx.getAsIncompleteArrayType(InitTy);
6432 InitTy = IAT->getElementType();
6433 }
6434 BuildFlattenedTypeList(InitTy, DestTypes);
6435 DstIt = DestTypes.begin();
6436 }
6437
6438 bool buildInitializerList(Expr *E) { return buildInitializerListImpl(E); }
6439
6440 Expr *generateInitLists() {
6441 assert(!ArgExprs.empty() &&
6442 "Call buildInitializerList to generate argument expressions.");
6443 ArgIt = ArgExprs.begin();
6444 if (!Wrap)
6445 return generateInitListsImpl(InitTy);
6446 llvm::SmallVector<Expr *> Inits;
6447 while (ArgIt != ArgExprs.end())
6448 Inits.push_back(generateInitListsImpl(InitTy));
6449
6450 auto *NewInit =
6451 new (Ctx) InitListExpr(Ctx, Inits.front()->getBeginLoc(), Inits,
6452 Inits.back()->getEndLoc(), /*isExplicit=*/false);
6453 llvm::APInt ArySize(64, Inits.size());
6454 NewInit->setType(Ctx.getConstantArrayType(InitTy, ArySize, nullptr,
6455 ArraySizeModifier::Normal, 0));
6456 return NewInit;
6457 }
6458};
6459} // namespace
6460
6461// Recursively detect any incomplete array anywhere in the type graph,
6462// including arrays, struct fields, and base classes.
6464 Ty = Ty.getCanonicalType();
6465
6466 // Array types
6467 if (const ArrayType *AT = dyn_cast<ArrayType>(Ty)) {
6469 return true;
6471 }
6472
6473 // Record (struct/class) types
6474 if (const auto *RT = Ty->getAs<RecordType>()) {
6475 const RecordDecl *RD = RT->getDecl();
6476
6477 // Walk base classes (for C++ / HLSL structs with inheritance)
6478 if (const auto *CXXRD = dyn_cast<CXXRecordDecl>(RD)) {
6479 for (const CXXBaseSpecifier &Base : CXXRD->bases()) {
6480 if (containsIncompleteArrayType(Base.getType()))
6481 return true;
6482 }
6483 }
6484
6485 // Walk fields
6486 for (const FieldDecl *F : RD->fields()) {
6487 if (containsIncompleteArrayType(F->getType()))
6488 return true;
6489 }
6490 }
6491
6492 return false;
6493}
6494
6496 InitListExpr *Init) {
6497 // If the initializer is a scalar, just return it.
6498 if (Init->getType()->isScalarType())
6499 return true;
6500 ASTContext &Ctx = SemaRef.getASTContext();
6501 InitListTransformer ILT(SemaRef, Entity);
6502
6503 for (unsigned I = 0; I < Init->getNumInits(); ++I) {
6504 Expr *E = Init->getInit(I);
6505 if (E->HasSideEffects(Ctx)) {
6506 QualType Ty = E->getType();
6507 if (Ty->isRecordType())
6508 E = new (Ctx) MaterializeTemporaryExpr(Ty, E, E->isLValue());
6509 E = new (Ctx) OpaqueValueExpr(E->getBeginLoc(), Ty, E->getValueKind(),
6510 E->getObjectKind(), E);
6511 Init->setInit(I, E);
6512 }
6513 if (!ILT.buildInitializerList(E))
6514 return false;
6515 }
6516 size_t ExpectedSize = ILT.DestTypes.size();
6517 size_t ActualSize = ILT.ArgExprs.size();
6518 if (ExpectedSize == 0 && ActualSize == 0)
6519 return true;
6520
6521 // Reject empty initializer if *any* incomplete array exists structurally
6522 if (ActualSize == 0 && containsIncompleteArrayType(Entity.getType())) {
6523 QualType InitTy = Entity.getType().getNonReferenceType();
6524 if (InitTy.hasAddressSpace())
6525 InitTy = SemaRef.getASTContext().removeAddrSpaceQualType(InitTy);
6526
6527 SemaRef.Diag(Init->getBeginLoc(), diag::err_hlsl_incorrect_num_initializers)
6528 << /*TooManyOrFew=*/(int)(ExpectedSize < ActualSize) << InitTy
6529 << /*ExpectedSize=*/ExpectedSize << /*ActualSize=*/ActualSize;
6530 return false;
6531 }
6532
6533 // We infer size after validating legality.
6534 // For incomplete arrays it is completely arbitrary to choose whether we think
6535 // the user intended fewer or more elements. This implementation assumes that
6536 // the user intended more, and errors that there are too few initializers to
6537 // complete the final element.
6538 if (Entity.getType()->isIncompleteArrayType()) {
6539 assert(ExpectedSize > 0 &&
6540 "The expected size of an incomplete array type must be at least 1.");
6541 ExpectedSize =
6542 ((ActualSize + ExpectedSize - 1) / ExpectedSize) * ExpectedSize;
6543 }
6544
6545 // An initializer list might be attempting to initialize a reference or
6546 // rvalue-reference. When checking the initializer we should look through
6547 // the reference.
6548 QualType InitTy = Entity.getType().getNonReferenceType();
6549 if (InitTy.hasAddressSpace())
6550 InitTy = SemaRef.getASTContext().removeAddrSpaceQualType(InitTy);
6551 if (ExpectedSize != ActualSize) {
6552 int TooManyOrFew = ActualSize > ExpectedSize ? 1 : 0;
6553 SemaRef.Diag(Init->getBeginLoc(), diag::err_hlsl_incorrect_num_initializers)
6554 << TooManyOrFew << InitTy << ExpectedSize << ActualSize;
6555 return false;
6556 }
6557
6558 // generateInitListsImpl will always return an InitListExpr here, because the
6559 // scalar case is handled above.
6560 auto *NewInit = cast<InitListExpr>(ILT.generateInitLists());
6561 Init->resizeInits(Ctx, NewInit->getNumInits());
6562 for (unsigned I = 0; I < NewInit->getNumInits(); ++I)
6563 Init->updateInit(Ctx, I, NewInit->getInit(I));
6564 return true;
6565}
6566
6567static QualType ReportMatrixInvalidMember(Sema &S, StringRef Name,
6568 StringRef Expected,
6569 SourceLocation OpLoc,
6570 SourceLocation CompLoc) {
6571 S.Diag(OpLoc, diag::err_builtin_matrix_invalid_member)
6572 << Name << Expected << SourceRange(CompLoc);
6573 return QualType();
6574}
6575
6578 const IdentifierInfo *CompName,
6579 SourceLocation CompLoc) {
6580 const auto *MT = baseType->castAs<ConstantMatrixType>();
6581 StringRef AccessorName = CompName->getName();
6582 assert(!AccessorName.empty() && "Matrix Accessor must have a name");
6583
6584 unsigned Rows = MT->getNumRows();
6585 unsigned Cols = MT->getNumColumns();
6586 bool IsZeroBasedAccessor = false;
6587 unsigned ChunkLen = 0;
6588 if (AccessorName.size() < 2)
6589 return ReportMatrixInvalidMember(S, AccessorName,
6590 "length 4 for zero based: \'_mRC\' or "
6591 "length 3 for one-based: \'_RC\' accessor",
6592 OpLoc, CompLoc);
6593
6594 if (AccessorName[0] == '_') {
6595 if (AccessorName[1] == 'm') {
6596 IsZeroBasedAccessor = true;
6597 ChunkLen = 4; // zero-based: "_mRC"
6598 } else {
6599 ChunkLen = 3; // one-based: "_RC"
6600 }
6601 } else
6603 S, AccessorName, "zero based: \'_mRC\' or one-based: \'_RC\' accessor",
6604 OpLoc, CompLoc);
6605
6606 if (AccessorName.size() % ChunkLen != 0) {
6607 const llvm::StringRef Expected = IsZeroBasedAccessor
6608 ? "zero based: '_mRC' accessor"
6609 : "one-based: '_RC' accessor";
6610
6611 return ReportMatrixInvalidMember(S, AccessorName, Expected, OpLoc, CompLoc);
6612 }
6613
6614 auto isDigit = [](char c) { return c >= '0' && c <= '9'; };
6615 auto isZeroBasedIndex = [](unsigned i) { return i <= 3; };
6616 auto isOneBasedIndex = [](unsigned i) { return i >= 1 && i <= 4; };
6617
6618 bool HasRepeated = false;
6619 SmallVector<bool, 16> Seen(Rows * Cols, false);
6620 unsigned NumComponents = 0;
6621 const char *Begin = AccessorName.data();
6622
6623 for (unsigned I = 0, E = AccessorName.size(); I < E; I += ChunkLen) {
6624 const char *Chunk = Begin + I;
6625 char RowChar = 0, ColChar = 0;
6626 if (IsZeroBasedAccessor) {
6627 // Zero-based: "_mRC"
6628 if (Chunk[0] != '_' || Chunk[1] != 'm') {
6629 char Bad = (Chunk[0] != '_') ? Chunk[0] : Chunk[1];
6631 S, StringRef(&Bad, 1), "\'_m\' prefix",
6632 OpLoc.getLocWithOffset(I + (Bad == Chunk[0] ? 1 : 2)), CompLoc);
6633 }
6634 RowChar = Chunk[2];
6635 ColChar = Chunk[3];
6636 } else {
6637 // One-based: "_RC"
6638 if (Chunk[0] != '_')
6640 S, StringRef(&Chunk[0], 1), "\'_\' prefix",
6641 OpLoc.getLocWithOffset(I + 1), CompLoc);
6642 RowChar = Chunk[1];
6643 ColChar = Chunk[2];
6644 }
6645
6646 // Must be digits.
6647 bool IsDigitsError = false;
6648 if (!isDigit(RowChar)) {
6649 unsigned BadPos = IsZeroBasedAccessor ? 2 : 1;
6650 ReportMatrixInvalidMember(S, StringRef(&RowChar, 1), "row as integer",
6651 OpLoc.getLocWithOffset(I + BadPos + 1),
6652 CompLoc);
6653 IsDigitsError = true;
6654 }
6655
6656 if (!isDigit(ColChar)) {
6657 unsigned BadPos = IsZeroBasedAccessor ? 3 : 2;
6658 ReportMatrixInvalidMember(S, StringRef(&ColChar, 1), "column as integer",
6659 OpLoc.getLocWithOffset(I + BadPos + 1),
6660 CompLoc);
6661 IsDigitsError = true;
6662 }
6663 if (IsDigitsError)
6664 return QualType();
6665
6666 unsigned Row = RowChar - '0';
6667 unsigned Col = ColChar - '0';
6668
6669 bool HasIndexingError = false;
6670 if (IsZeroBasedAccessor) {
6671 // 0-based [0..3]
6672 if (!isZeroBasedIndex(Row)) {
6673 S.Diag(OpLoc, diag::err_hlsl_matrix_element_not_in_bounds)
6674 << /*row*/ 0 << /*zero-based*/ 0 << SourceRange(CompLoc);
6675 HasIndexingError = true;
6676 }
6677 if (!isZeroBasedIndex(Col)) {
6678 S.Diag(OpLoc, diag::err_hlsl_matrix_element_not_in_bounds)
6679 << /*col*/ 1 << /*zero-based*/ 0 << SourceRange(CompLoc);
6680 HasIndexingError = true;
6681 }
6682 } else {
6683 // 1-based [1..4]
6684 if (!isOneBasedIndex(Row)) {
6685 S.Diag(OpLoc, diag::err_hlsl_matrix_element_not_in_bounds)
6686 << /*row*/ 0 << /*one-based*/ 1 << SourceRange(CompLoc);
6687 HasIndexingError = true;
6688 }
6689 if (!isOneBasedIndex(Col)) {
6690 S.Diag(OpLoc, diag::err_hlsl_matrix_element_not_in_bounds)
6691 << /*col*/ 1 << /*one-based*/ 1 << SourceRange(CompLoc);
6692 HasIndexingError = true;
6693 }
6694 // Convert to 0-based after range checking.
6695 --Row;
6696 --Col;
6697 }
6698
6699 if (HasIndexingError)
6700 return QualType();
6701
6702 // Note: matrix swizzle index is hard coded. That means Row and Col can
6703 // potentially be larger than Rows and Cols if matrix size is less than
6704 // the max index size.
6705 bool HasBoundsError = false;
6706 if (Row >= Rows) {
6707 Diag(OpLoc, diag::err_hlsl_matrix_index_out_of_bounds)
6708 << /*Row*/ 0 << Row << Rows << SourceRange(CompLoc);
6709 HasBoundsError = true;
6710 }
6711 if (Col >= Cols) {
6712 Diag(OpLoc, diag::err_hlsl_matrix_index_out_of_bounds)
6713 << /*Col*/ 1 << Col << Cols << SourceRange(CompLoc);
6714 HasBoundsError = true;
6715 }
6716 if (HasBoundsError)
6717 return QualType();
6718
6719 unsigned FlatIndex = Row * Cols + Col;
6720 if (Seen[FlatIndex])
6721 HasRepeated = true;
6722 Seen[FlatIndex] = true;
6723 ++NumComponents;
6724 }
6725 if (NumComponents == 0 || NumComponents > 4) {
6726 S.Diag(OpLoc, diag::err_hlsl_matrix_swizzle_invalid_length)
6727 << NumComponents << SourceRange(CompLoc);
6728 return QualType();
6729 }
6730
6731 QualType ElemTy = MT->getElementType();
6732 if (NumComponents == 1)
6733 return ElemTy;
6734 QualType VT = S.Context.getExtVectorType(ElemTy, NumComponents);
6735 if (HasRepeated)
6736 VK = VK_PRValue;
6737
6738 for (Sema::ExtVectorDeclsType::iterator
6740 E = S.ExtVectorDecls.end();
6741 I != E; ++I) {
6742 if ((*I)->getUnderlyingType() == VT)
6744 /*Qualifier=*/std::nullopt, *I);
6745 }
6746
6747 return VT;
6748}
6749
6751 // If initializing a local resource, track the resource binding it is using
6752 if (VDecl->getType()->isHLSLResourceRecord() && !VDecl->hasGlobalStorage())
6753 trackLocalResource(VDecl, Init);
6754
6755 const HLSLVkConstantIdAttr *ConstIdAttr =
6756 VDecl->getAttr<HLSLVkConstantIdAttr>();
6757 if (!ConstIdAttr)
6758 return true;
6759
6760 ASTContext &Context = SemaRef.getASTContext();
6761
6762 APValue InitValue;
6763 if (!Init->isCXX11ConstantExpr(Context, InitValue)) {
6764 Diag(VDecl->getLocation(), diag::err_specialization_const);
6765 VDecl->setInvalidDecl();
6766 return false;
6767 }
6768
6769 Builtin::ID BID =
6771
6772 // Argument 1: The ID from the attribute
6773 int ConstantID = ConstIdAttr->getId();
6774 llvm::APInt IDVal(Context.getIntWidth(Context.IntTy), ConstantID);
6775 Expr *IdExpr = IntegerLiteral::Create(Context, IDVal, Context.IntTy,
6776 ConstIdAttr->getLocation());
6777
6778 SmallVector<Expr *, 2> Args = {IdExpr, Init};
6779 Expr *C = SemaRef.BuildBuiltinCallExpr(Init->getExprLoc(), BID, Args);
6780 if (C->getType()->getCanonicalTypeUnqualified() !=
6782 C = SemaRef
6783 .BuildCStyleCastExpr(SourceLocation(),
6784 Context.getTrivialTypeSourceInfo(
6785 Init->getType(), Init->getExprLoc()),
6786 SourceLocation(), C)
6787 .get();
6788 }
6789 Init = C;
6790 return true;
6791}
6792
6794 SourceLocation NameLoc) {
6795 if (!Template)
6796 return QualType();
6797
6798 DeclContext *DC = Template->getDeclContext();
6799 if (!DC->isNamespace() || !cast<NamespaceDecl>(DC)->getIdentifier() ||
6800 cast<NamespaceDecl>(DC)->getName() != "hlsl")
6801 return QualType();
6802
6803 TemplateParameterList *Params = Template->getTemplateParameters();
6804 if (!Params || Params->size() != 1)
6805 return QualType();
6806
6807 if (!Template->isImplicit())
6808 return QualType();
6809
6810 // We manually extract default arguments here instead of letting
6811 // CheckTemplateIdType handle it. This ensures that for resource types that
6812 // lack a default argument (like Buffer), we return a null QualType, which
6813 // triggers the "requires template arguments" error rather than a less
6814 // descriptive "too few template arguments" error.
6815 TemplateArgumentListInfo TemplateArgs(NameLoc, NameLoc);
6816 for (NamedDecl *P : *Params) {
6817 if (auto *TTP = dyn_cast<TemplateTypeParmDecl>(P)) {
6818 if (TTP->hasDefaultArgument()) {
6819 TemplateArgs.addArgument(TTP->getDefaultArgument());
6820 continue;
6821 }
6822 } else if (auto *NTTP = dyn_cast<NonTypeTemplateParmDecl>(P)) {
6823 if (NTTP->hasDefaultArgument()) {
6824 TemplateArgs.addArgument(NTTP->getDefaultArgument());
6825 continue;
6826 }
6827 } else if (auto *TTPD = dyn_cast<TemplateTemplateParmDecl>(P)) {
6828 if (TTPD->hasDefaultArgument()) {
6829 TemplateArgs.addArgument(TTPD->getDefaultArgument());
6830 continue;
6831 }
6832 }
6833 return QualType();
6834 }
6835
6836 return SemaRef.CheckTemplateIdType(
6838 TemplateArgs, nullptr, /*ForNestedNameSpecifier=*/false);
6839}
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 isMatrixType(QualType QT)
static bool CheckUnsignedIntVecRepresentation(Sema *S, SourceLocation Loc, int ArgOrdinal, clang::QualType PassedType)
SampleKind
static bool isInvalidConstantBufferLeafElementType(const Type *Ty)
Definition SemaHLSL.cpp:415
static bool CheckCalculateLodBuiltin(Sema &S, CallExpr *TheCall)
static Builtin::ID getSpecConstBuiltinId(const Type *Type)
Definition SemaHLSL.cpp:166
static bool CheckNoDoubleElementType(Sema &S, CallExpr *TheCall, QualType ContainedType, StringRef DefaultName)
static bool CheckFloatingOrIntRepresentation(Sema *S, SourceLocation Loc, int ArgOrdinal, clang::QualType PassedType)
static const Type * createHostLayoutType(Sema &S, const Type *Ty)
Definition SemaHLSL.cpp:506
static bool CheckAnyScalarOrVector(Sema *S, CallExpr *TheCall, unsigned ArgIndex)
static bool CheckInterlockedBuiltin(Sema &S, CallExpr *TheCall, unsigned MinArgs, unsigned MaxArgs, InterlockedDest Dest, bool ReportsOriginalValue)
Check a call to an HLSL interlocked builtin.
static const HLSLAttributedResourceType * getResourceArrayHandleType(QualType QT)
Definition SemaHLSL.cpp:397
static IdentifierInfo * getHostLayoutStructName(Sema &S, NamedDecl *BaseDecl, bool MustBeUnique)
Definition SemaHLSL.cpp:471
static QualType createCounterHandleType(ASTContext &AST, QualType MainHandleTy)
static bool CheckMatrixSelect(Sema *S, CallExpr *TheCall)
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 QualType getTypedResourceElementType(QualType ContainedType)
static bool findExistingMatrixLayoutMarker(QualType T, attr::Kind &ExistingKind)
Walks the existing AttributedType sugar of T looking for a previously applied HLSLRowMajor/HLSLColumn...
static CXXRecordDecl * findRecordDeclInContext(IdentifierInfo *II, DeclContext *DC)
Definition SemaHLSL.cpp:454
static bool CheckWavePrefix(Sema *S, CallExpr *TheCall)
static bool CheckExpectedBitWidth(Sema *S, CallExpr *TheCall, unsigned ArgOrdinal, unsigned Width)
static LangAS getLangASFromResourceClass(ResourceClass RC)
Definition SemaHLSL.cpp:83
static bool CheckTextureSamplerAndLocation(Sema &S, CallExpr *TheCall, bool IncludeArraySlice=true)
static bool CheckVectorSelect(Sema *S, CallExpr *TheCall)
static QualType handleFloatVectorBinOpConversion(Sema &SemaRef, ExprResult &LHS, ExprResult &RHS, QualType LHSType, QualType RHSType, QualType LElTy, QualType RElTy, bool IsCompAssign)
static const Type * getHostLayoutFieldType(QualType QT)
Definition SemaHLSL.cpp:535
InterlockedDest
The dest types an interlocked operation accepts. Float is 32-bit only.
static ResourceClass getResourceClass(RegisterType RT)
Definition SemaHLSL.cpp:148
static CXXRecordDecl * createHostLayoutStruct(Sema &S, CXXRecordDecl *StructDecl)
Definition SemaHLSL.cpp:571
static bool CheckScalarOrVector(Sema *S, CallExpr *TheCall, QualType Scalar, unsigned ArgIndex)
static QualType getVectorOrScalarType(Sema &S, QualType BaseType, unsigned Count)
static bool CheckSamplingBuiltin(Sema &S, CallExpr *TheCall, SampleKind Kind)
static bool CheckScalarOrVectorOrMatrix(Sema *S, CallExpr *TheCall, QualType Scalar, unsigned ArgIndex)
static bool CheckFloatRepresentation(Sema *S, SourceLocation Loc, int ArgOrdinal, clang::QualType PassedType)
static bool CheckAnyDoubleRepresentation(Sema *S, SourceLocation Loc, int ArgOrdinal, clang::QualType PassedType)
static bool requiresImplicitBufferLayoutStructure(const CXXRecordDecl *RD)
Definition SemaHLSL.cpp:434
static bool CheckResourceHandle(Sema *S, CallExpr *TheCall, unsigned ArgIndex, llvm::function_ref< bool(const HLSLAttributedResourceType *ResType)> Check=nullptr)
static void validatePackoffset(Sema &S, HLSLBufferDecl *BufDecl)
Definition SemaHLSL.cpp:328
static StringRef getCurrentResourceMethodName(Sema &S, StringRef DefaultName)
static bool IsDefaultBufferConstantDecl(const ASTContext &Ctx, VarDecl *VD)
HLSLResourceBindingAttr::RegisterType RegisterType
Definition SemaHLSL.cpp:60
static CastKind getScalarCastKind(ASTContext &Ctx, QualType DestTy, QualType SrcTy)
static bool CheckGatherBuiltin(Sema &S, CallExpr *TheCall, bool IsCmp)
static bool isValidWaveSizeValue(unsigned Value)
static bool isResourceRecordTypeOrArrayOf(QualType Ty)
Definition SemaHLSL.cpp:388
static bool CheckLoadMSBuiltin(Sema &S, CallExpr *TheCall)
static bool AccumulateHLSLResourceSlots(QualType Ty, uint64_t &StartSlot, const uint64_t &Limit, const ResourceClass ResClass, ASTContext &Ctx, uint64_t ArrayCount=1)
static bool CheckNoDoubleVectors(Sema *S, SourceLocation Loc, int ArgOrdinal, clang::QualType PassedType)
static bool ValidateMultipleRegisterAnnotations(Sema &S, Decl *TheDecl, RegisterType regType)
static bool DiagnoseLocalRegisterBinding(Sema &S, SourceLocation &ArgLoc, Decl *D, RegisterType RegType, bool SpecifiedSpace)
static bool CheckIndexType(Sema *S, CallExpr *TheCall, unsigned IndexArgIndex)
This file declares semantic analysis for HLSL constructs.
Defines the clang::SourceLocation class and associated facilities.
Defines various enumerations that describe declaration and type specifiers.
C Language Family Type Representation.
Defines the clang::TypeLoc interface and its subclasses.
C Language Family Type Representation.
static const TypeInfo & getInfo(unsigned id)
Definition Types.cpp:44
return(__x > > __y)|(__x<<(32 - __y))
APValue - This class implements a discriminated union of [uninitialized] [APSInt] [APFloat],...
Definition APValue.h:123
virtual bool HandleTopLevelDecl(DeclGroupRef D)
HandleTopLevelDecl - Handle the specified top-level declaration.
Holds long-lived AST nodes (such as types and decls) that can be referred to throughout the semantic ...
Definition ASTContext.h:239
unsigned getIntWidth(QualType T) const
QualType getConstantMatrixType(QualType ElementType, unsigned NumRows, unsigned NumColumns, std::optional< MatrixType::LayoutKind > Layout=std::nullopt) const
Return the unique reference to the matrix type of the specified element type and size.
int getIntegerTypeOrder(QualType LHS, QualType RHS) const
Return the highest ranked integer type, see C99 6.3.1.8p1.
CanQualType FloatTy
QualType getPointerType(QualType T) const
Return the uniqued reference to the type for a pointer to the specified type.
const IncompleteArrayType * getAsIncompleteArrayType(QualType T) const
IdentifierTable & Idents
Definition ASTContext.h:850
QualType getConstantArrayType(QualType EltTy, const llvm::APInt &ArySize, const Expr *SizeExpr, ArraySizeModifier ASM, unsigned IndexTypeQuals) const
Return the unique reference to the type for a constant array of the specified element type.
QualType getBaseElementType(const ArrayType *VAT) const
Return the innermost element type of an array type.
int getFloatingTypeOrder(QualType LHS, QualType RHS) const
Compare the rank of the two specified floating point types, ignoring the domain of the type (i....
CanQualType BoolTy
TypeSourceInfo * getTrivialTypeSourceInfo(QualType T, SourceLocation Loc=SourceLocation()) const
Allocate a TypeSourceInfo where all locations have been initialized to a given location,...
QualType getStringLiteralArrayType(QualType EltTy, unsigned Length) const
Return a type for a constant array for a string literal of the specified element type and length.
CanQualType CharTy
CanQualType IntTy
uint64_t getTypeSize(QualType T) const
Return the size of the specified (complete) type T, in bits.
CharUnits getTypeSizeInChars(QualType T) const
Return the size of the specified (complete) type T, in characters.
CanQualType VoidTy
CanQualType UnsignedIntTy
QualType getTypedefType(ElaboratedTypeKeyword Keyword, NestedNameSpecifier Qualifier, const TypedefNameDecl *Decl, QualType UnderlyingType=QualType(), std::optional< bool > TypeMatchesDeclOrNone=std::nullopt) const
Return the unique reference to the type for the specified typedef-name decl.
llvm::StringRef backupStr(llvm::StringRef S) const
Definition ASTContext.h:932
QualType getSizeType() const
Return the unique type for "size_t" (C99 7.17), defined in <stddef.h>.
QualType getExtVectorType(QualType VectorType, unsigned NumElts) const
Return the unique reference to an extended vector type of the specified element type and size.
const TargetInfo & getTargetInfo() const
Definition ASTContext.h:969
QualType getHLSLAttributedResourceType(QualType Wrapped, QualType Contained, const HLSLAttributedResourceType::Attributes &Attrs)
QualType getAddrSpaceQualType(QualType T, LangAS AddressSpace) const
Return the uniqued reference to the type for an address space qualified type with the specified type ...
CanQualType getCanonicalTagType(const TagDecl *TD) const
static bool hasSameUnqualifiedType(QualType T1, QualType T2)
Determine whether the given types are equivalent after cvr-qualifiers have been removed.
unsigned getTypeAlign(QualType T) const
Return the ABI-specified alignment of a (complete) type T, in bits.
PtrTy get() const
Definition Ownership.h:171
bool isInvalid() const
Definition Ownership.h:167
Represents an array type, per C99 6.7.5.2 - Array Declarators.
Definition TypeBase.h:3819
QualType getElementType() const
Definition TypeBase.h:3831
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:3857
bool isZeroSize() const
Return true if the size is zero.
Definition TypeBase.h:3927
llvm::APInt getSize() const
Return the constant array size as an APInt.
Definition TypeBase.h:3913
uint64_t getZExtSize() const
Return the size zero-extended as a uint64_t.
Definition TypeBase.h:3933
Represents a concrete matrix type with constant number of rows and columns.
Definition TypeBase.h:4489
unsigned getNumColumns() const
Returns the number of columns in the matrix.
Definition TypeBase.h:4511
unsigned getNumRows() const
Returns the number of rows in the matrix.
Definition TypeBase.h:4508
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:4364
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:5983
void addLayoutStruct(CXXRecordDecl *LS)
Definition Decl.cpp:6023
void setHasValidPackoffset(bool PO)
Definition Decl.h:5377
static HLSLBufferDecl * CreateDefaultCBuffer(ASTContext &C, DeclContext *LexicalParent, ArrayRef< Decl * > DefaultCBufferDecls)
Definition Decl.cpp:6006
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:6069
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:4434
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:3816
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:8445
LangAS getAddressSpace() const
Return the address space of this type.
Definition TypeBase.h:8571
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:8630
QualType getCanonicalType() const
Definition TypeBase.h:8497
QualType getUnqualifiedType() const
Retrieve the unqualified variant of the given type, removing as little sugar as possible.
Definition TypeBase.h:8539
bool hasAddressSpace() const
Check if this type has any address space qualifier.
Definition TypeBase.h:8566
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:254
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)
bool CanPerformPackedTypeCast(Expr *Src, QualType DestTy)
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:9416
@ LookupMemberName
Member name lookup, which finds the names of class/struct/union members.
Definition Sema.h:9424
bool checkArgCountAtMost(CallExpr *Call, unsigned MaxArgCount)
Checks that a call expression's argument count is at most the desired number.
ExtVectorDeclsType ExtVectorDecls
ExtVectorDecls - This is a list all the extended vector types.
Definition Sema.h:5001
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:4977
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:8416
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:9067
bool isBooleanType() const
Definition TypeBase.h:9204
bool isIncompleteArrayType() const
Definition TypeBase.h:8789
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:8785
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:8781
CXXRecordDecl * castAsCXXRecordDecl() const
Definition Type.h:36
bool isArithmeticType() const
Definition Type.cpp:2548
bool isConstantMatrixType() const
Definition TypeBase.h:8849
bool isHLSLBuiltinIntangibleType() const
Definition TypeBase.h:9005
bool isPointerType() const
Definition TypeBase.h:8682
CanQualType getCanonicalTypeUnqualified() const
bool isIntegerType() const
isIntegerType() does not include complex integers (a GCC extension).
Definition TypeBase.h:9111
const T * castAs() const
Member-template castAs<specific type>.
Definition TypeBase.h:9361
bool isReferenceType() const
Definition TypeBase.h:8706
bool isHLSLIntangibleType() const
Definition Type.cpp:5716
bool isEnumeralType() const
Definition TypeBase.h:8813
bool isScalarType() const
Definition TypeBase.h:9173
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:9036
bool isDependentType() const
Whether this type is a dependent type, meaning that its definition somehow depends on a template para...
Definition TypeBase.h:2862
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:8845
bool isHLSLBuiltinPackedType() const
Definition TypeBase.h:9012
bool isHLSLResourceRecord() const
Definition Type.cpp:5703
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:8821
bool isRealFloatingType() const
Floating point categories.
Definition Type.cpp:2531
bool isHLSLAttributedResourceType() const
Definition TypeBase.h:9024
@ STK_FloatingComplex
Definition TypeBase.h:2844
@ STK_ObjCObjectPointer
Definition TypeBase.h:2838
@ STK_IntegralComplex
Definition TypeBase.h:2843
@ STK_MemberPointer
Definition TypeBase.h:2839
bool isFloatingType() const
Definition Type.cpp:2515
bool isSamplerT() const
Definition TypeBase.h:8926
const T * getAs() const
Member-template getAs<specific type>'.
Definition TypeBase.h:9294
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:8809
bool isHLSLResourceRecordArray() const
Definition Type.cpp:5707
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:4272
unsigned getNumElements() const
Definition TypeBase.h:4287
QualType getElementType() const
Definition TypeBase.h:4286
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:276
@ 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:155
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:128
@ AS_none
Definition Specifiers.h:131
@ SC_Extern
Definition Specifiers.h:255
@ SC_Static
Definition Specifiers.h:256
@ SC_None
Definition Specifiers.h:254
@ 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:384
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:136
@ VK_PRValue
A pr-value expression (in the C++11 taxonomy) produces a temporary value.
Definition Specifiers.h:139
@ VK_LValue
An l-value expression is a reference to an object with independent storage.
Definition Specifiers.h:143
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:6034
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.
int32_t 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