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 // Each array element occupies a separate semantic index.
1036 QualType T = D == FD ? FD->getReturnType() : D->getType();
1037 const ConstantArrayType *AT =
1038 getASTContext().getAsConstantArrayType(T.getNonReferenceType());
1039 if (isZeroSizedArray(AT)) {
1040 Diag(A->getLoc(), diag::err_hlsl_semantic_zero_sized_array)
1041 << A->getAttrName();
1042 return false;
1043 }
1044 unsigned ElementCount = AT ? ASTContext::getConstantArrayElementCount(AT) : 1;
1045
1046 checkSemanticAnnotation(FD, D, A, SC, ElementCount);
1047 OutputDecl->addAttr(A);
1048
1049 unsigned Location = ActiveSemantic.Index.value_or(0);
1050
1051 if (!isVkPipelineBuiltin(getASTContext(), FD, A,
1052 any(SC.CurrentIOType & IOType::In))) {
1053 bool HasVkLocation = false;
1054 if (auto *A = D->getAttr<HLSLVkLocationAttr>()) {
1055 HasVkLocation = true;
1056 Location = A->getLocation();
1057 }
1058
1059 if (SC.UsesExplicitVkLocations.value_or(HasVkLocation) != HasVkLocation) {
1060 Diag(D->getLocation(), diag::err_hlsl_semantic_partial_explicit_indexing);
1061 return false;
1062 }
1063 SC.UsesExplicitVkLocations = HasVkLocation;
1064 }
1065
1066 ActiveSemantic.Index = Location + ElementCount;
1067
1068 StringRef BaseName = ActiveSemantic.Semantic->getAttrName()->getName();
1069 std::string LowerName = BaseName.lower();
1070 for (unsigned I = 0; I < ElementCount; ++I) {
1071 auto [It, Inserted] = SC.ActiveSemantics.try_emplace(
1072 (Twine(LowerName) + Twine(Location + I)).str(), D->getLocation());
1073 if (!Inserted) {
1074 Diag(D->getLocation(), diag::err_hlsl_semantic_index_overlap)
1075 << (BaseName + Twine(Location + I)).str();
1076 Diag(It->second, diag::note_previous_use);
1077 return false;
1078 }
1079 }
1080
1081 return true;
1082}
1083
1084bool SemaHLSL::determineActiveSemantic(FunctionDecl *FD,
1085 DeclaratorDecl *OutputDecl,
1086 DeclaratorDecl *D,
1087 SemanticInfo &ActiveSemantic,
1088 SemaHLSL::SemanticContext &SC) {
1089 if (ActiveSemantic.Semantic == nullptr) {
1090 ActiveSemantic.Semantic = D->getAttr<HLSLParsedSemanticAttr>();
1091 if (ActiveSemantic.Semantic)
1092 ActiveSemantic.Index = ActiveSemantic.Semantic->getSemanticIndex();
1093 }
1094
1095 const Type *T = D == FD ? &*FD->getReturnType() : &*D->getType();
1097
1098 const RecordType *RT = dyn_cast<RecordType>(T);
1099 if (!RT)
1100 return determineActiveSemanticOnScalar(FD, OutputDecl, D, ActiveSemantic,
1101 SC);
1102
1103 const RecordDecl *RD = RT->getDecl();
1104 for (FieldDecl *Field : RD->fields()) {
1105 SemanticInfo Info = ActiveSemantic;
1106 if (!determineActiveSemantic(FD, OutputDecl, Field, Info, SC)) {
1107 Diag(Field->getLocation(), diag::note_hlsl_semantic_used_here) << Field;
1108 return false;
1109 }
1110 if (ActiveSemantic.Semantic)
1111 ActiveSemantic = Info;
1112 }
1113
1114 return true;
1115}
1116
1118 const auto *ShaderAttr = FD->getAttr<HLSLShaderAttr>();
1119 assert(ShaderAttr && "Entry point has no shader attribute");
1120 llvm::Triple::EnvironmentType ST = ShaderAttr->getType();
1122 VersionTuple Ver = TargetInfo.getTriple().getOSVersion();
1123 switch (ST) {
1124 case llvm::Triple::Pixel:
1125 case llvm::Triple::Vertex:
1126 case llvm::Triple::Geometry:
1127 case llvm::Triple::Hull:
1128 case llvm::Triple::Domain:
1129 case llvm::Triple::RayGeneration:
1130 case llvm::Triple::Intersection:
1131 case llvm::Triple::AnyHit:
1132 case llvm::Triple::ClosestHit:
1133 case llvm::Triple::Miss:
1134 case llvm::Triple::Callable:
1135 if (const auto *NT = FD->getAttr<HLSLNumThreadsAttr>()) {
1136 diagnoseAttrStageMismatch(NT, ST,
1137 {llvm::Triple::Compute,
1138 llvm::Triple::Amplification,
1139 llvm::Triple::Mesh});
1140 FD->setInvalidDecl();
1141 }
1142 if (const auto *WS = FD->getAttr<HLSLWaveSizeAttr>()) {
1143 diagnoseAttrStageMismatch(WS, ST,
1144 {llvm::Triple::Compute,
1145 llvm::Triple::Amplification,
1146 llvm::Triple::Mesh});
1147 FD->setInvalidDecl();
1148 }
1149 break;
1150
1151 case llvm::Triple::Compute:
1152 case llvm::Triple::Amplification:
1153 case llvm::Triple::Mesh:
1154 if (!FD->hasAttr<HLSLNumThreadsAttr>()) {
1155 Diag(FD->getLocation(), diag::err_hlsl_missing_numthreads)
1156 << llvm::Triple::getEnvironmentTypeName(ST);
1157 FD->setInvalidDecl();
1158 }
1159 if (const auto *WS = FD->getAttr<HLSLWaveSizeAttr>()) {
1160 if (TargetInfo.getTriple().isSPIRV()) {
1161 Diag(WS->getLocation(), diag::warn_hlsl_wavesize_unsupported_spirv);
1162 } else if (Ver < VersionTuple(6, 6)) {
1163 Diag(WS->getLocation(), diag::err_hlsl_attribute_in_wrong_shader_model)
1164 << WS << "6.6";
1165 FD->setInvalidDecl();
1166 } else if (WS->getSpelledArgsCount() > 1 && Ver < VersionTuple(6, 8)) {
1167 Diag(
1168 WS->getLocation(),
1169 diag::err_hlsl_attribute_number_arguments_insufficient_shader_model)
1170 << WS << WS->getSpelledArgsCount() << "6.8";
1171 FD->setInvalidDecl();
1172 }
1173 }
1174 break;
1175 case llvm::Triple::RootSignature:
1176 llvm_unreachable("rootsig environment has no function entry point");
1177 default:
1178 llvm_unreachable("Unhandled environment in triple");
1179 }
1180
1181 SemaHLSL::SemanticContext InputSC = {};
1182 InputSC.CurrentIOType = IOType::In;
1183 SemaHLSL::SemanticContext OutputSC = {};
1184 OutputSC.CurrentIOType = IOType::Out;
1185
1186 for (ParmVarDecl *Param : FD->parameters()) {
1187 SemanticInfo ActiveSemantic;
1188 ActiveSemantic.Semantic = Param->getAttr<HLSLParsedSemanticAttr>();
1189 if (ActiveSemantic.Semantic)
1190 ActiveSemantic.Index = ActiveSemantic.Semantic->getSemanticIndex();
1191
1192 // FIXME: An `inout` parameter is part of both signatures, but it is only
1193 // verified against the output one here.
1194 const auto *MA = Param->getAttr<HLSLParamModifierAttr>();
1195 SemanticContext &SC = MA && MA->isAnyOut() ? OutputSC : InputSC;
1196
1197 // Interpolation applies to pixel inputs and vertex outputs, including the
1198 // corresponding side of inout parameters.
1199 if (((ST == llvm::Triple::Pixel && (!MA || MA->isAnyIn())) ||
1200 (ST == llvm::Triple::Vertex && MA && MA->isAnyOut())) &&
1201 !checkInterpolationModifiers(Param, nullptr, nullptr))
1202 FD->setInvalidDecl();
1203
1204 if (!determineActiveSemantic(FD, Param, Param, ActiveSemantic, SC)) {
1205 Diag(Param->getLocation(), diag::note_previous_decl) << Param;
1206 FD->setInvalidDecl();
1207 }
1208 }
1209
1210 SemanticInfo ActiveSemantic;
1211 ActiveSemantic.Semantic = FD->getAttr<HLSLParsedSemanticAttr>();
1212 if (ActiveSemantic.Semantic)
1213 ActiveSemantic.Index = ActiveSemantic.Semantic->getSemanticIndex();
1214 if (!FD->getReturnType()->isVoidType()) {
1215 if (ST == llvm::Triple::Vertex &&
1216 !checkInterpolationModifiers(FD, nullptr, nullptr))
1217 FD->setInvalidDecl();
1218 determineActiveSemantic(FD, FD, FD, ActiveSemantic, OutputSC);
1219 }
1220}
1221
1222void SemaHLSL::checkSemanticAnnotation(
1223 FunctionDecl *EntryPoint, const Decl *Param,
1224 const HLSLAppliedSemanticAttr *SemanticAttr, const SemanticContext &SC,
1225 unsigned ElementCount) {
1226 auto *ShaderAttr = EntryPoint->getAttr<HLSLShaderAttr>();
1227 assert(ShaderAttr && "Entry point has no shader attribute");
1228 llvm::Triple::EnvironmentType ST = ShaderAttr->getType();
1229
1230 SemanticKind Kind =
1231 llvm::hlsl::getSemanticKind(SemanticAttr->getSemanticName());
1232 llvm::hlsl::SemanticInterpretation Interpretation =
1233 llvm::hlsl::getInterpretationKind(Kind, ST, SC.CurrentIOType);
1234 if (Interpretation == llvm::hlsl::SemanticInterpretation::Invalid) {
1235 diagnoseSemanticStageMismatch(SemanticAttr, ST, SC.CurrentIOType, Kind);
1236 return;
1237 }
1238
1239 // A system-value name can have an arbitrary interpretation, for example
1240 // SV_Position on a vertex input. Only the general type restrictions apply.
1241 if (Interpretation == llvm::hlsl::SemanticInterpretation::Arbitrary) {
1242 diagnoseSemanticType(Param, SemanticAttr, SemanticKind::Arbitrary);
1243 return;
1244 }
1245
1246 diagnoseSystemSemanticIndex(SemanticAttr, Kind, ElementCount);
1247 diagnoseSemanticType(Param, SemanticAttr, Kind);
1248}
1249
1250void SemaHLSL::diagnoseSystemSemanticIndex(const HLSLAppliedSemanticAttr *A,
1251 SemanticKind Kind,
1252 unsigned ElementCount) {
1253 assert(Kind != SemanticKind::Invalid && Kind != SemanticKind::Arbitrary &&
1254 "expected a recognized system semantic");
1255 assert(ElementCount > 0 && "a semantic covers at least one element");
1256 // The attribute stores the index in an int. Recover its unsigned value
1257 // before widening the arithmetic to detect overflow of the semantic range.
1258 uint32_t FirstIndex = A->getSemanticIndex();
1259 uint64_t LastIndex = uint64_t(FirstIndex) + ElementCount - 1;
1260 constexpr uint32_t MaxSemanticIndex = std::numeric_limits<uint32_t>::max();
1261 if (LastIndex > MaxSemanticIndex) {
1262 Diag(A->getLoc(), diag::err_hlsl_semantic_index_out_of_range)
1263 << A->getAttrName() << LastIndex << MaxSemanticIndex;
1264 return;
1265 }
1266 if (LastIndex == 0)
1267 return;
1268
1269 switch (Kind) {
1270 // These semantics are limited by signature packing, not semantic indices.
1271 case SemanticKind::ClipDistance:
1272 case SemanticKind::CullDistance:
1273 return;
1274 case SemanticKind::Target: {
1275 constexpr unsigned MaxTargetIndex = 7;
1276 if (LastIndex > MaxTargetIndex)
1277 Diag(A->getLoc(), diag::err_hlsl_semantic_index_out_of_range)
1278 << A->getAttrName() << LastIndex << MaxTargetIndex;
1279 return;
1280 }
1281 default:
1282 Diag(A->getLoc(), diag::err_hlsl_semantic_indexing_not_supported)
1283 << A->getAttrName();
1284 return;
1285 }
1286}
1287
1288static QualType getElementTypeOf(QualType T, bool IncludeMatrix) {
1289 if (const auto *VT = T->getAs<clang::VectorType>())
1290 return VT->getElementType();
1291 if (IncludeMatrix)
1292 if (const auto *MT = T->getAs<clang::MatrixType>())
1293 return MT->getElementType();
1294 return T;
1295}
1296
1298 if (const auto *VT = T->getAs<clang::VectorType>())
1299 return VT->getNumElements();
1300 return 1;
1301}
1302
1304 return Elem->isHalfType() || Elem->isFloat16Type() || Elem->isFloat32Type();
1305}
1306
1307// System-value integer types exclude bool.
1308static bool isIntElementOfWidth(const ASTContext &Ctx, QualType Elem,
1309 uint64_t Width) {
1310 if (!Elem->isIntegerType() || Elem->isBooleanType())
1311 return false;
1312 return Ctx.getTypeSize(Elem) == Width;
1313}
1314
1315static bool isIntUpTo32Element(const ASTContext &Ctx, QualType Elem) {
1316 return isIntElementOfWidth(Ctx, Elem, 16) ||
1317 isIntElementOfWidth(Ctx, Elem, 32);
1318}
1319
1320void SemaHLSL::diagnoseSemanticType(const Decl *D,
1321 const HLSLAppliedSemanticAttr *A,
1322 SemanticKind Kind) {
1323 assert(Kind != SemanticKind::Invalid && "expected a valid semantic");
1324 ASTContext &Ctx = getASTContext();
1325
1326 QualType T;
1327 if (const auto *FD = dyn_cast<FunctionDecl>(D))
1328 T = FD->getReturnType();
1329 else
1330 T = cast<ValueDecl>(D)->getType();
1331
1332 // `out` and `inout` parameters are passed by reference.
1333 T = T.getNonReferenceType();
1334
1335 // Array semantics constrain each element's type.
1336 QualType DeclaredTy = T;
1337 while (const ConstantArrayType *AT = Ctx.getAsConstantArrayType(T))
1338 T = AT->getElementType();
1339
1340 QualType ElemTy = getElementTypeOf(T, /*IncludeMatrix=*/false);
1341 unsigned Components = getComponentCountOf(T);
1342
1343 bool IsSPIRV = getASTContext().getTargetInfo().getTriple().isSPIRV();
1344
1345 switch (Kind) {
1346 case SemanticKind::DispatchThreadID:
1347 case SemanticKind::GroupID:
1348 case SemanticKind::GroupThreadID:
1349 if (!isIntUpTo32Element(Ctx, ElemTy) || Components > 3)
1350 Diag(A->getLoc(), diag::err_hlsl_semantic_invalid_type)
1351 << A->getAttrName() << /* scalar or vector of up to */ 1 << 3
1352 << /* 16 or 32 bit integer */ 0 << DeclaredTy;
1353 return;
1354 case SemanticKind::GroupIndex:
1355 if (!isIntElementOfWidth(Ctx, ElemTy, 32) || Components != 1)
1356 Diag(A->getLoc(), diag::err_hlsl_semantic_invalid_type)
1357 << A->getAttrName() << /* scalar */ 0 << 1 << /* 32 bit integer */ 1
1358 << DeclaredTy;
1359 return;
1360 case SemanticKind::VertexID:
1361 if (!isIntUpTo32Element(Ctx, ElemTy) || Components != 1)
1362 Diag(A->getLoc(), diag::err_hlsl_semantic_invalid_type)
1363 << A->getAttrName() << /* scalar */ 0 << 1
1364 << /* 16 or 32 bit integer */ 0 << DeclaredTy;
1365 return;
1366 case SemanticKind::Position:
1367 case SemanticKind::Target:
1368 if (!isFloatOrHalfElement(ElemTy) || Components > 4)
1369 Diag(A->getLoc(), diag::err_hlsl_semantic_invalid_type)
1370 << A->getAttrName() << /* scalar or vector of up to */ 1 << 4
1371 << /* 16 or 32 bit floating-point */ 2 << DeclaredTy;
1372 return;
1373 case SemanticKind::InstanceID:
1374 // DXIL permits U32 or U16. SPIR-V requires a 32-bit scalar per
1375 // VUID-InstanceIndex-InstanceIndex-04265.
1376 if (!T->isUnsignedIntegerType() ||
1377 !(isIntElementOfWidth(Ctx, T, 32) ||
1378 (!IsSPIRV && isIntElementOfWidth(Ctx, T, 16))))
1379 Diag(A->getLoc(), diag::err_hlsl_semantic_invalid_type)
1380 << A->getAttrName() << /* scalar */ 0 << 1
1381 << /* 16 or 32 bit unsigned integer / 32 bit unsigned integer */
1382 (IsSPIRV ? 4 : 3) << DeclaredTy;
1383 return;
1384 default:
1385 // Other semantics only have the general signature type restrictions.
1386 break;
1387 }
1388
1389 // DXIL signatures cannot carry 64-bit components, even for arbitrary
1390 // semantics. SPIR-V interfaces do support these types.
1391 QualType ScalarTy = getElementTypeOf(T, /*IncludeMatrix=*/true);
1392 if (Ctx.getTargetInfo().getTriple().isDXIL() &&
1393 (ScalarTy->isSpecificBuiltinType(BuiltinType::Double) ||
1394 isIntElementOfWidth(Ctx, ScalarTy, 64)))
1395 Diag(A->getLoc(), diag::err_hlsl_semantic_64bit_type)
1396 << A->getAttrName() << DeclaredTy;
1397}
1398
1399void SemaHLSL::diagnoseAttrStageMismatch(
1400 const Attr *A, llvm::Triple::EnvironmentType Stage,
1401 std::initializer_list<llvm::Triple::EnvironmentType> AllowedStages) {
1402 SmallVector<StringRef, 8> StageStrings;
1403 llvm::transform(AllowedStages, std::back_inserter(StageStrings),
1404 [](llvm::Triple::EnvironmentType ST) {
1405 return StringRef(
1406 HLSLShaderAttr::ConvertEnvironmentTypeToStr(ST));
1407 });
1408 Diag(A->getLoc(), diag::err_hlsl_attr_unsupported_in_stage)
1409 << A->getAttrName() << llvm::Triple::getEnvironmentTypeName(Stage)
1410 << (AllowedStages.size() != 1) << join(StageStrings, ", ");
1411}
1412
1413void SemaHLSL::diagnoseSemanticStageMismatch(
1414 const Attr *A, llvm::Triple::EnvironmentType Stage, IOType CurrentIOType,
1415 SemanticKind Kind) {
1416
1417 ArrayRef<SemanticStageInfo> Allowed = llvm::hlsl::getAvailableStages(Kind);
1418 auto It = llvm::find_if(Allowed, [&Stage](const SemanticStageInfo &Info) {
1419 return Info.Stage == Stage;
1420 });
1421
1422 StringRef CurrentIOTypeName = "patch constants or primitives";
1423 if (any(CurrentIOType & IOType::In))
1424 CurrentIOTypeName = "inputs";
1425 else if (any(CurrentIOType & IOType::Out))
1426 CurrentIOTypeName = "outputs";
1427
1428 // The semantic is not available in this shader stage at all.
1429 if (It == Allowed.end()) {
1430 Diag(A->getLoc(), diag::err_hlsl_semantic_unsupported_iotype_for_stage)
1431 << A->getAttrName() << llvm::Triple::getEnvironmentTypeName(Stage)
1432 << CurrentIOTypeName;
1433 return;
1434 }
1435
1436 IOType AllowedIOTypes = It->AllowedIOTypesMask;
1437 if (!(AllowedIOTypes & CurrentIOType)) {
1438 Diag(A->getLoc(), diag::err_hlsl_semantic_unsupported_iotype_for_stage)
1439 << A->getAttrName() << llvm::Triple::getEnvironmentTypeName(Stage)
1440 << CurrentIOTypeName;
1441 return;
1442 }
1443}
1444
1445template <CastKind Kind>
1446static void castVector(Sema &S, ExprResult &E, QualType &Ty, unsigned Sz) {
1447 if (const auto *VTy = Ty->getAs<VectorType>())
1448 Ty = VTy->getElementType();
1449 Ty = S.getASTContext().getExtVectorType(Ty, Sz);
1450 E = S.ImpCastExprToType(E.get(), Ty, Kind);
1451}
1452
1453template <CastKind Kind>
1455 E = S.ImpCastExprToType(E.get(), Ty, Kind);
1456 return Ty;
1457}
1458
1460 Sema &SemaRef, ExprResult &LHS, ExprResult &RHS, QualType LHSType,
1461 QualType RHSType, QualType LElTy, QualType RElTy, bool IsCompAssign) {
1462 bool LHSFloat = LElTy->isRealFloatingType();
1463 bool RHSFloat = RElTy->isRealFloatingType();
1464
1465 if (LHSFloat && RHSFloat) {
1466 if (IsCompAssign ||
1467 SemaRef.getASTContext().getFloatingTypeOrder(LElTy, RElTy) > 0)
1468 return castElement<CK_FloatingCast>(SemaRef, RHS, LHSType);
1469
1470 return castElement<CK_FloatingCast>(SemaRef, LHS, RHSType);
1471 }
1472
1473 if (LHSFloat)
1474 return castElement<CK_IntegralToFloating>(SemaRef, RHS, LHSType);
1475
1476 assert(RHSFloat);
1477 if (IsCompAssign)
1478 return castElement<clang::CK_FloatingToIntegral>(SemaRef, RHS, LHSType);
1479
1480 return castElement<CK_IntegralToFloating>(SemaRef, LHS, RHSType);
1481}
1482
1484 Sema &SemaRef, ExprResult &LHS, ExprResult &RHS, QualType LHSType,
1485 QualType RHSType, QualType LElTy, QualType RElTy, bool IsCompAssign) {
1486
1487 int IntOrder = SemaRef.Context.getIntegerTypeOrder(LElTy, RElTy);
1488 bool LHSSigned = LElTy->hasSignedIntegerRepresentation();
1489 bool RHSSigned = RElTy->hasSignedIntegerRepresentation();
1490 auto &Ctx = SemaRef.getASTContext();
1491
1492 // If both types have the same signedness, use the higher ranked type.
1493 if (LHSSigned == RHSSigned) {
1494 if (IsCompAssign || IntOrder >= 0)
1495 return castElement<CK_IntegralCast>(SemaRef, RHS, LHSType);
1496
1497 return castElement<CK_IntegralCast>(SemaRef, LHS, RHSType);
1498 }
1499
1500 // If the unsigned type has greater than or equal rank of the signed type, use
1501 // the unsigned type.
1502 if (IntOrder != (LHSSigned ? 1 : -1)) {
1503 if (IsCompAssign || RHSSigned)
1504 return castElement<CK_IntegralCast>(SemaRef, RHS, LHSType);
1505 return castElement<CK_IntegralCast>(SemaRef, LHS, RHSType);
1506 }
1507
1508 // At this point the signed type has higher rank than the unsigned type, which
1509 // means it will be the same size or bigger. If the signed type is bigger, it
1510 // can represent all the values of the unsigned type, so select it.
1511 if (Ctx.getIntWidth(LElTy) != Ctx.getIntWidth(RElTy)) {
1512 if (IsCompAssign || LHSSigned)
1513 return castElement<CK_IntegralCast>(SemaRef, RHS, LHSType);
1514 return castElement<CK_IntegralCast>(SemaRef, LHS, RHSType);
1515 }
1516
1517 // This is a bit of an odd duck case in HLSL. It shouldn't happen, but can due
1518 // to C/C++ leaking through. The place this happens today is long vs long
1519 // long. When arguments are vector<unsigned long, N> and vector<long long, N>,
1520 // the long long has higher rank than long even though they are the same size.
1521
1522 // If this is a compound assignment cast the right hand side to the left hand
1523 // side's type.
1524 if (IsCompAssign)
1525 return castElement<CK_IntegralCast>(SemaRef, RHS, LHSType);
1526
1527 // If this isn't a compound assignment we convert to unsigned long long.
1528 QualType ElTy = Ctx.getCorrespondingUnsignedType(LHSSigned ? LElTy : RElTy);
1529 QualType NewTy = Ctx.getExtVectorType(
1530 ElTy, RHSType->castAs<VectorType>()->getNumElements());
1531 (void)castElement<CK_IntegralCast>(SemaRef, RHS, NewTy);
1532
1533 return castElement<CK_IntegralCast>(SemaRef, LHS, NewTy);
1534}
1535
1537 QualType SrcTy) {
1538 if (DestTy->isRealFloatingType() && SrcTy->isRealFloatingType())
1539 return CK_FloatingCast;
1540 if (DestTy->isIntegralType(Ctx) && SrcTy->isIntegralType(Ctx))
1541 return CK_IntegralCast;
1542 if (DestTy->isRealFloatingType())
1543 return CK_IntegralToFloating;
1544 assert(SrcTy->isRealFloatingType() && DestTy->isIntegralType(Ctx));
1545 return CK_FloatingToIntegral;
1546}
1547
1549 QualType LHSType,
1550 QualType RHSType,
1551 bool IsCompAssign) {
1552 const auto *LVecTy = LHSType->getAs<VectorType>();
1553 const auto *RVecTy = RHSType->getAs<VectorType>();
1554 auto &Ctx = getASTContext();
1555
1556 // If the LHS is not a vector and this is a compound assignment, we truncate
1557 // the argument to a scalar then convert it to the LHS's type.
1558 if (!LVecTy && IsCompAssign) {
1559 QualType RElTy = RHSType->castAs<VectorType>()->getElementType();
1560 RHS = SemaRef.ImpCastExprToType(RHS.get(), RElTy, CK_HLSLVectorTruncation);
1561 RHSType = RHS.get()->getType();
1562 if (Ctx.hasSameUnqualifiedType(LHSType, RHSType))
1563 return LHSType;
1564 RHS = SemaRef.ImpCastExprToType(RHS.get(), LHSType,
1565 getScalarCastKind(Ctx, LHSType, RHSType));
1566 return LHSType;
1567 }
1568
1569 unsigned EndSz = std::numeric_limits<unsigned>::max();
1570 unsigned LSz = 0;
1571 if (LVecTy)
1572 LSz = EndSz = LVecTy->getNumElements();
1573 if (RVecTy)
1574 EndSz = std::min(RVecTy->getNumElements(), EndSz);
1575 assert(EndSz != std::numeric_limits<unsigned>::max() &&
1576 "one of the above should have had a value");
1577
1578 // In a compound assignment, the left operand does not change type, the right
1579 // operand is converted to the type of the left operand.
1580 if (IsCompAssign && LSz != EndSz) {
1581 Diag(LHS.get()->getBeginLoc(),
1582 diag::err_hlsl_vector_compound_assignment_truncation)
1583 << LHSType << RHSType;
1584 return QualType();
1585 }
1586
1587 if (RVecTy && RVecTy->getNumElements() > EndSz)
1588 castVector<CK_HLSLVectorTruncation>(SemaRef, RHS, RHSType, EndSz);
1589 if (!IsCompAssign && LVecTy && LVecTy->getNumElements() > EndSz)
1590 castVector<CK_HLSLVectorTruncation>(SemaRef, LHS, LHSType, EndSz);
1591
1592 if (!RVecTy)
1593 castVector<CK_VectorSplat>(SemaRef, RHS, RHSType, EndSz);
1594 if (!IsCompAssign && !LVecTy)
1595 castVector<CK_VectorSplat>(SemaRef, LHS, LHSType, EndSz);
1596
1597 // If we're at the same type after resizing we can stop here.
1598 if (Ctx.hasSameUnqualifiedType(LHSType, RHSType))
1599 return Ctx.getCommonSugaredType(LHSType, RHSType);
1600
1601 QualType LElTy = LHSType->castAs<VectorType>()->getElementType();
1602 QualType RElTy = RHSType->castAs<VectorType>()->getElementType();
1603
1604 // Handle conversion for floating point vectors.
1605 if (LElTy->isRealFloatingType() || RElTy->isRealFloatingType())
1606 return handleFloatVectorBinOpConversion(SemaRef, LHS, RHS, LHSType, RHSType,
1607 LElTy, RElTy, IsCompAssign);
1608
1609 assert(LElTy->isIntegralType(Ctx) && RElTy->isIntegralType(Ctx) &&
1610 "HLSL Vectors can only contain integer or floating point types");
1611 return handleIntegerVectorBinOpConversion(SemaRef, LHS, RHS, LHSType, RHSType,
1612 LElTy, RElTy, IsCompAssign);
1613}
1614
1616 BinaryOperatorKind Opc) {
1617 assert((Opc == BO_LOr || Opc == BO_LAnd) &&
1618 "Called with non-logical operator");
1620 llvm::raw_svector_ostream OS(Buff);
1621 PrintingPolicy PP(SemaRef.getLangOpts());
1622 StringRef NewFnName = Opc == BO_LOr ? "or" : "and";
1623 OS << NewFnName << "(";
1624 LHS->printPretty(OS, nullptr, PP);
1625 OS << ", ";
1626 RHS->printPretty(OS, nullptr, PP);
1627 OS << ")";
1628 SourceRange FullRange = SourceRange(LHS->getBeginLoc(), RHS->getEndLoc());
1629 SemaRef.Diag(LHS->getBeginLoc(), diag::note_function_suggestion)
1630 << NewFnName << FixItHint::CreateReplacement(FullRange, OS.str());
1631}
1632
1633std::pair<IdentifierInfo *, bool>
1635 llvm::hash_code Hash = llvm::hash_value(Signature);
1636 std::string IdStr = "__hlsl_rootsig_decl_" + std::to_string(Hash);
1637 IdentifierInfo *DeclIdent = &(getASTContext().Idents.get(IdStr));
1638
1639 // Check if we have already found a decl of the same name.
1640 LookupResult R(SemaRef, DeclIdent, SourceLocation(),
1642 bool Found = SemaRef.LookupQualifiedName(R, SemaRef.CurContext);
1643 return {DeclIdent, Found};
1644}
1645
1647 SourceLocation Loc, IdentifierInfo *DeclIdent,
1649
1650 if (handleRootSignatureElements(RootElements))
1651 return;
1652
1654 for (auto &RootSigElement : RootElements)
1655 Elements.push_back(RootSigElement.getElement());
1656
1657 auto *SignatureDecl = HLSLRootSignatureDecl::Create(
1658 SemaRef.getASTContext(), /*DeclContext=*/SemaRef.CurContext, Loc,
1659 DeclIdent, SemaRef.getLangOpts().HLSLRootSigVer, Elements);
1660
1661 SignatureDecl->setImplicit();
1662 SemaRef.PushOnScopeChains(SignatureDecl, SemaRef.getCurScope());
1663}
1664
1667 if (RootSigOverrideIdent) {
1668 LookupResult R(SemaRef, RootSigOverrideIdent, SourceLocation(),
1670 if (SemaRef.LookupQualifiedName(R, DC))
1671 return dyn_cast<HLSLRootSignatureDecl>(R.getFoundDecl());
1672 }
1673
1674 return nullptr;
1675}
1676
1677namespace {
1678
1679struct PerVisibilityBindingChecker {
1680 SemaHLSL *S;
1681 // We need one builder per `llvm::dxbc::ShaderVisibility` value.
1682 std::array<llvm::hlsl::BindingInfoBuilder, 8> Builders;
1683
1684 struct ElemInfo {
1685 const hlsl::RootSignatureElement *Elem;
1686 llvm::dxbc::ShaderVisibility Vis;
1687 bool Diagnosed;
1688 };
1689 llvm::SmallVector<ElemInfo> ElemInfoMap;
1690
1691 PerVisibilityBindingChecker(SemaHLSL *S) : S(S) {}
1692
1693 void trackBinding(llvm::dxbc::ShaderVisibility Visibility,
1694 llvm::dxil::ResourceClass RC, uint32_t Space,
1695 uint32_t LowerBound, uint32_t UpperBound,
1696 const hlsl::RootSignatureElement *Elem) {
1697 uint32_t BuilderIndex = llvm::to_underlying(Visibility);
1698 assert(BuilderIndex < Builders.size() &&
1699 "Not enough builders for visibility type");
1700 Builders[BuilderIndex].trackBinding(RC, Space, LowerBound, UpperBound,
1701 static_cast<const void *>(Elem));
1702
1703 static_assert(llvm::to_underlying(llvm::dxbc::ShaderVisibility::All) == 0,
1704 "'All' visibility must come first");
1705 if (Visibility == llvm::dxbc::ShaderVisibility::All)
1706 for (size_t I = 1, E = Builders.size(); I < E; ++I)
1707 Builders[I].trackBinding(RC, Space, LowerBound, UpperBound,
1708 static_cast<const void *>(Elem));
1709
1710 ElemInfoMap.push_back({Elem, Visibility, false});
1711 }
1712
1713 ElemInfo &getInfo(const hlsl::RootSignatureElement *Elem) {
1714 auto It = llvm::lower_bound(
1715 ElemInfoMap, Elem,
1716 [](const auto &LHS, const auto &RHS) { return LHS.Elem < RHS; });
1717 assert(It->Elem == Elem && "Element not in map");
1718 return *It;
1719 }
1720
1721 bool checkOverlap() {
1722 llvm::sort(ElemInfoMap, [](const auto &LHS, const auto &RHS) {
1723 return LHS.Elem < RHS.Elem;
1724 });
1725
1726 bool HadOverlap = false;
1727
1728 using llvm::hlsl::BindingInfoBuilder;
1729 auto ReportOverlap = [this,
1730 &HadOverlap](const BindingInfoBuilder &Builder,
1731 const llvm::hlsl::Binding &Reported) {
1732 HadOverlap = true;
1733
1734 const auto *Elem =
1735 static_cast<const hlsl::RootSignatureElement *>(Reported.Cookie);
1736 const llvm::hlsl::Binding &Previous = Builder.findOverlapping(Reported);
1737 const auto *PrevElem =
1738 static_cast<const hlsl::RootSignatureElement *>(Previous.Cookie);
1739
1740 ElemInfo &Info = getInfo(Elem);
1741 // We will have already diagnosed this binding if there's overlap in the
1742 // "All" visibility as well as any particular visibility.
1743 if (Info.Diagnosed)
1744 return;
1745 Info.Diagnosed = true;
1746
1747 ElemInfo &PrevInfo = getInfo(PrevElem);
1748 llvm::dxbc::ShaderVisibility CommonVis =
1749 Info.Vis == llvm::dxbc::ShaderVisibility::All ? PrevInfo.Vis
1750 : Info.Vis;
1751
1752 this->S->Diag(Elem->getLocation(), diag::err_hlsl_resource_range_overlap)
1753 << llvm::to_underlying(Reported.RC) << Reported.LowerBound
1754 << Reported.isUnbounded() << Reported.UpperBound
1755 << llvm::to_underlying(Previous.RC) << Previous.LowerBound
1756 << Previous.isUnbounded() << Previous.UpperBound << Reported.Space
1757 << CommonVis;
1758
1759 this->S->Diag(PrevElem->getLocation(),
1760 diag::note_hlsl_resource_range_here);
1761 };
1762
1763 for (BindingInfoBuilder &Builder : Builders)
1764 Builder.calculateBindingInfo(ReportOverlap);
1765
1766 return HadOverlap;
1767 }
1768};
1769
1770static CXXMethodDecl *lookupMethod(Sema &S, CXXRecordDecl *RecordDecl,
1771 StringRef Name, SourceLocation Loc) {
1772 DeclarationName DeclName(&S.getASTContext().Idents.get(Name));
1773 LookupResult Result(S, DeclName, Loc, Sema::LookupMemberName);
1774 if (!S.LookupQualifiedName(Result, static_cast<DeclContext *>(RecordDecl)))
1775 return nullptr;
1776 return cast<CXXMethodDecl>(Result.getFoundDecl());
1777}
1778
1779} // end anonymous namespace
1780
1783 // Define some common error handling functions
1784 bool HadError = false;
1785 auto ReportError = [this, &HadError](SourceLocation Loc, uint32_t LowerBound,
1786 uint32_t UpperBound) {
1787 HadError = true;
1788 this->Diag(Loc, diag::err_hlsl_invalid_rootsig_value)
1789 << LowerBound << UpperBound;
1790 };
1791
1792 auto ReportFloatError = [this, &HadError](SourceLocation Loc,
1793 float LowerBound,
1794 float UpperBound) {
1795 HadError = true;
1796 this->Diag(Loc, diag::err_hlsl_invalid_rootsig_value)
1797 << llvm::formatv("{0:f}", LowerBound).sstr<6>()
1798 << llvm::formatv("{0:f}", UpperBound).sstr<6>();
1799 };
1800
1801 auto VerifyRegister = [ReportError](SourceLocation Loc, uint32_t Register) {
1802 if (!llvm::hlsl::rootsig::verifyRegisterValue(Register))
1803 ReportError(Loc, 0, 0xfffffffe);
1804 };
1805
1806 auto VerifySpace = [ReportError](SourceLocation Loc, uint32_t Space) {
1807 if (!llvm::hlsl::rootsig::verifyRegisterSpace(Space))
1808 ReportError(Loc, 0, 0xffffffef);
1809 };
1810
1811 const uint32_t Version =
1812 llvm::to_underlying(SemaRef.getLangOpts().HLSLRootSigVer);
1813 const uint32_t VersionEnum = Version - 1;
1814 auto ReportFlagError = [this, &HadError, VersionEnum](SourceLocation Loc) {
1815 HadError = true;
1816 this->Diag(Loc, diag::err_hlsl_invalid_rootsig_flag)
1817 << /*version minor*/ VersionEnum;
1818 };
1819
1820 // Iterate through the elements and do basic validations
1821 for (const hlsl::RootSignatureElement &RootSigElem : Elements) {
1822 SourceLocation Loc = RootSigElem.getLocation();
1823 const llvm::hlsl::rootsig::RootElement &Elem = RootSigElem.getElement();
1824 if (const auto *Descriptor =
1825 std::get_if<llvm::hlsl::rootsig::RootDescriptor>(&Elem)) {
1826 VerifyRegister(Loc, Descriptor->Reg.Number);
1827 VerifySpace(Loc, Descriptor->Space);
1828
1829 if (!llvm::hlsl::rootsig::verifyRootDescriptorFlag(Version,
1830 Descriptor->Flags))
1831 ReportFlagError(Loc);
1832 } else if (const auto *Constants =
1833 std::get_if<llvm::hlsl::rootsig::RootConstants>(&Elem)) {
1834 VerifyRegister(Loc, Constants->Reg.Number);
1835 VerifySpace(Loc, Constants->Space);
1836 } else if (const auto *Sampler =
1837 std::get_if<llvm::hlsl::rootsig::StaticSampler>(&Elem)) {
1838 VerifyRegister(Loc, Sampler->Reg.Number);
1839 VerifySpace(Loc, Sampler->Space);
1840
1841 assert(!std::isnan(Sampler->MaxLOD) && !std::isnan(Sampler->MinLOD) &&
1842 "By construction, parseFloatParam can't produce a NaN from a "
1843 "float_literal token");
1844
1845 if (!llvm::hlsl::rootsig::verifyMaxAnisotropy(Sampler->MaxAnisotropy))
1846 ReportError(Loc, 0, 16);
1847 if (!llvm::hlsl::rootsig::verifyMipLODBias(Sampler->MipLODBias))
1848 ReportFloatError(Loc, -16.f, 15.99f);
1849 } else if (const auto *Clause =
1850 std::get_if<llvm::hlsl::rootsig::DescriptorTableClause>(
1851 &Elem)) {
1852 VerifyRegister(Loc, Clause->Reg.Number);
1853 VerifySpace(Loc, Clause->Space);
1854
1855 if (!llvm::hlsl::rootsig::verifyNumDescriptors(Clause->NumDescriptors)) {
1856 // NumDescriptor could techincally be ~0u but that is reserved for
1857 // unbounded, so the diagnostic will not report that as a valid int
1858 // value
1859 ReportError(Loc, 1, 0xfffffffe);
1860 }
1861
1862 if (!llvm::hlsl::rootsig::verifyDescriptorRangeFlag(Version, Clause->Type,
1863 Clause->Flags))
1864 ReportFlagError(Loc);
1865 }
1866 }
1867
1868 PerVisibilityBindingChecker BindingChecker(this);
1869 SmallVector<std::pair<const llvm::hlsl::rootsig::DescriptorTableClause *,
1871 UnboundClauses;
1872
1873 for (const hlsl::RootSignatureElement &RootSigElem : Elements) {
1874 const llvm::hlsl::rootsig::RootElement &Elem = RootSigElem.getElement();
1875 if (const auto *Descriptor =
1876 std::get_if<llvm::hlsl::rootsig::RootDescriptor>(&Elem)) {
1877 uint32_t LowerBound(Descriptor->Reg.Number);
1878 uint32_t UpperBound(LowerBound); // inclusive range
1879
1880 BindingChecker.trackBinding(
1881 Descriptor->Visibility,
1882 static_cast<llvm::dxil::ResourceClass>(Descriptor->Type),
1883 Descriptor->Space, LowerBound, UpperBound, &RootSigElem);
1884 } else if (const auto *Constants =
1885 std::get_if<llvm::hlsl::rootsig::RootConstants>(&Elem)) {
1886 uint32_t LowerBound(Constants->Reg.Number);
1887 uint32_t UpperBound(LowerBound); // inclusive range
1888
1889 BindingChecker.trackBinding(
1890 Constants->Visibility, llvm::dxil::ResourceClass::CBuffer,
1891 Constants->Space, LowerBound, UpperBound, &RootSigElem);
1892 } else if (const auto *Sampler =
1893 std::get_if<llvm::hlsl::rootsig::StaticSampler>(&Elem)) {
1894 uint32_t LowerBound(Sampler->Reg.Number);
1895 uint32_t UpperBound(LowerBound); // inclusive range
1896
1897 BindingChecker.trackBinding(
1898 Sampler->Visibility, llvm::dxil::ResourceClass::Sampler,
1899 Sampler->Space, LowerBound, UpperBound, &RootSigElem);
1900 } else if (const auto *Clause =
1901 std::get_if<llvm::hlsl::rootsig::DescriptorTableClause>(
1902 &Elem)) {
1903 // We'll process these once we see the table element.
1904 UnboundClauses.emplace_back(Clause, &RootSigElem);
1905 } else if (const auto *Table =
1906 std::get_if<llvm::hlsl::rootsig::DescriptorTable>(&Elem)) {
1907 assert(UnboundClauses.size() == Table->NumClauses &&
1908 "Number of unbound elements must match the number of clauses");
1909 bool HasAnySampler = false;
1910 bool HasAnyNonSampler = false;
1911 uint64_t Offset = 0;
1912 bool IsPrevUnbound = false;
1913 for (const auto &[Clause, ClauseElem] : UnboundClauses) {
1914 SourceLocation Loc = ClauseElem->getLocation();
1915 if (Clause->Type == llvm::dxil::ResourceClass::Sampler)
1916 HasAnySampler = true;
1917 else
1918 HasAnyNonSampler = true;
1919
1920 if (HasAnySampler && HasAnyNonSampler)
1921 Diag(Loc, diag::err_hlsl_invalid_mixed_resources);
1922
1923 // Relevant error will have already been reported above and needs to be
1924 // fixed before we can conduct further analysis, so shortcut error
1925 // return
1926 if (Clause->NumDescriptors == 0)
1927 return true;
1928
1929 bool IsAppending =
1930 Clause->Offset == llvm::hlsl::rootsig::DescriptorTableOffsetAppend;
1931 if (!IsAppending)
1932 Offset = Clause->Offset;
1933
1934 uint64_t RangeBound = llvm::hlsl::rootsig::computeRangeBound(
1935 Offset, Clause->NumDescriptors);
1936
1937 if (IsPrevUnbound && IsAppending)
1938 Diag(Loc, diag::err_hlsl_appending_onto_unbound);
1939 else if (!llvm::hlsl::rootsig::verifyNoOverflowedOffset(RangeBound))
1940 Diag(Loc, diag::err_hlsl_offset_overflow) << Offset << RangeBound;
1941
1942 // Update offset to be 1 past this range's bound
1943 Offset = RangeBound + 1;
1944 IsPrevUnbound = Clause->NumDescriptors ==
1945 llvm::hlsl::rootsig::NumDescriptorsUnbounded;
1946
1947 // Compute the register bounds and track resource binding
1948 uint32_t LowerBound(Clause->Reg.Number);
1949 uint32_t UpperBound = llvm::hlsl::rootsig::computeRangeBound(
1950 LowerBound, Clause->NumDescriptors);
1951
1952 BindingChecker.trackBinding(
1953 Table->Visibility,
1954 static_cast<llvm::dxil::ResourceClass>(Clause->Type), Clause->Space,
1955 LowerBound, UpperBound, ClauseElem);
1956 }
1957 UnboundClauses.clear();
1958 }
1959 }
1960
1961 return BindingChecker.checkOverlap();
1962}
1963
1965 if (AL.getNumArgs() != 1) {
1966 Diag(AL.getLoc(), diag::err_attribute_wrong_number_arguments) << AL << 1;
1967 return;
1968 }
1969
1971 if (auto *RS = D->getAttr<RootSignatureAttr>()) {
1972 if (RS->getSignatureIdent() != Ident) {
1973 Diag(AL.getLoc(), diag::err_disallowed_duplicate_attribute) << RS;
1974 return;
1975 }
1976
1977 Diag(AL.getLoc(), diag::warn_duplicate_attribute_exact) << RS;
1978 return;
1979 }
1980
1982 if (SemaRef.LookupQualifiedName(R, D->getDeclContext()))
1983 if (auto *SignatureDecl =
1984 dyn_cast<HLSLRootSignatureDecl>(R.getFoundDecl())) {
1985 D->addAttr(::new (getASTContext()) RootSignatureAttr(
1986 getASTContext(), AL, Ident, SignatureDecl));
1987 }
1988}
1989
1991 llvm::VersionTuple SMVersion =
1992 getASTContext().getTargetInfo().getTriple().getOSVersion();
1993 bool IsDXIL = getASTContext().getTargetInfo().getTriple().getArch() ==
1994 llvm::Triple::dxil;
1995
1996 uint32_t ZMax = 1024;
1997 uint32_t ThreadMax = 1024;
1998 if (IsDXIL && SMVersion.getMajor() <= 4) {
1999 ZMax = 1;
2000 ThreadMax = 768;
2001 } else if (IsDXIL && SMVersion.getMajor() == 5) {
2002 ZMax = 64;
2003 ThreadMax = 1024;
2004 }
2005
2006 uint32_t X;
2007 if (!SemaRef.checkUInt32Argument(AL, AL.getArgAsExpr(0), X))
2008 return;
2009 if (X > 1024) {
2010 Diag(AL.getArgAsExpr(0)->getExprLoc(),
2011 diag::err_hlsl_numthreads_argument_oor)
2012 << 0 << 1024;
2013 return;
2014 }
2015 uint32_t Y;
2016 if (!SemaRef.checkUInt32Argument(AL, AL.getArgAsExpr(1), Y))
2017 return;
2018 if (Y > 1024) {
2019 Diag(AL.getArgAsExpr(1)->getExprLoc(),
2020 diag::err_hlsl_numthreads_argument_oor)
2021 << 1 << 1024;
2022 return;
2023 }
2024 uint32_t Z;
2025 if (!SemaRef.checkUInt32Argument(AL, AL.getArgAsExpr(2), Z))
2026 return;
2027 if (Z > ZMax) {
2028 SemaRef.Diag(AL.getArgAsExpr(2)->getExprLoc(),
2029 diag::err_hlsl_numthreads_argument_oor)
2030 << 2 << ZMax;
2031 return;
2032 }
2033
2034 if (X * Y * Z > ThreadMax) {
2035 Diag(AL.getLoc(), diag::err_hlsl_numthreads_invalid) << ThreadMax;
2036 return;
2037 }
2038
2039 HLSLNumThreadsAttr *NewAttr = mergeNumThreadsAttr(D, AL, X, Y, Z);
2040 if (NewAttr)
2041 D->addAttr(NewAttr);
2042}
2043
2044static bool isValidWaveSizeValue(unsigned Value) {
2045 return llvm::isPowerOf2_32(Value) && Value >= 4 && Value <= 128;
2046}
2047
2049 // validate that the wavesize argument is a power of 2 between 4 and 128
2050 // inclusive
2051 unsigned SpelledArgsCount = AL.getNumArgs();
2052 if (SpelledArgsCount == 0 || SpelledArgsCount > 3)
2053 return;
2054
2055 uint32_t Min;
2056 if (!SemaRef.checkUInt32Argument(AL, AL.getArgAsExpr(0), Min))
2057 return;
2058
2059 uint32_t Max = 0;
2060 if (SpelledArgsCount > 1 &&
2061 !SemaRef.checkUInt32Argument(AL, AL.getArgAsExpr(1), Max))
2062 return;
2063
2064 uint32_t Preferred = 0;
2065 if (SpelledArgsCount > 2 &&
2066 !SemaRef.checkUInt32Argument(AL, AL.getArgAsExpr(2), Preferred))
2067 return;
2068
2069 if (SpelledArgsCount > 2) {
2070 if (!isValidWaveSizeValue(Preferred)) {
2071 Diag(AL.getArgAsExpr(2)->getExprLoc(),
2072 diag::err_attribute_power_of_two_in_range)
2073 << AL << llvm::dxil::MinWaveSize << llvm::dxil::MaxWaveSize
2074 << Preferred;
2075 return;
2076 }
2077 // Preferred not in range.
2078 if (Preferred < Min || Preferred > Max) {
2079 Diag(AL.getArgAsExpr(2)->getExprLoc(),
2080 diag::err_attribute_power_of_two_in_range)
2081 << AL << Min << Max << Preferred;
2082 return;
2083 }
2084 } else if (SpelledArgsCount > 1) {
2085 if (!isValidWaveSizeValue(Max)) {
2086 Diag(AL.getArgAsExpr(1)->getExprLoc(),
2087 diag::err_attribute_power_of_two_in_range)
2088 << AL << llvm::dxil::MinWaveSize << llvm::dxil::MaxWaveSize << Max;
2089 return;
2090 }
2091 if (Max < Min) {
2092 Diag(AL.getLoc(), diag::err_attribute_argument_invalid) << AL << 1;
2093 return;
2094 } else if (Max == Min) {
2095 Diag(AL.getLoc(), diag::warn_attr_min_eq_max) << AL;
2096 }
2097 } else {
2098 if (!isValidWaveSizeValue(Min)) {
2099 Diag(AL.getArgAsExpr(0)->getExprLoc(),
2100 diag::err_attribute_power_of_two_in_range)
2101 << AL << llvm::dxil::MinWaveSize << llvm::dxil::MaxWaveSize << Min;
2102 return;
2103 }
2104 }
2105
2106 HLSLWaveSizeAttr *NewAttr =
2107 mergeWaveSizeAttr(D, AL, Min, Max, Preferred, SpelledArgsCount);
2108 if (NewAttr)
2109 D->addAttr(NewAttr);
2110}
2111
2113 uint32_t ID;
2114 if (!SemaRef.checkUInt32Argument(AL, AL.getArgAsExpr(0), ID))
2115 return;
2116 D->addAttr(::new (getASTContext())
2117 HLSLVkExtBuiltinInputAttr(getASTContext(), AL, ID));
2118}
2119
2121 uint32_t ID;
2122 if (!SemaRef.checkUInt32Argument(AL, AL.getArgAsExpr(0), ID))
2123 return;
2124 D->addAttr(::new (getASTContext())
2125 HLSLVkExtBuiltinOutputAttr(getASTContext(), AL, ID));
2126}
2127
2129 D->addAttr(::new (getASTContext())
2130 HLSLVkPushConstantAttr(getASTContext(), AL));
2131}
2132
2134 uint32_t Id;
2135 if (!SemaRef.checkUInt32Argument(AL, AL.getArgAsExpr(0), Id))
2136 return;
2137 HLSLVkConstantIdAttr *NewAttr = mergeVkConstantIdAttr(D, AL, Id);
2138 if (NewAttr)
2139 D->addAttr(NewAttr);
2140}
2141
2143 uint32_t Binding = 0;
2144 if (!SemaRef.checkUInt32Argument(AL, AL.getArgAsExpr(0), Binding))
2145 return;
2146 uint32_t Set = 0;
2147 if (AL.getNumArgs() > 1 &&
2148 !SemaRef.checkUInt32Argument(AL, AL.getArgAsExpr(1), Set))
2149 return;
2150
2151 D->addAttr(::new (getASTContext())
2152 HLSLVkBindingAttr(getASTContext(), AL, Binding, Set));
2153}
2154
2156 uint32_t Location;
2157 if (!SemaRef.checkUInt32Argument(AL, AL.getArgAsExpr(0), Location))
2158 return;
2159
2160 D->addAttr(::new (getASTContext())
2161 HLSLVkLocationAttr(getASTContext(), AL, Location));
2162}
2163
2165 uint32_t IndexValue(0), ExplicitIndex(0);
2166 if (!SemaRef.checkUInt32Argument(AL, AL.getArgAsExpr(0), IndexValue) ||
2167 !SemaRef.checkUInt32Argument(AL, AL.getArgAsExpr(1), ExplicitIndex)) {
2168 assert(0 && "HLSLUnparsedSemantic is expected to have 2 int arguments.");
2169 }
2170 assert(IndexValue > 0 ? ExplicitIndex : true);
2171
2172 SemanticKind Kind = llvm::hlsl::getSemanticKind(AL.getAttrName()->getName());
2173 if (Kind == SemanticKind::Invalid) {
2174 Diag(AL.getLoc(), diag::err_hlsl_unknown_semantic) << AL;
2175 return;
2176 }
2177
2178 switch (Kind) {
2179 // FIXME: These semantics do not yet have CodeGen support.
2180 case SemanticKind::RenderTargetArrayIndex:
2181 case SemanticKind::ViewPortArrayIndex:
2182 case SemanticKind::ClipDistance:
2183 case SemanticKind::CullDistance:
2184 case SemanticKind::OutputControlPointID:
2185 case SemanticKind::DomainLocation:
2186 case SemanticKind::PrimitiveID:
2187 case SemanticKind::GSInstanceID:
2188 case SemanticKind::SampleIndex:
2189 case SemanticKind::IsFrontFace:
2190 case SemanticKind::Coverage:
2191 case SemanticKind::InnerCoverage:
2192 case SemanticKind::Depth:
2193 case SemanticKind::DepthLessEqual:
2194 case SemanticKind::DepthGreaterEqual:
2195 case SemanticKind::StencilRef:
2196 case SemanticKind::TessFactor:
2197 case SemanticKind::InsideTessFactor:
2198 case SemanticKind::ViewID:
2199 case SemanticKind::Barycentrics:
2200 case SemanticKind::ShadingRate:
2201 case SemanticKind::CullPrimitive:
2202 Diag(AL.getLoc(), diag::err_hlsl_unknown_semantic) << AL;
2203 return;
2204 default:
2205 break;
2206 }
2207
2208 D->addAttr(HLSLParsedSemanticAttr::Create(
2209 getASTContext(), AL.getAttrName()->getName(), IndexValue, AL));
2210}
2211
2214 Diag(AL.getLoc(), diag::err_hlsl_attr_invalid_ast_node)
2215 << AL << "shader constant in a constant buffer";
2216 return;
2217 }
2218
2219 uint32_t SubComponent;
2220 if (!SemaRef.checkUInt32Argument(AL, AL.getArgAsExpr(0), SubComponent))
2221 return;
2222 uint32_t Component;
2223 if (!SemaRef.checkUInt32Argument(AL, AL.getArgAsExpr(1), Component))
2224 return;
2225
2226 QualType T = cast<VarDecl>(D)->getType().getCanonicalType();
2227 // Check if T is an array or struct type.
2228 // TODO: mark matrix type as aggregate type.
2229 bool IsAggregateTy = (T->isArrayType() || T->isStructureType());
2230
2231 // Check Component is valid for T.
2232 if (Component) {
2233 unsigned Size = getASTContext().getTypeSize(T);
2234 if (IsAggregateTy) {
2235 Diag(AL.getLoc(), diag::err_hlsl_invalid_register_or_packoffset);
2236 return;
2237 } else {
2238 // Make sure Component + sizeof(T) <= 4.
2239 if ((Component * 32 + Size) > 128) {
2240 Diag(AL.getLoc(), diag::err_hlsl_packoffset_cross_reg_boundary);
2241 return;
2242 }
2243 QualType EltTy = T;
2244 if (const auto *VT = T->getAs<VectorType>())
2245 EltTy = VT->getElementType();
2246 unsigned Align = getASTContext().getTypeAlign(EltTy);
2247 if (Align > 32 && Component == 1) {
2248 // NOTE: Component 3 will hit err_hlsl_packoffset_cross_reg_boundary.
2249 // So we only need to check Component 1 here.
2250 Diag(AL.getLoc(), diag::err_hlsl_packoffset_alignment_mismatch)
2251 << Align << EltTy;
2252 return;
2253 }
2254 }
2255 }
2256
2257 D->addAttr(::new (getASTContext()) HLSLPackOffsetAttr(
2258 getASTContext(), AL, SubComponent, Component));
2259}
2260
2262 StringRef Str;
2263 SourceLocation ArgLoc;
2264 if (!SemaRef.checkStringLiteralArgumentAttr(AL, 0, Str, &ArgLoc))
2265 return;
2266
2267 llvm::Triple::EnvironmentType ShaderType;
2268 if (!HLSLShaderAttr::ConvertStrToEnvironmentType(Str, ShaderType)) {
2269 Diag(AL.getLoc(), diag::warn_attribute_type_not_supported)
2270 << AL << Str << ArgLoc;
2271 return;
2272 }
2273
2274 // FIXME: check function match the shader stage.
2275
2276 HLSLShaderAttr *NewAttr = mergeShaderAttr(D, AL, ShaderType);
2277 if (NewAttr)
2278 D->addAttr(NewAttr);
2279}
2280
2282 Sema &S, QualType Wrapped, ArrayRef<const Attr *> AttrList,
2283 QualType &ResType, HLSLAttributedResourceLocInfo *LocInfo,
2284 Expr *SampleCountExpr) {
2285 assert(AttrList.size() && "expected list of resource attributes");
2286
2287 QualType ContainedTy = QualType();
2288 TypeSourceInfo *ContainedTyInfo = nullptr;
2289 SourceLocation LocBegin = AttrList[0]->getRange().getBegin();
2290 SourceLocation LocEnd = AttrList[0]->getRange().getEnd();
2291
2292 HLSLAttributedResourceType::Attributes ResAttrs;
2293
2294 bool HasResourceClass = false;
2295 bool HasResourceDimension = false;
2296 for (const Attr *A : AttrList) {
2297 if (!A)
2298 continue;
2299 LocEnd = A->getRange().getEnd();
2300 switch (A->getKind()) {
2301 case attr::HLSLResourceClass: {
2302 ResourceClass RC = cast<HLSLResourceClassAttr>(A)->getResourceClass();
2303 if (HasResourceClass) {
2304 S.Diag(A->getLocation(), ResAttrs.ResourceClass == RC
2305 ? diag::warn_duplicate_attribute_exact
2306 : diag::warn_duplicate_attribute)
2307 << A;
2308 return false;
2309 }
2310 ResAttrs.ResourceClass = RC;
2311 HasResourceClass = true;
2312 break;
2313 }
2314 case attr::HLSLResourceDimension: {
2315 llvm::dxil::ResourceDimension RD =
2316 cast<HLSLResourceDimensionAttr>(A)->getDimension();
2317 if (HasResourceDimension) {
2318 S.Diag(A->getLocation(), ResAttrs.ResourceDimension == RD
2319 ? diag::warn_duplicate_attribute_exact
2320 : diag::warn_duplicate_attribute)
2321 << A;
2322 return false;
2323 }
2324 ResAttrs.ResourceDimension = RD;
2325 HasResourceDimension = true;
2326 break;
2327 }
2328 case attr::HLSLIsROV:
2329 if (ResAttrs.IsROV) {
2330 S.Diag(A->getLocation(), diag::warn_duplicate_attribute_exact) << A;
2331 return false;
2332 }
2333 ResAttrs.IsROV = true;
2334 break;
2335 case attr::HLSLRawBuffer:
2336 if (ResAttrs.RawBuffer) {
2337 S.Diag(A->getLocation(), diag::warn_duplicate_attribute_exact) << A;
2338 return false;
2339 }
2340 ResAttrs.RawBuffer = true;
2341 break;
2342 case attr::HLSLIsArray:
2343 if (ResAttrs.IsArray) {
2344 S.Diag(A->getLocation(), diag::warn_duplicate_attribute_exact) << A;
2345 return false;
2346 }
2347 ResAttrs.IsArray = true;
2348 break;
2349 case attr::HLSLIsMultiSampled:
2350 if (ResAttrs.SampleCountExpr) {
2351 S.Diag(A->getLocation(), diag::warn_duplicate_attribute_exact) << A;
2352 return false;
2353 }
2354 // A bare [[hlsl::is_ms]] carries no count, so default it to 0, the same
2355 // value Texture2DMS<T> gets from its template parameter.
2356 ResAttrs.SampleCountExpr =
2357 SampleCountExpr
2358 ? SampleCountExpr
2359 : IntegerLiteral::Create(S.Context, llvm::APInt(32, 0),
2360 S.Context.IntTy, A->getLocation());
2361 break;
2362 case attr::HLSLIsCounter:
2363 if (ResAttrs.IsCounter) {
2364 S.Diag(A->getLocation(), diag::warn_duplicate_attribute_exact) << A;
2365 return false;
2366 }
2367 ResAttrs.IsCounter = true;
2368 break;
2369 case attr::HLSLContainedType: {
2370 const HLSLContainedTypeAttr *CTAttr = cast<HLSLContainedTypeAttr>(A);
2371 QualType Ty = CTAttr->getType();
2372 if (!ContainedTy.isNull()) {
2373 S.Diag(A->getLocation(), ContainedTy == Ty
2374 ? diag::warn_duplicate_attribute_exact
2375 : diag::warn_duplicate_attribute)
2376 << A;
2377 return false;
2378 }
2379 ContainedTy = Ty;
2380 ContainedTyInfo = CTAttr->getTypeLoc();
2381 break;
2382 }
2383 default:
2384 llvm_unreachable("unhandled resource attribute type");
2385 }
2386 }
2387
2388 if (!HasResourceClass) {
2389 S.Diag(AttrList.back()->getRange().getEnd(),
2390 diag::err_hlsl_missing_resource_class);
2391 return false;
2392 }
2393
2395 Wrapped, ContainedTy, ResAttrs);
2396
2397 if (LocInfo && ContainedTyInfo) {
2398 LocInfo->Range = SourceRange(LocBegin, LocEnd);
2399 LocInfo->ContainedTyInfo = ContainedTyInfo;
2400 }
2401 return true;
2402}
2403
2404// Validates and creates an HLSL attribute that is applied as type attribute on
2405// HLSL resource. The attributes are collected in HLSLResourcesTypeAttrs and at
2406// the end of the declaration they are applied to the declaration type by
2407// wrapping it in HLSLAttributedResourceType.
2409 // only allow resource type attributes on intangible types
2410 if (!T->isHLSLResourceType()) {
2411 Diag(AL.getLoc(), diag::err_hlsl_attribute_needs_intangible_type)
2412 << AL << getASTContext().HLSLResourceTy;
2413 return false;
2414 }
2415
2416 // validate number of arguments
2417 if (!AL.checkExactlyNumArgs(SemaRef, AL.getMinArgs()))
2418 return false;
2419
2420 Attr *A = nullptr;
2421
2425 {
2426 AttributeCommonInfo::AS_CXX11, 0, false /*IsAlignas*/,
2427 false /*IsRegularKeywordAttribute*/
2428 });
2429
2430 switch (AL.getKind()) {
2431 case ParsedAttr::AT_HLSLResourceClass: {
2432 StringRef Identifier;
2433 SourceLocation ArgLoc;
2434 if (!SemaRef.checkStringLiteralArgumentAttr(AL, 0, Identifier, &ArgLoc))
2435 return false;
2436
2437 // Validate resource class value
2438 ResourceClass RC;
2439 if (!HLSLResourceClassAttr::ConvertStrToResourceClass(Identifier, RC)) {
2440 Diag(ArgLoc, diag::warn_attribute_type_not_supported)
2441 << "ResourceClass" << Identifier;
2442 return false;
2443 }
2444 A = HLSLResourceClassAttr::Create(getASTContext(), RC, ACI);
2445 break;
2446 }
2447
2448 case ParsedAttr::AT_HLSLResourceDimension: {
2449 StringRef Identifier;
2450 SourceLocation ArgLoc;
2451 if (!SemaRef.checkStringLiteralArgumentAttr(AL, 0, Identifier, &ArgLoc))
2452 return false;
2453
2454 // Validate resource dimension value
2455 llvm::dxil::ResourceDimension RD;
2456 if (!HLSLResourceDimensionAttr::ConvertStrToResourceDimension(Identifier,
2457 RD)) {
2458 Diag(ArgLoc, diag::warn_attribute_type_not_supported)
2459 << "ResourceDimension" << Identifier;
2460 return false;
2461 }
2462 A = HLSLResourceDimensionAttr::Create(getASTContext(), RD, ACI);
2463 break;
2464 }
2465
2466 case ParsedAttr::AT_HLSLIsROV:
2467 A = HLSLIsROVAttr::Create(getASTContext(), ACI);
2468 break;
2469
2470 case ParsedAttr::AT_HLSLRawBuffer:
2471 A = HLSLRawBufferAttr::Create(getASTContext(), ACI);
2472 break;
2473
2474 case ParsedAttr::AT_HLSLIsCounter:
2475 A = HLSLIsCounterAttr::Create(getASTContext(), ACI);
2476 break;
2477
2478 case ParsedAttr::AT_HLSLIsArray:
2479 A = HLSLIsArrayAttr::Create(getASTContext(), ACI);
2480 break;
2481
2482 case ParsedAttr::AT_HLSLIsMultiSampled:
2483 A = HLSLIsMultiSampledAttr::Create(getASTContext(), ACI);
2484 break;
2485
2486 case ParsedAttr::AT_HLSLContainedType: {
2487 if (AL.getNumArgs() != 1 && !AL.hasParsedType()) {
2488 Diag(AL.getLoc(), diag::err_attribute_wrong_number_arguments) << AL << 1;
2489 return false;
2490 }
2491
2492 TypeSourceInfo *TSI = nullptr;
2493 QualType QT = SemaRef.GetTypeFromParser(AL.getTypeArg(), &TSI);
2494 assert(TSI && "no type source info for attribute argument");
2495 if (SemaRef.RequireCompleteType(TSI->getTypeLoc().getBeginLoc(), QT,
2496 diag::err_incomplete_type))
2497 return false;
2498 A = HLSLContainedTypeAttr::Create(getASTContext(), TSI, ACI);
2499 break;
2500 }
2501
2502 default:
2503 llvm_unreachable("unhandled HLSL attribute");
2504 }
2505
2506 HLSLResourcesTypeAttrs.emplace_back(A);
2507 return true;
2508}
2509
2510// Combines all resource type attributes and creates HLSLAttributedResourceType.
2512 if (!HLSLResourcesTypeAttrs.size())
2513 return CurrentType;
2514
2515 QualType QT = CurrentType;
2518 HLSLResourcesTypeAttrs, QT, &LocInfo)) {
2519 const HLSLAttributedResourceType *RT =
2521
2522 // Temporarily store TypeLoc information for the new type.
2523 // It will be transferred to HLSLAttributesResourceTypeLoc
2524 // shortly after the type is created by TypeSpecLocFiller which
2525 // will call the TakeLocForHLSLAttribute method below.
2526 LocsForHLSLAttributedResources.insert(std::pair(RT, LocInfo));
2527 }
2528 HLSLResourcesTypeAttrs.clear();
2529 return QT;
2530}
2531
2532// Returns source location for the HLSLAttributedResourceType
2534SemaHLSL::TakeLocForHLSLAttribute(const HLSLAttributedResourceType *RT) {
2535 HLSLAttributedResourceLocInfo LocInfo = {};
2536 auto I = LocsForHLSLAttributedResources.find(RT);
2537 if (I != LocsForHLSLAttributedResources.end()) {
2538 LocInfo = I->second;
2539 LocsForHLSLAttributedResources.erase(I);
2540 return LocInfo;
2541 }
2542 LocInfo.Range = SourceRange();
2543 return LocInfo;
2544}
2545
2546// Walks though the global variable declaration, collects all resource binding
2547// requirements and adds them to Bindings
2548void SemaHLSL::collectResourceBindingsOnUserRecordDecl(const VarDecl *VD,
2549 const RecordType *RT) {
2550 const RecordDecl *RD = RT->getDecl()->getDefinitionOrSelf();
2551 for (FieldDecl *FD : RD->fields()) {
2552 const Type *Ty = FD->getType()->getUnqualifiedDesugaredType();
2553
2554 // Unwrap arrays
2555 // FIXME: Calculate array size while unwrapping
2556 assert(!Ty->isIncompleteArrayType() &&
2557 "incomplete arrays inside user defined types are not supported");
2558 while (Ty->isConstantArrayType()) {
2561 }
2562
2563 if (!Ty->isRecordType())
2564 continue;
2565
2566 if (const HLSLAttributedResourceType *AttrResType =
2567 HLSLAttributedResourceType::findHandleTypeOnResource(Ty)) {
2568 // Add a new DeclBindingInfo to Bindings if it does not already exist
2569 ResourceClass RC = AttrResType->getAttrs().ResourceClass;
2570 DeclBindingInfo *DBI = Bindings.getDeclBindingInfo(VD, RC);
2571 if (!DBI)
2572 Bindings.addDeclBindingInfo(VD, RC);
2573 } else if (const RecordType *RT = dyn_cast<RecordType>(Ty)) {
2574 // Recursively scan embedded struct or class; it would be nice to do this
2575 // without recursion, but tricky to correctly calculate the size of the
2576 // binding, which is something we are probably going to need to do later
2577 // on. Hopefully nesting of structs in structs too many levels is
2578 // unlikely.
2579 collectResourceBindingsOnUserRecordDecl(VD, RT);
2580 }
2581 }
2582}
2583
2584// Diagnose localized register binding errors for a single binding; does not
2585// diagnose resource binding on user record types, that will be done later
2586// in processResourceBindingOnDecl based on the information collected in
2587// collectResourceBindingsOnVarDecl.
2588// Returns false if the register binding is not valid.
2590 Decl *D, RegisterType RegType,
2591 bool SpecifiedSpace) {
2592 int RegTypeNum = static_cast<int>(RegType);
2593
2594 // check if the decl type is groupshared
2595 if (D->hasAttr<HLSLGroupSharedAddressSpaceAttr>()) {
2596 S.Diag(ArgLoc, diag::err_hlsl_binding_type_mismatch) << RegTypeNum;
2597 return false;
2598 }
2599
2600 // Cbuffers and Tbuffers are HLSLBufferDecl types
2601 if (HLSLBufferDecl *CBufferOrTBuffer = dyn_cast<HLSLBufferDecl>(D)) {
2602 ResourceClass RC = CBufferOrTBuffer->isCBuffer() ? ResourceClass::CBuffer
2603 : ResourceClass::SRV;
2604 if (RegType == getRegisterType(RC))
2605 return true;
2606
2607 S.Diag(D->getLocation(), diag::err_hlsl_binding_type_mismatch)
2608 << RegTypeNum;
2609 return false;
2610 }
2611
2612 // Samplers, UAVs, and SRVs are VarDecl types
2613 assert(isa<VarDecl>(D) && "D is expected to be VarDecl or HLSLBufferDecl");
2614 VarDecl *VD = cast<VarDecl>(D);
2615
2616 // Resource
2617 if (const HLSLAttributedResourceType *AttrResType =
2618 HLSLAttributedResourceType::findHandleTypeOnResource(
2619 VD->getType().getTypePtr())) {
2620 if (RegType == getRegisterType(AttrResType))
2621 return true;
2622
2623 S.Diag(D->getLocation(), diag::err_hlsl_binding_type_mismatch)
2624 << RegTypeNum;
2625 return false;
2626 }
2627
2628 const clang::Type *Ty = VD->getType().getTypePtr();
2629 while (Ty->isArrayType())
2631
2632 // Basic types
2633 if (Ty->isArithmeticType() || Ty->isVectorType()) {
2634 bool DeclaredInCOrTBuffer = isa<HLSLBufferDecl>(D->getDeclContext());
2635 if (SpecifiedSpace && !DeclaredInCOrTBuffer)
2636 S.Diag(ArgLoc, diag::err_hlsl_space_on_global_constant);
2637
2638 if (!DeclaredInCOrTBuffer && (Ty->isIntegralType(S.getASTContext()) ||
2639 Ty->isFloatingType() || Ty->isVectorType())) {
2640 // Register annotation on default constant buffer declaration ($Globals)
2641 if (RegType == RegisterType::CBuffer)
2642 S.Diag(ArgLoc, diag::warn_hlsl_deprecated_register_type_b);
2643 else if (RegType != RegisterType::C)
2644 S.Diag(ArgLoc, diag::err_hlsl_binding_type_mismatch) << RegTypeNum;
2645 else
2646 return true;
2647 } else {
2648 if (RegType == RegisterType::C)
2649 S.Diag(ArgLoc, diag::warn_hlsl_register_type_c_packoffset);
2650 else
2651 S.Diag(ArgLoc, diag::err_hlsl_binding_type_mismatch) << RegTypeNum;
2652 }
2653 return false;
2654 }
2655 if (Ty->isRecordType())
2656 // RecordTypes will be diagnosed in processResourceBindingOnDecl
2657 // that is called from ActOnVariableDeclarator
2658 return true;
2659
2660 // Anything else is an error
2661 S.Diag(ArgLoc, diag::err_hlsl_binding_type_mismatch) << RegTypeNum;
2662 return false;
2663}
2664
2666 RegisterType regType) {
2667 // make sure that there are no two register annotations
2668 // applied to the decl with the same register type
2669 bool RegisterTypesDetected[5] = {false};
2670 RegisterTypesDetected[static_cast<int>(regType)] = true;
2671
2672 for (auto it = TheDecl->attr_begin(); it != TheDecl->attr_end(); ++it) {
2673 if (HLSLResourceBindingAttr *attr =
2674 dyn_cast<HLSLResourceBindingAttr>(*it)) {
2675
2676 RegisterType otherRegType = attr->getRegisterType();
2677 if (RegisterTypesDetected[static_cast<int>(otherRegType)]) {
2678 int otherRegTypeNum = static_cast<int>(otherRegType);
2679 S.Diag(TheDecl->getLocation(),
2680 diag::err_hlsl_duplicate_register_annotation)
2681 << otherRegTypeNum;
2682 return false;
2683 }
2684 RegisterTypesDetected[static_cast<int>(otherRegType)] = true;
2685 }
2686 }
2687 return true;
2688}
2689
2691 Decl *D, RegisterType RegType,
2692 bool SpecifiedSpace) {
2693
2694 // exactly one of these two types should be set
2695 assert(((isa<VarDecl>(D) && !isa<HLSLBufferDecl>(D)) ||
2696 (!isa<VarDecl>(D) && isa<HLSLBufferDecl>(D))) &&
2697 "expecting VarDecl or HLSLBufferDecl");
2698
2699 // check if the declaration contains resource matching the register type
2700 if (!DiagnoseLocalRegisterBinding(S, ArgLoc, D, RegType, SpecifiedSpace))
2701 return false;
2702
2703 // next, if multiple register annotations exist, check that none conflict.
2704 return ValidateMultipleRegisterAnnotations(S, D, RegType);
2705}
2706
2707// return false if the slot count exceeds the limit, true otherwise
2708static bool AccumulateHLSLResourceSlots(QualType Ty, uint64_t &StartSlot,
2709 const uint64_t &Limit,
2710 const ResourceClass ResClass,
2711 ASTContext &Ctx,
2712 uint64_t ArrayCount = 1) {
2713 Ty = Ty.getCanonicalType();
2714 const Type *T = Ty.getTypePtr();
2715
2716 // Early exit if already overflowed
2717 if (StartSlot > Limit)
2718 return false;
2719
2720 // Case 1: array type
2721 if (const auto *AT = dyn_cast<ArrayType>(T)) {
2722 uint64_t Count = 1;
2723
2724 if (const auto *CAT = dyn_cast<ConstantArrayType>(AT))
2725 Count = CAT->getSize().getZExtValue();
2726
2727 QualType ElemTy = AT->getElementType();
2728 return AccumulateHLSLResourceSlots(ElemTy, StartSlot, Limit, ResClass, Ctx,
2729 ArrayCount * Count);
2730 }
2731
2732 // Case 2: resource leaf
2733 if (auto ResTy = dyn_cast<HLSLAttributedResourceType>(T)) {
2734 // First ensure this resource counts towards the corresponding
2735 // register type limit.
2736 if (ResTy->getAttrs().ResourceClass != ResClass)
2737 return true;
2738
2739 // Validate highest slot used
2740 uint64_t EndSlot = StartSlot + ArrayCount - 1;
2741 if (EndSlot > Limit)
2742 return false;
2743
2744 // Advance SlotCount past the consumed range
2745 StartSlot = EndSlot + 1;
2746 return true;
2747 }
2748
2749 // Case 3: struct / record
2750 if (const auto *RT = dyn_cast<RecordType>(T)) {
2751 const RecordDecl *RD = RT->getDecl();
2752
2753 if (const auto *CXXRD = dyn_cast<CXXRecordDecl>(RD)) {
2754 for (const CXXBaseSpecifier &Base : CXXRD->bases()) {
2755 if (!AccumulateHLSLResourceSlots(Base.getType(), StartSlot, Limit,
2756 ResClass, Ctx, ArrayCount))
2757 return false;
2758 }
2759 }
2760
2761 for (const FieldDecl *Field : RD->fields()) {
2762 if (!AccumulateHLSLResourceSlots(Field->getType(), StartSlot, Limit,
2763 ResClass, Ctx, ArrayCount))
2764 return false;
2765 }
2766
2767 return true;
2768 }
2769
2770 // Case 4: everything else
2771 return true;
2772}
2773
2774// return true if there is something invalid, false otherwise
2775static bool ValidateRegisterNumber(uint64_t SlotNum, Decl *TheDecl,
2776 ASTContext &Ctx, RegisterType RegTy) {
2777 const uint64_t Limit = UINT32_MAX;
2778 if (SlotNum > Limit)
2779 return true;
2780
2781 // after verifying the number doesn't exceed uint32max, we don't need
2782 // to look further into c or i register types
2783 if (RegTy == RegisterType::C || RegTy == RegisterType::I)
2784 return false;
2785
2786 if (VarDecl *VD = dyn_cast<VarDecl>(TheDecl)) {
2787 uint64_t BaseSlot = SlotNum;
2788
2789 if (!AccumulateHLSLResourceSlots(VD->getType(), SlotNum, Limit,
2790 getResourceClass(RegTy), Ctx))
2791 return true;
2792
2793 // After AccumulateHLSLResourceSlots runs, SlotNum is now
2794 // the first free slot; last used was SlotNum - 1
2795 return (BaseSlot > Limit);
2796 }
2797 // handle the cbuffer/tbuffer case
2798 if (isa<HLSLBufferDecl>(TheDecl))
2799 // resources cannot be put within a cbuffer, so no need
2800 // to analyze the structure since the register number
2801 // won't be pushed any higher.
2802 return (SlotNum > Limit);
2803
2804 // we don't expect any other decl type, so fail
2805 llvm_unreachable("unexpected decl type");
2806}
2807
2809 if (VarDecl *VD = dyn_cast<VarDecl>(TheDecl)) {
2810 QualType Ty = VD->getType();
2811 if (const auto *IAT = dyn_cast<IncompleteArrayType>(Ty))
2812 Ty = IAT->getElementType();
2813 if (SemaRef.RequireCompleteType(TheDecl->getBeginLoc(), Ty,
2814 diag::err_incomplete_type))
2815 return;
2816 }
2817
2818 StringRef Slot = "";
2819 StringRef Space = "";
2820 SourceLocation SlotLoc, SpaceLoc;
2821
2822 if (!AL.isArgIdent(0)) {
2823 Diag(AL.getLoc(), diag::err_attribute_argument_type)
2824 << AL << AANT_ArgumentIdentifier;
2825 return;
2826 }
2827 IdentifierLoc *Loc = AL.getArgAsIdent(0);
2828
2829 if (AL.getNumArgs() == 2) {
2830 Slot = Loc->getIdentifierInfo()->getName();
2831 SlotLoc = Loc->getLoc();
2832 if (!AL.isArgIdent(1)) {
2833 Diag(AL.getLoc(), diag::err_attribute_argument_type)
2834 << AL << AANT_ArgumentIdentifier;
2835 return;
2836 }
2837 Loc = AL.getArgAsIdent(1);
2838 Space = Loc->getIdentifierInfo()->getName();
2839 SpaceLoc = Loc->getLoc();
2840 } else {
2841 StringRef Str = Loc->getIdentifierInfo()->getName();
2842 if (Str.starts_with("space")) {
2843 Space = Str;
2844 SpaceLoc = Loc->getLoc();
2845 } else {
2846 Slot = Str;
2847 SlotLoc = Loc->getLoc();
2848 Space = "space0";
2849 }
2850 }
2851
2852 RegisterType RegType = RegisterType::SRV;
2853 std::optional<unsigned> SlotNum;
2854 unsigned SpaceNum = 0;
2855
2856 // Validate slot
2857 if (!Slot.empty()) {
2858 if (!convertToRegisterType(Slot, &RegType)) {
2859 Diag(SlotLoc, diag::err_hlsl_binding_type_invalid) << Slot.substr(0, 1);
2860 return;
2861 }
2862 if (RegType == RegisterType::I) {
2863 Diag(SlotLoc, diag::warn_hlsl_deprecated_register_type_i);
2864 return;
2865 }
2866 const StringRef SlotNumStr = Slot.substr(1);
2867
2868 uint64_t N;
2869
2870 // validate that the slot number is a non-empty number
2871 if (SlotNumStr.getAsInteger(10, N)) {
2872 Diag(SlotLoc, diag::err_hlsl_unsupported_register_number);
2873 return;
2874 }
2875
2876 // Validate register number. It should not exceed UINT32_MAX,
2877 // including if the resource type is an array that starts
2878 // before UINT32_MAX, but ends afterwards.
2879 if (ValidateRegisterNumber(N, TheDecl, getASTContext(), RegType)) {
2880 Diag(SlotLoc, diag::err_hlsl_register_number_too_large);
2881 return;
2882 }
2883
2884 // the slot number has been validated and does not exceed UINT32_MAX
2885 SlotNum = (unsigned)N;
2886 }
2887
2888 // Validate space
2889 if (!Space.starts_with("space")) {
2890 Diag(SpaceLoc, diag::err_hlsl_expected_space) << Space;
2891 return;
2892 }
2893 StringRef SpaceNumStr = Space.substr(5);
2894 if (SpaceNumStr.getAsInteger(10, SpaceNum)) {
2895 Diag(SpaceLoc, diag::err_hlsl_expected_space) << Space;
2896 return;
2897 }
2898
2899 // If we have slot, diagnose it is the right register type for the decl
2900 if (SlotNum.has_value())
2901 if (!DiagnoseHLSLRegisterAttribute(SemaRef, SlotLoc, TheDecl, RegType,
2902 !SpaceLoc.isInvalid()))
2903 return;
2904
2905 HLSLResourceBindingAttr *NewAttr =
2906 HLSLResourceBindingAttr::Create(getASTContext(), Slot, Space, AL);
2907 if (NewAttr) {
2908 NewAttr->setBinding(RegType, SlotNum, SpaceNum);
2909 TheDecl->addAttr(NewAttr);
2910 }
2911}
2912
2914 HLSLParamModifierAttr *NewAttr = mergeParamModifierAttr(
2915 D, AL,
2916 static_cast<HLSLParamModifierAttr::Spelling>(AL.getSemanticSpelling()));
2917 if (NewAttr)
2918 D->addAttr(NewAttr);
2919}
2920
2921static bool isMatrixType(QualType QT) {
2922 const Type *Ty = QT->getUnqualifiedDesugaredType();
2923 return Ty->isDependentType() || Ty->isConstantMatrixType();
2924}
2925
2926/// Walks the existing AttributedType sugar of \p T looking for a previously
2927/// applied HLSLRowMajor/HLSLColumnMajor marker. If one is found, populates
2928/// \p ExistingKind with its attr::Kind and returns true.
2930 attr::Kind &ExistingKind) {
2931 QualType Cur = T;
2932 while (const auto *AT = Cur->getAs<AttributedType>()) {
2933 attr::Kind K = AT->getAttrKind();
2934 if (K == attr::HLSLRowMajor || K == attr::HLSLColumnMajor) {
2935 ExistingKind = K;
2936 return true;
2937 }
2938 Cur = AT->getModifiedType();
2939 }
2940 return false;
2941}
2942
2944 if (T.isNull())
2945 return nullptr;
2946
2947 ASTContext &Ctx = getASTContext();
2948 attr::Kind AttrK = AL.getKind() == ParsedAttr::AT_HLSLRowMajor
2949 ? attr::HLSLRowMajor
2950 : attr::HLSLColumnMajor;
2951
2952 // For non-dependent types, the operand must be a matrix.
2953 if (!T->isDependentType() && !isMatrixType(T)) {
2954 Diag(AL.getLoc(), diag::err_hlsl_matrix_layout_non_matrix)
2955 << AL.getAttrName();
2956 AL.setInvalid();
2957 return nullptr;
2958 }
2959
2960 // Conflict / duplicate detection by walking existing sugar.
2961 attr::Kind ExistingKind;
2962 if (findExistingMatrixLayoutMarker(T, ExistingKind)) {
2963 if (ExistingKind == AttrK) {
2964 Diag(AL.getLoc(), diag::warn_duplicate_attribute_exact)
2965 << AL.getAttrName();
2966 Diag(AL.getLoc(), diag::note_previous_attribute);
2967 return nullptr;
2968 }
2969 IdentifierInfo *ExistingII = &Ctx.Idents.get(
2970 ExistingKind == attr::HLSLRowMajor ? "row_major" : "column_major");
2971 Diag(AL.getLoc(), diag::err_hlsl_matrix_layout_conflict)
2972 << AL.getAttrName() << ExistingII;
2973 Diag(AL.getLoc(), diag::note_conflicting_attribute);
2974 AL.setInvalid();
2975 return nullptr;
2976 }
2977
2978 if (AttrK == attr::HLSLRowMajor)
2979 return ::new (Ctx) HLSLRowMajorAttr(Ctx, AL);
2980 return ::new (Ctx) HLSLColumnMajorAttr(Ctx, AL);
2981}
2982
2983// Re-validates an HLSL `row_major` / `column_major` attribute after template
2984// substitution. The parse-time check in `buildMatrixLayoutTypeAttr` is skipped
2985// for dependent types; `TransformAttributedType` calls this once the type is
2986// concrete. Returns `true` (and emits a diagnostic) if the substituted type is
2987// not a matrix or array of matrices, signaling the caller to abort the
2988// transform.
2990 SourceLocation Loc) {
2991 if (K != attr::HLSLRowMajor && K != attr::HLSLColumnMajor)
2992 return false;
2993 if (T.isNull() || T->isDependentType())
2994 return false;
2995 if (isMatrixType(T))
2996 return false;
2998 K == attr::HLSLRowMajor ? "row_major" : "column_major");
2999 Diag(Loc, diag::err_hlsl_matrix_layout_non_matrix) << II;
3000 return true;
3001}
3002
3003// Transpose and matrix mul need to read the destination layout.
3004// Elementwise builtins reuse the operand layout instead.
3005namespace {
3006
3007using llvm::dxil::BarrierMemoryTypeFlag;
3008using llvm::dxil::BarrierSemanticFlag;
3009
3010template <typename T> constexpr uint64_t barrierFlagValue(T Flag) {
3011 return llvm::to_underlying(Flag);
3012}
3013
3014/// This class implements reachable HLSL diagnostics.
3015///
3016/// It diagnoses unavailable APIs in default and relaxed availability modes.
3017/// It also validates Barrier calls in all availability modes.
3018///
3019/// This is done by traversing the AST of all shader entry point functions
3020/// and of all exported functions, and any functions that are referenced
3021/// from this AST. In other words, any functions that are reachable from
3022/// the entry points.
3023class DiagnoseHLSLAvailability : public DynamicRecursiveASTVisitor {
3024 Sema &SemaRef;
3025 bool DiagnoseAvailability;
3026
3027 // Stack of functions to be scaned
3029
3030 // Tracks which environments functions have been scanned in.
3031 //
3032 // Maps FunctionDecl to an unsigned number that represents the set of shader
3033 // environments the function has been scanned for.
3034 // The llvm::Triple::EnvironmentType enum values for shader stages guaranteed
3035 // to be numbered from llvm::Triple::Pixel to llvm::Triple::Amplification
3036 // (verified by static_asserts in Triple.cpp), we can use it to index
3037 // individual bits in the set, as long as we shift the values to start with 0
3038 // by subtracting the value of llvm::Triple::Pixel first.
3039 //
3040 // The N'th bit in the set will be set if the function has been scanned
3041 // in shader environment whose llvm::Triple::EnvironmentType integer value
3042 // equals (llvm::Triple::Pixel + N).
3043 //
3044 // For example, if a function has been scanned in compute and pixel stage
3045 // environment, the value will be 0x21 (100001 binary) because:
3046 //
3047 // (int)(llvm::Triple::Pixel - llvm::Triple::Pixel) == 0
3048 // (int)(llvm::Triple::Compute - llvm::Triple::Pixel) == 5
3049 //
3050 // A FunctionDecl is mapped to 0 (or not included in the map) if it has not
3051 // been scanned in any environment.
3052 llvm::DenseMap<const FunctionDecl *, unsigned> ScannedDecls;
3053
3054 // Do not access these directly, use the get/set methods below to make
3055 // sure the values are in sync
3056 llvm::Triple::EnvironmentType CurrentShaderEnvironment;
3057 unsigned CurrentShaderStageBit;
3058
3059 // True if scanning a function that was already scanned in a different
3060 // shader stage context. Suppress stage-independent diagnostics because
3061 // they were reported during the first scan.
3062 bool ReportOnlyShaderStageIssues;
3063
3064 // Helper methods for dealing with current stage context / environment
3065 void SetShaderStageContext(llvm::Triple::EnvironmentType ShaderType) {
3066 static_assert(sizeof(unsigned) >= 4);
3067 assert(HLSLShaderAttr::isValidShaderType(ShaderType));
3068 assert((unsigned)(ShaderType - llvm::Triple::Pixel) < 31 &&
3069 "ShaderType is too big for this bitmap"); // 31 is reserved for
3070 // "unknown"
3071
3072 unsigned bitmapIndex = ShaderType - llvm::Triple::Pixel;
3073 CurrentShaderEnvironment = ShaderType;
3074 CurrentShaderStageBit = (1 << bitmapIndex);
3075 }
3076
3077 void SetUnknownShaderStageContext() {
3078 CurrentShaderEnvironment = llvm::Triple::UnknownEnvironment;
3079 CurrentShaderStageBit = (1 << 31);
3080 }
3081
3082 llvm::Triple::EnvironmentType GetCurrentShaderEnvironment() const {
3083 return CurrentShaderEnvironment;
3084 }
3085
3086 bool InUnknownShaderStageContext() const {
3087 return CurrentShaderEnvironment == llvm::Triple::UnknownEnvironment;
3088 }
3089
3090 // Helper methods for dealing with shader stage bitmap
3091 void AddToScannedFunctions(const FunctionDecl *FD) {
3092 unsigned &ScannedStages = ScannedDecls[FD];
3093 ScannedStages |= CurrentShaderStageBit;
3094 }
3095
3096 unsigned GetScannedStages(const FunctionDecl *FD) { return ScannedDecls[FD]; }
3097
3098 bool WasAlreadyScannedInCurrentStage(const FunctionDecl *FD) {
3099 return WasAlreadyScannedInCurrentStage(GetScannedStages(FD));
3100 }
3101
3102 bool WasAlreadyScannedInCurrentStage(unsigned ScannerStages) {
3103 return ScannerStages & CurrentShaderStageBit;
3104 }
3105
3106 static bool NeverBeenScanned(unsigned ScannedStages) {
3107 return ScannedStages == 0;
3108 }
3109
3110 // Scanning methods
3111 void HandleFunctionOrMethodRef(FunctionDecl *FD, Expr *RefExpr);
3112 void CheckDeclAvailability(NamedDecl *D, const AvailabilityAttr *AA,
3113 SourceRange Range);
3114 const AvailabilityAttr *FindAvailabilityAttr(const Decl *D);
3115 bool HasMatchingEnvironmentOrNone(const AvailabilityAttr *AA);
3116 void DiagnoseBarrierCall(CallExpr *CE);
3117 uint64_t DiagnoseBarrierGroupMemory(Expr *MemoryArg, uint64_t MemoryFlags,
3118 bool HasVisibleGroup, bool IsAllMemory);
3119 uint64_t DiagnoseBarrierNodeMemory(Expr *MemoryArg, uint64_t MemoryFlags,
3120 bool HasKnownStage, bool IsAllMemory);
3121 void DiagnoseBarrierGroupSemantic(Expr *SemanticArg, uint64_t SemanticFlags,
3122 bool HasVisibleGroup);
3123 void DiagnoseBarrierScope(Expr *SemanticArg, uint64_t MemoryFlags,
3124 uint64_t SemanticFlags);
3125
3126public:
3127 DiagnoseHLSLAvailability(Sema &SemaRef, bool DiagnoseAvailability)
3128 : SemaRef(SemaRef), DiagnoseAvailability(DiagnoseAvailability),
3129 CurrentShaderEnvironment(llvm::Triple::UnknownEnvironment),
3130 CurrentShaderStageBit(0), ReportOnlyShaderStageIssues(false) {}
3131
3132 // AST traversal methods
3133 void RunOnTranslationUnit(const TranslationUnitDecl *TU);
3134 void RunOnFunction(const FunctionDecl *FD);
3135
3136 bool VisitDeclRefExpr(DeclRefExpr *DRE) override {
3137 FunctionDecl *FD = llvm::dyn_cast<FunctionDecl>(DRE->getDecl());
3138 if (FD)
3139 HandleFunctionOrMethodRef(FD, DRE);
3140 return true;
3141 }
3142
3143 bool VisitMemberExpr(MemberExpr *ME) override {
3144 FunctionDecl *FD = llvm::dyn_cast<FunctionDecl>(ME->getMemberDecl());
3145 if (FD)
3146 HandleFunctionOrMethodRef(FD, ME);
3147 return true;
3148 }
3149
3150 bool VisitCallExpr(CallExpr *CE) override {
3151 DiagnoseBarrierCall(CE);
3152 return true;
3153 }
3154};
3155
3156uint64_t DiagnoseHLSLAvailability::DiagnoseBarrierGroupMemory(
3157 Expr *MemoryArg, uint64_t MemoryFlags, bool HasVisibleGroup,
3158 bool IsAllMemory) {
3159 const uint64_t GroupSharedMemory =
3160 barrierFlagValue(BarrierMemoryTypeFlag::GroupSharedMemory);
3161 if (HasVisibleGroup || (MemoryFlags & GroupSharedMemory) == 0)
3162 return MemoryFlags;
3163
3164 if (!IsAllMemory) {
3165 SemaRef.Diag(MemoryArg->getExprLoc(),
3166 diag::err_hlsl_barrier_flag_requires_group)
3167 << 0;
3168 return MemoryFlags;
3169 }
3170
3171 return MemoryFlags & ~GroupSharedMemory;
3172}
3173
3174uint64_t DiagnoseHLSLAvailability::DiagnoseBarrierNodeMemory(
3175 Expr *MemoryArg, uint64_t MemoryFlags, bool HasKnownStage,
3176 bool IsAllMemory) {
3177 const uint64_t NodeMemory =
3178 barrierFlagValue(BarrierMemoryTypeFlag::NodeMemory);
3179 if (!HasKnownStage || (MemoryFlags & NodeMemory) == 0)
3180 return MemoryFlags;
3181
3182 if (!IsAllMemory) {
3183 SemaRef.Diag(MemoryArg->getExprLoc(),
3184 diag::err_hlsl_barrier_node_memory_requires_node);
3185 return MemoryFlags;
3186 }
3187
3188 return MemoryFlags & ~NodeMemory;
3189}
3190
3191void DiagnoseHLSLAvailability::DiagnoseBarrierGroupSemantic(
3192 Expr *SemanticArg, uint64_t SemanticFlags, bool HasVisibleGroup) {
3193 if (HasVisibleGroup ||
3194 (SemanticFlags & barrierFlagValue(BarrierSemanticFlag::GroupFlags)) == 0)
3195 return;
3196
3197 SemaRef.Diag(SemanticArg->getExprLoc(),
3198 diag::err_hlsl_barrier_flag_requires_group)
3199 << ((SemanticFlags & barrierFlagValue(BarrierSemanticFlag::GroupSync)) !=
3200 0
3201 ? 1
3202 : 2);
3203}
3204
3205void DiagnoseHLSLAvailability::DiagnoseBarrierScope(Expr *SemanticArg,
3206 uint64_t MemoryFlags,
3207 uint64_t SemanticFlags) {
3208 if (ReportOnlyShaderStageIssues)
3209 return;
3210
3211 const uint64_t DeviceScopeMemory =
3212 barrierFlagValue(BarrierMemoryTypeFlag::UAVMemory) |
3213 barrierFlagValue(BarrierMemoryTypeFlag::NodeInputMemory);
3214 if ((SemanticFlags & barrierFlagValue(BarrierSemanticFlag::DeviceScope)) !=
3215 0 &&
3216 (MemoryFlags & DeviceScopeMemory) == 0)
3217 SemaRef.Diag(SemanticArg->getExprLoc(),
3218 diag::err_hlsl_barrier_scope_requires_memory)
3219 << 1;
3220 if ((SemanticFlags & barrierFlagValue(BarrierSemanticFlag::GroupScope)) !=
3221 0 &&
3222 MemoryFlags == 0)
3223 SemaRef.Diag(SemanticArg->getExprLoc(),
3224 diag::err_hlsl_barrier_scope_requires_memory)
3225 << 0;
3226}
3227
3228void DiagnoseHLSLAvailability::DiagnoseBarrierCall(CallExpr *CE) {
3229 const FunctionDecl *FD = CE->getDirectCallee();
3230 if (!FD || FD->getBuiltinID() != Builtin::BI__builtin_hlsl_barrier)
3231 return;
3232
3233 const llvm::Triple::EnvironmentType Stage = GetCurrentShaderEnvironment();
3234 const bool HasKnownStage = !InUnknownShaderStageContext();
3235 const bool HasVisibleGroup =
3236 !HasKnownStage || Stage == llvm::Triple::Compute ||
3237 Stage == llvm::Triple::Mesh || Stage == llvm::Triple::Amplification;
3238
3239 uint64_t MemoryFlags = barrierFlagValue(BarrierMemoryTypeFlag::ValidMask);
3240 Expr *MemoryArg = CE->getArg(0);
3241 if (MemoryArg->getType()->isUnsignedIntegerType()) {
3242 std::optional<llvm::APSInt> Value =
3243 MemoryArg->getIntegerConstantExpr(SemaRef.Context);
3244 if (!Value)
3245 return;
3246 MemoryFlags = Value->getZExtValue();
3247 const bool IsAllMemory =
3248 MemoryFlags == barrierFlagValue(BarrierMemoryTypeFlag::ValidMask);
3249
3250 MemoryFlags = DiagnoseBarrierGroupMemory(MemoryArg, MemoryFlags,
3251 HasVisibleGroup, IsAllMemory);
3252 MemoryFlags = DiagnoseBarrierNodeMemory(MemoryArg, MemoryFlags,
3253 HasKnownStage, IsAllMemory);
3254 } else if (!HasVisibleGroup) {
3255 SemaRef.Diag(MemoryArg->getExprLoc(),
3256 diag::err_hlsl_barrier_resource_requires_group);
3257 return;
3258 }
3259
3260 Expr *SemanticArg = CE->getArg(1);
3261 std::optional<llvm::APSInt> Value =
3262 SemanticArg->getIntegerConstantExpr(SemaRef.Context);
3263 if (!Value)
3264 return;
3265 const uint64_t SemanticFlags = Value->getZExtValue();
3266
3267 DiagnoseBarrierGroupSemantic(SemanticArg, SemanticFlags, HasVisibleGroup);
3268
3269 if (MemoryArg->getType()->isUnsignedIntegerType())
3270 DiagnoseBarrierScope(SemanticArg, MemoryFlags, SemanticFlags);
3271}
3272
3273void DiagnoseHLSLAvailability::HandleFunctionOrMethodRef(FunctionDecl *FD,
3274 Expr *RefExpr) {
3275 assert((isa<DeclRefExpr>(RefExpr) || isa<MemberExpr>(RefExpr)) &&
3276 "expected DeclRefExpr or MemberExpr");
3277
3278 if (DiagnoseAvailability)
3279 if (const AvailabilityAttr *AA = FindAvailabilityAttr(FD))
3280 CheckDeclAvailability(
3281 FD, AA, SourceRange(RefExpr->getBeginLoc(), RefExpr->getEndLoc()));
3282
3283 // has a definition -> add to stack to be scanned
3284 const FunctionDecl *FDWithBody = nullptr;
3285 if (FD->hasBody(FDWithBody) && !WasAlreadyScannedInCurrentStage(FDWithBody))
3286 DeclsToScan.push_back(FDWithBody);
3287}
3288
3289void DiagnoseHLSLAvailability::RunOnTranslationUnit(
3290 const TranslationUnitDecl *TU) {
3291 const TargetInfo &TargetInfo = SemaRef.getASTContext().getTargetInfo();
3292 std::string &EntryName = TargetInfo.getTargetOpts().HLSLEntry;
3293 bool IsLibraryShader = TargetInfo.getTriple().getEnvironment() ==
3294 llvm::Triple::EnvironmentType::Library;
3295 SourceLocation EntryLoc{};
3296
3297 // Iterate over all shader entry functions and library exports, and for those
3298 // that have a body (definiton), run diag scan on each, setting appropriate
3299 // shader environment context based on whether it is a shader entry function
3300 // or an exported function. Exported functions can be in namespaces and in
3301 // export declarations so we need to scan those declaration contexts as well.
3303 DeclContextsToScan.push_back(TU);
3304
3305 while (!DeclContextsToScan.empty()) {
3306 const DeclContext *DC = DeclContextsToScan.pop_back_val();
3307 for (auto &D : DC->decls()) {
3308 // do not scan implicit declaration generated by the implementation
3309 if (D->isImplicit())
3310 continue;
3311
3312 // for namespace or export declaration add the context to the list to be
3313 // scanned later
3314 if (llvm::dyn_cast<NamespaceDecl>(D) || llvm::dyn_cast<ExportDecl>(D)) {
3315 DeclContextsToScan.push_back(llvm::dyn_cast<DeclContext>(D));
3316 continue;
3317 }
3318
3319 // skip over other decls or function decls without body
3320 const FunctionDecl *FD = llvm::dyn_cast<FunctionDecl>(D);
3321 if (!FD || !FD->isThisDeclarationADefinition())
3322 continue;
3323
3324 // shader entry point
3325 if (HLSLShaderAttr *ShaderAttr = FD->getAttr<HLSLShaderAttr>()) {
3326 if (!IsLibraryShader && FD->getName() == EntryName) {
3327 if (EntryLoc.isValid()) {
3328 SemaRef.Diag(FD->getLocation(),
3329 diag::err_hlsl_ambiguous_entry_point)
3330 << EntryName;
3331 SemaRef.Diag(EntryLoc, diag::note_previous_declaration_as)
3332 << EntryName;
3333 return;
3334 }
3335 EntryLoc = FD->getLocation();
3336 }
3337 SetShaderStageContext(ShaderAttr->getType());
3338 RunOnFunction(FD);
3339 continue;
3340 }
3341 // exported library function
3342 // FIXME: replace this loop with external linkage check once issue #92071
3343 // is resolved
3344 bool isExport = FD->isInExportDeclContext();
3345 if (!isExport) {
3346 for (const auto *Redecl : FD->redecls()) {
3347 if (Redecl->isInExportDeclContext()) {
3348 isExport = true;
3349 break;
3350 }
3351 }
3352 }
3353 if (isExport) {
3354 SetUnknownShaderStageContext();
3355 RunOnFunction(FD);
3356 continue;
3357 }
3358 }
3359 }
3360
3361 if (!IsLibraryShader && EntryLoc.isInvalid()) {
3362 SemaRef.Diag(TU->getLocation(), diag::err_hlsl_missing_entry_point)
3363 << EntryName;
3364 return;
3365 }
3366}
3367
3368void DiagnoseHLSLAvailability::RunOnFunction(const FunctionDecl *FD) {
3369 assert(DeclsToScan.empty() && "DeclsToScan should be empty");
3370 DeclsToScan.push_back(FD);
3371
3372 while (!DeclsToScan.empty()) {
3373 // Take one decl from the stack and check it by traversing its AST.
3374 // For any CallExpr found during the traversal add it's callee to the top of
3375 // the stack to be processed next. Functions already processed are stored in
3376 // ScannedDecls.
3377 const FunctionDecl *FD = DeclsToScan.pop_back_val();
3378
3379 // Decl was already scanned
3380 const unsigned ScannedStages = GetScannedStages(FD);
3381 if (WasAlreadyScannedInCurrentStage(ScannedStages))
3382 continue;
3383
3384 ReportOnlyShaderStageIssues = !NeverBeenScanned(ScannedStages);
3385
3386 AddToScannedFunctions(FD);
3387 TraverseStmt(FD->getBody());
3388 }
3389}
3390
3391bool DiagnoseHLSLAvailability::HasMatchingEnvironmentOrNone(
3392 const AvailabilityAttr *AA) {
3393 const IdentifierInfo *IIEnvironment = AA->getEnvironment();
3394 if (!IIEnvironment)
3395 return true;
3396
3397 llvm::Triple::EnvironmentType CurrentEnv = GetCurrentShaderEnvironment();
3398 if (CurrentEnv == llvm::Triple::UnknownEnvironment)
3399 return false;
3400
3401 llvm::Triple::EnvironmentType AttrEnv =
3402 AvailabilityAttr::getEnvironmentType(IIEnvironment->getName());
3403
3404 return CurrentEnv == AttrEnv;
3405}
3406
3407const AvailabilityAttr *
3408DiagnoseHLSLAvailability::FindAvailabilityAttr(const Decl *D) {
3409 AvailabilityAttr const *PartialMatch = nullptr;
3410 // Check each AvailabilityAttr to find the one for this platform.
3411 // For multiple attributes with the same platform try to find one for this
3412 // environment.
3413 for (const auto *A : D->attrs()) {
3414 if (const auto *Avail = dyn_cast<AvailabilityAttr>(A)) {
3415 const AvailabilityAttr *EffectiveAvail = Avail->getEffectiveAttr();
3416 StringRef AttrPlatform = EffectiveAvail->getPlatform()->getName();
3417 StringRef TargetPlatform =
3419
3420 // Match the platform name.
3421 if (AttrPlatform == TargetPlatform) {
3422 // Find the best matching attribute for this environment
3423 if (HasMatchingEnvironmentOrNone(EffectiveAvail))
3424 return Avail;
3425 PartialMatch = Avail;
3426 }
3427 }
3428 }
3429 return PartialMatch;
3430}
3431
3432// Check availability against target shader model version and current shader
3433// stage and emit diagnostic
3434void DiagnoseHLSLAvailability::CheckDeclAvailability(NamedDecl *D,
3435 const AvailabilityAttr *AA,
3436 SourceRange Range) {
3437
3438 const IdentifierInfo *IIEnv = AA->getEnvironment();
3439
3440 if (!IIEnv) {
3441 // The availability attribute does not have environment -> it depends only
3442 // on shader model version and not on specific the shader stage.
3443
3444 // Skip emitting the diagnostics if the diagnostic mode is set to
3445 // strict (-fhlsl-strict-availability) because all relevant diagnostics
3446 // were already emitted in the DiagnoseUnguardedAvailability scan
3447 // (SemaAvailability.cpp).
3448 if (SemaRef.getLangOpts().HLSLStrictAvailability)
3449 return;
3450
3451 // Do not report shader-stage-independent issues if scanning a function
3452 // that was already scanned in a different shader stage context (they would
3453 // be duplicate)
3454 if (ReportOnlyShaderStageIssues)
3455 return;
3456
3457 } else {
3458 // The availability attribute has environment -> we need to know
3459 // the current stage context to property diagnose it.
3460 if (InUnknownShaderStageContext())
3461 return;
3462 }
3463
3464 // Check introduced version and if environment matches
3465 bool EnvironmentMatches = HasMatchingEnvironmentOrNone(AA);
3466 VersionTuple Introduced = AA->getIntroduced();
3467 VersionTuple TargetVersion =
3469
3470 if (TargetVersion >= Introduced && EnvironmentMatches)
3471 return;
3472
3473 // Emit diagnostic message
3474 const TargetInfo &TI = SemaRef.getASTContext().getTargetInfo();
3475 llvm::StringRef PlatformName(
3476 AvailabilityAttr::getPrettyPlatformName(TI.getPlatformName()));
3477
3478 llvm::StringRef CurrentEnvStr =
3479 llvm::Triple::getEnvironmentTypeName(GetCurrentShaderEnvironment());
3480
3481 llvm::StringRef AttrEnvStr =
3482 AA->getEnvironment() ? AA->getEnvironment()->getName() : "";
3483 bool UseEnvironment = !AttrEnvStr.empty();
3484
3485 if (EnvironmentMatches) {
3486 SemaRef.Diag(Range.getBegin(), diag::warn_hlsl_availability)
3487 << Range << D << PlatformName << Introduced.getAsString()
3488 << UseEnvironment << CurrentEnvStr;
3489 } else {
3490 SemaRef.Diag(Range.getBegin(), diag::warn_hlsl_availability_unavailable)
3491 << Range << D;
3492 }
3493
3494 SemaRef.Diag(D->getLocation(), diag::note_partial_availability_specified_here)
3495 << D << PlatformName << Introduced.getAsString()
3496 << SemaRef.Context.getTargetInfo().getPlatformMinVersion().getAsString()
3497 << UseEnvironment << AttrEnvStr << CurrentEnvStr;
3498}
3499
3500} // namespace
3501
3503 // process default CBuffer - create buffer layout struct and invoke codegenCGH
3504 if (!DefaultCBufferDecls.empty()) {
3506 SemaRef.getASTContext(), SemaRef.getCurLexicalContext(),
3507 DefaultCBufferDecls);
3508 addImplicitBindingAttrToDecl(SemaRef, DefaultCBuffer, RegisterType::CBuffer,
3510 SemaRef.getCurLexicalContext()->addDecl(DefaultCBuffer);
3512
3513 // Set HasValidPackoffset if any of the decls has a register(c#) annotation;
3514 for (const Decl *VD : DefaultCBufferDecls) {
3515 const HLSLResourceBindingAttr *RBA =
3516 VD->getAttr<HLSLResourceBindingAttr>();
3517 if (RBA && RBA->hasRegisterSlot() &&
3518 RBA->getRegisterType() == HLSLResourceBindingAttr::RegisterType::C) {
3519 DefaultCBuffer->setHasValidPackoffset(true);
3520 break;
3521 }
3522 }
3523
3524 DeclGroupRef DG(DefaultCBuffer);
3525 SemaRef.Consumer.HandleTopLevelDecl(DG);
3526 }
3527 diagnoseAvailabilityViolations(TU);
3528}
3529
3530// For resource member access through a global struct array, verify that the
3531// array index selecting the struct element is a constant integer expression.
3532// Returns false if the member expression is invalid.
3534 assert((ME->getType()->isHLSLResourceRecord() ||
3536 "expected member expr to have resource record type or array of them");
3537
3538 // Walk the AST from MemberExpr to the VarDecl of the parent struct instance
3539 // and take note of any non-constant array indexing along the way. If the
3540 // VarDecl we find is a global variable, report error if there was any
3541 // non-constant array index in the resource member access along the way.
3542 const Expr *NonConstIndexExpr = nullptr;
3543 const Expr *E = ME->getBase();
3544 while (E) {
3545 if (const DeclRefExpr *DRE = dyn_cast<DeclRefExpr>(E)) {
3546 if (!NonConstIndexExpr)
3547 return true;
3548
3549 const VarDecl *VD = cast<VarDecl>(DRE->getDecl());
3550 if (!VD->hasGlobalStorage())
3551 return true;
3552
3553 SemaRef.Diag(NonConstIndexExpr->getExprLoc(),
3554 diag::err_hlsl_resource_member_array_access_not_constant);
3555 return false;
3556 }
3557
3558 if (const auto *ASE = dyn_cast<ArraySubscriptExpr>(E)) {
3559 const Expr *IdxExpr = ASE->getIdx();
3560 if (!IdxExpr->isIntegerConstantExpr(SemaRef.getASTContext()))
3561 NonConstIndexExpr = IdxExpr;
3562 E = ASE->getBase();
3563 } else if (const auto *SubME = dyn_cast<MemberExpr>(E)) {
3564 E = SubME->getBase();
3565 } else if (const auto *ICE = dyn_cast<ImplicitCastExpr>(E)) {
3566 E = ICE->getSubExpr();
3567 } else {
3568 llvm_unreachable("unexpected expr type in resource member access");
3569 }
3570 }
3571 return true;
3572}
3573
3575 CXXRecordDecl *RD) {
3576 QualType AddrSpaceType =
3577 SemaRef.Context.getCanonicalType(SemaRef.Context.getAddrSpaceQualType(
3578 Type.withConst(), LangAS::hlsl_constant));
3579 QualType ReturnTy = SemaRef.Context.getCanonicalType(
3580 SemaRef.Context.getLValueReferenceType(AddrSpaceType));
3581
3582 DeclarationName ConvName =
3583 SemaRef.Context.DeclarationNames.getCXXConversionFunctionName(
3584 CanQualType::CreateUnsafe(ReturnTy));
3585 LookupResult ConvR(SemaRef, ConvName, SourceLocation(),
3587 [[maybe_unused]] bool LookupSucceeded =
3588 SemaRef.LookupQualifiedName(ConvR, RD);
3589 assert(LookupSucceeded);
3590
3591 for (NamedDecl *D : ConvR) {
3593 return D;
3594 }
3595 return nullptr;
3596}
3597
3598std::optional<ExprResult>
3600 QualType BaseType = BaseExpr->getType();
3601 const HLSLAttributedResourceType *ResTy =
3602 HLSLAttributedResourceType::findHandleTypeOnResource(
3603 BaseType.getTypePtr());
3604 if (!ResTy ||
3605 ResTy->getAttrs().ResourceClass != llvm::dxil::ResourceClass::CBuffer)
3606 return std::nullopt;
3607
3608 QualType TemplateType = ResTy->getContainedType();
3609
3610 NamedDecl *NamedConversionDecl = getConstantBufferConversionFunction(
3611 TemplateType, BaseType->getAsCXXRecordDecl());
3612 assert(NamedConversionDecl &&
3613 "Could not find conversion function for ConstantBuffer.");
3614 auto *ConversionDecl =
3615 cast<CXXConversionDecl>(NamedConversionDecl->getUnderlyingDecl());
3616
3617 return SemaRef.BuildCXXMemberCallExpr(BaseExpr, NamedConversionDecl,
3618 ConversionDecl,
3619 /*HadMultipleCandidates=*/false);
3620}
3621
3622void SemaHLSL::diagnoseAvailabilityViolations(TranslationUnitDecl *TU) {
3623 // Strict mode diagnoses availability during the
3624 // DiagnoseUnguardedAvailability scan in SemaAvailability.cpp. The reachable
3625 // function scan must still run to validate Barrier calls.
3627 const bool DiagnoseAvailability =
3628 !SemaRef.getLangOpts().HLSLStrictAvailability ||
3629 TI.getTriple().getEnvironment() == llvm::Triple::EnvironmentType::Library;
3630 DiagnoseHLSLAvailability(SemaRef, DiagnoseAvailability)
3631 .RunOnTranslationUnit(TU);
3632}
3633
3634static bool CheckAllArgsHaveSameType(Sema *S, CallExpr *TheCall) {
3635 assert(TheCall->getNumArgs() > 1);
3636 QualType ArgTy0 = TheCall->getArg(0)->getType();
3637
3638 for (unsigned I = 1, N = TheCall->getNumArgs(); I < N; ++I) {
3640 ArgTy0, TheCall->getArg(I)->getType())) {
3641 S->Diag(TheCall->getBeginLoc(), diag::err_vec_builtin_incompatible_vector)
3642 << TheCall->getDirectCallee() << /*useAllTerminology*/ true
3643 << SourceRange(TheCall->getArg(0)->getBeginLoc(),
3644 TheCall->getArg(N - 1)->getEndLoc());
3645 return true;
3646 }
3647 }
3648 return false;
3649}
3650
3652 QualType ArgType = Arg->getType();
3654 S->Diag(Arg->getBeginLoc(), diag::err_typecheck_convert_incompatible)
3655 << ArgType << ExpectedType << 1 << 0 << 0;
3656 return true;
3657 }
3658 return false;
3659}
3660
3662 Sema *S, CallExpr *TheCall,
3663 llvm::function_ref<bool(Sema *S, SourceLocation Loc, int ArgOrdinal,
3664 clang::QualType PassedType)>
3665 Check) {
3666 for (unsigned I = 0; I < TheCall->getNumArgs(); ++I) {
3667 Expr *Arg = TheCall->getArg(I);
3668 if (Check(S, Arg->getBeginLoc(), I + 1, Arg->getType()))
3669 return true;
3670 }
3671 return false;
3672}
3673
3675 int ArgOrdinal,
3676 clang::QualType PassedType) {
3677 clang::QualType BaseType =
3678 getElementTypeOf(PassedType, /*IncludeMatrix=*/false);
3679 if (!BaseType->isFloat32Type())
3680 return S->Diag(Loc, diag::err_builtin_invalid_arg_type)
3681 << ArgOrdinal << /* scalar or vector of */ 5 << /* no int */ 0
3682 << /* float */ 1 << PassedType;
3683 return false;
3684}
3685
3687 int ArgOrdinal,
3688 clang::QualType PassedType) {
3689 QualType BaseType = getScalarComponentType(PassedType);
3690
3691 if (!BaseType->isHalfType() && !BaseType->isFloat32Type())
3692 return S->Diag(Loc, diag::err_builtin_invalid_arg_type)
3693 << ArgOrdinal << /* scalar or vector of */ 5 << /* no int */ 0
3694 << /* half or float */ 2 << PassedType;
3695 return false;
3696}
3697
3699 int ArgOrdinal,
3700 clang::QualType PassedType) {
3701 QualType BaseType = getScalarComponentType(PassedType);
3702 if (!BaseType->isDoubleType()) {
3703 // FIXME: adopt standard `err_builtin_invalid_arg_type` instead of using
3704 // this custom error.
3705 return S->Diag(Loc, diag::err_builtin_requires_double_type)
3706 << ArgOrdinal << PassedType;
3707 }
3708
3709 return false;
3710}
3711
3712static bool CheckModifiableLValue(Sema *S, CallExpr *TheCall,
3713 unsigned ArgIndex) {
3714 auto *Arg = TheCall->getArg(ArgIndex);
3715 SourceLocation OrigLoc = Arg->getExprLoc();
3716 if (Arg->IgnoreCasts()->isModifiableLvalue(S->Context, &OrigLoc) ==
3718 return false;
3719 S->Diag(OrigLoc, diag::error_hlsl_inout_lvalue) << Arg << 0;
3720 return true;
3721}
3722
3723// Verifies that the argument at `ArgIndex` of `TheCall` refers to memory in
3724// one of `AllowedSpaces`. Intended for HLSL builtins (e.g. atomics).
3725static bool CheckArgAddrSpaceOneOf(Sema *S, CallExpr *TheCall,
3726 unsigned ArgIndex,
3727 ArrayRef<LangAS> AllowedSpaces) {
3728 Expr *Arg = TheCall->getArg(ArgIndex);
3729 QualType LValueTy = Arg->IgnoreCasts()->getType();
3730 if (llvm::is_contained(AllowedSpaces, LValueTy.getAddressSpace()))
3731 return false;
3732 S->Diag(Arg->getBeginLoc(), diag::err_hlsl_atomic_arg_addr_space)
3733 << (ArgIndex + 1) << LValueTy;
3734 return true;
3735}
3736
3737static bool CheckNoDoubleVectors(Sema *S, SourceLocation Loc, int ArgOrdinal,
3738 clang::QualType PassedType) {
3739 const auto *VecTy = PassedType->getAs<VectorType>();
3740 if (!VecTy)
3741 return false;
3742
3743 if (VecTy->getElementType()->isDoubleType())
3744 return S->Diag(Loc, diag::err_builtin_invalid_arg_type)
3745 << ArgOrdinal << /* scalar */ 1 << /* no int */ 0 << /* fp */ 1
3746 << PassedType;
3747 return false;
3748}
3749
3751 int ArgOrdinal,
3752 clang::QualType PassedType) {
3753 if (!PassedType->hasIntegerRepresentation() &&
3754 !PassedType->hasFloatingRepresentation())
3755 return S->Diag(Loc, diag::err_builtin_invalid_arg_type)
3756 << ArgOrdinal << /* scalar or vector of */ 5 << /* integer */ 1
3757 << /* fp */ 1 << PassedType;
3758 return false;
3759}
3760
3762 int ArgOrdinal,
3763 clang::QualType PassedType) {
3764 if (auto *VecTy = PassedType->getAs<VectorType>())
3765 if (VecTy->getElementType()->isUnsignedIntegerType())
3766 return false;
3767
3768 return S->Diag(Loc, diag::err_builtin_invalid_arg_type)
3769 << ArgOrdinal << /* vector of */ 4 << /* uint */ 3 << /* no fp */ 0
3770 << PassedType;
3771}
3772
3773// checks for unsigned ints of all sizes
3775 int ArgOrdinal,
3776 clang::QualType PassedType) {
3777 if (!PassedType->hasUnsignedIntegerRepresentation())
3778 return S->Diag(Loc, diag::err_builtin_invalid_arg_type)
3779 << ArgOrdinal << /* scalar or vector of */ 5 << /* unsigned int */ 3
3780 << /* no fp */ 0 << PassedType;
3781 return false;
3782}
3783
3784static bool CheckExpectedBitWidth(Sema *S, CallExpr *TheCall,
3785 unsigned ArgOrdinal, unsigned Width) {
3786 QualType ArgTy = TheCall->getArg(0)->getType();
3787 if (auto *VTy = ArgTy->getAs<VectorType>())
3788 ArgTy = VTy->getElementType();
3789 // ensure arg type has expected bit width
3790 uint64_t ElementBitCount =
3792 if (ElementBitCount != Width) {
3793 S->Diag(TheCall->getArg(0)->getBeginLoc(),
3794 diag::err_integer_incorrect_bit_count)
3795 << Width << ElementBitCount;
3796 return true;
3797 }
3798 return false;
3799}
3800
3802 QualType ReturnType) {
3803 if (auto *VecTyA = TheCall->getArg(0)->getType()->getAs<VectorType>())
3804 ReturnType =
3805 S->Context.getExtVectorType(ReturnType, VecTyA->getNumElements());
3806 else if (auto *MatTyA =
3807 TheCall->getArg(0)->getType()->getAs<ConstantMatrixType>())
3808 ReturnType = S->Context.getConstantMatrixType(
3809 ReturnType, MatTyA->getNumRows(), MatTyA->getNumColumns());
3810
3811 TheCall->setType(ReturnType);
3812}
3813
3814static bool CheckScalarOrVector(Sema *S, CallExpr *TheCall, QualType Scalar,
3815 unsigned ArgIndex) {
3816 assert(TheCall->getNumArgs() >= ArgIndex);
3817 QualType ArgType = TheCall->getArg(ArgIndex)->getType();
3818 auto *VTy = ArgType->getAs<VectorType>();
3819 // not the scalar or vector<scalar>
3820 if (!(S->Context.hasSameUnqualifiedType(ArgType, Scalar) ||
3821 (VTy &&
3822 S->Context.hasSameUnqualifiedType(VTy->getElementType(), Scalar)))) {
3823 S->Diag(TheCall->getArg(0)->getBeginLoc(),
3824 diag::err_typecheck_expect_scalar_or_vector)
3825 << ArgType << Scalar;
3826 return true;
3827 }
3828 return false;
3829}
3830
3832 QualType Scalar, unsigned ArgIndex) {
3833 assert(TheCall->getNumArgs() > ArgIndex);
3834
3835 Expr *Arg = TheCall->getArg(ArgIndex);
3836 QualType ArgType = Arg->getType();
3837
3838 // Scalar: T
3839 if (S->Context.hasSameUnqualifiedType(ArgType, Scalar))
3840 return false;
3841
3842 // Vector: vector<T>
3843 if (const auto *VTy = ArgType->getAs<VectorType>()) {
3844 if (S->Context.hasSameUnqualifiedType(VTy->getElementType(), Scalar))
3845 return false;
3846 }
3847
3848 // Matrix: ConstantMatrixType with element type T
3849 if (const auto *MTy = ArgType->getAs<ConstantMatrixType>()) {
3850 if (S->Context.hasSameUnqualifiedType(MTy->getElementType(), Scalar))
3851 return false;
3852 }
3853
3854 // Not a scalar/vector/matrix-of-scalar
3855 S->Diag(Arg->getBeginLoc(),
3856 diag::err_typecheck_expect_scalar_or_vector_or_matrix)
3857 << ArgType << Scalar;
3858 return true;
3859}
3860
3861static bool CheckAnyScalarOrVector(Sema *S, CallExpr *TheCall,
3862 unsigned ArgIndex) {
3863 assert(TheCall->getNumArgs() >= ArgIndex);
3864 QualType ArgType = TheCall->getArg(ArgIndex)->getType();
3865 auto *VTy = ArgType->getAs<VectorType>();
3866 // not the scalar or vector<scalar>
3867 if (!(ArgType->isScalarType() ||
3868 (VTy && VTy->getElementType()->isScalarType()))) {
3869 S->Diag(TheCall->getArg(0)->getBeginLoc(),
3870 diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
3871 << ArgType << 1;
3872 return true;
3873 }
3874 return false;
3875}
3876
3878 unsigned ArgIndex) {
3879 assert(TheCall->getNumArgs() > ArgIndex);
3880 QualType ArgType = TheCall->getArg(ArgIndex)->getType();
3881 if (ArgType->isDependentType())
3882 return false;
3883
3884 QualType ElementType = ArgType;
3885 if (const auto *VectorTy = ArgType->getAs<VectorType>())
3886 ElementType = VectorTy->getElementType();
3887 else if (const auto *MatrixTy = ArgType->getAs<ConstantMatrixType>())
3888 ElementType = MatrixTy->getElementType();
3889
3890 if (ElementType->isBooleanType())
3891 return false;
3892
3893 if (ElementType->isIntegerType() || ElementType->isRealFloatingType()) {
3894 unsigned BitWidth = S->Context.getTypeSize(ElementType);
3895 if (BitWidth == 16 || BitWidth == 32 || BitWidth == 64)
3896 return false;
3897 }
3898
3899 S->Diag(TheCall->getArg(ArgIndex)->getBeginLoc(),
3900 diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
3901 << ArgType << 2;
3902 return true;
3903}
3904
3905// Check that the argument is not a bool or vector<bool>
3906// Returns true on error
3908 unsigned ArgIndex) {
3909 QualType BoolType = S->getASTContext().BoolTy;
3910 assert(ArgIndex < TheCall->getNumArgs());
3911 QualType ArgType = TheCall->getArg(ArgIndex)->getType();
3912 auto *VTy = ArgType->getAs<VectorType>();
3913 // is the bool or vector<bool>
3914 if (S->Context.hasSameUnqualifiedType(ArgType, BoolType) ||
3915 (VTy &&
3916 S->Context.hasSameUnqualifiedType(VTy->getElementType(), BoolType))) {
3917 S->Diag(TheCall->getArg(0)->getBeginLoc(),
3918 diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
3919 << ArgType << 0;
3920 return true;
3921 }
3922 return false;
3923}
3924
3925static bool CheckWaveActive(Sema *S, CallExpr *TheCall) {
3926 if (CheckNotBoolScalarOrVector(S, TheCall, 0))
3927 return true;
3928 return false;
3929}
3930
3931static bool CheckWavePrefix(Sema *S, CallExpr *TheCall) {
3932 if (CheckNotBoolScalarOrVector(S, TheCall, 0))
3933 return true;
3934 return false;
3935}
3936
3937static bool CheckBoolSelect(Sema *S, CallExpr *TheCall) {
3938 assert(TheCall->getNumArgs() == 3);
3939 Expr *Arg1 = TheCall->getArg(1);
3940 Expr *Arg2 = TheCall->getArg(2);
3941 if (!S->Context.hasSameUnqualifiedType(Arg1->getType(), Arg2->getType())) {
3942 S->Diag(TheCall->getBeginLoc(),
3943 diag::err_typecheck_call_different_arg_types)
3944 << Arg1->getType() << Arg2->getType() << Arg1->getSourceRange()
3945 << Arg2->getSourceRange();
3946 return true;
3947 }
3948
3949 TheCall->setType(Arg1->getType());
3950 return false;
3951}
3952
3953static bool CheckVectorSelect(Sema *S, CallExpr *TheCall) {
3954 assert(TheCall->getNumArgs() == 3);
3955 Expr *Arg1 = TheCall->getArg(1);
3956 QualType Arg1Ty = Arg1->getType();
3957 Expr *Arg2 = TheCall->getArg(2);
3958 QualType Arg2Ty = Arg2->getType();
3959
3960 QualType Arg1ScalarTy = Arg1Ty;
3961 if (auto VTy = Arg1ScalarTy->getAs<VectorType>())
3962 Arg1ScalarTy = VTy->getElementType();
3963
3964 QualType Arg2ScalarTy = Arg2Ty;
3965 if (auto VTy = Arg2ScalarTy->getAs<VectorType>())
3966 Arg2ScalarTy = VTy->getElementType();
3967
3968 if (!S->Context.hasSameUnqualifiedType(Arg1ScalarTy, Arg2ScalarTy))
3969 S->Diag(Arg1->getBeginLoc(), diag::err_hlsl_builtin_scalar_vector_mismatch)
3970 << /* second and third */ 1 << TheCall->getCallee() << Arg1Ty << Arg2Ty;
3971
3972 QualType Arg0Ty = TheCall->getArg(0)->getType();
3973 unsigned Arg0Length = Arg0Ty->getAs<VectorType>()->getNumElements();
3974 unsigned Arg1Length = Arg1Ty->isVectorType()
3975 ? Arg1Ty->getAs<VectorType>()->getNumElements()
3976 : 0;
3977 unsigned Arg2Length = Arg2Ty->isVectorType()
3978 ? Arg2Ty->getAs<VectorType>()->getNumElements()
3979 : 0;
3980 if (Arg1Length > 0 && Arg0Length != Arg1Length) {
3981 S->Diag(TheCall->getBeginLoc(),
3982 diag::err_typecheck_vector_lengths_not_equal)
3983 << Arg0Ty << Arg1Ty << TheCall->getArg(0)->getSourceRange()
3984 << Arg1->getSourceRange();
3985 return true;
3986 }
3987
3988 if (Arg2Length > 0 && Arg0Length != Arg2Length) {
3989 S->Diag(TheCall->getBeginLoc(),
3990 diag::err_typecheck_vector_lengths_not_equal)
3991 << Arg0Ty << Arg2Ty << TheCall->getArg(0)->getSourceRange()
3992 << Arg2->getSourceRange();
3993 return true;
3994 }
3995
3996 TheCall->setType(
3997 S->getASTContext().getExtVectorType(Arg1ScalarTy, Arg0Length));
3998 return false;
3999}
4000
4001static bool CheckMatrixSelect(Sema *S, CallExpr *TheCall) {
4002 assert(TheCall->getNumArgs() == 3);
4003 Expr *Arg1 = TheCall->getArg(1);
4004 QualType Arg1Ty = Arg1->getType();
4005 Expr *Arg2 = TheCall->getArg(2);
4006 QualType Arg2Ty = Arg2->getType();
4007
4008 QualType Arg1ScalarTy = Arg1Ty;
4009 if (auto MTy = Arg1ScalarTy->getAs<ConstantMatrixType>())
4010 Arg1ScalarTy = MTy->getElementType();
4011
4012 QualType Arg2ScalarTy = Arg2Ty;
4013 if (auto MTy = Arg2ScalarTy->getAs<ConstantMatrixType>())
4014 Arg2ScalarTy = MTy->getElementType();
4015
4016 if (!S->Context.hasSameUnqualifiedType(Arg1ScalarTy, Arg2ScalarTy))
4017 S->Diag(Arg1->getBeginLoc(), diag::err_hlsl_builtin_scalar_vector_mismatch)
4018 << /* second and third */ 1 << TheCall->getCallee() << Arg1Ty << Arg2Ty;
4019
4020 QualType Arg0Ty = TheCall->getArg(0)->getType();
4021 auto *Arg0MatTy = Arg0Ty->getAs<ConstantMatrixType>();
4022 unsigned Arg0Rows = Arg0MatTy->getNumRows();
4023 unsigned Arg0Cols = Arg0MatTy->getNumColumns();
4024
4025 for (Expr *Arg : {Arg1, Arg2}) {
4026 auto *MTy = Arg->getType()->getAs<ConstantMatrixType>();
4027 if (MTy &&
4028 (MTy->getNumRows() != Arg0Rows || MTy->getNumColumns() != Arg0Cols)) {
4029 S->Diag(TheCall->getBeginLoc(),
4030 diag::err_typecheck_vector_lengths_not_equal)
4031 << Arg0Ty << Arg->getType() << TheCall->getArg(0)->getSourceRange()
4032 << Arg->getSourceRange();
4033 return true;
4034 }
4035 }
4036
4037 TheCall->setType(
4038 S->Context.getConstantMatrixType(Arg1ScalarTy, Arg0Rows, Arg0Cols));
4039 return false;
4040}
4041
4043 unsigned Count) {
4044 return Count > 1 ? S.Context.getExtVectorType(BaseType, Count) : BaseType;
4045}
4046
4047static bool CheckScalarFloatOperand(Sema &S, CallExpr *TheCall,
4048 unsigned ArgIndex) {
4049 return CheckArgTypeMatches(&S, TheCall->getArg(ArgIndex), S.Context.FloatTy);
4050}
4051
4052static bool CheckIndexType(Sema *S, CallExpr *TheCall, unsigned IndexArgIndex) {
4053 assert(TheCall->getNumArgs() > IndexArgIndex && "Index argument missing");
4054 QualType ArgType = TheCall->getArg(IndexArgIndex)->getType();
4055 QualType IndexTy = ArgType;
4056 unsigned int ActualDim = 1;
4057 if (const auto *VTy = IndexTy->getAs<VectorType>()) {
4058 ActualDim = VTy->getNumElements();
4059 IndexTy = VTy->getElementType();
4060 }
4061 if (!IndexTy->isIntegerType()) {
4062 S->Diag(TheCall->getArg(IndexArgIndex)->getBeginLoc(),
4063 diag::err_typecheck_expect_int)
4064 << ArgType;
4065 return true;
4066 }
4067
4068 QualType ResourceArgTy = TheCall->getArg(0)->getType();
4069 const HLSLAttributedResourceType *ResTy =
4070 ResourceArgTy.getTypePtr()->getAs<HLSLAttributedResourceType>();
4071 assert(ResTy && "Resource argument must be a resource");
4072 HLSLAttributedResourceType::Attributes ResAttrs = ResTy->getAttrs();
4073
4074 unsigned int ExpectedDim = 1;
4075 if (ResAttrs.ResourceDimension != llvm::dxil::ResourceDimension::Unknown)
4076 ExpectedDim = getResourceDimensions(ResAttrs.ResourceDimension) +
4077 (ResAttrs.IsArray ? 1 : 0);
4078
4079 if (ActualDim != ExpectedDim) {
4080 S->Diag(TheCall->getArg(IndexArgIndex)->getBeginLoc(),
4081 diag::err_hlsl_builtin_resource_coordinate_dimension_mismatch)
4082 << cast<NamedDecl>(TheCall->getCalleeDecl()) << ExpectedDim
4083 << ActualDim;
4084 return true;
4085 }
4086
4087 return false;
4088}
4089
4091 Sema *S, CallExpr *TheCall, unsigned ArgIndex,
4092 llvm::function_ref<bool(const HLSLAttributedResourceType *ResType)> Check =
4093 nullptr) {
4094 assert(TheCall->getNumArgs() >= ArgIndex);
4095 QualType ArgType = TheCall->getArg(ArgIndex)->getType();
4096 const HLSLAttributedResourceType *ResTy =
4097 ArgType.getTypePtr()->getAs<HLSLAttributedResourceType>();
4098 if (!ResTy) {
4099 S->Diag(TheCall->getArg(ArgIndex)->getBeginLoc(),
4100 diag::err_typecheck_expect_hlsl_resource)
4101 << ArgType;
4102 return true;
4103 }
4104 if (Check && Check(ResTy)) {
4105 S->Diag(TheCall->getArg(ArgIndex)->getExprLoc(),
4106 diag::err_invalid_hlsl_resource_type)
4107 << ArgType;
4108 return true;
4109 }
4110 return false;
4111}
4112
4114 QualType MainHandleTy) {
4115 assert(MainHandleTy->isHLSLAttributedResourceType() &&
4116 "expected resource handle type");
4117 auto *MainResType = MainHandleTy->getAs<HLSLAttributedResourceType>();
4118 auto MainAttrs = MainResType->getAttrs();
4119 assert(!MainAttrs.IsCounter && "cannot create a counter from a counter");
4120 MainAttrs.IsCounter = true;
4121 return AST.getHLSLAttributedResourceType(MainResType->getWrappedType(),
4122 MainResType->getContainedType(),
4123 MainAttrs);
4124}
4125
4126enum class SampleKind { Sample, Bias, Grad, Level, Cmp, CmpLevelZero };
4127
4128static StringRef getSampleMethodName(SampleKind Kind) {
4129 switch (Kind) {
4130 case SampleKind::Sample:
4131 return "Sample";
4132 case SampleKind::Bias:
4133 return "SampleBias";
4134 case SampleKind::Grad:
4135 return "SampleGrad";
4136 case SampleKind::Level:
4137 return "SampleLevel";
4138 case SampleKind::Cmp:
4139 return "SampleCmp";
4141 return "SampleCmpLevelZero";
4142 }
4143 llvm_unreachable("Invalid SampleKind");
4144}
4145
4146// Returns the name of the resource method whose body the sampling or gather
4147// builtin is being emitted into, which is the name the user called. This
4148// matters for methods that share a builtin, like 'Gather' and 'GatherRed'.
4149// Falls back to DefaultName if the builtin is used outside of a resource
4150// method.
4151static StringRef getCurrentResourceMethodName(Sema &S, StringRef DefaultName) {
4152 const auto *MD = dyn_cast_if_present<CXXMethodDecl>(S.getCurFunctionDecl());
4153 if (!MD || !MD->getDeclName().isIdentifier())
4154 return DefaultName;
4155
4156 QualType RecordTy = S.Context.getCanonicalTagType(MD->getParent());
4157 if (!RecordTy->isHLSLResourceRecord())
4158 return DefaultName;
4159
4160 return MD->getName();
4161}
4162
4163// Sampling from and gathering on resources with a 'double' element type is not
4164// supported. Such resources are still valid declarations whose contents can be
4165// accessed by other means, like Load or the subscript operator.
4166static bool CheckNoDoubleElementType(Sema &S, CallExpr *TheCall,
4167 QualType ContainedType,
4168 StringRef DefaultName) {
4169 QualType EltTy = getElementTypeOf(ContainedType, /*IncludeMatrix=*/false);
4170 if (!EltTy->isSpecificBuiltinType(BuiltinType::Double))
4171 return false;
4172
4173 S.Diag(TheCall->getBeginLoc(), diag::err_hlsl_sample_double_element_type)
4174 << getCurrentResourceMethodName(S, DefaultName) << ContainedType;
4175 return true;
4176}
4177
4178// Sampling textures with an integer element type was introduced in SM 6.7 as
4179// part of Advanced Texture Operations. The shader model only applies to DirectX
4180// targets; Vulkan has no such restriction.
4182 QualType ContainedType,
4183 SampleKind Kind) {
4184 // Comparison sampling requires a floating point element type at every shader
4185 // model, which the caller diagnoses.
4186 if (Kind == SampleKind::Cmp || Kind == SampleKind::CmpLevelZero)
4187 return false;
4188
4189 // 'bool' is an integer type in HLSL, but sampling bool resources is never
4190 // allowed, so it must not be reported as requiring shader model 6.7.
4191 QualType EltTy = getElementTypeOf(ContainedType, /*IncludeMatrix=*/false);
4192 if (!EltTy->isIntegerType() || EltTy->isBooleanType())
4193 return false;
4194
4195 const TargetInfo &TI = S.Context.getTargetInfo();
4196 if (!TI.getTriple().isDXIL())
4197 return false;
4198
4199 VersionTuple SMVersion = TI.getPlatformMinVersion();
4200 if (SMVersion >= VersionTuple(6, 7))
4201 return false;
4202
4203 S.Diag(TheCall->getBeginLoc(), diag::err_hlsl_sample_integer_element_type)
4205 << ContainedType << SMVersion.getAsString();
4206 return true;
4207}
4208
4210 bool IncludeArraySlice = true) {
4211 // Check the texture handle.
4212 if (CheckResourceHandle(&S, TheCall, 0,
4213 [](const HLSLAttributedResourceType *ResType) {
4214 return ResType->getAttrs().ResourceDimension ==
4215 llvm::dxil::ResourceDimension::Unknown;
4216 }))
4217 return true;
4218
4219 // Check the sampler handle.
4220 if (CheckResourceHandle(&S, TheCall, 1,
4221 [](const HLSLAttributedResourceType *ResType) {
4222 return ResType->getAttrs().ResourceClass !=
4223 llvm::hlsl::ResourceClass::Sampler;
4224 }))
4225 return true;
4226
4227 auto *ResourceTy =
4228 TheCall->getArg(0)->getType()->castAs<HLSLAttributedResourceType>();
4229
4230 // Check the location.
4231 unsigned ExpectedDim =
4232 getResourceDimensions(ResourceTy->getAttrs().ResourceDimension) +
4233 (IncludeArraySlice && ResourceTy->getAttrs().IsArray ? 1 : 0);
4235 &S, TheCall->getArg(2),
4236 getVectorOrScalarType(S, S.Context.FloatTy, ExpectedDim)))
4237 return true;
4238
4239 return false;
4240}
4241
4242static bool CheckCalculateLodBuiltin(Sema &S, CallExpr *TheCall) {
4243 if (S.checkArgCount(TheCall, 3))
4244 return true;
4245
4246 // CalculateLevelOfDetail location uses resource dimension only (e.g. float2
4247 // for 2D), not an extra array slice component like Sample/Gather.
4248 if (CheckTextureSamplerAndLocation(S, TheCall, /*IncludeArraySlice=*/false))
4249 return true;
4250
4251 TheCall->setType(S.Context.FloatTy);
4252 return false;
4253}
4254
4255static bool CheckGatherBuiltin(Sema &S, CallExpr *TheCall, bool IsCmp) {
4256 if (S.checkArgCountRange(TheCall, IsCmp ? 5 : 4, IsCmp ? 6 : 5))
4257 return true;
4258
4259 if (CheckTextureSamplerAndLocation(S, TheCall))
4260 return true;
4261
4262 unsigned NextIdx = 3;
4263 if (IsCmp) {
4264 // Check the compare value.
4265 if (CheckScalarFloatOperand(S, TheCall, NextIdx))
4266 return true;
4267 NextIdx++;
4268 }
4269
4270 // Check the component operand.
4271 if (CheckArgTypeMatches(&S, TheCall->getArg(NextIdx),
4273 return true;
4274 Expr *ComponentArg = TheCall->getArg(NextIdx);
4275
4276 // GatherCmp operations on Vulkan target must use component 0 (Red).
4277 if (IsCmp && S.getASTContext().getTargetInfo().getTriple().isSPIRV()) {
4278 std::optional<llvm::APSInt> ComponentOpt =
4279 ComponentArg->getIntegerConstantExpr(S.getASTContext());
4280 if (ComponentOpt) {
4281 int64_t ComponentVal = ComponentOpt->getSExtValue();
4282 if (ComponentVal != 0) {
4283 // Issue an error if the component is not 0 (Red).
4284 // 0 -> Red, 1 -> Green, 2 -> Blue, 3 -> Alpha
4285 assert(ComponentVal >= 0 && ComponentVal <= 3 &&
4286 "The component is not in the expected range.");
4287 S.Diag(ComponentArg->getBeginLoc(),
4288 diag::err_hlsl_gathercmp_invalid_component)
4289 << ComponentVal;
4290 return true;
4291 }
4292 }
4293 }
4294
4295 NextIdx++;
4296
4297 // Check the offset operand.
4298 const HLSLAttributedResourceType *ResourceTy =
4299 TheCall->getArg(0)->getType()->castAs<HLSLAttributedResourceType>();
4300 if (TheCall->getNumArgs() > NextIdx) {
4301 unsigned ExpectedDim =
4302 getResourceDimensions(ResourceTy->getAttrs().ResourceDimension);
4304 &S, TheCall->getArg(NextIdx),
4305 getVectorOrScalarType(S, S.Context.IntTy, ExpectedDim)))
4306 return true;
4307 NextIdx++;
4308 }
4309
4310 assert(ResourceTy->hasContainedType() &&
4311 "Expecting a contained type for resource with a dimension "
4312 "attribute.");
4313 QualType ReturnType = ResourceTy->getContainedType();
4314
4315 if (CheckNoDoubleElementType(S, TheCall, ReturnType,
4316 IsCmp ? "GatherCmp" : "Gather"))
4317 return true;
4318
4319 if (IsCmp) {
4320 if (!ReturnType->hasFloatingRepresentation()) {
4321 S.Diag(TheCall->getBeginLoc(), diag::err_hlsl_samplecmp_requires_float);
4322 return true;
4323 }
4324 }
4325
4326 if (const auto *VecTy = ReturnType->getAs<VectorType>())
4327 ReturnType = VecTy->getElementType();
4328 ReturnType = S.Context.getExtVectorType(ReturnType, 4);
4329
4330 TheCall->setType(ReturnType);
4331
4332 return false;
4333}
4334static bool CheckLoadLevelBuiltin(Sema &S, CallExpr *TheCall) {
4335 if (S.checkArgCountRange(TheCall, 2, 3))
4336 return true;
4337
4338 // Check the texture handle.
4339 if (CheckResourceHandle(&S, TheCall, 0,
4340 [](const HLSLAttributedResourceType *ResType) {
4341 return ResType->getAttrs().ResourceDimension ==
4342 llvm::dxil::ResourceDimension::Unknown;
4343 }))
4344 return true;
4345
4346 auto *ResourceTy =
4347 TheCall->getArg(0)->getType()->castAs<HLSLAttributedResourceType>();
4348
4349 // A UAV descriptor binds a single mip slice, so a RWTexture location has no
4350 // mip component to select, and TextureLoad on a UAV takes no offset.
4351 bool IsUAV =
4352 ResourceTy->getAttrs().ResourceClass == llvm::dxil::ResourceClass::UAV;
4353 if (IsUAV && S.checkArgCount(TheCall, 2))
4354 return true;
4355
4356 // Check the location: int3 for Texture2D and int4 for Texture2DArray, which
4357 // both carry a trailing mip level; int2 and int3 for the RWTexture forms,
4358 // which do not.
4359 unsigned ResourceDim =
4360 getResourceDimensions(ResourceTy->getAttrs().ResourceDimension);
4361 unsigned LocationDim = ResourceDim + (ResourceTy->getAttrs().IsArray ? 1 : 0);
4362 if (!IsUAV)
4363 ++LocationDim;
4365 &S, TheCall->getArg(1),
4366 getVectorOrScalarType(S, S.Context.IntTy, LocationDim)))
4367 return true;
4368
4369 // Check the offset operand (int2 for 2D textures; no array slice).
4370 if (TheCall->getNumArgs() > 2) {
4372 &S, TheCall->getArg(2),
4373 getVectorOrScalarType(S, S.Context.IntTy, ResourceDim)))
4374 return true;
4375 }
4376
4377 TheCall->setType(ResourceTy->getContainedType());
4378 return false;
4379}
4380
4381static bool CheckLoadMSBuiltin(Sema &S, CallExpr *TheCall) {
4382 if (S.checkArgCountRange(TheCall, 3, 4))
4383 return true;
4384
4385 // Check the multisampled texture handle.
4386 if (CheckResourceHandle(&S, TheCall, 0,
4387 [](const HLSLAttributedResourceType *ResType) {
4388 return !ResType->isMultiSampled();
4389 }))
4390 return true;
4391
4392 auto *ResourceTy =
4393 TheCall->getArg(0)->getType()->castAs<HLSLAttributedResourceType>();
4394
4395 // Check the location (int2 for Texture2DMS, int3 for Texture2DMSArray).
4396 // Unlike Load on regular textures, there is no mip/LOD component.
4397 unsigned ResourceDim =
4398 getResourceDimensions(ResourceTy->getAttrs().ResourceDimension);
4399 unsigned LocationDim = ResourceDim + (ResourceTy->getAttrs().IsArray ? 1 : 0);
4401 &S, TheCall->getArg(1),
4402 getVectorOrScalarType(S, S.Context.IntTy, LocationDim)))
4403 return true;
4404
4405 // Check the sample index operand (scalar int).
4406 if (CheckArgTypeMatches(&S, TheCall->getArg(2), S.Context.IntTy))
4407 return true;
4408
4409 // Check the offset operand (int2 for 2D textures; no array slice).
4410 if (TheCall->getNumArgs() > 3) {
4412 &S, TheCall->getArg(3),
4413 getVectorOrScalarType(S, S.Context.IntTy, ResourceDim)))
4414 return true;
4415 }
4416
4417 TheCall->setType(ResourceTy->getContainedType());
4418 return false;
4419}
4420
4421static bool CheckSamplingBuiltin(Sema &S, CallExpr *TheCall, SampleKind Kind) {
4422 unsigned MinArgs, MaxArgs;
4423 if (Kind == SampleKind::Sample) {
4424 MinArgs = 3;
4425 MaxArgs = 5;
4426 } else if (Kind == SampleKind::Bias) {
4427 MinArgs = 4;
4428 MaxArgs = 6;
4429 } else if (Kind == SampleKind::Grad) {
4430 MinArgs = 5;
4431 MaxArgs = 7;
4432 } else if (Kind == SampleKind::Level) {
4433 MinArgs = 4;
4434 MaxArgs = 5;
4435 } else if (Kind == SampleKind::Cmp) {
4436 MinArgs = 4;
4437 MaxArgs = 6;
4438 } else {
4439 assert(Kind == SampleKind::CmpLevelZero);
4440 MinArgs = 4;
4441 MaxArgs = 5;
4442 }
4443
4444 if (S.checkArgCountRange(TheCall, MinArgs, MaxArgs))
4445 return true;
4446
4447 if (CheckTextureSamplerAndLocation(S, TheCall))
4448 return true;
4449
4450 const HLSLAttributedResourceType *ResourceTy =
4451 TheCall->getArg(0)->getType()->castAs<HLSLAttributedResourceType>();
4452 unsigned ExpectedDim =
4453 getResourceDimensions(ResourceTy->getAttrs().ResourceDimension);
4454
4455 unsigned NextIdx = 3;
4456 if (Kind == SampleKind::Bias || Kind == SampleKind::Level ||
4457 Kind == SampleKind::Cmp || Kind == SampleKind::CmpLevelZero) {
4458 // Check the bias, lod level, or compare value, depending on the kind.
4459 // All of them must be a scalar float value.
4460 if (CheckScalarFloatOperand(S, TheCall, NextIdx))
4461 return true;
4462 NextIdx++;
4463 } else if (Kind == SampleKind::Grad) {
4464 QualType GradTy = getVectorOrScalarType(S, S.Context.FloatTy, ExpectedDim);
4465
4466 // Check the DDX operand.
4467 if (CheckArgTypeMatches(&S, TheCall->getArg(NextIdx), GradTy))
4468 return true;
4469
4470 // Check the DDY operand.
4471 if (CheckArgTypeMatches(&S, TheCall->getArg(NextIdx + 1), GradTy))
4472 return true;
4473 NextIdx += 2;
4474 }
4475
4476 // Check the offset operand (if applicable).
4477 if (hasResourceOffset(ResourceTy->getAttrs().ResourceDimension) &&
4478 TheCall->getNumArgs() > NextIdx) {
4480 &S, TheCall->getArg(NextIdx),
4481 getVectorOrScalarType(S, S.Context.IntTy, ExpectedDim)))
4482 return true;
4483 NextIdx++;
4484 }
4485
4486 // Check the clamp operand.
4487 if (Kind != SampleKind::Level && Kind != SampleKind::CmpLevelZero &&
4488 TheCall->getNumArgs() > NextIdx) {
4489 if (CheckScalarFloatOperand(S, TheCall, NextIdx))
4490 return true;
4491 }
4492
4493 assert(ResourceTy->hasContainedType() &&
4494 "Expecting a contained type for resource with a dimension "
4495 "attribute.");
4496 QualType ReturnType = ResourceTy->getContainedType();
4497
4498 if (CheckNoDoubleElementType(S, TheCall, ReturnType,
4499 getSampleMethodName(Kind)))
4500 return true;
4501
4502 if (CheckIntegerElementTypeShaderModel(S, TheCall, ReturnType, Kind))
4503 return true;
4504
4505 if (Kind == SampleKind::Cmp || Kind == SampleKind::CmpLevelZero) {
4506 if (!ReturnType->hasFloatingRepresentation()) {
4507 S.Diag(TheCall->getBeginLoc(), diag::err_hlsl_samplecmp_requires_float);
4508 return true;
4509 }
4510 ReturnType = S.Context.FloatTy;
4511 }
4512 TheCall->setType(ReturnType);
4513
4514 return false;
4515}
4516
4517/// The `dest` types an interlocked operation accepts. Float is 32-bit only.
4519
4520/// Check a call to an HLSL interlocked builtin. The builtins are variadic, so
4521/// this is the only check a direct call gets. Overload resolution checks the
4522/// calls that come through the `InterlockedOp` overload sets.
4523static bool CheckInterlockedBuiltin(Sema &S, CallExpr *TheCall,
4524 unsigned MinArgs, unsigned MaxArgs,
4525 InterlockedDest Dest,
4526 bool ReportsOriginalValue) {
4527 if (MinArgs == MaxArgs) {
4528 if (S.checkArgCount(TheCall, MinArgs))
4529 return true;
4530 } else if (TheCall->getNumArgs() < MinArgs) {
4531 S.Diag(TheCall->getEndLoc(), diag::err_typecheck_call_too_few_args_at_least)
4532 << /*callee_type=*/0 << /*min_arg_count=*/MinArgs
4533 << TheCall->getNumArgs() << /*is_non_object=*/0
4534 << TheCall->getSourceRange();
4535 return true;
4536 } else if (S.checkArgCountAtMost(TheCall, MaxArgs)) {
4537 return true;
4538 }
4539
4540 QualType DestTy = TheCall->getArg(0)->getType().getUnqualifiedType();
4541 const bool DestIsOK =
4542 DestTy->isSpecificBuiltinType(BuiltinType::Float)
4543 ? Dest != InterlockedDest::Int
4544 : Dest != InterlockedDest::Float && DestTy->isIntegerType();
4545 if (!DestIsOK) {
4546 S.Diag(TheCall->getArg(0)->getBeginLoc(),
4547 diag::err_builtin_invalid_arg_type)
4548 << /*ordinal=*/1 << /*scalar*/ 1
4549 << /*integer*/ (Dest == InterlockedDest::Float ? 0 : 1)
4550 << /*32 bit floating-point*/ (Dest == InterlockedDest::Int ? 0 : 3)
4551 << DestTy;
4552 return true;
4553 }
4554
4555 // 64-bit interlocked ops require SM 6.6 on DXIL. The synthesized wrapper
4556 // methods (e.g. RWByteAddressBuffer::InterlockedAdd64) are only declared on
4557 // SM 6.6+, so this defensive check only fires for direct builtin calls; skip
4558 // synthetic invocations (invalid source location).
4559 const TargetInfo &TI = S.Context.getTargetInfo();
4560 if (TheCall->getBeginLoc().isValid() &&
4561 TI.getTriple().getArch() == llvm::Triple::dxil &&
4562 S.Context.getTypeSize(DestTy) == 64 &&
4563 TI.getPlatformMinVersion() < VersionTuple(6, 6)) {
4564 S.Diag(TheCall->getBeginLoc(), diag::err_hlsl_builtin_requires_sm)
4565 << TheCall->getDirectCallee() << VersionTuple(6, 6).getAsString();
4566 return true;
4567 }
4568
4569 if (CheckModifiableLValue(&S, TheCall, 0))
4570 return true;
4571
4572 if (CheckArgAddrSpaceOneOf(&S, TheCall, 0,
4574 return true;
4575
4576 // Every argument after `dest` has the destination's type.
4577 for (unsigned I = 1, E = TheCall->getNumArgs(); I != E; ++I)
4578 if (CheckArgTypeMatches(&S, TheCall->getArg(I), DestTy))
4579 return true;
4580
4581 // Operations that report the previous value write it back through their last
4582 // argument.
4583 const unsigned NumArgs = TheCall->getNumArgs();
4584 if (ReportsOriginalValue && NumArgs == MaxArgs &&
4585 CheckModifiableLValue(&S, TheCall, NumArgs - 1))
4586 return true;
4587
4588 TheCall->setType(S.Context.VoidTy);
4589 return false;
4590}
4591
4592// Note: returning true in this case results in CheckBuiltinFunctionCall
4593// returning an ExprError
4594bool SemaHLSL::CheckBuiltinFunctionCall(unsigned BuiltinID, CallExpr *TheCall) {
4595 switch (BuiltinID) {
4596 case Builtin::BI__builtin_hlsl_barrier: {
4597 if (SemaRef.checkArgCount(TheCall, 2))
4598 return true;
4599
4600 if (SemaRef.Context.getTargetInfo().getTriple().getArch() !=
4601 llvm::Triple::dxil) {
4602 SemaRef.Diag(TheCall->getExprLoc(), diag::err_hlsl_dxil_only)
4603 << "Barrier";
4604 return true;
4605 }
4606
4607 Expr *MemoryArg = TheCall->getArg(0);
4608 if (MemoryArg->getType()->isUnsignedIntegerType()) {
4609 std::optional<llvm::APSInt> MemoryFlags =
4610 MemoryArg->getIntegerConstantExpr(SemaRef.Context);
4611 if (!MemoryFlags) {
4612 SemaRef.Diag(MemoryArg->getExprLoc(),
4613 diag::err_constant_integer_arg_type)
4614 << "Barrier";
4615 return true;
4616 }
4617 if ((MemoryFlags->getZExtValue() &
4618 ~barrierFlagValue(BarrierMemoryTypeFlag::ValidMask)) != 0) {
4619 SemaRef.Diag(MemoryArg->getExprLoc(),
4620 diag::err_hlsl_invalid_barrier_memory_flags);
4621 return true;
4622 }
4623 } else {
4624 const HLSLAttributedResourceType *ResTy =
4625 HLSLAttributedResourceType::findHandleTypeOnResource(
4626 MemoryArg->getType().getTypePtr());
4627 if (!ResTy) {
4628 SemaRef.Diag(MemoryArg->getExprLoc(),
4629 diag::err_typecheck_expect_hlsl_resource)
4630 << MemoryArg->getType();
4631 return true;
4632 }
4633 if (ResTy->getAttrs().ResourceClass != ResourceClass::UAV) {
4634 SemaRef.Diag(MemoryArg->getExprLoc(),
4635 diag::err_invalid_hlsl_resource_type)
4636 << MemoryArg->getType();
4637 return true;
4638 }
4639 }
4640
4641 Expr *SemanticArg = TheCall->getArg(1);
4642 std::optional<llvm::APSInt> SemanticFlags =
4643 SemanticArg->getIntegerConstantExpr(SemaRef.Context);
4644 if (!SemanticFlags) {
4645 SemaRef.Diag(SemanticArg->getExprLoc(),
4646 diag::err_constant_integer_arg_type)
4647 << "Barrier";
4648 return true;
4649 }
4650 if ((SemanticFlags->getZExtValue() &
4651 ~barrierFlagValue(BarrierSemanticFlag::ValidMask)) != 0) {
4652 SemaRef.Diag(SemanticArg->getExprLoc(),
4653 diag::err_hlsl_invalid_barrier_semantic_flags);
4654 return true;
4655 }
4656
4657 TheCall->setType(SemaRef.Context.VoidTy);
4658 break;
4659 }
4660 case Builtin::BI__builtin_hlsl_adduint64: {
4661 if (SemaRef.checkArgCount(TheCall, 2))
4662 return true;
4663
4664 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
4666 return true;
4667
4668 // ensure arg integers are 32-bits
4669 if (CheckExpectedBitWidth(&SemaRef, TheCall, 0, 32))
4670 return true;
4671
4672 // ensure both args are vectors of total bit size of a multiple of 64
4673 auto *VTy = TheCall->getArg(0)->getType()->getAs<VectorType>();
4674 int NumElementsArg = VTy->getNumElements();
4675 if (NumElementsArg != 2 && NumElementsArg != 4) {
4676 SemaRef.Diag(TheCall->getBeginLoc(), diag::err_vector_incorrect_bit_count)
4677 << 1 /*a multiple of*/ << 64 << NumElementsArg * 32;
4678 return true;
4679 }
4680
4681 // ensure first arg and second arg have the same type
4682 if (CheckAllArgsHaveSameType(&SemaRef, TheCall))
4683 return true;
4684
4685 ExprResult A = TheCall->getArg(0);
4686 QualType ArgTyA = A.get()->getType();
4687 // return type is the same as the input type
4688 TheCall->setType(ArgTyA);
4689 break;
4690 }
4691 case Builtin::BI__builtin_hlsl_resource_getpointer: {
4692 if (SemaRef.checkArgCountRange(TheCall, 1, 2) ||
4693 CheckResourceHandle(&SemaRef, TheCall, 0) ||
4694 (TheCall->getNumArgs() == 2 && CheckIndexType(&SemaRef, TheCall, 1)))
4695 return true;
4696
4697 auto *ResourceTy =
4698 TheCall->getArg(0)->getType()->castAs<HLSLAttributedResourceType>();
4699 QualType ContainedTy = ResourceTy->getContainedType();
4700 auto ReturnType = SemaRef.Context.getAddrSpaceQualType(
4701 ContainedTy,
4702 getLangASFromResourceClass(ResourceTy->getAttrs().ResourceClass));
4703 ReturnType = SemaRef.Context.getPointerType(ReturnType);
4704 TheCall->setType(ReturnType);
4705
4706 break;
4707 }
4708 case Builtin::BI__builtin_hlsl_resource_getpointer_typed: {
4709 if (SemaRef.checkArgCount(TheCall, 3) ||
4710 CheckResourceHandle(&SemaRef, TheCall, 0) ||
4711 CheckIndexType(&SemaRef, TheCall, 1))
4712 return true;
4713
4714 QualType ElementTy = TheCall->getArg(2)->getType();
4715 assert(ElementTy->isPointerType() &&
4716 "expected pointer type for second argument");
4717 ElementTy = ElementTy->getPointeeType();
4718
4719 // Reject array types
4720 if (ElementTy->isArrayType())
4721 return SemaRef.Diag(
4722 cast<FunctionDecl>(SemaRef.CurContext)->getPointOfInstantiation(),
4723 diag::err_invalid_use_of_array_type);
4724
4725 auto *ResourceTy =
4726 TheCall->getArg(0)->getType()->castAs<HLSLAttributedResourceType>();
4727 auto ReturnType = SemaRef.Context.getAddrSpaceQualType(
4728 ElementTy,
4729 getLangASFromResourceClass(ResourceTy->getAttrs().ResourceClass));
4730 ReturnType = SemaRef.Context.getPointerType(ReturnType);
4731 TheCall->setType(ReturnType);
4732
4733 break;
4734 }
4735 case Builtin::BI__builtin_hlsl_transpose_if_memory_is_row_major: {
4736 if (SemaRef.checkArgCount(TheCall, 2) ||
4737 CheckArgTypeMatches(&SemaRef, TheCall->getArg(1),
4738 SemaRef.getASTContext().IntTy))
4739 return true;
4740
4741 TheCall->setType(TheCall->getArg(0)->getType());
4742
4743 break;
4744 }
4745 case Builtin::BI__builtin_hlsl_resource_load_with_status: {
4746 if (SemaRef.checkArgCount(TheCall, 3) ||
4747 CheckResourceHandle(&SemaRef, TheCall, 0) ||
4748 CheckArgTypeMatches(&SemaRef, TheCall->getArg(1),
4749 SemaRef.getASTContext().UnsignedIntTy) ||
4750 CheckArgTypeMatches(&SemaRef, TheCall->getArg(2),
4751 SemaRef.getASTContext().UnsignedIntTy) ||
4752 CheckModifiableLValue(&SemaRef, TheCall, 2))
4753 return true;
4754
4755 auto *ResourceTy =
4756 TheCall->getArg(0)->getType()->castAs<HLSLAttributedResourceType>();
4757 QualType ReturnType = ResourceTy->getContainedType();
4758 TheCall->setType(ReturnType);
4759
4760 break;
4761 }
4762 case Builtin::BI__builtin_hlsl_resource_load_with_status_typed: {
4763 if (SemaRef.checkArgCount(TheCall, 4) ||
4764 CheckResourceHandle(&SemaRef, TheCall, 0) ||
4765 CheckArgTypeMatches(&SemaRef, TheCall->getArg(1),
4766 SemaRef.getASTContext().UnsignedIntTy) ||
4767 CheckArgTypeMatches(&SemaRef, TheCall->getArg(2),
4768 SemaRef.getASTContext().UnsignedIntTy) ||
4769 CheckModifiableLValue(&SemaRef, TheCall, 2))
4770 return true;
4771
4772 QualType ReturnType = TheCall->getArg(3)->getType();
4773 assert(ReturnType->isPointerType() &&
4774 "expected pointer type for second argument");
4775 ReturnType = ReturnType->getPointeeType();
4776
4777 // Reject array types
4778 if (ReturnType->isArrayType())
4779 return SemaRef.Diag(
4780 cast<FunctionDecl>(SemaRef.CurContext)->getPointOfInstantiation(),
4781 diag::err_invalid_use_of_array_type);
4782
4783 TheCall->setType(ReturnType);
4784
4785 break;
4786 }
4787 case Builtin::BI__builtin_hlsl_resource_load_level:
4788 return CheckLoadLevelBuiltin(SemaRef, TheCall);
4789 case Builtin::BI__builtin_hlsl_resource_load_ms:
4790 return CheckLoadMSBuiltin(SemaRef, TheCall);
4791 case Builtin::BI__builtin_hlsl_resource_sample:
4793 case Builtin::BI__builtin_hlsl_resource_sample_bias:
4795 case Builtin::BI__builtin_hlsl_resource_sample_grad:
4797 case Builtin::BI__builtin_hlsl_resource_sample_level:
4799 case Builtin::BI__builtin_hlsl_resource_sample_cmp:
4801 case Builtin::BI__builtin_hlsl_resource_sample_cmp_level_zero:
4803 case Builtin::BI__builtin_hlsl_resource_calculate_lod:
4804 case Builtin::BI__builtin_hlsl_resource_calculate_lod_unclamped:
4805 return CheckCalculateLodBuiltin(SemaRef, TheCall);
4806 case Builtin::BI__builtin_hlsl_resource_gather:
4807 return CheckGatherBuiltin(SemaRef, TheCall, /*IsCmp=*/false);
4808 case Builtin::BI__builtin_hlsl_resource_gather_cmp:
4809 return CheckGatherBuiltin(SemaRef, TheCall, /*IsCmp=*/true);
4810 case Builtin::BI__builtin_hlsl_resource_uninitializedhandle: {
4811 assert(TheCall->getNumArgs() == 1 && "expected 1 arg");
4812 // Update return type to be the attributed resource type from arg0.
4813 QualType ResourceTy = TheCall->getArg(0)->getType();
4814 TheCall->setType(ResourceTy);
4815 break;
4816 }
4817 case Builtin::BI__builtin_hlsl_resource_handlefrombinding: {
4818 assert(TheCall->getNumArgs() == 6 && "expected 6 args");
4819 // Update return type to be the attributed resource type from arg0.
4820 QualType ResourceTy = TheCall->getArg(0)->getType();
4821 TheCall->setType(ResourceTy);
4822 break;
4823 }
4824 case Builtin::BI__builtin_hlsl_resource_handlefromimplicitbinding: {
4825 assert(TheCall->getNumArgs() == 6 && "expected 6 args");
4826 // Update return type to be the attributed resource type from arg0.
4827 QualType ResourceTy = TheCall->getArg(0)->getType();
4828 TheCall->setType(ResourceTy);
4829 break;
4830 }
4831 case Builtin::BI__builtin_hlsl_resource_counterhandlefromimplicitbinding: {
4832 assert(TheCall->getNumArgs() == 3 && "expected 3 args");
4833 // Update return type to be the attributed resource type from arg0
4834 // with added IsCounter flag.
4835 QualType MainHandleTy = TheCall->getArg(0)->getType();
4836 QualType CounterHandleTy =
4837 createCounterHandleType(SemaRef.getASTContext(), MainHandleTy);
4838 TheCall->setType(CounterHandleTy);
4839 break;
4840 }
4841 case Builtin::BI__builtin_hlsl_resource_handlefromheap: {
4842 if (SemaRef.checkArgCount(TheCall, 2) ||
4843 CheckResourceHandle(&SemaRef, TheCall, 0) ||
4844 CheckArgTypeMatches(&SemaRef, TheCall->getArg(1),
4845 SemaRef.getASTContext().UnsignedIntTy))
4846 return true;
4847
4848 // Update return type to be the attributed resource type from arg0.
4849 QualType ResourceTy = TheCall->getArg(0)->getType();
4850 TheCall->setType(ResourceTy);
4851 break;
4852 }
4853 case Builtin::BI__builtin_hlsl_resource_counterhandlefromheap: {
4854 if (SemaRef.checkArgCount(TheCall, 1) ||
4855 CheckResourceHandle(&SemaRef, TheCall, 0))
4856 return true;
4857 // Update return type to be the attributed resource type from arg0
4858 // with added IsCounter flag.
4859 QualType MainHandleTy = TheCall->getArg(0)->getType();
4860 QualType CounterHandleTy =
4861 createCounterHandleType(SemaRef.getASTContext(), MainHandleTy);
4862 TheCall->setType(CounterHandleTy);
4863 break;
4864 }
4865 case Builtin::BI__builtin_hlsl_and:
4866 case Builtin::BI__builtin_hlsl_or: {
4867 if (SemaRef.checkArgCount(TheCall, 2))
4868 return true;
4869 if (CheckScalarOrVectorOrMatrix(&SemaRef, TheCall, getASTContext().BoolTy,
4870 0))
4871 return true;
4872 if (CheckAllArgsHaveSameType(&SemaRef, TheCall))
4873 return true;
4874
4875 ExprResult A = TheCall->getArg(0);
4876 QualType ArgTyA = A.get()->getType();
4877 // return type is the same as the input type
4878 TheCall->setType(ArgTyA);
4879 break;
4880 }
4881 case Builtin::BI__builtin_hlsl_all:
4882 case Builtin::BI__builtin_hlsl_any: {
4883 if (SemaRef.checkArgCount(TheCall, 1))
4884 return true;
4885 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
4886 return true;
4887 break;
4888 }
4889 case Builtin::BI__builtin_hlsl_asdouble: {
4890 if (SemaRef.checkArgCount(TheCall, 2))
4891 return true;
4893 &SemaRef, TheCall,
4894 /*only check for uint*/ SemaRef.Context.UnsignedIntTy,
4895 /* arg index */ 0))
4896 return true;
4898 &SemaRef, TheCall,
4899 /*only check for uint*/ SemaRef.Context.UnsignedIntTy,
4900 /* arg index */ 1))
4901 return true;
4902 if (CheckAllArgsHaveSameType(&SemaRef, TheCall))
4903 return true;
4904
4905 SetElementTypeAsReturnType(&SemaRef, TheCall, getASTContext().DoubleTy);
4906 break;
4907 }
4908 case Builtin::BI__builtin_hlsl_elementwise_clamp: {
4909 if (SemaRef.BuiltinElementwiseTernaryMath(
4910 TheCall, /*ArgTyRestr=*/
4912 return true;
4913 break;
4914 }
4915 case Builtin::BI__builtin_hlsl_dot: {
4916 // arg count is checked by BuiltinVectorToScalarMath
4917 if (SemaRef.BuiltinVectorToScalarMath(TheCall))
4918 return true;
4920 return true;
4921 break;
4922 }
4923 case Builtin::BI__builtin_hlsl_elementwise_firstbithigh:
4924 case Builtin::BI__builtin_hlsl_elementwise_firstbitlow: {
4925 if (SemaRef.PrepareBuiltinElementwiseMathOneArgCall(TheCall))
4926 return true;
4927
4928 const Expr *Arg = TheCall->getArg(0);
4929 QualType ArgTy = Arg->getType();
4930 QualType EltTy = ArgTy;
4931
4932 QualType ResTy = SemaRef.Context.UnsignedIntTy;
4933
4934 if (auto *VecTy = EltTy->getAs<VectorType>()) {
4935 EltTy = VecTy->getElementType();
4936 ResTy = SemaRef.Context.getExtVectorType(ResTy, VecTy->getNumElements());
4937 }
4938
4939 if (!EltTy->isIntegerType()) {
4940 Diag(Arg->getBeginLoc(), diag::err_builtin_invalid_arg_type)
4941 << 1 << /* scalar or vector of */ 5 << /* integer ty */ 1
4942 << /* no fp */ 0 << ArgTy;
4943 return true;
4944 }
4945
4946 TheCall->setType(ResTy);
4947 break;
4948 }
4949 case Builtin::BI__builtin_hlsl_select: {
4950 if (SemaRef.checkArgCount(TheCall, 3))
4951 return true;
4952 if (CheckScalarOrVectorOrMatrix(&SemaRef, TheCall, getASTContext().BoolTy,
4953 0))
4954 return true;
4955 QualType ArgTy = TheCall->getArg(0)->getType();
4956 if (ArgTy->isBooleanType() && CheckBoolSelect(&SemaRef, TheCall))
4957 return true;
4958 auto *VTy = ArgTy->getAs<VectorType>();
4959 if (VTy && VTy->getElementType()->isBooleanType() &&
4960 CheckVectorSelect(&SemaRef, TheCall))
4961 return true;
4962 auto *MTy = ArgTy->getAs<ConstantMatrixType>();
4963 if (MTy && MTy->getElementType()->isBooleanType() &&
4964 CheckMatrixSelect(&SemaRef, TheCall))
4965 return true;
4966 break;
4967 }
4968 case Builtin::BI__builtin_hlsl_elementwise_saturate:
4969 case Builtin::BI__builtin_hlsl_elementwise_rcp: {
4970 if (SemaRef.checkArgCount(TheCall, 1))
4971 return true;
4972 if (!TheCall->getArg(0)
4973 ->getType()
4974 ->hasFloatingRepresentation()) // half or float or double
4975 return SemaRef.Diag(TheCall->getArg(0)->getBeginLoc(),
4976 diag::err_builtin_invalid_arg_type)
4977 << /* ordinal */ 1 << /* scalar or vector */ 5 << /* no int */ 0
4978 << /* fp */ 1 << TheCall->getArg(0)->getType();
4979 if (SemaRef.PrepareBuiltinElementwiseMathOneArgCall(TheCall))
4980 return true;
4981 break;
4982 }
4983 case Builtin::BI__builtin_hlsl_elementwise_rsqrt:
4984 case Builtin::BI__builtin_hlsl_elementwise_frac:
4985 case Builtin::BI__builtin_hlsl_elementwise_ddx_coarse:
4986 case Builtin::BI__builtin_hlsl_elementwise_ddy_coarse:
4987 case Builtin::BI__builtin_hlsl_elementwise_ddx_fine:
4988 case Builtin::BI__builtin_hlsl_elementwise_ddy_fine: {
4989 if (SemaRef.checkArgCount(TheCall, 1))
4990 return true;
4991 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
4993 return true;
4994 if (SemaRef.PrepareBuiltinElementwiseMathOneArgCall(TheCall))
4995 return true;
4996 break;
4997 }
4998 case Builtin::BI__builtin_hlsl_elementwise_isfinite:
4999 case Builtin::BI__builtin_hlsl_elementwise_isinf:
5000 case Builtin::BI__builtin_hlsl_elementwise_isnan: {
5001 if (SemaRef.checkArgCount(TheCall, 1))
5002 return true;
5003 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
5005 return true;
5006 if (SemaRef.PrepareBuiltinElementwiseMathOneArgCall(TheCall))
5007 return true;
5009 break;
5010 }
5011 case Builtin::BI__builtin_hlsl_mad: {
5012 if (SemaRef.BuiltinElementwiseTernaryMath(
5013 TheCall, /*ArgTyRestr=*/
5015 return true;
5016 break;
5017 }
5018 case Builtin::BI__builtin_hlsl_mul: {
5019 if (SemaRef.checkArgCount(TheCall, 2))
5020 return true;
5021
5022 Expr *Arg0 = TheCall->getArg(0);
5023 Expr *Arg1 = TheCall->getArg(1);
5024 QualType Ty0 = Arg0->getType();
5025 QualType Ty1 = Arg1->getType();
5026
5027 auto getElemType = [](QualType T) -> QualType {
5028 if (const auto *VTy = T->getAs<VectorType>())
5029 return VTy->getElementType();
5030 if (const auto *MTy = T->getAs<ConstantMatrixType>())
5031 return MTy->getElementType();
5032 return T;
5033 };
5034
5035 QualType EltTy0 = getElemType(Ty0);
5036
5037 bool IsVec0 = Ty0->isVectorType();
5038 bool IsMat0 = Ty0->isConstantMatrixType();
5039 bool IsVec1 = Ty1->isVectorType();
5040 bool IsMat1 = Ty1->isConstantMatrixType();
5041
5042 QualType RetTy;
5043
5044 if (IsVec0 && IsMat1) {
5045 auto *MatTy = Ty1->castAs<ConstantMatrixType>();
5046 RetTy = getASTContext().getExtVectorType(EltTy0, MatTy->getNumColumns());
5047 } else if (IsMat0 && IsVec1) {
5048 auto *MatTy = Ty0->castAs<ConstantMatrixType>();
5049 RetTy = getASTContext().getExtVectorType(EltTy0, MatTy->getNumRows());
5050 } else {
5051 assert(IsMat0 && IsMat1);
5052 auto *MatTy0 = Ty0->castAs<ConstantMatrixType>();
5053 auto *MatTy1 = Ty1->castAs<ConstantMatrixType>();
5055 EltTy0, MatTy0->getNumRows(), MatTy1->getNumColumns());
5056 }
5057
5058 TheCall->setType(RetTy);
5059 break;
5060 }
5061 case Builtin::BI__builtin_elementwise_fma: {
5062 if (SemaRef.checkArgCount(TheCall, 3) ||
5063 CheckAllArgsHaveSameType(&SemaRef, TheCall)) {
5064 return true;
5065 }
5066
5067 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
5069 return true;
5070
5071 ExprResult A = TheCall->getArg(0);
5072 QualType ArgTyA = A.get()->getType();
5073 // return type is the same as input type
5074 TheCall->setType(ArgTyA);
5075 break;
5076 }
5077 case Builtin::BI__builtin_hlsl_transpose: {
5078 if (SemaRef.checkArgCount(TheCall, 1))
5079 return true;
5080
5081 Expr *Arg = TheCall->getArg(0);
5082 QualType ArgTy = Arg->getType();
5083
5084 const auto *MatTy = ArgTy->getAs<ConstantMatrixType>();
5085 if (!MatTy) {
5086 SemaRef.Diag(Arg->getBeginLoc(), diag::err_builtin_invalid_arg_type)
5087 << 1 << /* matrix */ 3 << /* no int */ 0 << /* no fp */ 0 << ArgTy;
5088 return true;
5089 }
5090
5092 MatTy->getElementType(), MatTy->getNumColumns(), MatTy->getNumRows());
5093 TheCall->setType(RetTy);
5094 break;
5095 }
5096 case Builtin::BI__builtin_hlsl_elementwise_sign: {
5097 if (SemaRef.PrepareBuiltinElementwiseMathOneArgCall(TheCall))
5098 return true;
5099 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
5101 return true;
5103 break;
5104 }
5105 case Builtin::BI__builtin_hlsl_wave_active_all_equal: {
5106 if (SemaRef.checkArgCount(TheCall, 1))
5107 return true;
5108
5109 // Ensure input expr type is a scalar/vector
5110 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
5111 return true;
5112
5113 QualType InputTy = TheCall->getArg(0)->getType();
5114 ASTContext &Ctx = getASTContext();
5115
5116 QualType RetTy;
5117
5118 // If vector, construct bool vector of same size
5119 if (const auto *VecTy = InputTy->getAs<ExtVectorType>()) {
5120 unsigned NumElts = VecTy->getNumElements();
5121 RetTy = Ctx.getExtVectorType(Ctx.BoolTy, NumElts);
5122 } else {
5123 // Scalar case
5124 RetTy = Ctx.BoolTy;
5125 }
5126
5127 TheCall->setType(RetTy);
5128 break;
5129 }
5130 case Builtin::BI__builtin_hlsl_wave_active_max:
5131 case Builtin::BI__builtin_hlsl_wave_active_min:
5132 case Builtin::BI__builtin_hlsl_wave_active_sum:
5133 case Builtin::BI__builtin_hlsl_wave_active_product: {
5134 if (SemaRef.checkArgCount(TheCall, 1))
5135 return true;
5136
5137 // Ensure input expr type is a scalar/vector and the same as the return type
5138 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
5139 return true;
5140 if (CheckWaveActive(&SemaRef, TheCall))
5141 return true;
5142 ExprResult Expr = TheCall->getArg(0);
5143 QualType ArgTyExpr = Expr.get()->getType();
5144 TheCall->setType(ArgTyExpr);
5145 break;
5146 }
5147 case Builtin::BI__builtin_hlsl_wave_active_bit_or:
5148 case Builtin::BI__builtin_hlsl_wave_active_bit_xor:
5149 case Builtin::BI__builtin_hlsl_wave_active_bit_and: {
5150 if (SemaRef.checkArgCount(TheCall, 1))
5151 return true;
5152
5153 // Ensure input expr type is a scalar/vector
5154 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
5155 return true;
5156
5157 if (CheckWaveActive(&SemaRef, TheCall))
5158 return true;
5159
5160 // Ensure the expr type is interpretable as a uint or vector<uint>
5161 ExprResult Expr = TheCall->getArg(0);
5162 QualType ArgTyExpr = Expr.get()->getType();
5163 auto *VTy = ArgTyExpr->getAs<VectorType>();
5164 if (!(ArgTyExpr->isIntegerType() ||
5165 (VTy && VTy->getElementType()->isIntegerType()))) {
5166 SemaRef.Diag(TheCall->getArg(0)->getBeginLoc(),
5167 diag::err_builtin_invalid_arg_type)
5168 << ArgTyExpr << SemaRef.Context.UnsignedIntTy << 1 << 0 << 0;
5169 return true;
5170 }
5171
5172 // Ensure input expr type is the same as the return type
5173 TheCall->setType(ArgTyExpr);
5174 break;
5175 }
5176 case Builtin::BI__builtin_hlsl_interlocked_add:
5177 case Builtin::BI__builtin_hlsl_interlocked_and:
5178 case Builtin::BI__builtin_hlsl_interlocked_max:
5179 case Builtin::BI__builtin_hlsl_interlocked_min:
5180 case Builtin::BI__builtin_hlsl_interlocked_or:
5181 case Builtin::BI__builtin_hlsl_interlocked_xor:
5182 if (CheckInterlockedBuiltin(SemaRef, TheCall, /*MinArgs=*/2, /*MaxArgs=*/3,
5184 /*ReportsOriginalValue=*/true))
5185 return true;
5186 break;
5187 case Builtin::BI__builtin_hlsl_interlocked_exchange:
5188 if (CheckInterlockedBuiltin(SemaRef, TheCall, /*MinArgs=*/3, /*MaxArgs=*/3,
5190 /*ReportsOriginalValue=*/true))
5191 return true;
5192 break;
5193 case Builtin::BI__builtin_hlsl_interlocked_compare_store:
5194 if (CheckInterlockedBuiltin(SemaRef, TheCall, /*MinArgs=*/3, /*MaxArgs=*/3,
5196 /*ReportsOriginalValue=*/false))
5197 return true;
5198 break;
5199 case Builtin::BI__builtin_hlsl_interlocked_compare_store_float_bitwise:
5200 if (CheckInterlockedBuiltin(SemaRef, TheCall, /*MinArgs=*/3, /*MaxArgs=*/3,
5202 /*ReportsOriginalValue=*/false))
5203 return true;
5204 break;
5205 case Builtin::BI__builtin_hlsl_interlocked_compare_exchange:
5206 if (CheckInterlockedBuiltin(SemaRef, TheCall, /*MinArgs=*/4, /*MaxArgs=*/4,
5208 /*ReportsOriginalValue=*/true))
5209 return true;
5210 break;
5211 case Builtin::BI__builtin_hlsl_interlocked_compare_exchange_float_bitwise:
5212 if (CheckInterlockedBuiltin(SemaRef, TheCall, /*MinArgs=*/4, /*MaxArgs=*/4,
5214 /*ReportsOriginalValue=*/true))
5215 return true;
5216 break;
5217 // Note these are llvm builtins that we want to catch invalid intrinsic
5218 // generation. Normal handling of these builtins will occur elsewhere.
5219 case Builtin::BI__builtin_elementwise_bitreverse: {
5220 // does not include a check for number of arguments
5221 // because that is done previously
5222 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
5224 return true;
5225 break;
5226 }
5227 case Builtin::BI__builtin_hlsl_wave_prefix_count_bits: {
5228 if (SemaRef.checkArgCount(TheCall, 1))
5229 return true;
5230
5231 QualType ArgType = TheCall->getArg(0)->getType();
5232
5233 if (!(ArgType->isScalarType())) {
5234 SemaRef.Diag(TheCall->getArg(0)->getBeginLoc(),
5235 diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
5236 << ArgType << 0;
5237 return true;
5238 }
5239
5240 if (!(ArgType->isBooleanType())) {
5241 SemaRef.Diag(TheCall->getArg(0)->getBeginLoc(),
5242 diag::err_typecheck_expect_any_scalar_or_vector_or_matrix)
5243 << ArgType << 0;
5244 return true;
5245 }
5246
5247 break;
5248 }
5249 case Builtin::BI__builtin_hlsl_wave_read_lane_at: {
5250 if (SemaRef.checkArgCount(TheCall, 2))
5251 return true;
5252
5253 // Ensure index parameter type can be interpreted as a uint
5254 ExprResult Index = TheCall->getArg(1);
5255 QualType ArgTyIndex = Index.get()->getType();
5256 if (!ArgTyIndex->isIntegerType()) {
5257 SemaRef.Diag(TheCall->getArg(1)->getBeginLoc(),
5258 diag::err_typecheck_convert_incompatible)
5259 << ArgTyIndex << SemaRef.Context.UnsignedIntTy << 1 << 0 << 0;
5260 return true;
5261 }
5262
5263 // Ensure input expr type is a scalar/vector and the same as the return type
5264 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
5265 return true;
5266
5267 ExprResult Expr = TheCall->getArg(0);
5268 QualType ArgTyExpr = Expr.get()->getType();
5269 TheCall->setType(ArgTyExpr);
5270 break;
5271 }
5272 case Builtin::BI__builtin_hlsl_wave_read_lane_first: {
5273 if (SemaRef.checkArgCount(TheCall, 1))
5274 return true;
5275
5276 if (CheckAnyScalarOrVectorOrMatrix(&SemaRef, TheCall, 0))
5277 return true;
5278
5279 TheCall->setType(TheCall->getArg(0)->getType());
5280 break;
5281 }
5282 case Builtin::BI__builtin_hlsl_wave_get_lane_index: {
5283 if (SemaRef.checkArgCount(TheCall, 0))
5284 return true;
5285 break;
5286 }
5287 case Builtin::BI__builtin_hlsl_wave_prefix_sum:
5288 case Builtin::BI__builtin_hlsl_wave_prefix_product: {
5289 if (SemaRef.checkArgCount(TheCall, 1))
5290 return true;
5291
5292 // Ensure input expr type is a scalar/vector and the same as the return type
5293 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
5294 return true;
5295 if (CheckWavePrefix(&SemaRef, TheCall))
5296 return true;
5297 ExprResult Expr = TheCall->getArg(0);
5298 QualType ArgTyExpr = Expr.get()->getType();
5299 TheCall->setType(ArgTyExpr);
5300 break;
5301 }
5302 case Builtin::BI__builtin_hlsl_quad_read_across_x:
5303 case Builtin::BI__builtin_hlsl_quad_read_across_y:
5304 case Builtin::BI__builtin_hlsl_quad_read_across_diagonal: {
5305 if (SemaRef.checkArgCount(TheCall, 1))
5306 return true;
5307
5308 if (CheckAnyScalarOrVector(&SemaRef, TheCall, 0))
5309 return true;
5310 if (CheckNotBoolScalarOrVector(&SemaRef, TheCall, 0))
5311 return true;
5312 ExprResult Expr = TheCall->getArg(0);
5313 QualType ArgTyExpr = Expr.get()->getType();
5314 TheCall->setType(ArgTyExpr);
5315 break;
5316 }
5317 case Builtin::BI__builtin_hlsl_elementwise_splitdouble: {
5318 if (SemaRef.checkArgCount(TheCall, 3))
5319 return true;
5320
5321 if (CheckScalarOrVectorOrMatrix(&SemaRef, TheCall, SemaRef.Context.DoubleTy,
5322 0) ||
5324 SemaRef.Context.UnsignedIntTy, 1) ||
5326 SemaRef.Context.UnsignedIntTy, 2))
5327 return true;
5328
5329 if (CheckModifiableLValue(&SemaRef, TheCall, 1) ||
5330 CheckModifiableLValue(&SemaRef, TheCall, 2))
5331 return true;
5332 break;
5333 }
5334 case Builtin::BI__builtin_hlsl_elementwise_clip: {
5335 if (SemaRef.checkArgCount(TheCall, 1))
5336 return true;
5337
5338 if (CheckScalarOrVector(&SemaRef, TheCall, SemaRef.Context.FloatTy, 0))
5339 return true;
5340 break;
5341 }
5342 case Builtin::BI__builtin_elementwise_acos:
5343 case Builtin::BI__builtin_elementwise_asin:
5344 case Builtin::BI__builtin_elementwise_atan:
5345 case Builtin::BI__builtin_elementwise_atan2:
5346 case Builtin::BI__builtin_elementwise_ceil:
5347 case Builtin::BI__builtin_elementwise_cos:
5348 case Builtin::BI__builtin_elementwise_cosh:
5349 case Builtin::BI__builtin_elementwise_exp:
5350 case Builtin::BI__builtin_elementwise_exp2:
5351 case Builtin::BI__builtin_elementwise_exp10:
5352 case Builtin::BI__builtin_elementwise_floor:
5353 case Builtin::BI__builtin_elementwise_fmod:
5354 case Builtin::BI__builtin_elementwise_log:
5355 case Builtin::BI__builtin_elementwise_log2:
5356 case Builtin::BI__builtin_elementwise_log10:
5357 case Builtin::BI__builtin_elementwise_pow:
5358 case Builtin::BI__builtin_elementwise_roundeven:
5359 case Builtin::BI__builtin_elementwise_sin:
5360 case Builtin::BI__builtin_elementwise_sinh:
5361 case Builtin::BI__builtin_elementwise_sqrt:
5362 case Builtin::BI__builtin_elementwise_tan:
5363 case Builtin::BI__builtin_elementwise_tanh:
5364 case Builtin::BI__builtin_elementwise_trunc: {
5365 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
5367 return true;
5368 break;
5369 }
5370 case Builtin::BI__builtin_hlsl_buffer_update_counter: {
5371 assert(TheCall->getNumArgs() == 2 && "expected 2 args");
5372 auto checkResTy = [](const HLSLAttributedResourceType *ResTy) -> bool {
5373 return !(ResTy->getAttrs().ResourceClass == ResourceClass::UAV &&
5374 ResTy->getAttrs().RawBuffer && ResTy->hasContainedType());
5375 };
5376 if (CheckResourceHandle(&SemaRef, TheCall, 0, checkResTy))
5377 return true;
5378 Expr *OffsetExpr = TheCall->getArg(1);
5379 std::optional<llvm::APSInt> Offset =
5380 OffsetExpr->getIntegerConstantExpr(SemaRef.getASTContext());
5381 if (!Offset.has_value() || std::abs(Offset->getExtValue()) != 1) {
5382 SemaRef.Diag(TheCall->getArg(1)->getBeginLoc(),
5383 diag::err_hlsl_expect_arg_const_int_one_or_neg_one)
5384 << 1;
5385 return true;
5386 }
5387 break;
5388 }
5389 case Builtin::BI__builtin_hlsl_elementwise_f16tof32: {
5390 if (SemaRef.checkArgCount(TheCall, 1))
5391 return true;
5392 if (CheckAllArgTypesAreCorrect(&SemaRef, TheCall,
5394 return true;
5395 // ensure arg integers are 32 bits
5396 if (CheckExpectedBitWidth(&SemaRef, TheCall, 0, 32))
5397 return true;
5398 // check it wasn't a bool type
5399 QualType ArgTy = TheCall->getArg(0)->getType();
5400 if (auto *VTy = ArgTy->getAs<VectorType>())
5401 ArgTy = VTy->getElementType();
5402 if (ArgTy->isBooleanType()) {
5403 SemaRef.Diag(TheCall->getArg(0)->getBeginLoc(),
5404 diag::err_builtin_invalid_arg_type)
5405 << 1 << /* scalar or vector of */ 5 << /* unsigned int */ 3
5406 << /* no fp */ 0 << TheCall->getArg(0)->getType();
5407 return true;
5408 }
5409
5410 SetElementTypeAsReturnType(&SemaRef, TheCall, getASTContext().FloatTy);
5411 break;
5412 }
5413 case Builtin::BI__builtin_hlsl_elementwise_f32tof16: {
5414 if (SemaRef.checkArgCount(TheCall, 1))
5415 return true;
5417 return true;
5419 getASTContext().UnsignedIntTy);
5420 break;
5421 }
5422 }
5423 return false;
5424}
5425
5429 WorkList.push_back(BaseTy);
5430 while (!WorkList.empty()) {
5431 QualType T = WorkList.pop_back_val();
5432 T = T.getCanonicalType().getUnqualifiedType();
5433 if (const auto *AT = dyn_cast<ConstantArrayType>(T)) {
5434 llvm::SmallVector<QualType, 16> ElementFields;
5435 // Generally I've avoided recursion in this algorithm, but arrays of
5436 // structs could be time-consuming to flatten and churn through on the
5437 // work list. Hopefully nesting arrays of structs containing arrays
5438 // of structs too many levels deep is unlikely.
5439 BuildFlattenedTypeList(AT->getElementType(), ElementFields);
5440 // Repeat the element's field list n times.
5441 for (uint64_t Ct = 0; Ct < AT->getZExtSize(); ++Ct)
5442 llvm::append_range(List, ElementFields);
5443 continue;
5444 }
5445 // Vectors can only have element types that are builtin types, so this can
5446 // add directly to the list instead of to the WorkList.
5447 if (const auto *VT = dyn_cast<VectorType>(T)) {
5448 List.insert(List.end(), VT->getNumElements(), VT->getElementType());
5449 continue;
5450 }
5451 if (const auto *MT = dyn_cast<ConstantMatrixType>(T)) {
5452 List.insert(List.end(), MT->getNumElementsFlattened(),
5453 MT->getElementType());
5454 continue;
5455 }
5456 if (const auto *RD = T->getAsCXXRecordDecl()) {
5457 if (RD->isStandardLayout())
5458 RD = RD->getStandardLayoutBaseWithFields();
5459
5460 // For types that we shouldn't decompose (unions and non-aggregates), just
5461 // add the type itself to the list.
5462 if (RD->isUnion() || !RD->isAggregate()) {
5463 List.push_back(T);
5464 continue;
5465 }
5466
5468 for (const auto *FD : RD->fields())
5469 if (!FD->isUnnamedBitField())
5470 FieldTypes.push_back(FD->getType());
5471 // Reverse the newly added sub-range.
5472 std::reverse(FieldTypes.begin(), FieldTypes.end());
5473 llvm::append_range(WorkList, FieldTypes);
5474
5475 // If this wasn't a standard layout type we may also have some base
5476 // classes to deal with.
5477 if (!RD->isStandardLayout()) {
5478 FieldTypes.clear();
5479 for (const auto &Base : RD->bases())
5480 FieldTypes.push_back(Base.getType());
5481 std::reverse(FieldTypes.begin(), FieldTypes.end());
5482 llvm::append_range(WorkList, FieldTypes);
5483 }
5484 continue;
5485 }
5486 List.push_back(T);
5487 }
5488}
5489
5491 if (QT.isNull())
5492 return false;
5493
5494 // Must be a class/struct.
5495 const auto *RD = QT->getAsCXXRecordDecl();
5496 if (!RD || RD->isUnion())
5497 return false;
5498
5499 // Cannot be a resource type or contain one.
5500 return !QT->isHLSLIntangibleType();
5501}
5502
5504 // null and array types are not allowed.
5505 if (QT.isNull() || QT->isArrayType())
5506 return false;
5507
5508 // UDT types are not allowed
5509 if (QT->isRecordType())
5510 return false;
5511
5512 if (QT->isBooleanType() || QT->isEnumeralType())
5513 return false;
5514
5515 // the only other valid builtin types are scalars or vectors
5516 if (QT->isArithmeticType()) {
5517 if (SemaRef.Context.getTypeSize(QT) / 8 > 16)
5518 return false;
5519 return true;
5520 }
5521
5522 if (const VectorType *VT = QT->getAs<VectorType>()) {
5523 int ArraySize = VT->getNumElements();
5524
5525 if (ArraySize > 4)
5526 return false;
5527
5528 QualType ElTy = VT->getElementType();
5529 if (ElTy->isBooleanType())
5530 return false;
5531
5532 if (SemaRef.Context.getTypeSize(QT) / 8 > 16)
5533 return false;
5534 return true;
5535 }
5536
5537 return false;
5538}
5539
5541 if (T1.isNull() || T2.isNull())
5542 return false;
5543
5546
5547 // If both types are the same canonical type, they're obviously compatible.
5548 if (SemaRef.getASTContext().hasSameType(T1, T2))
5549 return true;
5550
5552 BuildFlattenedTypeList(T1, T1Types);
5554 BuildFlattenedTypeList(T2, T2Types);
5555
5556 // Check the flattened type list
5557 return llvm::equal(T1Types, T2Types,
5558 [this](QualType LHS, QualType RHS) -> bool {
5559 return SemaRef.IsLayoutCompatible(LHS, RHS);
5560 });
5561}
5562
5564 FunctionDecl *Old) {
5565 if (New->getNumParams() != Old->getNumParams())
5566 return true;
5567
5568 bool HadError = false;
5569
5570 for (unsigned i = 0, e = New->getNumParams(); i != e; ++i) {
5571 ParmVarDecl *NewParam = New->getParamDecl(i);
5572 ParmVarDecl *OldParam = Old->getParamDecl(i);
5573
5574 // HLSL parameter declarations for inout and out must match between
5575 // declarations. In HLSL inout and out are ambiguous at the call site,
5576 // but have different calling behavior, so you cannot overload a
5577 // method based on a difference between inout and out annotations.
5578 const auto *NDAttr = NewParam->getAttr<HLSLParamModifierAttr>();
5579 unsigned NSpellingIdx = (NDAttr ? NDAttr->getSpellingListIndex() : 0);
5580 const auto *ODAttr = OldParam->getAttr<HLSLParamModifierAttr>();
5581 unsigned OSpellingIdx = (ODAttr ? ODAttr->getSpellingListIndex() : 0);
5582
5583 if (NSpellingIdx != OSpellingIdx) {
5584 SemaRef.Diag(NewParam->getLocation(),
5585 diag::err_hlsl_param_qualifier_mismatch)
5586 << NDAttr << NewParam;
5587 SemaRef.Diag(OldParam->getLocation(), diag::note_previous_declaration_as)
5588 << ODAttr;
5589 HadError = true;
5590 }
5591 }
5592 return HadError;
5593}
5594
5595// Generally follows PerformScalarCast, with cases reordered for
5596// clarity of what types are supported
5598
5599 if (!SrcTy->isScalarType() || !DestTy->isScalarType())
5600 return false;
5601
5602 if (SemaRef.getASTContext().hasSameUnqualifiedType(SrcTy, DestTy))
5603 return true;
5604
5605 switch (SrcTy->getScalarTypeKind()) {
5606 case Type::STK_Bool: // casting from bool is like casting from an integer
5607 case Type::STK_Integral:
5608 switch (DestTy->getScalarTypeKind()) {
5609 case Type::STK_Bool:
5610 case Type::STK_Integral:
5611 case Type::STK_Floating:
5612 return true;
5613 case Type::STK_CPointer:
5617 llvm_unreachable("HLSL doesn't support pointers.");
5620 llvm_unreachable("HLSL doesn't support complex types.");
5622 llvm_unreachable("HLSL doesn't support fixed point types.");
5623 }
5624 llvm_unreachable("Should have returned before this");
5625
5626 case Type::STK_Floating:
5627 switch (DestTy->getScalarTypeKind()) {
5628 case Type::STK_Floating:
5629 case Type::STK_Bool:
5630 case Type::STK_Integral:
5631 return true;
5634 llvm_unreachable("HLSL doesn't support complex types.");
5636 llvm_unreachable("HLSL doesn't support fixed point types.");
5637 case Type::STK_CPointer:
5641 llvm_unreachable("HLSL doesn't support pointers.");
5642 }
5643 llvm_unreachable("Should have returned before this");
5644
5646 case Type::STK_CPointer:
5649 llvm_unreachable("HLSL doesn't support pointers.");
5650
5652 llvm_unreachable("HLSL doesn't support fixed point types.");
5653
5656 llvm_unreachable("HLSL doesn't support complex types.");
5657 }
5658
5659 llvm_unreachable("Unhandled scalar cast");
5660}
5661
5662// Can perform an HLSL Aggregate splat cast if the Dest is an aggregate and the
5663// Src is a scalar, a vector of length 1, or a 1x1 matrix
5664// Or if Dest is a vector and Src is a vector of length 1 or a 1x1 matrix
5666
5667 QualType SrcTy = Src->getType();
5668 // Not a valid HLSL Aggregate Splat cast if Dest is a scalar or if this is
5669 // going to be a vector splat from a scalar.
5670 if ((SrcTy->isScalarType() && DestTy->isVectorType()) ||
5671 DestTy->isScalarType())
5672 return false;
5673
5674 const VectorType *SrcVecTy = SrcTy->getAs<VectorType>();
5675 const ConstantMatrixType *SrcMatTy = SrcTy->getAs<ConstantMatrixType>();
5676
5677 // Src isn't a scalar, a vector of length 1, or a 1x1 matrix
5678 if (!SrcTy->isScalarType() &&
5679 !(SrcVecTy && SrcVecTy->getNumElements() == 1) &&
5680 !(SrcMatTy && SrcMatTy->getNumElementsFlattened() == 1))
5681 return false;
5682
5683 if (SrcVecTy)
5684 SrcTy = SrcVecTy->getElementType();
5685 else if (SrcMatTy)
5686 SrcTy = SrcMatTy->getElementType();
5687
5689 BuildFlattenedTypeList(DestTy, DestTypes);
5690
5691 for (unsigned I = 0, Size = DestTypes.size(); I < Size; ++I) {
5692 if (DestTypes[I]->isUnionType())
5693 return false;
5694 if (!CanPerformScalarCast(SrcTy, DestTypes[I]))
5695 return false;
5696 }
5697 return true;
5698}
5699
5700// Can we perform an HLSL Elementwise cast?
5702
5703 // Don't handle casts where LHS and RHS are any combination of scalar/vector
5704 // There must be an aggregate somewhere
5705 QualType SrcTy = Src->getType();
5706 if (SrcTy->isScalarType()) // always a splat and this cast doesn't handle that
5707 return false;
5708
5709 if (SrcTy->isVectorType() &&
5710 (DestTy->isScalarType() || DestTy->isVectorType()))
5711 return false;
5712
5713 if (SrcTy->isConstantMatrixType() &&
5714 (DestTy->isScalarType() || DestTy->isConstantMatrixType()))
5715 return false;
5716
5718 BuildFlattenedTypeList(DestTy, DestTypes);
5720 BuildFlattenedTypeList(SrcTy, SrcTypes);
5721
5722 // Usually the size of SrcTypes must be greater than or equal to the size of
5723 // DestTypes.
5724 if (SrcTypes.size() < DestTypes.size())
5725 return false;
5726
5727 unsigned SrcSize = SrcTypes.size();
5728 unsigned DstSize = DestTypes.size();
5729 unsigned I;
5730 for (I = 0; I < DstSize && I < SrcSize; I++) {
5731 if (SrcTypes[I]->isUnionType() || DestTypes[I]->isUnionType())
5732 return false;
5733 if (!CanPerformScalarCast(SrcTypes[I], DestTypes[I])) {
5734 return false;
5735 }
5736 }
5737
5738 // check the rest of the source type for unions.
5739 for (; I < SrcSize; I++) {
5740 if (SrcTypes[I]->isUnionType())
5741 return false;
5742 }
5743 return true;
5744}
5745
5747 ASTContext &Ctx = SemaRef.getASTContext();
5748 QualType UIntTy = Ctx.UnsignedIntTy;
5749 QualType SrcTy = Src->getType();
5750
5751 return (SrcTy->isHLSLBuiltinPackedType() &&
5752 DestTy->isHLSLBuiltinPackedType()) ||
5753 (SrcTy->isHLSLBuiltinPackedType() &&
5754 Ctx.hasSameUnqualifiedType(DestTy, UIntTy)) ||
5755 (DestTy->isHLSLBuiltinPackedType() &&
5756 Ctx.hasSameUnqualifiedType(SrcTy, UIntTy));
5757}
5758
5760 assert(Param->hasAttr<HLSLParamModifierAttr>() &&
5761 "We should not get here without a parameter modifier expression");
5762 const auto *Attr = Param->getAttr<HLSLParamModifierAttr>();
5763 if (Attr->getABI() == ParameterABI::Ordinary)
5764 return ExprResult(Arg);
5765
5766 bool IsInOut = Attr->getABI() == ParameterABI::HLSLInOut;
5767 if (!Arg->isLValue()) {
5768 SemaRef.Diag(Arg->getBeginLoc(), diag::error_hlsl_inout_lvalue)
5769 << Arg << (IsInOut ? 1 : 0);
5770 return ExprError();
5771 }
5772
5773 ASTContext &Ctx = SemaRef.getASTContext();
5774
5775 QualType Ty = Param->getType().getNonLValueExprType(Ctx);
5776
5777 // HLSL allows implicit conversions from scalars to vectors, but not the
5778 // inverse, so we need to disallow `inout` with scalar->vector or
5779 // scalar->matrix conversions.
5780 if (Arg->getType()->isScalarType() != Ty->isScalarType()) {
5781 SemaRef.Diag(Arg->getBeginLoc(), diag::error_hlsl_inout_scalar_extension)
5782 << Arg << (IsInOut ? 1 : 0);
5783 return ExprError();
5784 }
5785
5786 auto *ArgOpV = new (Ctx) OpaqueValueExpr(Param->getBeginLoc(), Arg->getType(),
5787 VK_LValue, OK_Ordinary, Arg);
5788
5789 // Parameters are initialized via copy initialization. This allows for
5790 // overload resolution of argument constructors.
5791 InitializedEntity Entity =
5793 ExprResult Res =
5794 SemaRef.PerformCopyInitialization(Entity, Param->getBeginLoc(), ArgOpV);
5795 if (Res.isInvalid())
5796 return ExprError();
5797 Expr *Base = Res.get();
5798 // After the cast, drop the reference type when creating the exprs.
5799 Ty = Ty.getNonLValueExprType(Ctx);
5800 auto *OpV = new (Ctx)
5801 OpaqueValueExpr(Param->getBeginLoc(), Ty, VK_LValue, OK_Ordinary, Base);
5802
5803 // Writebacks are performed with `=` binary operator, which allows for
5804 // overload resolution on writeback result expressions.
5805 Res = SemaRef.ActOnBinOp(SemaRef.getCurScope(), Arg->getBeginLoc(),
5806 tok::equal, ArgOpV, OpV);
5807
5808 if (Res.isInvalid())
5809 return ExprError();
5810 Expr *Writeback = Res.get();
5811 auto *OutExpr =
5812 HLSLOutArgExpr::Create(Ctx, Ty, ArgOpV, OpV, Writeback, IsInOut);
5813
5814 return ExprResult(OutExpr);
5815}
5816
5818 // If HLSL gains support for references, all the cites that use this will need
5819 // to be updated with semantic checking to produce errors for
5820 // pointers/references.
5821 assert(!Ty->isReferenceType() &&
5822 "Pointer and reference types cannot be inout or out parameters");
5823 Ty = SemaRef.getASTContext().getLValueReferenceType(Ty);
5824 Ty.addRestrict();
5825 return Ty;
5826}
5827
5828// Returns true if the type has a non-empty constant buffer layout (if it is
5829// scalar, vector or matrix, or if it contains any of these.
5831 const Type *Ty = QT->getUnqualifiedDesugaredType();
5832 if (Ty->isScalarType() || Ty->isVectorType() || Ty->isMatrixType())
5833 return true;
5834
5836 return false;
5837
5838 if (const auto *RD = Ty->getAsCXXRecordDecl()) {
5839 for (const auto *FD : RD->fields()) {
5841 return true;
5842 }
5843 assert(RD->getNumBases() <= 1 &&
5844 "HLSL doesn't support multiple inheritance");
5845 return RD->getNumBases()
5846 ? hasConstantBufferLayout(RD->bases_begin()->getType())
5847 : false;
5848 }
5849
5850 if (const auto *AT = dyn_cast<ArrayType>(Ty)) {
5851 if (const auto *CAT = dyn_cast<ConstantArrayType>(AT))
5852 if (isZeroSizedArray(CAT))
5853 return false;
5854 return hasConstantBufferLayout(AT->getElementType());
5855 }
5856
5857 return false;
5858}
5859
5860static bool IsDefaultBufferConstantDecl(const ASTContext &Ctx, VarDecl *VD) {
5861 bool IsVulkan =
5862 Ctx.getTargetInfo().getTriple().getOS() == llvm::Triple::Vulkan;
5863 bool IsVKPushConstant = IsVulkan && VD->hasAttr<HLSLVkPushConstantAttr>();
5864 QualType QT = VD->getType();
5865 return VD->getDeclContext()->isTranslationUnit() &&
5866 QT.getAddressSpace() == LangAS::Default &&
5867 VD->getStorageClass() != SC_Static &&
5868 !VD->hasAttr<HLSLVkConstantIdAttr>() && !IsVKPushConstant &&
5870}
5871
5873 // The variable already has an address space (groupshared for ex).
5874 if (Decl->getType().hasAddressSpace())
5875 return;
5876
5877 if (Decl->getType()->isDependentType())
5878 return;
5879
5880 QualType Type = Decl->getType();
5881
5882 if (Decl->hasAttr<HLSLVkExtBuiltinInputAttr>()) {
5883 LangAS ImplAS = LangAS::hlsl_input;
5884 Type = SemaRef.getASTContext().getAddrSpaceQualType(Type, ImplAS);
5885 Decl->setType(Type);
5886 return;
5887 }
5888
5889 if (Decl->hasAttr<HLSLVkExtBuiltinOutputAttr>()) {
5890 LangAS ImplAS = LangAS::hlsl_output;
5891 Type = SemaRef.getASTContext().getAddrSpaceQualType(Type, ImplAS);
5892 Decl->setType(Type);
5893
5894 // HLSL uses `static` differently than C++. For BuiltIn output, the static
5895 // does not imply private to the module scope.
5896 // Marking it as external to reflect the semantic this attribute brings.
5897 // See https://github.com/microsoft/hlsl-specs/issues/350
5898 Decl->setStorageClass(SC_Extern);
5899 return;
5900 }
5901
5902 bool IsVulkan = getASTContext().getTargetInfo().getTriple().getOS() ==
5903 llvm::Triple::Vulkan;
5904 if (IsVulkan && Decl->hasAttr<HLSLVkPushConstantAttr>()) {
5905 if (HasDeclaredAPushConstant)
5906 SemaRef.Diag(Decl->getLocation(), diag::err_hlsl_push_constant_unique);
5907
5909 Type = SemaRef.getASTContext().getAddrSpaceQualType(Type, ImplAS);
5910 Decl->setType(Type);
5911 HasDeclaredAPushConstant = true;
5912 return;
5913 }
5914
5915 if (Type->isSamplerT() || Type->isVoidType())
5916 return;
5917
5918 // Resource handles.
5920 return;
5921
5922 // Only static globals belong to the Private address space.
5923 // Non-static globals belongs to the cbuffer.
5924 if (Decl->getStorageClass() != SC_Static && !Decl->isStaticDataMember())
5925 return;
5926
5928 Type = SemaRef.getASTContext().getAddrSpaceQualType(Type, ImplAS);
5929 Decl->setType(Type);
5930}
5931
5932namespace {
5933
5934// Helper class for assigning bindings to resources declared within a struct.
5935// It keeps track of all binding attributes declared on a struct instance, and
5936// the offsets for each register type that have been assigned so far.
5937// Handles both explicit and implicit bindings.
5938class StructBindingContext {
5939 // Bindings and offsets per register type. We only need to support four
5940 // register types - SRV (u), UAV (t), CBuffer (c), and Sampler (s).
5941 HLSLResourceBindingAttr *RegBindingsAttrs[4];
5942 unsigned RegBindingOffset[4];
5943
5944 // Make sure the RegisterType values are what we expect
5945 static_assert(static_cast<unsigned>(RegisterType::SRV) == 0 &&
5946 static_cast<unsigned>(RegisterType::UAV) == 1 &&
5947 static_cast<unsigned>(RegisterType::CBuffer) == 2 &&
5948 static_cast<unsigned>(RegisterType::Sampler) == 3,
5949 "unexpected register type values");
5950
5951 // Vulkan binding attribute does not vary by register type.
5952 HLSLVkBindingAttr *VkBindingAttr;
5953 unsigned VkBindingOffset;
5954
5955public:
5956 // Constructor: gather all binding attributes on a struct instance and
5957 // initialize offsets.
5958 StructBindingContext(VarDecl *VD) {
5959 for (unsigned i = 0; i < 4; ++i) {
5960 RegBindingsAttrs[i] = nullptr;
5961 RegBindingOffset[i] = 0;
5962 }
5963 VkBindingAttr = nullptr;
5964 VkBindingOffset = 0;
5965
5966 ASTContext &AST = VD->getASTContext();
5967 bool IsSpirv = AST.getTargetInfo().getTriple().isSPIRV();
5968
5969 for (Attr *A : VD->attrs()) {
5970 if (auto *RBA = dyn_cast<HLSLResourceBindingAttr>(A)) {
5971 RegisterType RegType = RBA->getRegisterType();
5972 unsigned RegTypeIdx = static_cast<unsigned>(RegType);
5973 // Ignore unsupported register annotations, such as 'c' or 'i'.
5974 if (RegTypeIdx < 4)
5975 RegBindingsAttrs[RegTypeIdx] = RBA;
5976 continue;
5977 }
5978 // Gather the Vulkan binding attributes only if the target is SPIR-V.
5979 if (IsSpirv) {
5980 if (auto *VBA = dyn_cast<HLSLVkBindingAttr>(A))
5981 VkBindingAttr = VBA;
5982 }
5983 }
5984 }
5985
5986 // Creates a binding attribute for a resource based on the gathered attributes
5987 // and the required register type and range.
5988 Attr *createBindingAttr(SemaHLSL &S, ASTContext &AST, RegisterType RegType,
5989 unsigned Range, bool HasCounter) {
5990 assert(static_cast<unsigned>(RegType) < 4 && "unexpected register type");
5991
5992 if (VkBindingAttr) {
5993 unsigned Offset = VkBindingOffset;
5994 VkBindingOffset += Range;
5995 return HLSLVkBindingAttr::CreateImplicit(
5996 AST, VkBindingAttr->getBinding() + Offset, VkBindingAttr->getSet(),
5997 VkBindingAttr->getRange());
5998 }
5999
6000 HLSLResourceBindingAttr *RBA =
6001 RegBindingsAttrs[static_cast<unsigned>(RegType)];
6002 HLSLResourceBindingAttr *NewAttr = nullptr;
6003
6004 if (RBA && RBA->hasRegisterSlot()) {
6005 // Explicit binding - create a new attribute with offseted slot number
6006 // based on the required register type.
6007 unsigned Offset = RegBindingOffset[static_cast<unsigned>(RegType)];
6008 RegBindingOffset[static_cast<unsigned>(RegType)] += Range;
6009
6010 unsigned NewSlotNumber = RBA->getSlotNumber() + Offset;
6011 StringRef NewSlotNumberStr =
6012 createRegisterString(AST, RBA->getRegisterType(), NewSlotNumber);
6013 NewAttr = HLSLResourceBindingAttr::CreateImplicit(
6014 AST, NewSlotNumberStr, RBA->getSpace(), RBA->getRange());
6015 NewAttr->setBinding(RegType, NewSlotNumber, RBA->getSpaceNumber());
6016 } else {
6017 // No binding attribute or space-only binding - create a binding
6018 // attribute for implicit binding.
6019 NewAttr = HLSLResourceBindingAttr::CreateImplicit(AST, "", "0", {});
6020 NewAttr->setBinding(RegType, std::nullopt,
6021 RBA ? RBA->getSpaceNumber() : 0);
6022 NewAttr->setImplicitBindingOrderID(S.getNextImplicitBindingOrderID());
6023 }
6024 if (HasCounter)
6025 NewAttr->setImplicitCounterBindingOrderID(
6027 return NewAttr;
6028 }
6029};
6030
6031// Creates a global variable declaration for a resource field embedded in a
6032// struct, assigns it a binding, initializes it, and associates it with the
6033// struct declaration via an HLSLAssociatedResourceDeclAttr.
6034static void createGlobalResourceDeclForStruct(
6035 Sema &S, VarDecl *ParentVD, SourceLocation Loc, IdentifierInfo *Id,
6036 QualType ResTy, StructBindingContext &BindingCtx) {
6037 assert(isResourceRecordTypeOrArrayOf(ResTy) &&
6038 "expected resource type or array of resources");
6039
6040 DeclContext *DC = ParentVD->getNonTransparentDeclContext();
6041 assert(DC->isTranslationUnit() && "expected translation unit decl context");
6042
6043 ASTContext &AST = S.getASTContext();
6044 VarDecl *ResDecl =
6045 VarDecl::Create(AST, DC, Loc, Loc, Id, ResTy, nullptr, SC_None);
6046
6047 unsigned Range = 1;
6048 const Type *SingleResTy = ResTy.getTypePtr()->getUnqualifiedDesugaredType();
6049 while (const auto *AT = dyn_cast<ArrayType>(SingleResTy)) {
6050 const auto *CAT = dyn_cast<ConstantArrayType>(AT);
6051 Range = CAT ? (Range * CAT->getSize().getZExtValue()) : 0;
6052 SingleResTy =
6054 }
6055 const HLSLAttributedResourceType *ResHandleTy =
6056 HLSLAttributedResourceType::findHandleTypeOnResource(SingleResTy);
6057
6058 // Add a binding attribute to the global resource declaration.
6059 bool HasCounter = hasCounterHandle(SingleResTy->getAsCXXRecordDecl());
6060 Attr *BindingAttr = BindingCtx.createBindingAttr(
6061 S.HLSL(), AST, getRegisterType(ResHandleTy), Range, HasCounter);
6062 ResDecl->addAttr(BindingAttr);
6063 ResDecl->addAttr(InternalLinkageAttr::CreateImplicit(AST));
6064 ResDecl->setImplicit();
6065
6066 if (Range == 1)
6067 S.HLSL().initGlobalResourceDecl(ResDecl);
6068 else
6069 S.HLSL().initGlobalResourceArrayDecl(ResDecl);
6070
6071 ParentVD->addAttr(
6072 HLSLAssociatedResourceDeclAttr::CreateImplicit(AST, ResDecl));
6073 DC->addDecl(ResDecl);
6074
6075 DeclGroupRef DG(ResDecl);
6077}
6078
6079static void handleArrayOfStructWithResources(
6080 Sema &S, VarDecl *ParentVD, const ConstantArrayType *CAT,
6081 EmbeddedResourceNameBuilder &NameBuilder, StructBindingContext &BindingCtx);
6082
6083// Scans base and all fields of a struct/class type to find all embedded
6084// resources or resource arrays. Creates a global variable for each resource
6085// found.
6086static void handleStructWithResources(Sema &S, VarDecl *ParentVD,
6087 const CXXRecordDecl *RD,
6088 EmbeddedResourceNameBuilder &NameBuilder,
6089 StructBindingContext &BindingCtx) {
6090
6091 // Scan the base classes.
6092 assert(RD->getNumBases() <= 1 && "HLSL doesn't support multiple inheritance");
6093 const auto *BasesIt = RD->bases_begin();
6094 if (BasesIt != RD->bases_end()) {
6095 QualType QT = BasesIt->getType();
6096 if (QT->isHLSLIntangibleType()) {
6097 CXXRecordDecl *BaseRD = QT->getAsCXXRecordDecl();
6098 NameBuilder.pushBaseName(BaseRD->getName());
6099 handleStructWithResources(S, ParentVD, BaseRD, NameBuilder, BindingCtx);
6100 NameBuilder.pop();
6101 }
6102 }
6103 // Process this class fields.
6104 for (const FieldDecl *FD : RD->fields()) {
6105 QualType FDTy = FD->getType().getCanonicalType();
6106 if (!FDTy->isHLSLIntangibleType())
6107 continue;
6108
6109 NameBuilder.pushName(FD->getName());
6110
6112 IdentifierInfo *II = NameBuilder.getNameAsIdentifier(S.getASTContext());
6113 createGlobalResourceDeclForStruct(S, ParentVD, FD->getLocation(), II,
6114 FDTy, BindingCtx);
6115 } else if (const auto *RD = FDTy->getAsCXXRecordDecl()) {
6116 handleStructWithResources(S, ParentVD, RD, NameBuilder, BindingCtx);
6117
6118 } else if (const auto *ArrayTy = dyn_cast<ConstantArrayType>(FDTy)) {
6119 assert(!FDTy->isHLSLResourceRecordArray() &&
6120 "resource arrays should have been already handled");
6121 handleArrayOfStructWithResources(S, ParentVD, ArrayTy, NameBuilder,
6122 BindingCtx);
6123 }
6124 NameBuilder.pop();
6125 }
6126}
6127
6128// Processes array of structs with resources.
6129static void
6130handleArrayOfStructWithResources(Sema &S, VarDecl *ParentVD,
6131 const ConstantArrayType *CAT,
6132 EmbeddedResourceNameBuilder &NameBuilder,
6133 StructBindingContext &BindingCtx) {
6134
6135 QualType ElementTy = CAT->getElementType().getCanonicalType();
6136 assert(ElementTy->isHLSLIntangibleType() && "Expected HLSL intangible type");
6137
6138 const ConstantArrayType *SubCAT = dyn_cast<ConstantArrayType>(ElementTy);
6139 const CXXRecordDecl *ElementRD = ElementTy->getAsCXXRecordDecl();
6140
6141 if (!SubCAT && !ElementRD)
6142 return;
6143
6144 for (unsigned I = 0, E = CAT->getSize().getZExtValue(); I < E; ++I) {
6145 NameBuilder.pushArrayIndex(I);
6146 if (ElementRD)
6147 handleStructWithResources(S, ParentVD, ElementRD, NameBuilder,
6148 BindingCtx);
6149 else
6150 handleArrayOfStructWithResources(S, ParentVD, SubCAT, NameBuilder,
6151 BindingCtx);
6152 NameBuilder.pop();
6153 }
6154}
6155
6156} // namespace
6157
6158// Scans all fields of a user-defined struct (or array of structs)
6159// to find all embedded resources or resource arrays. For each resource
6160// a global variable of the resource type is created and associated
6161// with the parent declaration (VD) through a HLSLAssociatedResourceDeclAttr
6162// attribute.
6163void SemaHLSL::handleGlobalStructOrArrayOfWithResources(VarDecl *VD) {
6164 EmbeddedResourceNameBuilder NameBuilder(VD->getName());
6165 StructBindingContext BindingCtx(VD);
6166
6167 const Type *VDTy = VD->getType().getTypePtr();
6168 assert(VDTy->isHLSLIntangibleType() && !isResourceRecordTypeOrArrayOf(VD) &&
6169 "Expected non-resource struct or array type");
6170
6171 if (const CXXRecordDecl *RD = VDTy->getAsCXXRecordDecl()) {
6172 handleStructWithResources(SemaRef, VD, RD, NameBuilder, BindingCtx);
6173 return;
6174 }
6175
6176 if (const auto *CAT = dyn_cast<ConstantArrayType>(VDTy)) {
6177 handleArrayOfStructWithResources(SemaRef, VD, CAT, NameBuilder, BindingCtx);
6178 return;
6179 }
6180}
6181
6183 if (VD->hasGlobalStorage()) {
6184 // make sure the declaration has a complete type
6185 if (SemaRef.RequireCompleteType(
6186 VD->getLocation(),
6187 SemaRef.getASTContext().getBaseElementType(VD->getType()),
6188 diag::err_typecheck_decl_incomplete_type)) {
6189 VD->setInvalidDecl();
6191 return;
6192 }
6193
6194 // Global variables outside a cbuffer block that are not a resource, static,
6195 // groupshared, or an empty array or struct belong to the default constant
6196 // buffer $Globals (to be created at the end of the translation unit).
6198 // update address space to hlsl_constant
6201 VD->setType(NewTy);
6202 DefaultCBufferDecls.push_back(VD);
6203 }
6204
6205 // find all resources bindings on decl
6206 if (VD->getType()->isHLSLIntangibleType())
6207 collectResourceBindingsOnVarDecl(VD);
6208
6209 if (VD->hasAttr<HLSLVkConstantIdAttr>())
6211
6213 VD->getStorageClass() != SC_Static) {
6214 // Add internal linkage attribute to non-static resource variables. The
6215 // global externally visible storage is accessed through the handle, which
6216 // is a member. The variable itself is not externally visible.
6217 VD->addAttr(InternalLinkageAttr::CreateImplicit(getASTContext()));
6218 }
6219
6220 // process explicit bindings
6221 processExplicitBindingsOnDecl(VD);
6222
6223 // Add implicit binding attribute to non-static resource arrays.
6224 if (VD->getType()->isHLSLResourceRecordArray() &&
6225 VD->getStorageClass() != SC_Static) {
6226 // If the resource array does not have an explicit binding attribute,
6227 // create an implicit one. It will be used to transfer implicit binding
6228 // order_ID to codegen.
6229 ResourceBindingAttrs Binding(VD);
6230 if (!Binding.isExplicit()) {
6231 uint32_t OrderID = getNextImplicitBindingOrderID();
6232 if (Binding.hasBinding())
6233 Binding.setImplicitOrderID(OrderID);
6234 else {
6237 OrderID);
6238 // Re-create the binding object to pick up the new attribute.
6239 Binding = ResourceBindingAttrs(VD);
6240 }
6241 }
6242
6243 // Get to the base type of a potentially multi-dimensional array.
6245
6246 const CXXRecordDecl *RD = Ty->getAsCXXRecordDecl();
6247 if (hasCounterHandle(RD)) {
6248 if (!Binding.hasCounterImplicitOrderID()) {
6249 uint32_t OrderID = getNextImplicitBindingOrderID();
6250 Binding.setCounterImplicitOrderID(OrderID);
6251 }
6252 }
6253 }
6254
6255 // Process resources in user-defined structs, or arrays of such structs.
6256 const Type *VDTy = VD->getType().getTypePtr();
6257 if (VD->getStorageClass() != SC_Static && VDTy->isHLSLIntangibleType() &&
6259 handleGlobalStructOrArrayOfWithResources(VD);
6260
6261 // Mark groupshared variables as extern so they will have
6262 // external storage and won't be default initialized
6263 if (VD->hasAttr<HLSLGroupSharedAddressSpaceAttr>())
6265 }
6266
6268}
6269
6271 assert(VD->getType()->isHLSLResourceRecord() &&
6272 "expected resource record type");
6273
6274 ASTContext &AST = SemaRef.getASTContext();
6275 uint64_t UIntTySize = AST.getTypeSize(AST.UnsignedIntTy);
6276 uint64_t IntTySize = AST.getTypeSize(AST.IntTy);
6277
6278 // Gather resource binding attributes.
6279 ResourceBindingAttrs Binding(VD);
6280
6281 // Find correct initialization method and create its arguments.
6282 QualType ResourceTy = VD->getType();
6283 CXXRecordDecl *ResourceDecl = ResourceTy->getAsCXXRecordDecl();
6284 CXXMethodDecl *CreateMethod = nullptr;
6286
6287 bool HasCounter = hasCounterHandle(ResourceDecl);
6288 const char *CreateMethodName;
6289 if (Binding.isExplicit())
6290 CreateMethodName = HasCounter ? "__createFromBindingWithImplicitCounter"
6291 : "__createFromBinding";
6292 else
6293 CreateMethodName = HasCounter
6294 ? "__createFromImplicitBindingWithImplicitCounter"
6295 : "__createFromImplicitBinding";
6296
6297 CreateMethod =
6298 lookupMethod(SemaRef, ResourceDecl, CreateMethodName, VD->getLocation());
6299
6300 if (!CreateMethod) {
6301 // This can happen if someone creates a struct that looks like an HLSL
6302 // resource record but does not have the required static create method.
6303 // No binding will be generated for it.
6304 assert(!ResourceDecl->isImplicit() &&
6305 "create method lookup should always succeed for built-in resource "
6306 "records");
6307 return false;
6308 }
6309
6310 if (Binding.isExplicit()) {
6311 IntegerLiteral *RegSlot =
6312 IntegerLiteral::Create(AST, llvm::APInt(UIntTySize, Binding.getSlot()),
6314 Args.push_back(RegSlot);
6315 } else {
6316 uint32_t OrderID = (Binding.hasImplicitOrderID())
6317 ? Binding.getImplicitOrderID()
6319 IntegerLiteral *OrderId =
6320 IntegerLiteral::Create(AST, llvm::APInt(UIntTySize, OrderID),
6322 Args.push_back(OrderId);
6323 }
6324
6325 IntegerLiteral *Space =
6326 IntegerLiteral::Create(AST, llvm::APInt(UIntTySize, Binding.getSpace()),
6328 Args.push_back(Space);
6329
6331 AST, llvm::APInt(IntTySize, 1), AST.IntTy, SourceLocation());
6332 Args.push_back(RangeSize);
6333
6335 AST, llvm::APInt(UIntTySize, 0), AST.UnsignedIntTy, SourceLocation());
6336 Args.push_back(Index);
6337
6338 StringRef VarName = VD->getName();
6340 AST, VarName, StringLiteralKind::Ordinary, false,
6341 AST.getStringLiteralArrayType(AST.CharTy.withConst(), VarName.size()),
6342 SourceLocation());
6344 AST, AST.getPointerType(AST.CharTy.withConst()), CK_ArrayToPointerDecay,
6345 Name, nullptr, VK_PRValue, FPOptionsOverride());
6346 Args.push_back(NameCast);
6347
6348 if (HasCounter) {
6349 // Will this be in the correct order?
6350 uint32_t CounterOrderID = getNextImplicitBindingOrderID();
6351 IntegerLiteral *CounterId =
6352 IntegerLiteral::Create(AST, llvm::APInt(UIntTySize, CounterOrderID),
6354 Args.push_back(CounterId);
6355 }
6356
6357 // Make sure the create method template is instantiated and emitted.
6358 if (!CreateMethod->isDefined() && CreateMethod->isTemplateInstantiation())
6359 SemaRef.InstantiateFunctionDefinition(VD->getLocation(), CreateMethod,
6360 true);
6361
6362 // Create CallExpr with a call to the static method and set it as the decl
6363 // initialization.
6365 AST, NestedNameSpecifierLoc(), SourceLocation(), CreateMethod, false,
6366 CreateMethod->getNameInfo(), CreateMethod->getType(), VK_PRValue);
6367
6368 auto *ImpCast = ImplicitCastExpr::Create(
6369 AST, AST.getPointerType(CreateMethod->getType()),
6370 CK_FunctionToPointerDecay, DRE, nullptr, VK_PRValue, FPOptionsOverride());
6371
6372 CallExpr *InitExpr =
6373 CallExpr::Create(AST, ImpCast, Args, ResourceTy, VK_PRValue,
6375 VD->setInit(InitExpr);
6377 SemaRef.CheckCompleteVariableDeclaration(VD);
6378 return true;
6379}
6380
6382 assert(VD->getType()->isHLSLResourceRecordArray() &&
6383 "expected array of resource records");
6384
6385 // Individual resources in a resource array are not initialized here. They
6386 // are initialized later on during codegen when the individual resources are
6387 // accessed. Codegen will emit a call to the resource initialization method
6388 // with the specified array index. We need to make sure though that the method
6389 // for the specific resource type is instantiated, so codegen can emit a call
6390 // to it when the array element is accessed.
6391
6392 // Find correct initialization method based on the resource binding
6393 // information.
6394 ASTContext &AST = SemaRef.getASTContext();
6395 QualType ResElementTy = AST.getBaseElementType(VD->getType());
6396 CXXRecordDecl *ResourceDecl = ResElementTy->getAsCXXRecordDecl();
6397 CXXMethodDecl *CreateMethod = nullptr;
6398
6399 bool HasCounter = hasCounterHandle(ResourceDecl);
6400 ResourceBindingAttrs ResourceAttrs(VD);
6401 if (ResourceAttrs.isExplicit())
6402 // Resource has explicit binding.
6403 CreateMethod =
6404 lookupMethod(SemaRef, ResourceDecl,
6405 HasCounter ? "__createFromBindingWithImplicitCounter"
6406 : "__createFromBinding",
6407 VD->getLocation());
6408 else
6409 // Resource has implicit binding.
6410 CreateMethod = lookupMethod(
6411 SemaRef, ResourceDecl,
6412 HasCounter ? "__createFromImplicitBindingWithImplicitCounter"
6413 : "__createFromImplicitBinding",
6414 VD->getLocation());
6415
6416 if (!CreateMethod)
6417 return false;
6418
6419 // Make sure the create method template is instantiated and emitted.
6420 if (!CreateMethod->isDefined() && CreateMethod->isTemplateInstantiation())
6421 SemaRef.InstantiateFunctionDefinition(VD->getLocation(), CreateMethod,
6422 true);
6423 return true;
6424}
6425
6426// Returns true if the initialization has been handled.
6427// Returns false to use default initialization.
6429 // Objects in the hlsl_constant address space are initialized
6430 // externally, so don't synthesize an implicit initializer.
6432 return true;
6433
6434 if (VD->hasGlobalStorage() && VD->getStorageClass() != SC_Static) {
6435 const Type *Ty = VD->getType().getTypePtr();
6437 return true;
6439 return true;
6440 }
6441
6442 // User-defined structs/classes do not have constructors.
6443 // When declared at a global scope, they are part of the constant buffer
6444 // and should not be initialized by the compiler.
6445 // When declared at a local scope, they are not initialized.
6446 // Also applies to arrays of user-defined structs/classes.
6447 const Type *Ty = VD->getType()->getUnqualifiedDesugaredType();
6448 while (Ty->isArrayType())
6450 if (CXXRecordDecl *RD = Ty->getAsCXXRecordDecl())
6451 return !RD->isHLSLBuiltinRecord();
6452
6453 return false;
6454}
6455
6456std::optional<const DeclBindingInfo *> SemaHLSL::inferGlobalBinding(Expr *E) {
6457 if (auto *Ternary = dyn_cast<ConditionalOperator>(E)) {
6458 auto TrueInfo = inferGlobalBinding(Ternary->getTrueExpr());
6459 auto FalseInfo = inferGlobalBinding(Ternary->getFalseExpr());
6460 if (!TrueInfo || !FalseInfo)
6461 return std::nullopt;
6462 if (*TrueInfo != *FalseInfo)
6463 return std::nullopt;
6464 return TrueInfo;
6465 }
6466
6467 if (auto *ASE = dyn_cast<ArraySubscriptExpr>(E))
6468 E = ASE->getBase()->IgnoreParenImpCasts();
6469
6470 if (DeclRefExpr *DRE = dyn_cast<DeclRefExpr>(E->IgnoreParens()))
6471 if (VarDecl *VD = dyn_cast<VarDecl>(DRE->getDecl())) {
6472 const Type *Ty = VD->getType()->getUnqualifiedDesugaredType();
6473 if (Ty->isArrayType())
6475
6476 if (const auto *AttrResType =
6477 HLSLAttributedResourceType::findHandleTypeOnResource(Ty)) {
6478 ResourceClass RC = AttrResType->getAttrs().ResourceClass;
6479 return Bindings.getDeclBindingInfo(VD, RC);
6480 }
6481 }
6482
6483 return nullptr;
6484}
6485
6486void SemaHLSL::trackLocalResource(VarDecl *VD, Expr *E) {
6487 std::optional<const DeclBindingInfo *> ExprBinding = inferGlobalBinding(E);
6488 if (!ExprBinding) {
6489 SemaRef.Diag(E->getBeginLoc(),
6490 diag::warn_hlsl_assigning_local_resource_is_not_unique)
6491 << E << VD;
6492 return; // Expr use multiple resources
6493 }
6494
6495 if (*ExprBinding == nullptr)
6496 return; // No binding could be inferred to track, return without error
6497
6498 auto PrevBinding = Assigns.find(VD);
6499 if (PrevBinding == Assigns.end()) {
6500 // No previous binding recorded, simply record the new assignment
6501 Assigns.insert({VD, *ExprBinding});
6502 return;
6503 }
6504
6505 // Otherwise, warn if the assignment implies different resource bindings
6506 if (*ExprBinding != PrevBinding->second) {
6507 SemaRef.Diag(E->getBeginLoc(),
6508 diag::warn_hlsl_assigning_local_resource_is_not_unique)
6509 << E << VD;
6510 SemaRef.Diag(VD->getLocation(), diag::note_var_declared_here) << VD;
6511 return;
6512 }
6513
6514 return;
6515}
6516
6518 Expr *RHSExpr, SourceLocation Loc) {
6519 assert((LHSExpr->getType()->isHLSLResourceRecord() ||
6520 LHSExpr->getType()->isHLSLResourceRecordArray()) &&
6521 "expected LHS to be a resource record or array of resource records");
6522 if (Opc != BO_Assign)
6523 return true;
6524
6525 // If LHS is an array subscript, get the underlying declaration.
6526 Expr *E = LHSExpr;
6527 while (auto *ASE = dyn_cast<ArraySubscriptExpr>(E))
6528 E = ASE->getBase()->IgnoreParenImpCasts();
6529
6530 // Report error if LHS is a non-static resource declared at a global scope.
6531 if (DeclRefExpr *DRE = dyn_cast<DeclRefExpr>(E->IgnoreParens())) {
6532 if (VarDecl *VD = dyn_cast<VarDecl>(DRE->getDecl())) {
6533 if (VD->hasGlobalStorage() && VD->getStorageClass() != SC_Static) {
6534 // assignment to global resource is not allowed
6535 SemaRef.Diag(Loc, diag::err_hlsl_assign_to_global_resource) << VD;
6536 SemaRef.Diag(VD->getLocation(), diag::note_var_declared_here) << VD;
6537 return false;
6538 }
6539
6540 trackLocalResource(VD, RHSExpr);
6541 }
6542 }
6543 return true;
6544}
6545
6546// Returns true if the given type can have an overload of the given
6547// binary operator.
6549 CXXRecordDecl *RD = LHSTy->getAsCXXRecordDecl();
6550 if (!RD)
6551 return true;
6552 return RD->isHLSLBuiltinRecord() || Opc != BO_Assign;
6553}
6554
6555// Walks though the global variable declaration, collects all resource binding
6556// requirements and adds them to Bindings
6557void SemaHLSL::collectResourceBindingsOnVarDecl(VarDecl *VD) {
6558 assert(VD->hasGlobalStorage() && VD->getType()->isHLSLIntangibleType() &&
6559 "expected global variable that contains HLSL resource");
6560
6561 // Cbuffers and Tbuffers are HLSLBufferDecl types
6562 if (const HLSLBufferDecl *CBufferOrTBuffer = dyn_cast<HLSLBufferDecl>(VD)) {
6563 Bindings.addDeclBindingInfo(VD, CBufferOrTBuffer->isCBuffer()
6564 ? ResourceClass::CBuffer
6565 : ResourceClass::SRV);
6566 return;
6567 }
6568
6569 // Unwrap arrays
6570 // FIXME: Calculate array size while unwrapping
6571 const Type *Ty = VD->getType()->getUnqualifiedDesugaredType();
6572 while (Ty->isArrayType()) {
6573 const ArrayType *AT = cast<ArrayType>(Ty);
6575 }
6576
6577 // Resource (or array of resources)
6578 if (const HLSLAttributedResourceType *AttrResType =
6579 HLSLAttributedResourceType::findHandleTypeOnResource(Ty)) {
6580 Bindings.addDeclBindingInfo(VD, AttrResType->getAttrs().ResourceClass);
6581 return;
6582 }
6583
6584 // User defined record type
6585 if (const RecordType *RT = dyn_cast<RecordType>(Ty))
6586 collectResourceBindingsOnUserRecordDecl(VD, RT);
6587}
6588
6589// Walks though the explicit resource binding attributes on the declaration,
6590// and makes sure there is a resource that matched the binding and updates
6591// DeclBindingInfoLists
6592void SemaHLSL::processExplicitBindingsOnDecl(VarDecl *VD) {
6593 assert(VD->hasGlobalStorage() && "expected global variable");
6594
6595 bool HasBinding = false;
6596 for (Attr *A : VD->attrs()) {
6597 if (isa<HLSLVkBindingAttr>(A)) {
6598 HasBinding = true;
6599 if (auto PA = VD->getAttr<HLSLVkPushConstantAttr>())
6600 Diag(PA->getLoc(), diag::err_hlsl_attr_incompatible) << A << PA;
6601 }
6602
6603 HLSLResourceBindingAttr *RBA = dyn_cast<HLSLResourceBindingAttr>(A);
6604 if (!RBA || !RBA->hasRegisterSlot())
6605 continue;
6606 HasBinding = true;
6607
6608 RegisterType RT = RBA->getRegisterType();
6609 assert(RT != RegisterType::I && "invalid or obsolete register type should "
6610 "never have an attribute created");
6611
6612 if (RT == RegisterType::C) {
6613 if (Bindings.hasBindingInfoForDecl(VD))
6614 SemaRef.Diag(VD->getLocation(),
6615 diag::warn_hlsl_user_defined_type_missing_member)
6616 << static_cast<int>(RT);
6617 continue;
6618 }
6619
6620 // Find DeclBindingInfo for this binding and update it, or report error
6621 // if it does not exist (user type does to contain resources with the
6622 // expected resource class).
6624 if (DeclBindingInfo *BI = Bindings.getDeclBindingInfo(VD, RC)) {
6625 // update binding info
6626 BI->setBindingAttribute(RBA, BindingType::Explicit);
6627 } else {
6628 SemaRef.Diag(VD->getLocation(),
6629 diag::warn_hlsl_user_defined_type_missing_member)
6630 << static_cast<int>(RT);
6631 }
6632 }
6633
6634 if (!HasBinding && isResourceRecordTypeOrArrayOf(VD))
6635 SemaRef.Diag(VD->getLocation(), diag::warn_hlsl_implicit_binding);
6636}
6637namespace {
6638class InitListTransformer {
6639 Sema &S;
6640 ASTContext &Ctx;
6641 QualType InitTy;
6642 QualType *DstIt = nullptr;
6643 Expr **ArgIt = nullptr;
6644 // Is wrapping the destination type iterator required? This is only used for
6645 // incomplete array types where we loop over the destination type since we
6646 // don't know the full number of elements from the declaration.
6647 bool Wrap;
6648
6649 bool castInitializer(Expr *E) {
6650 assert(DstIt && "This should always be something!");
6651 if (DstIt == DestTypes.end()) {
6652 if (!Wrap) {
6653 ArgExprs.push_back(E);
6654 // This is odd, but it isn't technically a failure due to conversion, we
6655 // handle mismatched counts of arguments differently.
6656 return true;
6657 }
6658 DstIt = DestTypes.begin();
6659 }
6660 InitializedEntity Entity = InitializedEntity::InitializeParameter(
6661 Ctx, *DstIt, /* Consumed (ObjC) */ false);
6662 ExprResult Res = S.PerformCopyInitialization(Entity, E->getBeginLoc(), E);
6663 if (Res.isInvalid())
6664 return false;
6665 Expr *Init = Res.get();
6666 ArgExprs.push_back(Init);
6667 DstIt++;
6668 return true;
6669 }
6670
6671 bool buildInitializerListImpl(Expr *E) {
6672 // If this is an initialization list, traverse the sub initializers.
6673 if (auto *Init = dyn_cast<InitListExpr>(E)) {
6674 for (auto *SubInit : Init->inits())
6675 if (!buildInitializerListImpl(SubInit))
6676 return false;
6677 return true;
6678 }
6679
6680 // If this is a scalar type, just enqueue the expression.
6681 QualType Ty = E->getType().getDesugaredType(Ctx);
6682
6683 if (Ty->isScalarType() || (Ty->isRecordType() && !Ty->isAggregateType()) ||
6685 return castInitializer(E);
6686
6687 // If this is an aggregate type and a prvalue, create an xvalue temporary
6688 // so the member accesses will be xvalues. Wrap it in OpaqueExpr to make
6689 // sure codegen will not generate duplicate copies.
6690 if (E->isPRValue() && Ty->isAggregateType()) {
6692 if (TmpExpr.isInvalid())
6693 return false;
6694 E = TmpExpr.get();
6695 E = new (Ctx) OpaqueValueExpr(E->getBeginLoc(), E->getType(),
6696 E->getValueKind(), E->getObjectKind(), E);
6697 }
6698
6699 if (auto *VecTy = Ty->getAs<VectorType>()) {
6700 uint64_t Size = VecTy->getNumElements();
6701
6702 QualType SizeTy = Ctx.getSizeType();
6703 uint64_t SizeTySize = Ctx.getTypeSize(SizeTy);
6704 for (uint64_t I = 0; I < Size; ++I) {
6705 auto *Idx = IntegerLiteral::Create(Ctx, llvm::APInt(SizeTySize, I),
6706 SizeTy, SourceLocation());
6707
6709 E, E->getBeginLoc(), Idx, E->getEndLoc());
6710 if (ElExpr.isInvalid())
6711 return false;
6712 if (!castInitializer(ElExpr.get()))
6713 return false;
6714 }
6715 return true;
6716 }
6717 if (auto *MTy = Ty->getAs<ConstantMatrixType>()) {
6718 unsigned Rows = MTy->getNumRows();
6719 unsigned Cols = MTy->getNumColumns();
6720 QualType ElemTy = MTy->getElementType();
6721
6722 for (unsigned R = 0; R < Rows; ++R) {
6723 for (unsigned C = 0; C < Cols; ++C) {
6724 // row index literal
6725 Expr *RowIdx = IntegerLiteral::Create(
6726 Ctx, llvm::APInt(Ctx.getIntWidth(Ctx.IntTy), R), Ctx.IntTy,
6727 E->getBeginLoc());
6728 // column index literal
6729 Expr *ColIdx = IntegerLiteral::Create(
6730 Ctx, llvm::APInt(Ctx.getIntWidth(Ctx.IntTy), C), Ctx.IntTy,
6731 E->getBeginLoc());
6733 E, RowIdx, ColIdx, E->getEndLoc());
6734 if (ElExpr.isInvalid())
6735 return false;
6736 if (!castInitializer(ElExpr.get()))
6737 return false;
6738 ElExpr.get()->setType(ElemTy);
6739 }
6740 }
6741 return true;
6742 }
6743
6744 if (auto *ArrTy = dyn_cast<ConstantArrayType>(Ty.getTypePtr())) {
6745 uint64_t Size = ArrTy->getZExtSize();
6746 QualType SizeTy = Ctx.getSizeType();
6747 uint64_t SizeTySize = Ctx.getTypeSize(SizeTy);
6748 for (uint64_t I = 0; I < Size; ++I) {
6749 auto *Idx = IntegerLiteral::Create(Ctx, llvm::APInt(SizeTySize, I),
6750 SizeTy, SourceLocation());
6752 E, E->getBeginLoc(), Idx, E->getEndLoc());
6753 if (ElExpr.isInvalid())
6754 return false;
6755 if (!buildInitializerListImpl(ElExpr.get()))
6756 return false;
6757 }
6758 return true;
6759 }
6760
6761 if (auto *RD = Ty->getAsCXXRecordDecl()) {
6762 llvm::SmallVector<CXXRecordDecl *> RecordDecls;
6763 RecordDecls.push_back(RD);
6764 while (RecordDecls.back()->getNumBases()) {
6765 CXXRecordDecl *D = RecordDecls.back();
6766 assert(D->getNumBases() == 1 &&
6767 "HLSL doesn't support multiple inheritance");
6768 RecordDecls.push_back(
6770 }
6771 while (!RecordDecls.empty()) {
6772 CXXRecordDecl *RD = RecordDecls.pop_back_val();
6773 for (auto *FD : RD->fields()) {
6774 if (FD->isUnnamedBitField())
6775 continue;
6776 DeclAccessPair Found = DeclAccessPair::make(FD, FD->getAccess());
6777 DeclarationNameInfo NameInfo(FD->getDeclName(), E->getBeginLoc());
6779 E, false, E->getBeginLoc(), CXXScopeSpec(), FD, Found, NameInfo);
6780 if (Res.isInvalid())
6781 return false;
6782 if (!buildInitializerListImpl(Res.get()))
6783 return false;
6784 }
6785 }
6786 }
6787 return true;
6788 }
6789
6790 Expr *generateInitListsImpl(QualType Ty) {
6791 Ty = Ty.getDesugaredType(Ctx);
6792 assert(ArgIt != ArgExprs.end() && "Something is off in iteration!");
6793 if (Ty->isScalarType() || (Ty->isRecordType() && !Ty->isAggregateType()) ||
6795 return *(ArgIt++);
6796
6797 llvm::SmallVector<Expr *> Inits;
6798 if (Ty->isVectorType() || Ty->isConstantArrayType() ||
6799 Ty->isConstantMatrixType()) {
6800 QualType ElTy;
6801 uint64_t Size = 0;
6802 if (auto *ATy = Ty->getAs<VectorType>()) {
6803 ElTy = ATy->getElementType();
6804 Size = ATy->getNumElements();
6805 } else if (auto *CMTy = Ty->getAs<ConstantMatrixType>()) {
6806 ElTy = CMTy->getElementType();
6807 Size = CMTy->getNumElementsFlattened();
6808 } else {
6809 auto *VTy = cast<ConstantArrayType>(Ty.getTypePtr());
6810 ElTy = VTy->getElementType();
6811 Size = VTy->getZExtSize();
6812 }
6813 for (uint64_t I = 0; I < Size; ++I)
6814 Inits.push_back(generateInitListsImpl(ElTy));
6815 }
6816 if (auto *RD = Ty->getAsCXXRecordDecl()) {
6817 llvm::SmallVector<CXXRecordDecl *> RecordDecls;
6818 RecordDecls.push_back(RD);
6819 while (RecordDecls.back()->getNumBases()) {
6820 CXXRecordDecl *D = RecordDecls.back();
6821 assert(D->getNumBases() == 1 &&
6822 "HLSL doesn't support multiple inheritance");
6823 RecordDecls.push_back(
6825 }
6826 while (!RecordDecls.empty()) {
6827 CXXRecordDecl *RD = RecordDecls.pop_back_val();
6828 for (auto *FD : RD->fields())
6829 if (!FD->isUnnamedBitField())
6830 Inits.push_back(generateInitListsImpl(FD->getType()));
6831 }
6832 }
6833 auto *NewInit =
6834 new (Ctx) InitListExpr(Ctx, Inits.front()->getBeginLoc(), Inits,
6835 Inits.back()->getEndLoc(), /*isExplicit=*/false);
6836 NewInit->setType(Ty);
6837 return NewInit;
6838 }
6839
6840public:
6841 llvm::SmallVector<QualType, 16> DestTypes;
6842 llvm::SmallVector<Expr *, 16> ArgExprs;
6843 InitListTransformer(Sema &SemaRef, const InitializedEntity &Entity)
6844 : S(SemaRef), Ctx(SemaRef.getASTContext()),
6845 Wrap(Entity.getType()->isIncompleteArrayType()) {
6846 InitTy = Entity.getType().getNonReferenceType();
6847 // When we're generating initializer lists for incomplete array types we
6848 // need to wrap around both when building the initializers and when
6849 // generating the final initializer lists.
6850 if (Wrap) {
6851 assert(InitTy->isIncompleteArrayType());
6852 const IncompleteArrayType *IAT = Ctx.getAsIncompleteArrayType(InitTy);
6853 InitTy = IAT->getElementType();
6854 }
6855 BuildFlattenedTypeList(InitTy, DestTypes);
6856 DstIt = DestTypes.begin();
6857 }
6858
6859 bool buildInitializerList(Expr *E) { return buildInitializerListImpl(E); }
6860
6861 Expr *generateInitLists() {
6862 assert(!ArgExprs.empty() &&
6863 "Call buildInitializerList to generate argument expressions.");
6864 ArgIt = ArgExprs.begin();
6865 if (!Wrap)
6866 return generateInitListsImpl(InitTy);
6867 llvm::SmallVector<Expr *> Inits;
6868 while (ArgIt != ArgExprs.end())
6869 Inits.push_back(generateInitListsImpl(InitTy));
6870
6871 auto *NewInit =
6872 new (Ctx) InitListExpr(Ctx, Inits.front()->getBeginLoc(), Inits,
6873 Inits.back()->getEndLoc(), /*isExplicit=*/false);
6874 llvm::APInt ArySize(64, Inits.size());
6875 NewInit->setType(Ctx.getConstantArrayType(InitTy, ArySize, nullptr,
6876 ArraySizeModifier::Normal, 0));
6877 return NewInit;
6878 }
6879};
6880} // namespace
6881
6882// Recursively detect any incomplete array anywhere in the type graph,
6883// including arrays, struct fields, and base classes.
6885 Ty = Ty.getCanonicalType();
6886
6887 // Array types
6888 if (const ArrayType *AT = dyn_cast<ArrayType>(Ty)) {
6890 return true;
6892 }
6893
6894 // Record (struct/class) types
6895 if (const auto *RT = Ty->getAs<RecordType>()) {
6896 const RecordDecl *RD = RT->getDecl();
6897
6898 // Walk base classes (for C++ / HLSL structs with inheritance)
6899 if (const auto *CXXRD = dyn_cast<CXXRecordDecl>(RD)) {
6900 for (const CXXBaseSpecifier &Base : CXXRD->bases()) {
6901 if (containsIncompleteArrayType(Base.getType()))
6902 return true;
6903 }
6904 }
6905
6906 // Walk fields
6907 for (const FieldDecl *F : RD->fields()) {
6908 if (containsIncompleteArrayType(F->getType()))
6909 return true;
6910 }
6911 }
6912
6913 return false;
6914}
6915
6917 InitListExpr *Init) {
6918 // If the initializer is a scalar, just return it.
6919 if (Init->getType()->isScalarType())
6920 return true;
6921 ASTContext &Ctx = SemaRef.getASTContext();
6922 InitListTransformer ILT(SemaRef, Entity);
6923
6924 for (unsigned I = 0; I < Init->getNumInits(); ++I) {
6925 Expr *E = Init->getInit(I);
6926 if (E->HasSideEffects(Ctx)) {
6927 QualType Ty = E->getType();
6928 if (Ty->isRecordType())
6929 E = new (Ctx) MaterializeTemporaryExpr(Ty, E, E->isLValue());
6930 E = new (Ctx) OpaqueValueExpr(E->getBeginLoc(), Ty, E->getValueKind(),
6931 E->getObjectKind(), E);
6932 Init->setInit(I, E);
6933 }
6934 if (!ILT.buildInitializerList(E))
6935 return false;
6936 }
6937 size_t ExpectedSize = ILT.DestTypes.size();
6938 size_t ActualSize = ILT.ArgExprs.size();
6939 if (ExpectedSize == 0 && ActualSize == 0)
6940 return true;
6941
6942 // Reject empty initializer if *any* incomplete array exists structurally
6943 if (ActualSize == 0 && containsIncompleteArrayType(Entity.getType())) {
6944 QualType InitTy = Entity.getType().getNonReferenceType();
6945 if (InitTy.hasAddressSpace())
6946 InitTy = SemaRef.getASTContext().removeAddrSpaceQualType(InitTy);
6947
6948 SemaRef.Diag(Init->getBeginLoc(), diag::err_hlsl_incorrect_num_initializers)
6949 << /*TooManyOrFew=*/(int)(ExpectedSize < ActualSize) << InitTy
6950 << /*ExpectedSize=*/ExpectedSize << /*ActualSize=*/ActualSize;
6951 return false;
6952 }
6953
6954 // We infer size after validating legality.
6955 // For incomplete arrays it is completely arbitrary to choose whether we think
6956 // the user intended fewer or more elements. This implementation assumes that
6957 // the user intended more, and errors that there are too few initializers to
6958 // complete the final element.
6959 if (Entity.getType()->isIncompleteArrayType()) {
6960 assert(ExpectedSize > 0 &&
6961 "The expected size of an incomplete array type must be at least 1.");
6962 ExpectedSize =
6963 ((ActualSize + ExpectedSize - 1) / ExpectedSize) * ExpectedSize;
6964 }
6965
6966 // An initializer list might be attempting to initialize a reference or
6967 // rvalue-reference. When checking the initializer we should look through
6968 // the reference.
6969 QualType InitTy = Entity.getType().getNonReferenceType();
6970 if (InitTy.hasAddressSpace())
6971 InitTy = SemaRef.getASTContext().removeAddrSpaceQualType(InitTy);
6972 if (ExpectedSize != ActualSize) {
6973 int TooManyOrFew = ActualSize > ExpectedSize ? 1 : 0;
6974 SemaRef.Diag(Init->getBeginLoc(), diag::err_hlsl_incorrect_num_initializers)
6975 << TooManyOrFew << InitTy << ExpectedSize << ActualSize;
6976 return false;
6977 }
6978
6979 // generateInitListsImpl will always return an InitListExpr here, because the
6980 // scalar case is handled above.
6981 auto *NewInit = cast<InitListExpr>(ILT.generateInitLists());
6982 Init->resizeInits(Ctx, NewInit->getNumInits());
6983 for (unsigned I = 0; I < NewInit->getNumInits(); ++I)
6984 Init->updateInit(Ctx, I, NewInit->getInit(I));
6985 return true;
6986}
6987
6988static QualType ReportMatrixInvalidMember(Sema &S, StringRef Name,
6989 StringRef Expected,
6990 SourceLocation OpLoc,
6991 SourceLocation CompLoc) {
6992 S.Diag(OpLoc, diag::err_builtin_matrix_invalid_member)
6993 << Name << Expected << SourceRange(CompLoc);
6994 return QualType();
6995}
6996
6999 const IdentifierInfo *CompName,
7000 SourceLocation CompLoc) {
7001 const auto *MT = baseType->castAs<ConstantMatrixType>();
7002 StringRef AccessorName = CompName->getName();
7003 assert(!AccessorName.empty() && "Matrix Accessor must have a name");
7004
7005 unsigned Rows = MT->getNumRows();
7006 unsigned Cols = MT->getNumColumns();
7007 bool IsZeroBasedAccessor = false;
7008 unsigned ChunkLen = 0;
7009 if (AccessorName.size() < 2)
7010 return ReportMatrixInvalidMember(S, AccessorName,
7011 "length 4 for zero based: \'_mRC\' or "
7012 "length 3 for one-based: \'_RC\' accessor",
7013 OpLoc, CompLoc);
7014
7015 if (AccessorName[0] == '_') {
7016 if (AccessorName[1] == 'm') {
7017 IsZeroBasedAccessor = true;
7018 ChunkLen = 4; // zero-based: "_mRC"
7019 } else {
7020 ChunkLen = 3; // one-based: "_RC"
7021 }
7022 } else
7024 S, AccessorName, "zero based: \'_mRC\' or one-based: \'_RC\' accessor",
7025 OpLoc, CompLoc);
7026
7027 if (AccessorName.size() % ChunkLen != 0) {
7028 const llvm::StringRef Expected = IsZeroBasedAccessor
7029 ? "zero based: '_mRC' accessor"
7030 : "one-based: '_RC' accessor";
7031
7032 return ReportMatrixInvalidMember(S, AccessorName, Expected, OpLoc, CompLoc);
7033 }
7034
7035 auto isDigit = [](char c) { return c >= '0' && c <= '9'; };
7036 auto isZeroBasedIndex = [](unsigned i) { return i <= 3; };
7037 auto isOneBasedIndex = [](unsigned i) { return i >= 1 && i <= 4; };
7038
7039 bool HasRepeated = false;
7040 SmallVector<bool, 16> Seen(Rows * Cols, false);
7041 unsigned NumComponents = 0;
7042 const char *Begin = AccessorName.data();
7043
7044 for (unsigned I = 0, E = AccessorName.size(); I < E; I += ChunkLen) {
7045 const char *Chunk = Begin + I;
7046 char RowChar = 0, ColChar = 0;
7047 if (IsZeroBasedAccessor) {
7048 // Zero-based: "_mRC"
7049 if (Chunk[0] != '_' || Chunk[1] != 'm') {
7050 char Bad = (Chunk[0] != '_') ? Chunk[0] : Chunk[1];
7052 S, StringRef(&Bad, 1), "\'_m\' prefix",
7053 OpLoc.getLocWithOffset(I + (Bad == Chunk[0] ? 1 : 2)), CompLoc);
7054 }
7055 RowChar = Chunk[2];
7056 ColChar = Chunk[3];
7057 } else {
7058 // One-based: "_RC"
7059 if (Chunk[0] != '_')
7061 S, StringRef(&Chunk[0], 1), "\'_\' prefix",
7062 OpLoc.getLocWithOffset(I + 1), CompLoc);
7063 RowChar = Chunk[1];
7064 ColChar = Chunk[2];
7065 }
7066
7067 // Must be digits.
7068 bool IsDigitsError = false;
7069 if (!isDigit(RowChar)) {
7070 unsigned BadPos = IsZeroBasedAccessor ? 2 : 1;
7071 ReportMatrixInvalidMember(S, StringRef(&RowChar, 1), "row as integer",
7072 OpLoc.getLocWithOffset(I + BadPos + 1),
7073 CompLoc);
7074 IsDigitsError = true;
7075 }
7076
7077 if (!isDigit(ColChar)) {
7078 unsigned BadPos = IsZeroBasedAccessor ? 3 : 2;
7079 ReportMatrixInvalidMember(S, StringRef(&ColChar, 1), "column as integer",
7080 OpLoc.getLocWithOffset(I + BadPos + 1),
7081 CompLoc);
7082 IsDigitsError = true;
7083 }
7084 if (IsDigitsError)
7085 return QualType();
7086
7087 unsigned Row = RowChar - '0';
7088 unsigned Col = ColChar - '0';
7089
7090 bool HasIndexingError = false;
7091 if (IsZeroBasedAccessor) {
7092 // 0-based [0..3]
7093 if (!isZeroBasedIndex(Row)) {
7094 S.Diag(OpLoc, diag::err_hlsl_matrix_element_not_in_bounds)
7095 << /*row*/ 0 << /*zero-based*/ 0 << SourceRange(CompLoc);
7096 HasIndexingError = true;
7097 }
7098 if (!isZeroBasedIndex(Col)) {
7099 S.Diag(OpLoc, diag::err_hlsl_matrix_element_not_in_bounds)
7100 << /*col*/ 1 << /*zero-based*/ 0 << SourceRange(CompLoc);
7101 HasIndexingError = true;
7102 }
7103 } else {
7104 // 1-based [1..4]
7105 if (!isOneBasedIndex(Row)) {
7106 S.Diag(OpLoc, diag::err_hlsl_matrix_element_not_in_bounds)
7107 << /*row*/ 0 << /*one-based*/ 1 << SourceRange(CompLoc);
7108 HasIndexingError = true;
7109 }
7110 if (!isOneBasedIndex(Col)) {
7111 S.Diag(OpLoc, diag::err_hlsl_matrix_element_not_in_bounds)
7112 << /*col*/ 1 << /*one-based*/ 1 << SourceRange(CompLoc);
7113 HasIndexingError = true;
7114 }
7115 // Convert to 0-based after range checking.
7116 --Row;
7117 --Col;
7118 }
7119
7120 if (HasIndexingError)
7121 return QualType();
7122
7123 // Note: matrix swizzle index is hard coded. That means Row and Col can
7124 // potentially be larger than Rows and Cols if matrix size is less than
7125 // the max index size.
7126 bool HasBoundsError = false;
7127 if (Row >= Rows) {
7128 Diag(OpLoc, diag::err_hlsl_matrix_index_out_of_bounds)
7129 << /*Row*/ 0 << Row << Rows << SourceRange(CompLoc);
7130 HasBoundsError = true;
7131 }
7132 if (Col >= Cols) {
7133 Diag(OpLoc, diag::err_hlsl_matrix_index_out_of_bounds)
7134 << /*Col*/ 1 << Col << Cols << SourceRange(CompLoc);
7135 HasBoundsError = true;
7136 }
7137 if (HasBoundsError)
7138 return QualType();
7139
7140 unsigned FlatIndex = Row * Cols + Col;
7141 if (Seen[FlatIndex])
7142 HasRepeated = true;
7143 Seen[FlatIndex] = true;
7144 ++NumComponents;
7145 }
7146 if (NumComponents == 0 || NumComponents > 4) {
7147 S.Diag(OpLoc, diag::err_hlsl_matrix_swizzle_invalid_length)
7148 << NumComponents << SourceRange(CompLoc);
7149 return QualType();
7150 }
7151
7152 QualType ElemTy = MT->getElementType();
7153 if (NumComponents == 1)
7154 return ElemTy;
7155 QualType VT = S.Context.getExtVectorType(ElemTy, NumComponents);
7156 if (HasRepeated)
7157 VK = VK_PRValue;
7158
7159 for (Sema::ExtVectorDeclsType::iterator
7161 E = S.ExtVectorDecls.end();
7162 I != E; ++I) {
7163 if ((*I)->getUnderlyingType() == VT)
7165 /*Qualifier=*/std::nullopt, *I);
7166 }
7167
7168 return VT;
7169}
7170
7172 // If initializing a local resource, track the resource binding it is using
7173 if (VDecl->getType()->isHLSLResourceRecord() && !VDecl->hasGlobalStorage())
7174 trackLocalResource(VDecl, Init);
7175
7176 const HLSLVkConstantIdAttr *ConstIdAttr =
7177 VDecl->getAttr<HLSLVkConstantIdAttr>();
7178 if (!ConstIdAttr)
7179 return true;
7180
7181 ASTContext &Context = SemaRef.getASTContext();
7182
7183 APValue InitValue;
7184 if (!Init->isCXX11ConstantExpr(Context, InitValue)) {
7185 Diag(VDecl->getLocation(), diag::err_specialization_const);
7186 VDecl->setInvalidDecl();
7187 return false;
7188 }
7189
7190 Builtin::ID BID =
7192
7193 // Argument 1: The ID from the attribute
7194 int ConstantID = ConstIdAttr->getId();
7195 llvm::APInt IDVal(Context.getIntWidth(Context.IntTy), ConstantID);
7196 Expr *IdExpr = IntegerLiteral::Create(Context, IDVal, Context.IntTy,
7197 ConstIdAttr->getLocation());
7198
7199 SmallVector<Expr *, 2> Args = {IdExpr, Init};
7200 Expr *C = SemaRef.BuildBuiltinCallExpr(Init->getExprLoc(), BID, Args);
7201 if (C->getType()->getCanonicalTypeUnqualified() !=
7203 C = SemaRef
7204 .BuildCStyleCastExpr(SourceLocation(),
7205 Context.getTrivialTypeSourceInfo(
7206 Init->getType(), Init->getExprLoc()),
7207 SourceLocation(), C)
7208 .get();
7209 }
7210 Init = C;
7211 return true;
7212}
7213
7215 SourceLocation NameLoc) {
7216 if (!Template)
7217 return QualType();
7218
7219 DeclContext *DC = Template->getDeclContext();
7220 if (!DC->isNamespace() || !cast<NamespaceDecl>(DC)->getIdentifier() ||
7221 cast<NamespaceDecl>(DC)->getName() != "hlsl")
7222 return QualType();
7223
7224 TemplateParameterList *Params = Template->getTemplateParameters();
7225 if (!Params || Params->size() != 1)
7226 return QualType();
7227
7228 if (!Template->isImplicit())
7229 return QualType();
7230
7231 // We manually extract default arguments here instead of letting
7232 // CheckTemplateIdType handle it. This ensures that for resource types that
7233 // lack a default argument (like Buffer), we return a null QualType, which
7234 // triggers the "requires template arguments" error rather than a less
7235 // descriptive "too few template arguments" error.
7236 TemplateArgumentListInfo TemplateArgs(NameLoc, NameLoc);
7237 for (NamedDecl *P : *Params) {
7238 if (auto *TTP = dyn_cast<TemplateTypeParmDecl>(P)) {
7239 if (TTP->hasDefaultArgument()) {
7240 TemplateArgs.addArgument(TTP->getDefaultArgument());
7241 continue;
7242 }
7243 } else if (auto *NTTP = dyn_cast<NonTypeTemplateParmDecl>(P)) {
7244 if (NTTP->hasDefaultArgument()) {
7245 TemplateArgs.addArgument(NTTP->getDefaultArgument());
7246 continue;
7247 }
7248 } else if (auto *TTPD = dyn_cast<TemplateTemplateParmDecl>(P)) {
7249 if (TTPD->hasDefaultArgument()) {
7250 TemplateArgs.addArgument(TTPD->getDefaultArgument());
7251 continue;
7252 }
7253 }
7254 return QualType();
7255 }
7256
7257 return SemaRef.CheckTemplateIdType(
7259 TemplateArgs, nullptr, /*ForNestedNameSpecifier=*/false);
7260}
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 bool isIntUpTo32Element(const ASTContext &Ctx, QualType Elem)
static unsigned calculateLegacyCbufferFieldAlign(const ASTContext &Context, QualType T)
Definition SemaHLSL.cpp:272
static bool CheckScalarFloatOperand(Sema &S, CallExpr *TheCall, unsigned ArgIndex)
static bool isIntElementOfWidth(const ASTContext &Ctx, QualType Elem, uint64_t Width)
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 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 unsigned getComponentCountOf(QualType T)
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 QualType getElementTypeOf(QualType T, bool IncludeMatrix)
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)
static bool isFloatOrHalfElement(QualType Elem)
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:158
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
const ConstantArrayType * getAsConstantArrayType(QualType T) const
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.
static uint64_t getConstantArrayElementCount(const ConstantArrayType *CA)
Return number of (potentially nested) constant array elements.
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 getCorrespondingUnsignedType(QualType T) const
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:2991
Expr * getArg(unsigned Arg)
getArg - Return the specified argument.
Definition Expr.h:3195
SourceLocation getBeginLoc() const
Definition Expr.h:3325
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:3174
Expr * getCallee()
Definition Expr.h:3138
unsigned getNumArgs() const
getNumArgs - Return the number of actual arguments to this call.
Definition Expr.h:3182
SourceLocation getEndLoc() const
Definition Expr.h:3344
Decl * getCalleeDecl()
Definition Expr.h:3168
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
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:1294
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:1362
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:3901
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:5356
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:3412
ValueDecl * getMemberDecl() const
Retrieve the member declaration to which this expression refers.
Definition Expr.h:3495
Expr * getBase() const
Definition Expr.h:3489
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:1202
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:239
void CheckEntryPoint(FunctionDecl *FD)
void handleVkExtBuiltinOutputAttr(Decl *D, const ParsedAttr &AL)
void emitLogicalOperatorFixIt(Expr *LHS, Expr *RHS, BinaryOperatorKind Opc)
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 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
QualType getInoutParameterType(QualType Ty)
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:1823
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
bool isFloat16Type() const
Definition TypeBase.h:9081
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
bool isFloat32Type() const
Definition TypeBase.h:9085
bool isHalfType() const
Definition TypeBase.h:9076
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