Move the repeated command-stream buffer range validation into a
helper in preparation for validating all weight and scale streams.

Assisted-by: LLM
Signed-off-by: Rob Herring (Arm) <[email protected]>
---
v2:
 - Adjust due to previous patch
---
 drivers/accel/ethosu/ethosu_gem.c | 37 +++++++++++++++++++++++++------------
 1 file changed, 25 insertions(+), 12 deletions(-)

diff --git a/drivers/accel/ethosu/ethosu_gem.c 
b/drivers/accel/ethosu/ethosu_gem.c
index 1e9ffacbc394..2707b7df5dbe 100644
--- a/drivers/accel/ethosu/ethosu_gem.c
+++ b/drivers/accel/ethosu/ethosu_gem.c
@@ -433,6 +433,25 @@ static u64 feat_matrix_length(struct ethosu_device *edev,
        return addr;
 }
 
+static int buffer_size(struct ethosu_validated_cmdstream_info *info,
+                      struct cmd_state *st, struct buffer *buf, s8 region,
+                      u16 region_cmd, u16 base_cmd, u16 length_cmd)
+{
+       u64 end;
+
+       if (region < 0 || !cmd_state_reg_is_set(st, region_cmd) ||
+           !cmd_state_reg_is_set(st, base_cmd) ||
+           !cmd_state_reg_is_set(st, length_cmd))
+               return -EINVAL;
+
+       if (check_add_overflow(buf->base, (u64)buf->length, &end))
+               return -EINVAL;
+
+       info->region_size[region] = max(info->region_size[region], end);
+
+       return 0;
+}
+
 static int calc_sizes(struct drm_device *ddev,
                      struct ethosu_validated_cmdstream_info *info,
                      u16 op, struct cmd_state *st,
@@ -485,26 +504,20 @@ static int calc_sizes(struct drm_device *ddev,
                dev_dbg(ddev->dev, "op %d: W:%d:0x%llx-0x%llx\n",
                        op, st->weight[0].region, st->weight[0].base,
                        st->weight[0].base + st->weight[0].length - 1);
-               if (!cmd_state_reg_is_set(st, NPU_SET_WEIGHT_REGION) ||
-                   !cmd_state_reg_is_set(st, NPU_SET_WEIGHT_BASE) ||
-                   !cmd_state_reg_is_set(st, NPU_SET_WEIGHT_LENGTH))
+               if (buffer_size(info, st, &st->weight[0], st->weight[0].region,
+                               NPU_SET_WEIGHT_REGION, NPU_SET_WEIGHT_BASE,
+                               NPU_SET_WEIGHT_LENGTH))
                        return -EINVAL;
-               info->region_size[st->weight[0].region] =
-                       max(info->region_size[st->weight[0].region],
-                           st->weight[0].base + st->weight[0].length);
        }
 
        if (scale) {
                dev_dbg(ddev->dev, "op %d: S:%d:0x%llx-0x%llx\n",
                        op, st->scale[0].region, st->scale[0].base,
                        st->scale[0].base + st->scale[0].length - 1);
-               if (!cmd_state_reg_is_set(st, NPU_SET_SCALE_REGION) ||
-                   !cmd_state_reg_is_set(st, NPU_SET_SCALE_BASE) ||
-                   !cmd_state_reg_is_set(st, NPU_SET_SCALE_LENGTH))
+               if (buffer_size(info, st, &st->scale[0], st->scale[0].region,
+                               NPU_SET_SCALE_REGION, NPU_SET_SCALE_BASE,
+                               NPU_SET_SCALE_LENGTH))
                        return -EINVAL;
-               info->region_size[st->scale[0].region] =
-                       max(info->region_size[st->scale[0].region],
-                           st->scale[0].base + st->scale[0].length);
        }
 
        len = feat_matrix_length(edev, info, st, &st->ofm, FEAT_MATRIX_OFM,

-- 
2.53.0

Reply via email to