Fix crash when calling `GetBindGroupLayout()` on null bindGroupLayout This patch fixes a crash issue when calling `GetBindGroupLayout()` on a null bind group layout item in the pipeline layout used to create the pipeline. According to the latest WebGPU SPEC, the parameter `index` is valid as long as `index < the size of this.[[layout]].[[bindGroupLayouts]]`, so calling `GetBindGroupLayout()` on a null bind group layout item in the pipeline layout should be legal and return an empty bind group layout. Bug: 377836524, 42241530 Test: dawn_unittests Change-Id: I57bc9a136e3da01e8442f6752c98bfa09eef6f2f Reviewed-on: https://dawn-review.googlesource.com/c/dawn/+/220455 Reviewed-by: Loko Kung <lokokung@google.com> Commit-Queue: Jiawei Shao <jiawei.shao@intel.com> Reviewed-by: Kai Ninomiya <kainino@chromium.org>
diff --git a/src/dawn/native/Pipeline.cpp b/src/dawn/native/Pipeline.cpp index a11b852..1508876 100644 --- a/src/dawn/native/Pipeline.cpp +++ b/src/dawn/native/Pipeline.cpp
@@ -312,7 +312,7 @@ "Bind group layout index (%u) exceeds the maximum number of bind groups (%u).", groupIndex, kMaxBindGroups); DAWN_INVALID_IF( - !mLayout->GetBindGroupLayoutsMask()[groupIndex], + static_cast<uint32_t>(groupIndex) >= mLayout->GetExplicitBindGroupLayoutsCount(), "Bind group layout index (%u) doesn't correspond to a bind group for this pipeline.", groupIndex); return {};
diff --git a/src/dawn/native/PipelineLayout.cpp b/src/dawn/native/PipelineLayout.cpp index e094a6a..737e040 100644 --- a/src/dawn/native/PipelineLayout.cpp +++ b/src/dawn/native/PipelineLayout.cpp
@@ -139,6 +139,7 @@ const UnpackedPtr<PipelineLayoutDescriptor>& descriptor, ApiObjectBase::UntrackedByDeviceTag tag) : ApiObjectBase(device, descriptor->label), + mExplicitBindGroupLayoutsCount(static_cast<uint32_t>(descriptor->bindGroupLayoutCount)), mImmediateDataRangeByteSize(descriptor->immediateDataRangeByteSize) { DAWN_ASSERT(descriptor->bindGroupLayoutCount <= kMaxBindGroups); @@ -447,20 +448,26 @@ const BindGroupLayoutBase* PipelineLayoutBase::GetFrontendBindGroupLayout( BindGroupIndex group) const { DAWN_ASSERT(!IsError()); - DAWN_ASSERT(group < kMaxBindGroupsTyped); - DAWN_ASSERT(mMask[group]); - const BindGroupLayoutBase* bgl = mBindGroupLayouts[group].Get(); - DAWN_ASSERT(bgl != nullptr); - return bgl; + DAWN_ASSERT(static_cast<uint32_t>(group) < mExplicitBindGroupLayoutsCount); + if (mMask[group]) { + const BindGroupLayoutBase* bgl = mBindGroupLayouts[group].Get(); + DAWN_ASSERT(bgl != nullptr); + return bgl; + } else { + return GetDevice()->GetEmptyBindGroupLayout(); + } } BindGroupLayoutBase* PipelineLayoutBase::GetFrontendBindGroupLayout(BindGroupIndex group) { DAWN_ASSERT(!IsError()); - DAWN_ASSERT(group < kMaxBindGroupsTyped); - DAWN_ASSERT(mMask[group]); - BindGroupLayoutBase* bgl = mBindGroupLayouts[group].Get(); - DAWN_ASSERT(bgl != nullptr); - return bgl; + DAWN_ASSERT(static_cast<uint32_t>(group) < mExplicitBindGroupLayoutsCount); + if (mMask[group]) { + BindGroupLayoutBase* bgl = mBindGroupLayouts[group].Get(); + DAWN_ASSERT(bgl != nullptr); + return bgl; + } else { + return GetDevice()->GetEmptyBindGroupLayout(); + } } const BindGroupLayoutInternalBase* PipelineLayoutBase::GetBindGroupLayout( @@ -510,6 +517,7 @@ ObjectContentHasher recorder; recorder.Record(mMask); + recorder.Record(mExplicitBindGroupLayoutsCount); for (BindGroupIndex group : IterateBitSet(mMask)) { recorder.Record(GetBindGroupLayout(group)->GetContentHash()); } @@ -532,6 +540,10 @@ return false; } + if (a->mExplicitBindGroupLayoutsCount != b->mExplicitBindGroupLayoutsCount) { + return false; + } + for (BindGroupIndex group : IterateBitSet(a->mMask)) { if (a->GetBindGroupLayout(group) != b->GetBindGroupLayout(group)) { return false; @@ -559,6 +571,10 @@ return true; } +uint32_t PipelineLayoutBase::GetExplicitBindGroupLayoutsCount() const { + return mExplicitBindGroupLayoutsCount; +} + uint32_t PipelineLayoutBase::GetImmediateDataRangeByteSize() const { return mImmediateDataRangeByteSize; }
diff --git a/src/dawn/native/PipelineLayout.h b/src/dawn/native/PipelineLayout.h index e742c76..1646b2c 100644 --- a/src/dawn/native/PipelineLayout.h +++ b/src/dawn/native/PipelineLayout.h
@@ -117,10 +117,13 @@ uint32_t GetImmediateDataRangeByteSize() const; + uint32_t GetExplicitBindGroupLayoutsCount() const; + protected: PipelineLayoutBase(DeviceBase* device, ObjectBase::ErrorTag tag, StringView label); void DestroyImpl() override; + uint32_t mExplicitBindGroupLayoutsCount = 0; PerBindGroup<Ref<BindGroupLayoutBase>> mBindGroupLayouts; BindGroupMask mMask; bool mHasPLS = false;
diff --git a/src/dawn/tests/unittests/validation/GetBindGroupLayoutValidationTests.cpp b/src/dawn/tests/unittests/validation/GetBindGroupLayoutValidationTests.cpp index d186ee8..63295b3 100644 --- a/src/dawn/tests/unittests/validation/GetBindGroupLayoutValidationTests.cpp +++ b/src/dawn/tests/unittests/validation/GetBindGroupLayoutValidationTests.cpp
@@ -1280,5 +1280,27 @@ EXPECT_THAT(pipeline.GetBindGroupLayout(3), BindGroupLayoutEq(emptyBGL)); } +// Test that a pipeline full of explicitly null BGLs correctly reflects empty BGLs. +TEST_F(GetBindGroupLayoutTests, NullBGLs) { + DAWN_SKIP_TEST_IF(UsesWire()); + + wgpu::PipelineLayout pl = + utils::MakePipelineLayout(device, {nullptr, nullptr, nullptr, nullptr}); + + wgpu::ComputePipelineDescriptor pipelineDesc; + pipelineDesc.layout = pl; + pipelineDesc.compute.module = utils::CreateShaderModule(device, R"( + @compute @workgroup_size(1) fn main() { + } + )"); + wgpu::ComputePipeline pipeline = device.CreateComputePipeline(&pipelineDesc); + + wgpu::BindGroupLayout emptyBGL = utils::MakeBindGroupLayout(device, {}); + EXPECT_THAT(pipeline.GetBindGroupLayout(0), BindGroupLayoutEq(emptyBGL)); + EXPECT_THAT(pipeline.GetBindGroupLayout(1), BindGroupLayoutEq(emptyBGL)); + EXPECT_THAT(pipeline.GetBindGroupLayout(2), BindGroupLayoutEq(emptyBGL)); + EXPECT_THAT(pipeline.GetBindGroupLayout(3), BindGroupLayoutEq(emptyBGL)); +} + } // anonymous namespace } // namespace dawn