1//===--- HLSLExternalSemaSource.cpp - HLSL Sema Source --------------------===//
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//
9//
10//===----------------------------------------------------------------------===//
11
12#include "clang/Sema/HLSLExternalSemaSource.h"
13#include "HLSLBuiltinTypeDeclBuilder.h"
14#include "clang/AST/ASTContext.h"
15#include "clang/AST/Attr.h"
16#include "clang/AST/Decl.h"
17#include "clang/AST/DeclCXX.h"
18#include "clang/AST/DeclTemplate.h"
19#include "clang/AST/Expr.h"
20#include "clang/AST/Type.h"
21#include "clang/Basic/AddressSpaces.h"
22#include "clang/Basic/SourceLocation.h"
23#include "clang/Lex/Preprocessor.h"
24#include "clang/Sema/Lookup.h"
25#include "clang/Sema/Sema.h"
26#include "clang/Sema/SemaHLSL.h"
27#include "llvm/ADT/BitmaskEnum.h"
28#include "llvm/ADT/STLExtras.h"
29#include "llvm/ADT/SmallVector.h"
30
31using namespace clang;
32using namespace llvm::hlsl;
33
34using clang::hlsl::BuiltinTypeDeclBuilder;
35
36static NamespaceDecl *createImplicitNamespace(Sema &S, StringRef Name,
37 DeclContext *DC) {
38 ASTContext &AST = S.getASTContext();
39 IdentifierInfo &II = AST.Idents.get(Name, TokenCode: tok::TokenKind::identifier);
40 LookupResult Result(S, &II, SourceLocation(), Sema::LookupNamespaceName);
41 NamespaceDecl *PrevDecl = nullptr;
42 if (S.LookupQualifiedName(R&: Result, LookupCtx: DC))
43 PrevDecl = Result.getAsSingle<NamespaceDecl>();
44
45 NamespaceDecl *NS =
46 NamespaceDecl::Create(C&: AST, DC, /*Inline=*/false, StartLoc: SourceLocation(),
47 IdLoc: SourceLocation(), Id: &II, PrevDecl, /*Nested=*/false);
48 NS->setImplicit(true);
49 NS->setHasExternalLexicalStorage();
50 DC->addDecl(D: NS);
51
52 // Force external decls in the namespace to load from the PCH.
53 (void)NS->getCanonicalDecl()->decls_begin();
54
55 return NS;
56}
57
58void HLSLExternalSemaSource::InitializeSema(Sema &S) {
59 SemaPtr = &S;
60 ASTContext &AST = SemaPtr->getASTContext();
61 // If the translation unit has external storage force external decls to load.
62 if (AST.getTranslationUnitDecl()->hasExternalLexicalStorage())
63 (void)AST.getTranslationUnitDecl()->decls_begin();
64
65 HLSLNamespace = createImplicitNamespace(
66 S, Name: "hlsl", DC: cast<DeclContext>(Val: AST.getTranslationUnitDecl()));
67 HLSLDetailNamespace = createImplicitNamespace(S, Name: "__detail", DC: HLSLNamespace);
68
69 defineTrivialHLSLTypes();
70 defineInternalHLSLTypes();
71 defineHLSLTypesWithForwardDeclarations();
72 defineHLSLAtomicIntrinsics();
73
74 // This adds a `using namespace hlsl` directive. In DXC, we don't put HLSL's
75 // built in types inside a namespace, but we are planning to change that in
76 // the near future. In order to be source compatible older versions of HLSL
77 // will need to implicitly use the hlsl namespace. For now in clang everything
78 // will get added to the namespace, and we can remove the using directive for
79 // future language versions to match HLSL's evolution.
80 auto *UsingDecl = UsingDirectiveDecl::Create(
81 C&: AST, DC: AST.getTranslationUnitDecl(), UsingLoc: SourceLocation(), NamespaceLoc: SourceLocation(),
82 QualifierLoc: NestedNameSpecifierLoc(), IdentLoc: SourceLocation(), Nominated: HLSLNamespace,
83 CommonAncestor: AST.getTranslationUnitDecl());
84
85 AST.getTranslationUnitDecl()->addDecl(D: UsingDecl);
86}
87
88void HLSLExternalSemaSource::defineHLSLVectorAlias() {
89 ASTContext &AST = SemaPtr->getASTContext();
90
91 llvm::SmallVector<NamedDecl *> TemplateParams;
92
93 auto *TypeParam = TemplateTypeParmDecl::Create(
94 C: AST, DC: HLSLNamespace, KeyLoc: SourceLocation(), NameLoc: SourceLocation(), D: 0, P: 0,
95 Id: &AST.Idents.get(Name: "element", TokenCode: tok::TokenKind::identifier), Typename: false, ParameterPack: false);
96 TypeParam->setDefaultArgument(
97 C: AST, DefArg: SemaPtr->getTrivialTemplateArgumentLoc(
98 Arg: TemplateArgument(AST.FloatTy), NTTPType: QualType(), Loc: SourceLocation()));
99
100 TemplateParams.emplace_back(Args&: TypeParam);
101
102 auto *SizeParam = NonTypeTemplateParmDecl::Create(
103 C: AST, DC: HLSLNamespace, StartLoc: SourceLocation(), IdLoc: SourceLocation(), D: 0, P: 1,
104 Id: &AST.Idents.get(Name: "element_count", TokenCode: tok::TokenKind::identifier), T: AST.IntTy,
105 ParameterPack: false, TInfo: AST.getTrivialTypeSourceInfo(T: AST.IntTy));
106 llvm::APInt Val(AST.getIntWidth(T: AST.IntTy), 4);
107 TemplateArgument Default(AST, llvm::APSInt(std::move(Val)), AST.IntTy,
108 /*IsDefaulted=*/true);
109 SizeParam->setDefaultArgument(C: AST, DefArg: SemaPtr->getTrivialTemplateArgumentLoc(
110 Arg: Default, NTTPType: AST.IntTy, Loc: SourceLocation()));
111 TemplateParams.emplace_back(Args&: SizeParam);
112
113 auto *ParamList =
114 TemplateParameterList::Create(C: AST, TemplateLoc: SourceLocation(), LAngleLoc: SourceLocation(),
115 Params: TemplateParams, RAngleLoc: SourceLocation(), RequiresClause: nullptr);
116
117 IdentifierInfo &II = AST.Idents.get(Name: "vector", TokenCode: tok::TokenKind::identifier);
118
119 QualType AliasType = AST.getDependentSizedExtVectorType(
120 VectorType: AST.getTemplateTypeParmType(Depth: 0, Index: 0, ParameterPack: false, ParmDecl: TypeParam),
121 SizeExpr: DeclRefExpr::Create(
122 Context: AST, QualifierLoc: NestedNameSpecifierLoc(), TemplateKWLoc: SourceLocation(), D: SizeParam, RefersToEnclosingVariableOrCapture: false,
123 NameInfo: DeclarationNameInfo(SizeParam->getDeclName(), SourceLocation()),
124 T: AST.IntTy, VK: VK_LValue),
125 AttrLoc: SourceLocation());
126
127 auto *Record = TypeAliasDecl::Create(C&: AST, DC: HLSLNamespace, StartLoc: SourceLocation(),
128 IdLoc: SourceLocation(), Id: &II,
129 TInfo: AST.getTrivialTypeSourceInfo(T: AliasType));
130 Record->setImplicit(true);
131
132 auto *Template =
133 TypeAliasTemplateDecl::Create(C&: AST, DC: HLSLNamespace, L: SourceLocation(),
134 Name: Record->getIdentifier(), Params: ParamList, Decl: Record);
135
136 Record->setDescribedAliasTemplate(Template);
137 Template->setImplicit(true);
138 Template->setLexicalDeclContext(Record->getDeclContext());
139 HLSLNamespace->addDecl(D: Template);
140}
141
142void HLSLExternalSemaSource::defineHLSLMatrixAlias() {
143 ASTContext &AST = SemaPtr->getASTContext();
144 llvm::SmallVector<NamedDecl *> TemplateParams;
145
146 auto *TypeParam = TemplateTypeParmDecl::Create(
147 C: AST, DC: HLSLNamespace, KeyLoc: SourceLocation(), NameLoc: SourceLocation(), D: 0, P: 0,
148 Id: &AST.Idents.get(Name: "element", TokenCode: tok::TokenKind::identifier), Typename: false, ParameterPack: false);
149 TypeParam->setDefaultArgument(
150 C: AST, DefArg: SemaPtr->getTrivialTemplateArgumentLoc(
151 Arg: TemplateArgument(AST.FloatTy), NTTPType: QualType(), Loc: SourceLocation()));
152
153 TemplateParams.emplace_back(Args&: TypeParam);
154
155 // these should be 64 bit to be consistent with other clang matrices.
156 auto *RowsParam = NonTypeTemplateParmDecl::Create(
157 C: AST, DC: HLSLNamespace, StartLoc: SourceLocation(), IdLoc: SourceLocation(), D: 0, P: 1,
158 Id: &AST.Idents.get(Name: "rows_count", TokenCode: tok::TokenKind::identifier), T: AST.IntTy,
159 ParameterPack: false, TInfo: AST.getTrivialTypeSourceInfo(T: AST.IntTy));
160 llvm::APInt RVal(AST.getIntWidth(T: AST.IntTy), 4);
161 TemplateArgument RDefault(AST, llvm::APSInt(std::move(RVal)), AST.IntTy,
162 /*IsDefaulted=*/true);
163 RowsParam->setDefaultArgument(
164 C: AST, DefArg: SemaPtr->getTrivialTemplateArgumentLoc(Arg: RDefault, NTTPType: AST.IntTy,
165 Loc: SourceLocation()));
166 TemplateParams.emplace_back(Args&: RowsParam);
167
168 auto *ColsParam = NonTypeTemplateParmDecl::Create(
169 C: AST, DC: HLSLNamespace, StartLoc: SourceLocation(), IdLoc: SourceLocation(), D: 0, P: 2,
170 Id: &AST.Idents.get(Name: "cols_count", TokenCode: tok::TokenKind::identifier), T: AST.IntTy,
171 ParameterPack: false, TInfo: AST.getTrivialTypeSourceInfo(T: AST.IntTy));
172 llvm::APInt CVal(AST.getIntWidth(T: AST.IntTy), 4);
173 TemplateArgument CDefault(AST, llvm::APSInt(std::move(CVal)), AST.IntTy,
174 /*IsDefaulted=*/true);
175 ColsParam->setDefaultArgument(
176 C: AST, DefArg: SemaPtr->getTrivialTemplateArgumentLoc(Arg: CDefault, NTTPType: AST.IntTy,
177 Loc: SourceLocation()));
178 TemplateParams.emplace_back(Args&: ColsParam);
179
180 const unsigned MaxMatDim = SemaPtr->getLangOpts().MaxMatrixDimension;
181
182 auto *MaxRow = IntegerLiteral::Create(
183 C: AST, V: llvm::APInt(AST.getIntWidth(T: AST.IntTy), MaxMatDim), type: AST.IntTy,
184 l: SourceLocation());
185 auto *MaxCol = IntegerLiteral::Create(
186 C: AST, V: llvm::APInt(AST.getIntWidth(T: AST.IntTy), MaxMatDim), type: AST.IntTy,
187 l: SourceLocation());
188
189 auto *RowsRef = DeclRefExpr::Create(
190 Context: AST, QualifierLoc: NestedNameSpecifierLoc(), TemplateKWLoc: SourceLocation(), D: RowsParam,
191 /*RefersToEnclosingVariableOrCapture*/ false,
192 NameInfo: DeclarationNameInfo(RowsParam->getDeclName(), SourceLocation()),
193 T: AST.IntTy, VK: VK_LValue);
194 auto *ColsRef = DeclRefExpr::Create(
195 Context: AST, QualifierLoc: NestedNameSpecifierLoc(), TemplateKWLoc: SourceLocation(), D: ColsParam,
196 /*RefersToEnclosingVariableOrCapture*/ false,
197 NameInfo: DeclarationNameInfo(ColsParam->getDeclName(), SourceLocation()),
198 T: AST.IntTy, VK: VK_LValue);
199
200 auto *RowsLE = BinaryOperator::Create(C: AST, lhs: RowsRef, rhs: MaxRow, opc: BO_LE, ResTy: AST.BoolTy,
201 VK: VK_PRValue, OK: OK_Ordinary,
202 opLoc: SourceLocation(), FPFeatures: FPOptionsOverride());
203 auto *ColsLE = BinaryOperator::Create(C: AST, lhs: ColsRef, rhs: MaxCol, opc: BO_LE, ResTy: AST.BoolTy,
204 VK: VK_PRValue, OK: OK_Ordinary,
205 opLoc: SourceLocation(), FPFeatures: FPOptionsOverride());
206
207 auto *RequiresExpr = BinaryOperator::Create(
208 C: AST, lhs: RowsLE, rhs: ColsLE, opc: BO_LAnd, ResTy: AST.BoolTy, VK: VK_PRValue, OK: OK_Ordinary,
209 opLoc: SourceLocation(), FPFeatures: FPOptionsOverride());
210
211 auto *ParamList = TemplateParameterList::Create(
212 C: AST, TemplateLoc: SourceLocation(), LAngleLoc: SourceLocation(), Params: TemplateParams, RAngleLoc: SourceLocation(),
213 RequiresClause: RequiresExpr);
214
215 IdentifierInfo &II = AST.Idents.get(Name: "matrix", TokenCode: tok::TokenKind::identifier);
216
217 QualType AliasType = AST.getDependentSizedMatrixType(
218 ElementType: AST.getTemplateTypeParmType(Depth: 0, Index: 0, ParameterPack: false, ParmDecl: TypeParam),
219 RowExpr: DeclRefExpr::Create(
220 Context: AST, QualifierLoc: NestedNameSpecifierLoc(), TemplateKWLoc: SourceLocation(), D: RowsParam, RefersToEnclosingVariableOrCapture: false,
221 NameInfo: DeclarationNameInfo(RowsParam->getDeclName(), SourceLocation()),
222 T: AST.IntTy, VK: VK_LValue),
223 ColumnExpr: DeclRefExpr::Create(
224 Context: AST, QualifierLoc: NestedNameSpecifierLoc(), TemplateKWLoc: SourceLocation(), D: ColsParam, RefersToEnclosingVariableOrCapture: false,
225 NameInfo: DeclarationNameInfo(ColsParam->getDeclName(), SourceLocation()),
226 T: AST.IntTy, VK: VK_LValue),
227 AttrLoc: SourceLocation());
228
229 auto *Record = TypeAliasDecl::Create(C&: AST, DC: HLSLNamespace, StartLoc: SourceLocation(),
230 IdLoc: SourceLocation(), Id: &II,
231 TInfo: AST.getTrivialTypeSourceInfo(T: AliasType));
232 Record->setImplicit(true);
233
234 auto *Template =
235 TypeAliasTemplateDecl::Create(C&: AST, DC: HLSLNamespace, L: SourceLocation(),
236 Name: Record->getIdentifier(), Params: ParamList, Decl: Record);
237
238 Record->setDescribedAliasTemplate(Template);
239 Template->setImplicit(true);
240 Template->setLexicalDeclContext(Record->getDeclContext());
241 HLSLNamespace->addDecl(D: Template);
242}
243
244void HLSLExternalSemaSource::defineTrivialHLSLTypes() {
245 defineHLSLVectorAlias();
246 defineHLSLMatrixAlias();
247}
248
249void HLSLExternalSemaSource::defineHeapResourceInfoTypes() {
250 ASTContext &AST = SemaPtr->getASTContext();
251 CXXRecordDecl *ResDecl = BuiltinTypeDeclBuilder(*SemaPtr, HLSLDetailNamespace,
252 "heap_resource_info")
253 .finalizeForwardDeclaration();
254 if (!ResDecl->isCompleteDefinition())
255 BuiltinTypeDeclBuilder(*SemaPtr, ResDecl)
256 .addMemberVariable(Name: "Index", Type: AST.UnsignedIntTy, Attrs: {})
257 .completeDefinition();
258
259 CXXRecordDecl *SampDecl =
260 BuiltinTypeDeclBuilder(*SemaPtr, HLSLDetailNamespace, "heap_sampler_info")
261 .finalizeForwardDeclaration();
262 if (!SampDecl->isCompleteDefinition())
263 BuiltinTypeDeclBuilder(*SemaPtr, SampDecl)
264 .addMemberVariable(Name: "Index", Type: AST.UnsignedIntTy, Attrs: {})
265 .completeDefinition();
266}
267
268void HLSLExternalSemaSource::defineInternalHLSLTypes() {
269 defineHeapResourceInfoTypes();
270}
271
272/// Set up common members and attributes for buffer types
273static BuiltinTypeDeclBuilder setupBufferType(CXXRecordDecl *Decl, Sema &S,
274 ResourceClass RC, bool IsROV,
275 bool RawBuffer, bool HasCounter) {
276 return BuiltinTypeDeclBuilder(S, Decl)
277 .addBufferHandles(RC, IsROV, RawBuffer, HasCounter)
278 .addDefaultHandleConstructor()
279 .addCopyConstructor()
280 .addCopyAssignmentOperator()
281 .addHeapResourceInfoConstructor(HasCounter)
282 .addStaticInitializationFunctions(HasCounter);
283}
284
285/// Set up common members and attributes for sampler types
286static BuiltinTypeDeclBuilder setupSamplerType(CXXRecordDecl *Decl, Sema &S) {
287 return BuiltinTypeDeclBuilder(S, Decl)
288 .addSamplerHandle()
289 .addDefaultHandleConstructor()
290 .addCopyConstructor()
291 .addCopyAssignmentOperator()
292 .addHeapSamplerInfoConstructor()
293 .addStaticInitializationFunctions(HasCounter: false);
294}
295
296namespace {
297LLVM_ENABLE_BITMASK_ENUMS_IN_NAMESPACE();
298
299/// Which members a texture type has. Overloads within a member family
300/// (e.g., offset overloads for samplers) follow from ResourceDimension.
301enum class TexCap : uint32_t {
302 Load = 1u << 0, // Load(int<N+1>) taking a mip level
303 LoadMS = 1u << 1, // Load(int<N>, int sampleIndex) on a multisampled type
304 LoadRW = 1u << 2, // Load(int<N>) on a writable texture
305 Subscript = 1u << 3, // operator[]
306 Mips = 1u << 4, // mips[]
307 Sample = 1u << 5, // Sample, SampleBias, SampleGrad, SampleLevel
308 SampleCmp = 1u << 6, // SampleCmp, SampleCmpLevelZero
309 Gather = 1u << 7, // Gather*, GatherCmp*
310 CalcLOD = 1u << 8, // CalculateLevelOfDetail, ...Unclamped
311 GetDims = 1u << 9, // GetDimensions
312
313 // TODO: multisampled types need an MS-specific GetDimensions
314 // https://github.com/llvm/wg-hlsl/issues/347
315
316 LLVM_MARK_AS_BITMASK_ENUM(/*LargestValue=*/GetDims)
317};
318
319/// How a type's template parameters are spelled. Independent of its
320/// capabilities; also decides which types get a vector partial specialization.
321enum class TemplateShape {
322 ElementType, // template<typename T = float4>
323 ElementTypeAndSampleCount, // template<typename T, uint N>
324};
325
326struct TextureTypeInfo {
327 const char *Name;
328 ResourceClass RC;
329 ResourceDimension Dim;
330 bool IsArray;
331 bool IsROV;
332 TemplateShape Shape;
333 TexCap Caps;
334
335 bool has(TexCap C) const { return (Caps & C) != TexCap{}; }
336 bool hasSampleCount() const {
337 return Shape == TemplateShape::ElementTypeAndSampleCount;
338 }
339};
340} // namespace
341
342static const TextureTypeInfo TextureTypes[] = {
343 {.Name: "Texture1D", .RC: ResourceClass::SRV, .Dim: ResourceDimension::Dim1D,
344 /*IsArray=*/false, /*IsROV=*/false, .Shape: TemplateShape::ElementType,
345 .Caps: TexCap::Load | TexCap::Subscript | TexCap::Mips | TexCap::Sample |
346 TexCap::SampleCmp | TexCap::CalcLOD},
347 {.Name: "RWTexture1D", .RC: ResourceClass::UAV, .Dim: ResourceDimension::Dim1D,
348 /*IsArray=*/false, /*IsROV=*/false, .Shape: TemplateShape::ElementType,
349 .Caps: TexCap::LoadRW | TexCap::Subscript},
350 {.Name: "Texture1DArray", .RC: ResourceClass::SRV, .Dim: ResourceDimension::Dim1D,
351 /*IsArray=*/true, /*IsROV=*/false, .Shape: TemplateShape::ElementType,
352 .Caps: TexCap::Load | TexCap::Subscript | TexCap::Mips | TexCap::Sample |
353 TexCap::SampleCmp | TexCap::CalcLOD},
354 {.Name: "RWTexture1DArray", .RC: ResourceClass::UAV, .Dim: ResourceDimension::Dim1D,
355 /*IsArray=*/true, /*IsROV=*/false, .Shape: TemplateShape::ElementType,
356 .Caps: TexCap::LoadRW | TexCap::Subscript},
357 {.Name: "Texture2D", .RC: ResourceClass::SRV, .Dim: ResourceDimension::Dim2D,
358 /*IsArray=*/false, /*IsROV=*/false, .Shape: TemplateShape::ElementType,
359 .Caps: TexCap::Load | TexCap::Subscript | TexCap::Mips | TexCap::Sample |
360 TexCap::SampleCmp | TexCap::CalcLOD | TexCap::Gather |
361 TexCap::GetDims},
362 {.Name: "RWTexture2D", .RC: ResourceClass::UAV, .Dim: ResourceDimension::Dim2D,
363 /*IsArray=*/false, /*IsROV=*/false, .Shape: TemplateShape::ElementType,
364 .Caps: TexCap::LoadRW | TexCap::Subscript | TexCap::GetDims},
365 {.Name: "Texture2DArray", .RC: ResourceClass::SRV, .Dim: ResourceDimension::Dim2D,
366 /*IsArray=*/true, /*IsROV=*/false, .Shape: TemplateShape::ElementType,
367 .Caps: TexCap::Load | TexCap::Subscript | TexCap::Mips | TexCap::Sample |
368 TexCap::SampleCmp | TexCap::CalcLOD | TexCap::Gather |
369 TexCap::GetDims},
370 {.Name: "RWTexture2DArray", .RC: ResourceClass::UAV, .Dim: ResourceDimension::Dim2D,
371 /*IsArray=*/true, /*IsROV=*/false, .Shape: TemplateShape::ElementType,
372 .Caps: TexCap::LoadRW | TexCap::Subscript | TexCap::GetDims},
373 {.Name: "Texture2DMS", .RC: ResourceClass::SRV, .Dim: ResourceDimension::Dim2D,
374 /*IsArray=*/false, /*IsROV=*/false,
375 .Shape: TemplateShape::ElementTypeAndSampleCount,
376 .Caps: TexCap::LoadMS | TexCap::Subscript},
377 {.Name: "Texture3D", .RC: ResourceClass::SRV, .Dim: ResourceDimension::Dim3D,
378 /*IsArray=*/false, /*IsROV=*/false, .Shape: TemplateShape::ElementType,
379 .Caps: TexCap::Load | TexCap::Subscript | TexCap::Mips | TexCap::Sample |
380 TexCap::CalcLOD | TexCap::GetDims},
381 {.Name: "RWTexture3D", .RC: ResourceClass::UAV, .Dim: ResourceDimension::Dim3D,
382 /*IsArray=*/false, /*IsROV=*/false, .Shape: TemplateShape::ElementType,
383 .Caps: TexCap::LoadRW | TexCap::Subscript | TexCap::GetDims},
384 {.Name: "TextureCube", .RC: ResourceClass::SRV, .Dim: ResourceDimension::Cube,
385 /*IsArray=*/false, /*IsROV=*/false, .Shape: TemplateShape::ElementType,
386 .Caps: TexCap::Sample | TexCap::SampleCmp | TexCap::CalcLOD | TexCap::Gather |
387 TexCap::GetDims},
388 {.Name: "TextureCubeArray", .RC: ResourceClass::SRV, .Dim: ResourceDimension::Cube,
389 /*IsArray=*/true, /*IsROV=*/false, .Shape: TemplateShape::ElementType,
390 .Caps: TexCap::Sample | TexCap::SampleCmp | TexCap::CalcLOD | TexCap::Gather |
391 TexCap::GetDims},
392};
393
394static BuiltinTypeDeclBuilder setupTextureType(CXXRecordDecl *Decl, Sema &S,
395 const TextureTypeInfo &T) {
396 const ResourceDimension Dim = T.Dim;
397 const bool IsArray = T.IsArray;
398
399 Expr *SampleCountExpr = nullptr;
400 if (T.hasSampleCount()) {
401 ClassTemplateDecl *CTD = Decl->getDescribedClassTemplate();
402 assert(CTD && "multisampled texture must be a class template");
403 // Parameter 1 is the N in Texture2DMS<T, N>.
404 auto *NTTP = cast<NonTypeTemplateParmDecl>(
405 Val: CTD->getTemplateParameters()->getParam(Idx: 1));
406 SampleCountExpr =
407 S.BuildDeclRefExpr(D: NTTP, Ty: NTTP->getType(), VK: VK_PRValue, Loc: SourceLocation());
408 }
409
410 BuiltinTypeDeclBuilder B(S, Decl);
411 B.addTextureHandle(RC: T.RC, IsROV: T.IsROV, IsArray, RD: Dim, SampleCountExpr);
412
413 // The `mips` member holds a second copy of the resource handle.
414 // addCopyConstructor, addCopyAssignmentOperator and
415 // addStaticInitializationFunctions are what initialize that copy, and they
416 // look the member up by name, so it has to exist before they run.
417 if (T.has(C: TexCap::Mips))
418 B.addMipsMember(Dim);
419
420 B.addDefaultHandleConstructor()
421 .addCopyConstructor()
422 .addCopyAssignmentOperator()
423 .addHeapResourceInfoConstructor()
424 .addStaticInitializationFunctions(HasCounter: false);
425
426 if (T.has(C: TexCap::Load))
427 B.addTextureLoadMethods(Dim, IsArray);
428 if (T.has(C: TexCap::LoadMS))
429 B.addTextureLoadMSMethods(Dim, IsArray);
430 if (T.has(C: TexCap::LoadRW))
431 B.addRWTextureLoadMethods(Dim, IsArray);
432 if (T.has(C: TexCap::Subscript))
433 B.addArraySubscriptOperators(Dim, IsArray);
434
435 if (T.has(C: TexCap::Sample))
436 B.addSampleMethods(Dim, IsArray)
437 .addSampleBiasMethods(Dim, IsArray)
438 .addSampleGradMethods(Dim, IsArray)
439 .addSampleLevelMethods(Dim, IsArray);
440 if (T.has(C: TexCap::SampleCmp))
441 B.addSampleCmpMethods(Dim, IsArray)
442 .addSampleCmpLevelZeroMethods(Dim, IsArray);
443 if (T.has(C: TexCap::CalcLOD))
444 B.addCalculateLodMethods(Dim);
445 if (T.has(C: TexCap::GetDims))
446 B.addGetDimensionsMethods(Dim);
447 if (T.has(C: TexCap::Gather))
448 B.addGatherMethods(Dim, IsArray).addGatherCmpMethods(Dim, IsArray);
449
450 return B;
451}
452
453// Add a partial specialization for a template. The `TextureTemplate` is
454// `Texture<element_type>`, and it will be specialized for vectors:
455// `Texture<vector<element_type, element_count>>`.
456static ClassTemplatePartialSpecializationDecl *
457addVectorTexturePartialSpecialization(Sema &S, NamespaceDecl *HLSLNamespace,
458 ClassTemplateDecl *TextureTemplate) {
459 ASTContext &AST = S.getASTContext();
460
461 // Create the template parameters: element_type and element_count.
462 auto *ElementType = TemplateTypeParmDecl::Create(
463 C: AST, DC: HLSLNamespace, KeyLoc: SourceLocation(), NameLoc: SourceLocation(), D: 0, P: 0,
464 Id: &AST.Idents.get(Name: "element_type"), Typename: false, ParameterPack: false);
465 auto *ElementCount = NonTypeTemplateParmDecl::Create(
466 C: AST, DC: HLSLNamespace, StartLoc: SourceLocation(), IdLoc: SourceLocation(), D: 0, P: 1,
467 Id: &AST.Idents.get(Name: "element_count"), T: AST.IntTy, ParameterPack: false,
468 TInfo: AST.getTrivialTypeSourceInfo(T: AST.IntTy));
469
470 auto *TemplateParams = TemplateParameterList::Create(
471 C: AST, TemplateLoc: SourceLocation(), LAngleLoc: SourceLocation(), Params: {ElementType, ElementCount},
472 RAngleLoc: SourceLocation(), RequiresClause: nullptr);
473
474 // Create the dependent vector type: vector<element_type, element_count>.
475 QualType VectorType = AST.getDependentSizedExtVectorType(
476 VectorType: AST.getTemplateTypeParmType(Depth: 0, Index: 0, ParameterPack: false, ParmDecl: ElementType),
477 SizeExpr: DeclRefExpr::Create(
478 Context: AST, QualifierLoc: NestedNameSpecifierLoc(), TemplateKWLoc: SourceLocation(), D: ElementCount, RefersToEnclosingVariableOrCapture: false,
479 NameInfo: DeclarationNameInfo(ElementCount->getDeclName(), SourceLocation()),
480 T: AST.IntTy, VK: VK_LValue),
481 AttrLoc: SourceLocation());
482
483 // Create the partial specialization declaration.
484 QualType CanonInjectedTST =
485 AST.getCanonicalType(T: AST.getTemplateSpecializationType(
486 Keyword: ElaboratedTypeKeyword::Class, T: TemplateName(TextureTemplate),
487 SpecifiedArgs: {TemplateArgument(VectorType)}, CanonicalArgs: {}));
488
489 auto *PartialSpec = ClassTemplatePartialSpecializationDecl::Create(
490 Context&: AST, TK: TagDecl::TagKind::Class, DC: HLSLNamespace, StartLoc: SourceLocation(),
491 IdLoc: SourceLocation(), Params: TemplateParams, SpecializedTemplate: TextureTemplate,
492 Args: {TemplateArgument(VectorType)},
493 CanonInjectedTST: CanQualType::CreateUnsafe(Other: CanonInjectedTST), PrevDecl: nullptr);
494
495 // Set the template arguments as written.
496 TemplateArgument Arg(VectorType);
497 TemplateArgumentLoc ArgLoc =
498 S.getTrivialTemplateArgumentLoc(Arg, NTTPType: QualType(), Loc: SourceLocation());
499 TemplateArgumentListInfo ArgsInfo =
500 TemplateArgumentListInfo(SourceLocation(), SourceLocation());
501 ArgsInfo.addArgument(Loc: ArgLoc);
502 PartialSpec->setTemplateArgsAsWritten(
503 ASTTemplateArgumentListInfo::Create(C: AST, List: ArgsInfo));
504
505 PartialSpec->setImplicit(true);
506 PartialSpec->setLexicalDeclContext(HLSLNamespace);
507 PartialSpec->setHasExternalLexicalStorage();
508
509 // Add the partial specialization to the namespace and the class template.
510 HLSLNamespace->addDecl(D: PartialSpec);
511 TextureTemplate->AddPartialSpecialization(D: PartialSpec, InsertToken: {});
512
513 return PartialSpec;
514}
515
516// This function is responsible for constructing the constraint expression for
517// this concept:
518// template<typename T> concept is_typed_resource_element_compatible =
519// __is_typed_resource_element_compatible<T>;
520static Expr *constructTypedBufferConstraintExpr(Sema &S, SourceLocation NameLoc,
521 TemplateTypeParmDecl *T) {
522 ASTContext &Context = S.getASTContext();
523
524 // Obtain the QualType for 'bool'
525 QualType BoolTy = Context.BoolTy;
526
527 // Create a QualType that points to this TemplateTypeParmDecl
528 QualType TType = Context.getTypeDeclType(Decl: T);
529
530 // Create a TypeSourceInfo for the template type parameter 'T'
531 TypeSourceInfo *TTypeSourceInfo =
532 Context.getTrivialTypeSourceInfo(T: TType, Loc: NameLoc);
533
534 TypeTraitExpr *TypedResExpr = TypeTraitExpr::Create(
535 C: Context, T: BoolTy, Loc: NameLoc, Kind: UTT_IsTypedResourceElementCompatible,
536 Args: {TTypeSourceInfo}, RParenLoc: NameLoc, Value: true);
537
538 return TypedResExpr;
539}
540
541// This function is responsible for constructing the constraint expression for
542// this concept:
543// template<typename T> concept is_constant_buffer_element_compatible =
544// std::is_class_v<T> && !__is_intangible(T);
545static Expr *constructConstantBufferConstraintExpr(Sema &S,
546 SourceLocation NameLoc,
547 TemplateTypeParmDecl *T) {
548 ASTContext &Context = S.getASTContext();
549
550 // Obtain the QualType for 'bool'
551 QualType BoolTy = Context.BoolTy;
552
553 // Create a QualType that points to this TemplateTypeParmDecl
554 QualType TType = Context.getTypeDeclType(Decl: T);
555
556 // Create a TypeSourceInfo for the template type parameter 'T'
557 TypeSourceInfo *TTypeSourceInfo =
558 Context.getTrivialTypeSourceInfo(T: TType, Loc: NameLoc);
559
560 TypeTraitExpr *ResExpr = TypeTraitExpr::Create(
561 C: Context, T: BoolTy, Loc: NameLoc, Kind: UTT_IsConstantBufferElementCompatible,
562 Args: {TTypeSourceInfo}, RParenLoc: NameLoc, Value: true);
563
564 return ResExpr;
565}
566
567// This function is responsible for constructing the constraint expression for
568// this concept:
569// template<typename T> concept is_structured_resource_element_compatible =
570// !__is_intangible<T> && sizeof(T) >= 1;
571static Expr *constructStructuredBufferConstraintExpr(Sema &S,
572 SourceLocation NameLoc,
573 TemplateTypeParmDecl *T) {
574 ASTContext &Context = S.getASTContext();
575
576 // Obtain the QualType for 'bool'
577 QualType BoolTy = Context.BoolTy;
578
579 // Create a QualType that points to this TemplateTypeParmDecl
580 QualType TType = Context.getTypeDeclType(Decl: T);
581
582 // Create a TypeSourceInfo for the template type parameter 'T'
583 TypeSourceInfo *TTypeSourceInfo =
584 Context.getTrivialTypeSourceInfo(T: TType, Loc: NameLoc);
585
586 TypeTraitExpr *IsIntangibleExpr =
587 TypeTraitExpr::Create(C: Context, T: BoolTy, Loc: NameLoc, Kind: UTT_IsIntangibleType,
588 Args: {TTypeSourceInfo}, RParenLoc: NameLoc, Value: true);
589
590 // negate IsIntangibleExpr
591 UnaryOperator *NotIntangibleExpr = UnaryOperator::Create(
592 C: Context, input: IsIntangibleExpr, opc: UO_LNot, type: BoolTy, VK: VK_LValue, OK: OK_Ordinary,
593 l: NameLoc, CanOverflow: false, FPFeatures: FPOptionsOverride());
594
595 // element types also may not be of 0 size
596 UnaryExprOrTypeTraitExpr *SizeOfExpr = new (Context) UnaryExprOrTypeTraitExpr(
597 UETT_SizeOf, TTypeSourceInfo, BoolTy, NameLoc, NameLoc);
598
599 // Create a BinaryOperator that checks if the size of the type is not equal to
600 // 1 Empty structs have a size of 1 in HLSL, so we need to check for that
601 IntegerLiteral *rhs = IntegerLiteral::Create(
602 C: Context, V: llvm::APInt(Context.getTypeSize(T: Context.getSizeType()), 1, true),
603 type: Context.getSizeType(), l: NameLoc);
604
605 BinaryOperator *SizeGEQOneExpr =
606 BinaryOperator::Create(C: Context, lhs: SizeOfExpr, rhs, opc: BO_GE, ResTy: BoolTy, VK: VK_LValue,
607 OK: OK_Ordinary, opLoc: NameLoc, FPFeatures: FPOptionsOverride());
608
609 // Combine the two constraints
610 BinaryOperator *CombinedExpr = BinaryOperator::Create(
611 C: Context, lhs: NotIntangibleExpr, rhs: SizeGEQOneExpr, opc: BO_LAnd, ResTy: BoolTy, VK: VK_LValue,
612 OK: OK_Ordinary, opLoc: NameLoc, FPFeatures: FPOptionsOverride());
613
614 return CombinedExpr;
615}
616
617enum class HLSLBufferType { Typed, Structured, Constant };
618
619static ConceptDecl *constructBufferConceptDecl(Sema &S, NamespaceDecl *NSD,
620 HLSLBufferType BT) {
621 ASTContext &Context = S.getASTContext();
622 DeclContext *DC = NSD->getDeclContext();
623 SourceLocation DeclLoc = SourceLocation();
624
625 IdentifierInfo &ElementTypeII = Context.Idents.get(Name: "element_type");
626 TemplateTypeParmDecl *T = TemplateTypeParmDecl::Create(
627 C: Context, DC: NSD->getDeclContext(), KeyLoc: DeclLoc, NameLoc: DeclLoc,
628 /*D=*/0,
629 /*P=*/0,
630 /*Id=*/&ElementTypeII,
631 /*Typename=*/true,
632 /*ParameterPack=*/false);
633
634 T->setDeclContext(DC);
635 T->setReferenced();
636
637 // Create and Attach Template Parameter List to ConceptDecl
638 TemplateParameterList *ConceptParams = TemplateParameterList::Create(
639 C: Context, TemplateLoc: DeclLoc, LAngleLoc: DeclLoc, Params: {T}, RAngleLoc: DeclLoc, RequiresClause: nullptr);
640
641 DeclarationName DeclName;
642 Expr *ConstraintExpr = nullptr;
643
644 switch (BT) {
645 case HLSLBufferType::Typed:
646 DeclName = DeclarationName(
647 &Context.Idents.get(Name: "__is_typed_resource_element_compatible"));
648 ConstraintExpr = constructTypedBufferConstraintExpr(S, NameLoc: DeclLoc, T);
649 break;
650 case HLSLBufferType::Structured:
651 DeclName = DeclarationName(
652 &Context.Idents.get(Name: "__is_structured_resource_element_compatible"));
653 ConstraintExpr = constructStructuredBufferConstraintExpr(S, NameLoc: DeclLoc, T);
654 break;
655 case HLSLBufferType::Constant:
656 DeclName = DeclarationName(
657 &Context.Idents.get(Name: "__is_constant_buffer_element_compatible"));
658 ConstraintExpr = constructConstantBufferConstraintExpr(S, NameLoc: DeclLoc, T);
659 break;
660 }
661
662 // Create a ConceptDecl
663 ConceptDecl *CD =
664 ConceptDecl::Create(C&: Context, DC: NSD->getDeclContext(), L: DeclLoc, Name: DeclName,
665 Params: ConceptParams, ConstraintExpr);
666
667 // Attach the template parameter list to the ConceptDecl
668 CD->setTemplateParameters(ConceptParams);
669
670 // Add the concept declaration to the Translation Unit Decl
671 NSD->getDeclContext()->addDecl(D: CD);
672
673 return CD;
674}
675
676void HLSLExternalSemaSource::defineHLSLTypesWithForwardDeclarations() {
677 ASTContext &AST = SemaPtr->getASTContext();
678 CXXRecordDecl *Decl;
679 ConceptDecl *TypedBufferConcept = constructBufferConceptDecl(
680 S&: *SemaPtr, NSD: HLSLNamespace, BT: HLSLBufferType::Typed);
681 ConceptDecl *StructuredBufferConcept = constructBufferConceptDecl(
682 S&: *SemaPtr, NSD: HLSLNamespace, BT: HLSLBufferType::Structured);
683 ConceptDecl *ConstantBufferConcept = constructBufferConceptDecl(
684 S&: *SemaPtr, NSD: HLSLNamespace, BT: HLSLBufferType::Constant);
685
686 Decl = BuiltinTypeDeclBuilder(*SemaPtr, HLSLNamespace, "ConstantBuffer")
687 .addSimpleTemplateParams(Names: {"element_type"}, CD: ConstantBufferConcept)
688 .finalizeForwardDeclaration();
689
690 onCompletion(Record: Decl, Fn: [this](CXXRecordDecl *Decl) {
691 setupBufferType(Decl, S&: *SemaPtr, RC: ResourceClass::CBuffer, /*IsROV=*/false,
692 /*RawBuffer=*/false, /*HasCounter=*/false)
693 .addConstantBufferConversionToType()
694 .completeDefinition();
695 });
696
697 Decl = BuiltinTypeDeclBuilder(*SemaPtr, HLSLNamespace, "Buffer")
698 .addSimpleTemplateParams(Names: {"element_type"}, CD: TypedBufferConcept)
699 .finalizeForwardDeclaration();
700
701 onCompletion(Record: Decl, Fn: [this](CXXRecordDecl *Decl) {
702 setupBufferType(Decl, S&: *SemaPtr, RC: ResourceClass::SRV, /*IsROV=*/false,
703 /*RawBuffer=*/false, /*HasCounter=*/false)
704 .addArraySubscriptOperators()
705 .addLoadMethods()
706 .addGetDimensionsMethodForBuffer()
707 .completeDefinition();
708 });
709
710 Decl = BuiltinTypeDeclBuilder(*SemaPtr, HLSLNamespace, "RWBuffer")
711 .addSimpleTemplateParams(Names: {"element_type"}, CD: TypedBufferConcept)
712 .finalizeForwardDeclaration();
713
714 onCompletion(Record: Decl, Fn: [this](CXXRecordDecl *Decl) {
715 setupBufferType(Decl, S&: *SemaPtr, RC: ResourceClass::UAV, /*IsROV=*/false,
716 /*RawBuffer=*/false, /*HasCounter=*/false)
717 .addArraySubscriptOperators()
718 .addLoadMethods()
719 .addGetDimensionsMethodForBuffer()
720 .completeDefinition();
721 });
722
723 Decl =
724 BuiltinTypeDeclBuilder(*SemaPtr, HLSLNamespace, "RasterizerOrderedBuffer")
725 .addSimpleTemplateParams(Names: {"element_type"}, CD: StructuredBufferConcept)
726 .finalizeForwardDeclaration();
727 onCompletion(Record: Decl, Fn: [this](CXXRecordDecl *Decl) {
728 setupBufferType(Decl, S&: *SemaPtr, RC: ResourceClass::UAV, /*IsROV=*/true,
729 /*RawBuffer=*/false, /*HasCounter=*/false)
730 .addArraySubscriptOperators()
731 .addLoadMethods()
732 .addGetDimensionsMethodForBuffer()
733 .completeDefinition();
734 });
735
736 Decl = BuiltinTypeDeclBuilder(*SemaPtr, HLSLNamespace, "StructuredBuffer")
737 .addSimpleTemplateParams(Names: {"element_type"}, CD: StructuredBufferConcept)
738 .finalizeForwardDeclaration();
739 onCompletion(Record: Decl, Fn: [this](CXXRecordDecl *Decl) {
740 setupBufferType(Decl, S&: *SemaPtr, RC: ResourceClass::SRV, /*IsROV=*/false,
741 /*RawBuffer=*/true, /*HasCounter=*/false)
742 .addArraySubscriptOperators()
743 .addLoadMethods()
744 .addGetDimensionsMethodForBuffer()
745 .completeDefinition();
746 });
747
748 Decl = BuiltinTypeDeclBuilder(*SemaPtr, HLSLNamespace, "RWStructuredBuffer")
749 .addSimpleTemplateParams(Names: {"element_type"}, CD: StructuredBufferConcept)
750 .finalizeForwardDeclaration();
751 onCompletion(Record: Decl, Fn: [this](CXXRecordDecl *Decl) {
752 setupBufferType(Decl, S&: *SemaPtr, RC: ResourceClass::UAV, /*IsROV=*/false,
753 /*RawBuffer=*/true, /*HasCounter=*/true)
754 .addArraySubscriptOperators()
755 .addLoadMethods()
756 .addIncrementCounterMethod()
757 .addDecrementCounterMethod()
758 .addGetDimensionsMethodForBuffer()
759 .completeDefinition();
760 });
761
762 Decl =
763 BuiltinTypeDeclBuilder(*SemaPtr, HLSLNamespace, "AppendStructuredBuffer")
764 .addSimpleTemplateParams(Names: {"element_type"}, CD: StructuredBufferConcept)
765 .finalizeForwardDeclaration();
766 onCompletion(Record: Decl, Fn: [this](CXXRecordDecl *Decl) {
767 setupBufferType(Decl, S&: *SemaPtr, RC: ResourceClass::UAV, /*IsROV=*/false,
768 /*RawBuffer=*/true, /*HasCounter=*/true)
769 .addAppendMethod()
770 .addGetDimensionsMethodForBuffer()
771 .completeDefinition();
772 });
773
774 Decl =
775 BuiltinTypeDeclBuilder(*SemaPtr, HLSLNamespace, "ConsumeStructuredBuffer")
776 .addSimpleTemplateParams(Names: {"element_type"}, CD: StructuredBufferConcept)
777 .finalizeForwardDeclaration();
778 onCompletion(Record: Decl, Fn: [this](CXXRecordDecl *Decl) {
779 setupBufferType(Decl, S&: *SemaPtr, RC: ResourceClass::UAV, /*IsROV=*/false,
780 /*RawBuffer=*/true, /*HasCounter=*/true)
781 .addConsumeMethod()
782 .addGetDimensionsMethodForBuffer()
783 .completeDefinition();
784 });
785
786 Decl = BuiltinTypeDeclBuilder(*SemaPtr, HLSLNamespace,
787 "RasterizerOrderedStructuredBuffer")
788 .addSimpleTemplateParams(Names: {"element_type"}, CD: StructuredBufferConcept)
789 .finalizeForwardDeclaration();
790 onCompletion(Record: Decl, Fn: [this](CXXRecordDecl *Decl) {
791 setupBufferType(Decl, S&: *SemaPtr, RC: ResourceClass::UAV, /*IsROV=*/true,
792 /*RawBuffer=*/true, /*HasCounter=*/true)
793 .addArraySubscriptOperators()
794 .addLoadMethods()
795 .addIncrementCounterMethod()
796 .addDecrementCounterMethod()
797 .addGetDimensionsMethodForBuffer()
798 .completeDefinition();
799 });
800
801 Decl = BuiltinTypeDeclBuilder(*SemaPtr, HLSLNamespace, "ByteAddressBuffer")
802 .finalizeForwardDeclaration();
803 onCompletion(Record: Decl, Fn: [this](CXXRecordDecl *Decl) {
804 setupBufferType(Decl, S&: *SemaPtr, RC: ResourceClass::SRV, /*IsROV=*/false,
805 /*RawBuffer=*/true, /*HasCounter=*/false)
806 .addByteAddressBufferLoadMethods()
807 .addGetDimensionsMethodForBuffer()
808 .completeDefinition();
809 });
810 Decl = BuiltinTypeDeclBuilder(*SemaPtr, HLSLNamespace, "RWByteAddressBuffer")
811 .finalizeForwardDeclaration();
812 onCompletion(Record: Decl, Fn: [this](CXXRecordDecl *Decl) {
813 setupBufferType(Decl, S&: *SemaPtr, RC: ResourceClass::UAV, /*IsROV=*/false,
814 /*RawBuffer=*/true, /*HasCounter=*/false)
815 .addByteAddressBufferLoadMethods()
816 .addByteAddressBufferStoreMethods()
817 .addByteAddressBufferInterlockedMethods()
818 .addGetDimensionsMethodForBuffer()
819 .completeDefinition();
820 });
821 Decl = BuiltinTypeDeclBuilder(*SemaPtr, HLSLNamespace,
822 "RasterizerOrderedByteAddressBuffer")
823 .finalizeForwardDeclaration();
824 onCompletion(Record: Decl, Fn: [this](CXXRecordDecl *Decl) {
825 setupBufferType(Decl, S&: *SemaPtr, RC: ResourceClass::UAV, /*IsROV=*/true,
826 /*RawBuffer=*/true, /*HasCounter=*/false)
827 .addByteAddressBufferInterlockedMethods()
828 .addGetDimensionsMethodForBuffer()
829 .completeDefinition();
830 });
831
832 Decl = BuiltinTypeDeclBuilder(*SemaPtr, HLSLNamespace, "SamplerState")
833 .finalizeForwardDeclaration();
834 onCompletion(Record: Decl, Fn: [this](CXXRecordDecl *Decl) {
835 setupSamplerType(Decl, S&: *SemaPtr).completeDefinition();
836 });
837
838 Decl =
839 BuiltinTypeDeclBuilder(*SemaPtr, HLSLNamespace, "SamplerComparisonState")
840 .finalizeForwardDeclaration();
841 onCompletion(Record: Decl, Fn: [this](CXXRecordDecl *Decl) {
842 setupSamplerType(Decl, S&: *SemaPtr).completeDefinition();
843 });
844
845 QualType Float4Ty = AST.getExtVectorType(VectorType: AST.FloatTy, NumElts: 4);
846 for (const TextureTypeInfo &T : TextureTypes) {
847 BuiltinTypeDeclBuilder TexBuilder(*SemaPtr, HLSLNamespace, T.Name);
848 switch (T.Shape) {
849 case TemplateShape::ElementType:
850 TexBuilder.addSimpleTemplateParams(Names: {"element_type"}, DefaultTypes: {Float4Ty},
851 CD: TypedBufferConcept);
852 break;
853 case TemplateShape::ElementTypeAndSampleCount:
854 TexBuilder.addMSTextureTemplateParams(ElementName: "element_type", SampleCountName: "sample_count",
855 CD: TypedBufferConcept);
856 break;
857 }
858 Decl = TexBuilder.finalizeForwardDeclaration();
859
860 onCompletion(Record: Decl, Fn: [this, &T](CXXRecordDecl *Decl) {
861 setupTextureType(Decl, S&: *SemaPtr, T).completeDefinition();
862 });
863
864 if (T.Shape != TemplateShape::ElementType)
865 continue;
866
867 CXXRecordDecl *PartialSpec = addVectorTexturePartialSpecialization(
868 S&: *SemaPtr, HLSLNamespace, TextureTemplate: Decl->getDescribedClassTemplate());
869 onCompletion(Record: PartialSpec, Fn: [this, &T](CXXRecordDecl *Decl) {
870 setupTextureType(Decl, S&: *SemaPtr, T).completeDefinition();
871 });
872 }
873}
874
875// Shapes of the synthesized atomic overloads. The read-modify-write operations
876// can report the previous value through a trailing reference. Compare-store
877// takes only inputs; compare-exchange adds the trailing reference.
878enum class AtomicOverloadShape {
879 Binary, // (dest, value)
880 BinaryWithOriginal, // (dest, value, original_value)
881 CompareStore, // (dest, compare_value, value)
882 CompareExchange, // (dest, compare_value, value, original_value)
883};
884
885// Build a single overload of an HLSL atomic intrinsic in the hlsl namespace.
886// `dest` is an address-space-qualified reference; `original_value` (when
887// present) is an HLSL `out` parameter. The synthesized FunctionDecl aliases
888// the underlying clang builtin via BuiltinAliasAttr.
889static void buildAtomicOverload(Sema &S, NamespaceDecl *NS, StringRef FuncName,
890 StringRef BuiltinName, QualType ElemTy,
891 LangAS DestAS, AtomicOverloadShape Shape) {
892 ASTContext &AST = S.getASTContext();
893
894 QualType DestTy =
895 AST.getLValueReferenceType(T: AST.getAddrSpaceQualType(T: ElemTy, AddressSpace: DestAS));
896 QualType OrigRefTy = S.HLSL().getInoutParameterType(Ty: ElemTy);
897
898 SmallVector<QualType, 4> ParamTypes = {DestTy, ElemTy};
899 switch (Shape) {
900 case AtomicOverloadShape::Binary:
901 break;
902 case AtomicOverloadShape::BinaryWithOriginal:
903 ParamTypes.push_back(Elt: OrigRefTy);
904 break;
905 case AtomicOverloadShape::CompareStore:
906 ParamTypes.push_back(Elt: ElemTy);
907 break;
908 case AtomicOverloadShape::CompareExchange:
909 ParamTypes.push_back(Elt: ElemTy);
910 ParamTypes.push_back(Elt: OrigRefTy);
911 break;
912 }
913
914 // The loop below stops at the end of ParamTypes, so a shorter overload
915 // ignores the trailing names.
916 constexpr const char *BinaryNames[] = {"dest", "value", "original_value"};
917 constexpr const char *CompareNames[] = {"dest", "compare_value", "value",
918 "original_value"};
919 const bool IsCompare = Shape == AtomicOverloadShape::CompareStore ||
920 Shape == AtomicOverloadShape::CompareExchange;
921 const bool HasOriginalValue =
922 Shape == AtomicOverloadShape::BinaryWithOriginal ||
923 Shape == AtomicOverloadShape::CompareExchange;
924 ArrayRef<const char *> ParamNames = IsCompare
925 ? ArrayRef<const char *>(CompareNames)
926 : ArrayRef<const char *>(BinaryNames);
927
928 FunctionProtoType::ExtProtoInfo EPI;
929 SmallVector<FunctionType::ExtParameterInfo, 4> ParamExtInfos(
930 ParamTypes.size());
931 if (HasOriginalValue) {
932 ParamExtInfos.back() = ParamExtInfos.back().withABI(kind: ParameterABI::HLSLOut);
933 EPI.ExtParameterInfos = ParamExtInfos.data();
934 }
935 QualType FuncTy = AST.getFunctionType(ResultTy: AST.VoidTy, Args: ParamTypes, EPI);
936 auto *TSInfo = AST.getTrivialTypeSourceInfo(T: FuncTy, Loc: SourceLocation());
937
938 IdentifierInfo &FuncII = AST.Idents.get(Name: FuncName, TokenCode: tok::TokenKind::identifier);
939 DeclarationName FuncDeclName(&FuncII);
940
941 FunctionDecl *FD = FunctionDecl::Create(
942 C&: AST, DC: NS, StartLoc: SourceLocation(), NLoc: SourceLocation(), N: FuncDeclName, T: FuncTy, TInfo: TSInfo,
943 SC: SC_Extern, /*UsesFPIntrin=*/false, /*isInlineSpecified=*/false,
944 /*hasWrittenPrototype=*/true);
945
946 SmallVector<ParmVarDecl *, 4> ParmDecls;
947 unsigned I = 0;
948 for (auto [ParamType, ParamName] : llvm::zip(t&: ParamTypes, u&: ParamNames)) {
949 IdentifierInfo &PII = AST.Idents.get(Name: ParamName, TokenCode: tok::TokenKind::identifier);
950 ParmVarDecl *Parm = ParmVarDecl::Create(
951 C&: AST, DC: FD, StartLoc: SourceLocation(), IdLoc: SourceLocation(), Id: &PII, T: ParamType,
952 TInfo: AST.getTrivialTypeSourceInfo(T: ParamType, Loc: SourceLocation()), S: SC_None,
953 DefArg: nullptr);
954 if (HasOriginalValue && I == ParamTypes.size() - 1)
955 Parm->addAttr(A: HLSLParamModifierAttr::Create(
956 Ctx&: AST, Range: SourceRange(), S: HLSLParamModifierAttr::Keyword_out));
957 Parm->setScopeInfo(scopeDepth: 0, parameterIndex: I++);
958 ParmDecls.push_back(Elt: Parm);
959 }
960 FD->setParams(ParmDecls);
961
962 IdentifierInfo &BuiltinII =
963 S.getPreprocessor().getIdentifierTable().get(Name: BuiltinName);
964 FD->addAttr(A: BuiltinAliasAttr::CreateImplicit(Ctx&: AST, BuiltinName: &BuiltinII));
965 FD->setImplicit();
966 NS->addDecl(D: FD);
967}
968
969// Synthesize the InterlockedFunc overload set: {int, uint, int64_t, uint64_t}
970// x {groupshared, device} x {2-arg, 3-arg}. Operations that always report the
971// previous value, such as InterlockedExchange, only get the 3-arg form.
972// InterlockedExchange also accepts float, which lowers to a bitwise exchange
973// of the 32-bit pattern.
974static void defineHLSLInterlockedFunc(Sema &S, NamespaceDecl *NS,
975 StringRef FuncName, StringRef BuiltinName,
976 bool RequiresOriginalValue = false,
977 bool SupportsFloat = false) {
978 ASTContext &AST = S.getASTContext();
979 // HLSL: int64_t == long, uint64_t == unsigned long (see hlsl_basic_types.h).
980 SmallVector<QualType, 5> Elems = {AST.IntTy, AST.UnsignedIntTy, AST.LongTy,
981 AST.UnsignedLongTy};
982 if (SupportsFloat)
983 Elems.push_back(Elt: AST.FloatTy);
984 LangAS AddrSpaces[] = {LangAS::hlsl_groupshared, LangAS::hlsl_device};
985
986 for (QualType ElemTy : Elems)
987 for (LangAS AS : AddrSpaces) {
988 if (!RequiresOriginalValue)
989 buildAtomicOverload(S, NS, FuncName, BuiltinName, ElemTy, DestAS: AS,
990 Shape: AtomicOverloadShape::Binary);
991 buildAtomicOverload(S, NS, FuncName, BuiltinName, ElemTy, DestAS: AS,
992 Shape: AtomicOverloadShape::BinaryWithOriginal);
993 }
994}
995
996// Synthesize the compare-and-swap overload sets: {int, uint, int64_t,
997// uint64_t} x {groupshared, device}. Each function has one form only.
998static void defineHLSLInterlockedCompareFunc(Sema &S, NamespaceDecl *NS,
999 StringRef FuncName,
1000 StringRef BuiltinName,
1001 AtomicOverloadShape Shape) {
1002 ASTContext &AST = S.getASTContext();
1003 QualType Elems[] = {AST.IntTy, AST.UnsignedIntTy, AST.LongTy,
1004 AST.UnsignedLongTy};
1005
1006 for (QualType ElemTy : Elems)
1007 for (LangAS AS : {LangAS::hlsl_groupshared, LangAS::hlsl_device})
1008 buildAtomicOverload(S, NS, FuncName, BuiltinName, ElemTy, DestAS: AS, Shape);
1009}
1010
1011// The float-bitwise compare-and-swap functions have their own names and take
1012// float only.
1013static void defineHLSLInterlockedCompareFuncFloat(Sema &S, NamespaceDecl *NS,
1014 StringRef FuncName,
1015 StringRef BuiltinName,
1016 AtomicOverloadShape Shape) {
1017 ASTContext &AST = S.getASTContext();
1018 buildAtomicOverload(S, NS, FuncName, BuiltinName, ElemTy: AST.FloatTy,
1019 DestAS: LangAS::hlsl_groupshared, Shape);
1020 buildAtomicOverload(S, NS, FuncName, BuiltinName, ElemTy: AST.FloatTy,
1021 DestAS: LangAS::hlsl_device, Shape);
1022}
1023
1024void HLSLExternalSemaSource::defineHLSLAtomicIntrinsics() {
1025 defineHLSLInterlockedFunc(S&: *SemaPtr, NS: HLSLNamespace, FuncName: "InterlockedAdd",
1026 BuiltinName: "__builtin_hlsl_interlocked_add");
1027 defineHLSLInterlockedFunc(S&: *SemaPtr, NS: HLSLNamespace, FuncName: "InterlockedAnd",
1028 BuiltinName: "__builtin_hlsl_interlocked_and");
1029 defineHLSLInterlockedCompareFunc(
1030 S&: *SemaPtr, NS: HLSLNamespace, FuncName: "InterlockedCompareExchange",
1031 BuiltinName: "__builtin_hlsl_interlocked_compare_exchange",
1032 Shape: AtomicOverloadShape::CompareExchange);
1033 defineHLSLInterlockedCompareFuncFloat(
1034 S&: *SemaPtr, NS: HLSLNamespace, FuncName: "InterlockedCompareExchangeFloatBitwise",
1035 BuiltinName: "__builtin_hlsl_interlocked_compare_exchange_float_bitwise",
1036 Shape: AtomicOverloadShape::CompareExchange);
1037 defineHLSLInterlockedCompareFunc(S&: *SemaPtr, NS: HLSLNamespace,
1038 FuncName: "InterlockedCompareStore",
1039 BuiltinName: "__builtin_hlsl_interlocked_compare_store",
1040 Shape: AtomicOverloadShape::CompareStore);
1041 defineHLSLInterlockedCompareFuncFloat(
1042 S&: *SemaPtr, NS: HLSLNamespace, FuncName: "InterlockedCompareStoreFloatBitwise",
1043 BuiltinName: "__builtin_hlsl_interlocked_compare_store_float_bitwise",
1044 Shape: AtomicOverloadShape::CompareStore);
1045 defineHLSLInterlockedFunc(S&: *SemaPtr, NS: HLSLNamespace, FuncName: "InterlockedExchange",
1046 BuiltinName: "__builtin_hlsl_interlocked_exchange",
1047 /*RequiresOriginalValue=*/true,
1048 /*SupportsFloat=*/true);
1049 defineHLSLInterlockedFunc(S&: *SemaPtr, NS: HLSLNamespace, FuncName: "InterlockedMax",
1050 BuiltinName: "__builtin_hlsl_interlocked_max");
1051 defineHLSLInterlockedFunc(S&: *SemaPtr, NS: HLSLNamespace, FuncName: "InterlockedMin",
1052 BuiltinName: "__builtin_hlsl_interlocked_min");
1053 defineHLSLInterlockedFunc(S&: *SemaPtr, NS: HLSLNamespace, FuncName: "InterlockedOr",
1054 BuiltinName: "__builtin_hlsl_interlocked_or");
1055 defineHLSLInterlockedFunc(S&: *SemaPtr, NS: HLSLNamespace, FuncName: "InterlockedXor",
1056 BuiltinName: "__builtin_hlsl_interlocked_xor");
1057}
1058
1059void HLSLExternalSemaSource::onCompletion(CXXRecordDecl *Record,
1060 CompletionFunction Fn) {
1061 if (!Record->isCompleteDefinition())
1062 Completions.insert(KV: std::make_pair(x: Record->getCanonicalDecl(), y&: Fn));
1063}
1064
1065void HLSLExternalSemaSource::CompleteType(TagDecl *Tag) {
1066 if (!isa<CXXRecordDecl>(Val: Tag))
1067 return;
1068 auto *Record = cast<CXXRecordDecl>(Val: Tag);
1069 Record = Record->getCanonicalDecl();
1070 auto It = Completions.find(Val: Record);
1071 if (It == Completions.end())
1072 return;
1073 // Move out the callback and erase before invoking it: the callback can
1074 // re-enter CompleteType and mutate Completions, which invalidates It under
1075 // backward-shift deletion.
1076 CompletionFunction Fn = std::move(It->second);
1077 Completions.erase(I: It);
1078 Fn(Record);
1079}
1080