[wgsl] Add subgroupMatrixLoad/Store resolver validation

* Add WGSL validation for subgroupMatrixLoad/Store
  * stride >= min stride
  * matrix required size <= pointer size

Bug: 527993420
Change-Id: I209b78b8984b8cc2d4ba319ad1bd57abed5f2f6d
Reviewed-on: https://dawn-review.googlesource.com/c/dawn/+/320915
Commit-Queue: Alan Baker <alanbaker@google.com>
Auto-Submit: Alan Baker <alanbaker@google.com>
Reviewed-by: dan sinclair <dsinclair@chromium.org>
Commit-Queue: dan sinclair <dsinclair@chromium.org>
diff --git a/src/tint/lang/wgsl/resolver/resolver.cc b/src/tint/lang/wgsl/resolver/resolver.cc
index 4c4a834..0e24018 100644
--- a/src/tint/lang/wgsl/resolver/resolver.cc
+++ b/src/tint/lang/wgsl/resolver/resolver.cc
@@ -2347,9 +2347,11 @@
             break;
 
         case wgsl::BuiltinFn::kSubgroupMatrixLoad:
+            TINT_RET_IF(!validator_.SubgroupMatrixLoadStore(call));
             RegisterLoad(args[0]);
             break;
         case wgsl::BuiltinFn::kSubgroupMatrixStore:
+            TINT_RET_IF(!validator_.SubgroupMatrixLoadStore(call));
             RegisterStore(args[0]);
             break;
 
diff --git a/src/tint/lang/wgsl/resolver/subgroup_matrix_test.cc b/src/tint/lang/wgsl/resolver/subgroup_matrix_test.cc
index ae8a2ff..7d17e95 100644
--- a/src/tint/lang/wgsl/resolver/subgroup_matrix_test.cc
+++ b/src/tint/lang/wgsl/resolver/subgroup_matrix_test.cc
@@ -1384,5 +1384,90 @@
 )");
 }
 
+TEST_F(ResolverSubgroupMatrixTest, Load_ColMajor_StrideLessThanMinStride_U32) {
+    ExpectError(
+        R"(
+enable chromium_experimental_subgroup_matrix;
+@group(0) @binding(0) var<storage> in : array<u32>;
+fn foo() {
+  _ = subgroupMatrixLoad<subgroup_matrix_left<u32, 8, 16>, col_major>(&in, 0, 15);
+})",
+        R"(input.wgsl:5:79 error: the stride argument (15, 60 bytes) of subgroupMatrixLoad must be greater than the minimum stride (64 bytes)
+  _ = subgroupMatrixLoad<subgroup_matrix_left<u32, 8, 16>, col_major>(&in, 0, 15);
+                                                                              ^^
+)");
+}
+
+TEST_F(ResolverSubgroupMatrixTest, Store_RowMajor_StrideLessThanMinStride_U8) {
+    ExpectError(
+        R"(
+enable chromium_experimental_subgroup_matrix;
+@group(0) @binding(0) var<storage, read_write> out : array<u32>;
+fn foo(m : subgroup_matrix_result<u8, 8, 8>) {
+  subgroupMatrixStore<row_major>(&out, 0, m, 1);
+})",
+        R"(input.wgsl:5:46 error: the stride argument (1, 4 bytes) of subgroupMatrixStore must be greater than the minimum stride (8 bytes)
+  subgroupMatrixStore<row_major>(&out, 0, m, 1);
+                                             ^
+)");
+}
+
+TEST_F(ResolverSubgroupMatrixTest, Load_ColMajor_PointerTooSmall_F32) {
+    ExpectError(
+        R"(
+enable chromium_experimental_subgroup_matrix;
+@group(0) @binding(0) var<storage> in : array<f32, 127>;
+fn foo(offset: u32, stride: u32) {
+  _ = subgroupMatrixLoad<subgroup_matrix_left<f32, 8, 16>, col_major>(&in, offset, stride);
+})",
+        R"(input.wgsl:5:7 error: the pointer operand of subgroupMatrixLoad is too small (508 bytes) for the matrix access (512 bytes)
+  _ = subgroupMatrixLoad<subgroup_matrix_left<f32, 8, 16>, col_major>(&in, offset, stride);
+      ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
+)");
+}
+
+TEST_F(ResolverSubgroupMatrixTest, Load_ColMajor_PointerTooSmall_F32_ConstOffset) {
+    ExpectError(
+        R"(
+enable chromium_experimental_subgroup_matrix;
+@group(0) @binding(0) var<storage> in : array<f32, 128>;
+fn foo(stride: u32) {
+  _ = subgroupMatrixLoad<subgroup_matrix_left<f32, 8, 16>, col_major>(&in, 1, stride);
+})",
+        R"(input.wgsl:5:7 error: the pointer operand of subgroupMatrixLoad is too small (512 bytes) for the matrix access (516 bytes)
+  _ = subgroupMatrixLoad<subgroup_matrix_left<f32, 8, 16>, col_major>(&in, 1, stride);
+      ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
+)");
+}
+
+TEST_F(ResolverSubgroupMatrixTest, Store_RowMajor_PointerTooSmall_F16_ConstStride) {
+    ExpectError(
+        R"(
+enable f16;
+enable chromium_experimental_subgroup_matrix;
+@group(0) @binding(0) var<storage, read_write> out : array<f16, 64>;
+fn foo(m : subgroup_matrix_right<f16, 8, 8>, offset: u32) {
+  subgroupMatrixStore<row_major>(&out, offset, m, 9);
+})",
+        R"(input.wgsl:6:3 error: the pointer operand of subgroupMatrixStore is too small (128 bytes) for the matrix access (142 bytes)
+  subgroupMatrixStore<row_major>(&out, offset, m, 9);
+  ^^^^^^^^^^^^^^^^^^^
+)");
+}
+
+TEST_F(ResolverSubgroupMatrixTest, Store_RowMajor_PointerTooSmall_U8_ConstOffsetAndStride) {
+    ExpectError(
+        R"(
+enable chromium_experimental_subgroup_matrix;
+@group(0) @binding(0) var<storage, read_write> out : array<u32, 63>;
+fn foo(m : subgroup_matrix_right<u8, 16, 8>) {
+  subgroupMatrixStore<row_major>(&out, 5, m, 8);
+})",
+        R"(input.wgsl:5:3 error: the pointer operand of subgroupMatrixStore is too small (252 bytes) for the matrix access (260 bytes)
+  subgroupMatrixStore<row_major>(&out, 5, m, 8);
+  ^^^^^^^^^^^^^^^^^^^
+)");
+}
+
 }  // namespace
 }  // namespace tint::resolver
diff --git a/src/tint/lang/wgsl/resolver/validator.cc b/src/tint/lang/wgsl/resolver/validator.cc
index 1eb276a..b8a6fb0 100644
--- a/src/tint/lang/wgsl/resolver/validator.cc
+++ b/src/tint/lang/wgsl/resolver/validator.cc
@@ -2264,6 +2264,116 @@
     return true;
 }
 
+bool Validator::SubgroupMatrixLoadStore(const sem::Call* call) const {
+    auto* builtin = call->Target()->As<sem::BuiltinFn>();
+    if (!builtin) {
+        return false;
+    }
+
+    const bool is_load = builtin->Fn() == wgsl::BuiltinFn::kSubgroupMatrixLoad;
+    auto* ptr_arg = call->Arguments()[0];
+    auto* ptr_arr_ty = ptr_arg->Type()->UnwrapPtr()->As<core::type::Array>();
+    auto* offset_arg = call->Arguments()[1];
+    const sem::ValueExpression* stride_arg = nullptr;
+    auto* templated_ident = call->Declaration()->target->identifier->As<ast::TemplatedIdentifier>();
+    bool col_major = false;
+    const core::type::SubgroupMatrix* mat_ty = nullptr;
+    if (is_load) {
+        TINT_ASSERT(templated_ident);
+        // Don't validate deprecated variant.
+        // TODO(b/529415904): remove this after deprecated variant is removed.
+        if (templated_ident->arguments.Length() != 2) {
+            return true;
+        }
+        auto* sem_expr = sem_.Get(templated_ident->arguments[1]);
+        auto* arg_major = sem_expr->As<sem::BuiltinEnumExpression<core::Majorness>>();
+        TINT_ASSERT(arg_major);
+        col_major = arg_major->Value() == core::Majorness::kColMajor;
+        stride_arg = call->Arguments()[2];
+        mat_ty = call->Target()->ReturnType()->As<core::type::SubgroupMatrix>();
+    } else {
+        // Don't validate deprecated variant.
+        // TODO(b/529415904): remove this after deprecated variant is removed.
+        if (!templated_ident) {
+            return true;
+        }
+        TINT_ASSERT(templated_ident->arguments.Length() == 1);
+        auto* sem_expr = sem_.Get(templated_ident->arguments[0]);
+        auto* arg_major = sem_expr->As<sem::BuiltinEnumExpression<core::Majorness>>();
+        TINT_ASSERT(arg_major);
+        col_major = arg_major->Value() == core::Majorness::kColMajor;
+        stride_arg = call->Arguments()[3];
+        mat_ty = call->Arguments()[2]->Type()->As<core::type::SubgroupMatrix>();
+    }
+
+    uint32_t major_size = col_major ? mat_ty->Columns() : mat_ty->Rows();
+    uint32_t minor_size = col_major ? mat_ty->Rows() : mat_ty->Columns();
+    auto* ele_ty = mat_ty->Type();
+
+    const uint32_t min_stride = ele_ty->Size() * minor_size;
+
+    uint64_t stride_value = 0;
+    if (stride_arg->ConstantValue()) {
+        if (stride_arg->Type()->IsUnsignedIntegerScalar()) {
+            stride_value = stride_arg->ConstantValue()->ValueAs<uint64_t>();
+        } else {
+            TINT_ASSERT(offset_arg->Type()->IsSignedIntegerScalar());
+            int32_t ivalue = stride_arg->ConstantValue()->ValueAs<int32_t>();
+            if (ivalue < 0) {
+                AddError(offset_arg->Declaration()->source)
+                    << "the stride argument of " << builtin->str() << " must be non-negative";
+                return false;
+            }
+            stride_value = static_cast<uint64_t>(ivalue);
+        }
+        stride_value *= ptr_arr_ty->ElemType()->Size();
+        if (stride_value < min_stride) {
+            AddError(stride_arg->Declaration()->source)
+                << "the stride argument (" << stride_value / ptr_arr_ty->ElemType()->Size() << ", "
+                << stride_value << " bytes) of " << builtin->str()
+                << " must be greater than the minimum stride (" << min_stride << " bytes)";
+            return false;
+        }
+    } else {
+        // Use the minimum stride value so we can validate the required matrix size below. Minimum
+        // stride (and 0 offset) allow us to avoid predication.
+        stride_value = min_stride;
+    }
+
+    if (!ptr_arr_ty->ConstantCount()) {
+        return true;
+    }
+
+    uint64_t offset_value = 0;
+    if (offset_arg->ConstantValue()) {
+        if (offset_arg->Type()->IsUnsignedIntegerScalar()) {
+            offset_value = offset_arg->ConstantValue()->ValueAs<uint64_t>();
+        } else {
+            TINT_ASSERT(offset_arg->Type()->IsSignedIntegerScalar());
+            int32_t ivalue = offset_arg->ConstantValue()->ValueAs<int32_t>();
+            if (ivalue < 0) {
+                AddError(offset_arg->Declaration()->source)
+                    << "the offset argument of " << builtin->str() << " must be non-negative";
+                return false;
+            }
+            offset_value = static_cast<uint64_t>(ivalue);
+        }
+        offset_value *= ptr_arr_ty->ElemType()->Size();
+    }
+
+    uint64_t mat_required_size =
+        stride_value * (major_size - 1) + minor_size * ele_ty->Size() + offset_value;
+    uint64_t arr_size = ptr_arr_ty->Size();
+    if (arr_size < mat_required_size) {
+        AddError(call->Declaration()->source)
+            << "the pointer operand of " << builtin->str() << " is too small (" << arr_size
+            << " bytes) for the matrix access (" << mat_required_size << " bytes)";
+        return false;
+    }
+
+    return true;
+}
+
 bool Validator::TextureBuiltinFn(const sem::Call* call) const {
     auto* builtin = call->Target()->As<sem::BuiltinFn>();
     if (!builtin) {
diff --git a/src/tint/lang/wgsl/resolver/validator.h b/src/tint/lang/wgsl/resolver/validator.h
index e7b8a05..a7be70c 100644
--- a/src/tint/lang/wgsl/resolver/validator.h
+++ b/src/tint/lang/wgsl/resolver/validator.h
@@ -551,6 +551,11 @@
     /// @returns true on success, false otherwise
     bool BufferView(const sem::Call* call) const;
 
+    /// Validate subgroupMatrixLoad and subgroupMatrixStore builtin functions
+    /// @param call the builtin call to validate
+    /// @returns true on success, false otherwise
+    bool SubgroupMatrixLoadStore(const sem::Call* call) const;
+
     /// Validates an optional builtin function and its required extensions and language features.
     /// @param call the builtin call to validate
     /// @returns true on success, false otherwise
@@ -658,6 +663,7 @@
     /// @param p_arg the pointer argument
     /// @param offset_arg the offset argument
     /// @returns true on success, false if an error was raised.
+    /// TODO(b/529415904): remove this when deprecated load/store variants are removed.
     bool CheckSubgroupMatrixOpOffset(const sem::BuiltinFn* fn,
                                      const sem::ValueExpression* p_arg,
                                      const sem::ValueExpression* offset_arg) const;