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