Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions tools/clang/unittests/HLSL/PixTest.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -3389,6 +3389,8 @@ void RaygenInternalName()
}

TEST_F(PixTest, DebugBreakInstrumentation_Basic) {
if (m_ver.SkipDxilVersion(1, 10))
return;

const char *source = R"x(
[numthreads(1, 1, 1)]
Expand Down Expand Up @@ -3426,6 +3428,8 @@ void main() {
}

TEST_F(PixTest, DebugBreakInstrumentation_Multiple) {
if (m_ver.SkipDxilVersion(1, 10))
return;

const char *source = R"x(
RWByteAddressBuffer buf : register(u0);
Expand Down
62 changes: 38 additions & 24 deletions tools/clang/unittests/HLSL/ValidationTest.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -4882,6 +4882,15 @@ TEST_F(ValidationTest, CacheInitWithLowPrec) {
TestCheck(L"..\\DXILValidation\\val-dx-type-lowprec.ll");
}

// PSVRuntimeInfo4 adds NumBytesGroupSharedMemory and is only emitted for
// validator version >= 1.10; earlier validators emit PSVRuntimeInfo3.
static uint32_t GetExpectedPSVRuntimeInfoSize(const VersionSupportInfo &ver) {
bool HasV4 =
ver.m_ValMajor > 1 || (ver.m_ValMajor == 1 && ver.m_ValMinor >= 10);
return HasV4 ? static_cast<uint32_t>(sizeof(PSVRuntimeInfo4))
: static_cast<uint32_t>(sizeof(PSVRuntimeInfo3));
}

TEST_F(ValidationTest, PSVStringTableReorder) {
if (!m_ver.m_InternalValidator)
if (m_ver.SkipDxilVersion(1, 8))
Expand Down Expand Up @@ -4916,7 +4925,7 @@ TEST_F(ValidationTest, PSVStringTableReorder) {
const uint32_t *PSVPtr = (const uint32_t *)GetDxilPartData(pPSVPart);

uint32_t PSVRuntimeInfo_size = *(PSVPtr++);
VERIFY_ARE_EQUAL(sizeof(PSVRuntimeInfo4), PSVRuntimeInfo_size);
VERIFY_ARE_EQUAL(GetExpectedPSVRuntimeInfoSize(m_ver), PSVRuntimeInfo_size);
PSVRuntimeInfo4 *PSVInfo =
const_cast<PSVRuntimeInfo4 *>((const PSVRuntimeInfo4 *)PSVPtr);
VERIFY_ARE_EQUAL(2u, PSVInfo->SigInputElements);
Expand Down Expand Up @@ -5108,7 +5117,7 @@ TEST_F(ValidationTest, PSVSemanticIndexTableReorder) {
const uint32_t *PSVPtr = (const uint32_t *)GetDxilPartData(pPSVPart);

uint32_t PSVRuntimeInfo_size = *(PSVPtr++);
VERIFY_ARE_EQUAL(sizeof(PSVRuntimeInfo4), PSVRuntimeInfo_size);
VERIFY_ARE_EQUAL(GetExpectedPSVRuntimeInfoSize(m_ver), PSVRuntimeInfo_size);
PSVRuntimeInfo4 *PSVInfo =
const_cast<PSVRuntimeInfo4 *>((const PSVRuntimeInfo4 *)PSVPtr);
VERIFY_ARE_EQUAL(PSVInfo->SigInputElements, 3u);
Expand Down Expand Up @@ -5443,17 +5452,18 @@ struct SimplePSV {
llvm::MutableArrayRef<uint32_t> PCInputToOutputTable;
llvm::MutableArrayRef<uint32_t> ViewIDOutputMask[DXIL::kNumOutputStreams];
llvm::MutableArrayRef<uint32_t> ViewIDPCOutputMask;
SimplePSV(const DxilPartHeader *pPSVPart);
SimplePSV(const DxilPartHeader *pPSVPart, uint32_t ExpectedRuntimeInfoSize);
};

SimplePSV::SimplePSV(const DxilPartHeader *pPSVPart) {
SimplePSV::SimplePSV(const DxilPartHeader *pPSVPart,
uint32_t ExpectedRuntimeInfoSize) {
uint32_t PartSize = pPSVPart->PartSize;
uint32_t *PSVPtr =
const_cast<uint32_t *>((const uint32_t *)GetDxilPartData(pPSVPart));
const uint32_t *PSVPtrEnd = PSVPtr + PartSize / 4;

uint32_t PSVRuntimeInfoSize = *(PSVPtr++);
VERIFY_ARE_EQUAL(sizeof(PSVRuntimeInfo4), PSVRuntimeInfoSize);
VERIFY_ARE_EQUAL(ExpectedRuntimeInfoSize, PSVRuntimeInfoSize);
PSVRuntimeInfo4 *PSVInfo4 =
const_cast<PSVRuntimeInfo4 *>((const PSVRuntimeInfo4 *)PSVPtr);
PSVInfo = PSVInfo4;
Expand Down Expand Up @@ -5581,7 +5591,7 @@ TEST_F(ValidationTest, PSVContentValidationVS) {
VERIFY_ARE_NOT_EQUAL(hlsl::end(pHeader), pPartIter);

const DxilPartHeader *pPSVPart = (const DxilPartHeader *)(*pPartIter);
SimplePSV PSV(pPSVPart);
SimplePSV PSV(pPSVPart, GetExpectedPSVRuntimeInfoSize(m_ver));

// Update PSV.
PSV.SigInput[0].InterpolationMode = 20;
Expand Down Expand Up @@ -5737,7 +5747,7 @@ TEST_F(ValidationTest, PSVContentValidationHS) {
VERIFY_ARE_NOT_EQUAL(hlsl::end(pHeader), pPartIter);

const DxilPartHeader *pPSVPart = (const DxilPartHeader *)(*pPartIter);
SimplePSV PSV(pPSVPart);
SimplePSV PSV(pPSVPart, GetExpectedPSVRuntimeInfoSize(m_ver));

// Update PSV.
PSV.SigPatchConstOrPrim[0].InterpolationMode = 20;
Expand Down Expand Up @@ -5887,7 +5897,7 @@ TEST_F(ValidationTest, PSVContentValidationDS) {
VERIFY_ARE_NOT_EQUAL(hlsl::end(pHeader), pPartIter);

const DxilPartHeader *pPSVPart = (const DxilPartHeader *)(*pPartIter);
SimplePSV PSV(pPSVPart);
SimplePSV PSV(pPSVPart, GetExpectedPSVRuntimeInfoSize(m_ver));

// Update PSV.
PSV.SigPatchConstOrPrim[0].InterpolationMode = 20;
Expand Down Expand Up @@ -6044,7 +6054,7 @@ TEST_F(ValidationTest, PSVContentValidationGS) {
VERIFY_ARE_NOT_EQUAL(hlsl::end(pHeader), pPartIter);

const DxilPartHeader *pPSVPart = (const DxilPartHeader *)(*pPartIter);
SimplePSV PSV(pPSVPart);
SimplePSV PSV(pPSVPart, GetExpectedPSVRuntimeInfoSize(m_ver));
// Update PSV.
PSV.PSVInfo->MaxVertexCount = 2;

Expand Down Expand Up @@ -6132,7 +6142,7 @@ TEST_F(ValidationTest, PSVContentValidationPS) {
VERIFY_ARE_NOT_EQUAL(hlsl::end(pHeader), pPartIter);

const DxilPartHeader *pPSVPart = (const DxilPartHeader *)(*pPartIter);
SimplePSV PSV(pPSVPart);
SimplePSV PSV(pPSVPart, GetExpectedPSVRuntimeInfoSize(m_ver));

// Update PSV.
PSV.PSVInfo->PS.DepthOutput = 1;
Expand Down Expand Up @@ -6217,7 +6227,7 @@ TEST_F(ValidationTest, PSVContentValidationCS) {
VERIFY_ARE_NOT_EQUAL(hlsl::end(pHeader), pPartIter);

const DxilPartHeader *pPSVPart = (const DxilPartHeader *)(*pPartIter);
SimplePSV PSV(pPSVPart);
SimplePSV PSV(pPSVPart, GetExpectedPSVRuntimeInfoSize(m_ver));
// Update PSV.
PSV.PSVInfo->NumThreadsX = 1;

Expand Down Expand Up @@ -6299,7 +6309,7 @@ TEST_F(ValidationTest, PSVContentValidationMS) {
VERIFY_ARE_NOT_EQUAL(hlsl::end(pHeader), pPartIter);

const DxilPartHeader *pPSVPart = (const DxilPartHeader *)(*pPartIter);
SimplePSV PSV(pPSVPart);
SimplePSV PSV(pPSVPart, GetExpectedPSVRuntimeInfoSize(m_ver));
// Update PSV.
memset(PSV.ViewIDOutputMask[0].data(), 0,
PSV.ViewIDOutputMask[0].size() * sizeof(uint32_t));
Expand Down Expand Up @@ -6366,7 +6376,7 @@ TEST_F(ValidationTest, PSVContentValidationAS) {
VERIFY_ARE_NOT_EQUAL(hlsl::end(pHeader), pPartIter);

const DxilPartHeader *pPSVPart = (const DxilPartHeader *)(*pPartIter);
SimplePSV PSV(pPSVPart);
SimplePSV PSV(pPSVPart, GetExpectedPSVRuntimeInfoSize(m_ver));

// Update PSV.
PSV.PSVInfo->AS.PayloadSizeInBytes = 0;
Expand Down Expand Up @@ -6560,7 +6570,7 @@ TEST_F(ValidationTest, WrongPSVSizeOnZeros) {
const uint32_t *PSVPtr = (const uint32_t *)GetDxilPartData(pPSVPart);

uint32_t PSVRuntimeInfo_size = *(PSVPtr++);
VERIFY_ARE_EQUAL(sizeof(PSVRuntimeInfo4), PSVRuntimeInfo_size);
VERIFY_ARE_EQUAL(GetExpectedPSVRuntimeInfoSize(m_ver), PSVRuntimeInfo_size);
PSVRuntimeInfo4 *PSVInfo =
const_cast<PSVRuntimeInfo4 *>((const PSVRuntimeInfo4 *)PSVPtr);
VERIFY_ARE_EQUAL(2u, PSVInfo->SigInputElements);
Expand Down Expand Up @@ -6791,11 +6801,14 @@ TEST_F(ValidationTest, WrongPSVVersion) {
VERIFY_IS_NOT_NULL(p60WithPSV68Result);
VERIFY_SUCCEEDED(p60WithPSV68Result->GetStatus(&status));
VERIFY_FAILED(status);
CheckOperationResultMsgs(
p60WithPSV68Result,
{"DXIL container mismatch for 'PSVRuntimeInfoSize' between 'PSV0' "
"part:('56') and DXIL module:('24')"},
/*maySucceedAnyway*/ false, /*bRegex*/ false);
std::string ExpectedPSVSizeStr =
std::to_string(GetExpectedPSVRuntimeInfoSize(m_ver));
std::string Msg60WithPSV68 =
"DXIL container mismatch for 'PSVRuntimeInfoSize' between 'PSV0' "
"part:('" +
ExpectedPSVSizeStr + "') and DXIL module:('24')";
CheckOperationResultMsgs(p60WithPSV68Result, {Msg60WithPSV68.c_str()},
/*maySucceedAnyway*/ false, /*bRegex*/ false);

// Create a new Blob.
CComPtr<IDxcBlobEncoding> pProgram68WithPSV60;
Expand All @@ -6809,9 +6822,10 @@ TEST_F(ValidationTest, WrongPSVVersion) {
VERIFY_IS_NOT_NULL(p68WithPSV60Result);
VERIFY_SUCCEEDED(p68WithPSV60Result->GetStatus(&status));
VERIFY_FAILED(status);
CheckOperationResultMsgs(
p68WithPSV60Result,
{"DXIL container mismatch for 'PSVRuntimeInfoSize' between 'PSV0' "
"part:('24') and DXIL module:('56')"},
/*maySucceedAnyway*/ false, /*bRegex*/ false);
std::string Msg68WithPSV60 =
"DXIL container mismatch for 'PSVRuntimeInfoSize' between 'PSV0' "
"part:('24') and DXIL module:('" +
ExpectedPSVSizeStr + "')";
CheckOperationResultMsgs(p68WithPSV60Result, {Msg68WithPSV60.c_str()},
/*maySucceedAnyway*/ false, /*bRegex*/ false);
}
Loading