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