[llvm] [Offload] Avoid writing past the destination in host dataFill (PR #223331)

via llvm-commits llvm-commits at lists.llvm.org
Wed Sep 23 23:02:14 PDT 2026


https://github.com/StevenYangCC updated https://github.com/llvm/llvm-project/pull/223331

>From e4206f9cc6df7016ec0c0044b6d216ef3a7711f3 Mon Sep 17 00:00:00 2001
From: "chengcang.yang" <yangchengcang at gmail.com>
Date: Mon, 14 Sep 2026 16:51:01 +0800
Subject: [PATCH] [Offload] Avoid writing past the destination in host dataFill

The host plugin copied PatternSize bytes on every loop iteration, so a
fill size that was not a multiple of the pattern size overran the
destination. Cap each copy with std::min to the remaining bytes.

Add unit tests for the remainder path and for not writing past the
fill size.
---
 offload/plugins-nextgen/host/src/rtl.cpp      |  4 +-
 offload/unittests/OffloadAPI/CMakeLists.txt   |  4 +
 .../unittests/OffloadAPI/memory/olMemFill.cpp | 40 ++++++++
 .../OffloadAPI/memory/olMemFillHost.cpp       | 94 +++++++++++++++++++
 4 files changed, 141 insertions(+), 1 deletion(-)
 create mode 100644 offload/unittests/OffloadAPI/memory/olMemFillHost.cpp

diff --git a/offload/plugins-nextgen/host/src/rtl.cpp b/offload/plugins-nextgen/host/src/rtl.cpp
index 55ada2f82c360..1155ec5c9bd80 100644
--- a/offload/plugins-nextgen/host/src/rtl.cpp
+++ b/offload/plugins-nextgen/host/src/rtl.cpp
@@ -10,6 +10,7 @@
 //
 //===----------------------------------------------------------------------===//
 
+#include <algorithm>
 #include <cassert>
 #include <cstddef>
 #include <string>
@@ -308,7 +309,8 @@ struct GenELF64DeviceTy : public GenericDeviceTy {
     } else {
       for (unsigned int Step = 0; Step < Size; Step += PatternSize) {
         auto *Dst = static_cast<char *>(TgtPtr) + Step;
-        std::memcpy(Dst, PatternPtr, PatternSize);
+        std::memcpy(Dst, PatternPtr,
+                    std::min<int64_t>(PatternSize, Size - Step));
       }
     }
 
diff --git a/offload/unittests/OffloadAPI/CMakeLists.txt b/offload/unittests/OffloadAPI/CMakeLists.txt
index 292ea1eb4852f..f51097c74d317 100644
--- a/offload/unittests/OffloadAPI/CMakeLists.txt
+++ b/offload/unittests/OffloadAPI/CMakeLists.txt
@@ -42,6 +42,10 @@ add_offload_unittest("memory"
     memory/olGetMemInfoSize.cpp
     memory/olMemRegister.cpp)
 
+add_offload_unittest("memory_host"
+    memory/olMemFillHost.cpp)
+target_compile_definitions("memory_host.unittests" PRIVATE DISABLE_WRAPPER)
+
 add_offload_unittest("platform"
     platform/olGetPlatformInfo.cpp
     platform/olGetPlatformInfoSize.cpp)
diff --git a/offload/unittests/OffloadAPI/memory/olMemFill.cpp b/offload/unittests/OffloadAPI/memory/olMemFill.cpp
index 467a551c48c94..be41c1075fca1 100644
--- a/offload/unittests/OffloadAPI/memory/olMemFill.cpp
+++ b/offload/unittests/OffloadAPI/memory/olMemFill.cpp
@@ -48,6 +48,46 @@ struct olMemFillTest : OffloadQueueTest {
 };
 OFFLOAD_TESTS_INSTANTIATE_DEVICE_FIXTURE(olMemFillTest);
 
+using olMemFillHostDeviceTest = OffloadQueueTest;
+OFFLOAD_TESTS_INSTANTIATE_HOST_DEVICE_FIXTURE(olMemFillHostDeviceTest);
+
+TEST_P(olMemFillHostDeviceTest, DoesNotWritePastFillSize) {
+  constexpr size_t AllocSize = 16;
+  constexpr size_t FillSize = 8;
+  constexpr unsigned char Canary = 0xAA;
+  const unsigned char Pattern[4] = {1, 2, 3, 4};
+
+  void *Alloc;
+  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED, AllocSize, &Alloc));
+  ASSERT_SUCCESS(olMemFill(Queue, Alloc, 1, &Canary, AllocSize));
+  ASSERT_SUCCESS(olMemFill(Queue, Alloc, sizeof(Pattern), Pattern, FillSize));
+  olSyncQueue(Queue);
+
+  auto *Bytes = static_cast<unsigned char *>(Alloc);
+  for (size_t I = 0; I < FillSize; ++I)
+    ASSERT_EQ(Bytes[I], Pattern[I % sizeof(Pattern)]);
+  for (size_t I = FillSize; I < AllocSize; ++I)
+    ASSERT_EQ(Bytes[I], Canary) << "wrote past FillSize at byte " << I;
+
+  olMemFree(Alloc);
+}
+
+TEST_P(olMemFillHostDeviceTest, SuccessSingleByteNotMultipleOfFour) {
+  constexpr size_t FillSize = 7;
+  constexpr unsigned char Pattern = 0x5A;
+
+  void *Alloc;
+  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED, FillSize, &Alloc));
+  ASSERT_SUCCESS(olMemFill(Queue, Alloc, 1, &Pattern, FillSize));
+  olSyncQueue(Queue);
+
+  auto *Bytes = static_cast<unsigned char *>(Alloc);
+  for (size_t I = 0; I < FillSize; ++I)
+    ASSERT_EQ(Bytes[I], Pattern);
+
+  olMemFree(Alloc);
+}
+
 TEST_P(olMemFillTest, Success8) { test_body<uint8_t, 0x42, 1024>(); }
 TEST_P(olMemFillTest, Success8NotMultiple4) {
   test_body<uint8_t, 0x42, 1023>();
diff --git a/offload/unittests/OffloadAPI/memory/olMemFillHost.cpp b/offload/unittests/OffloadAPI/memory/olMemFillHost.cpp
new file mode 100644
index 0000000000000..e26b8672e35ce
--- /dev/null
+++ b/offload/unittests/OffloadAPI/memory/olMemFillHost.cpp
@@ -0,0 +1,94 @@
+//===------- Offload API tests - host olMemFill remainder -----------------===//
+//
+// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+//
+//===----------------------------------------------------------------------===//
+
+// olMemFill rejects a fill size that is not a multiple of the pattern size.
+// This suite skips implicit olInit so it can disable that check and reach the
+// host plugin's dataFill with a remainder, which used to write past the
+// destination.
+
+#include "../common/Fixtures.hpp"
+#include <OffloadAPI.h>
+#include <cstdlib>
+#include <gtest/gtest.h>
+
+struct olMemFillHostRemainderTest : ::testing::Test {
+  void SetUp() override {
+    ASSERT_EQ(setenv("OFFLOAD_DISABLE_VALIDATION", "1", 1), 0);
+
+    ol_init_args_t Args = OL_INIT_ARGS_INIT;
+    ol_platform_backend_t Backends[] = {OL_PLATFORM_BACKEND_HOST};
+    Args.NumPlatforms = 1;
+    Args.Platforms = Backends;
+    ASSERT_SUCCESS(olInit(&Args));
+    Initialized = true;
+
+    Device = TestEnvironment::getHostDevice();
+    if (!Device)
+      GTEST_SKIP() << "No host device.";
+
+    ASSERT_SUCCESS(olCreateContext(1, &Device, &Context));
+    ASSERT_SUCCESS(olCreateQueue(Context, Device, &Queue));
+  }
+
+  void TearDown() override {
+    if (Queue)
+      olDestroyQueue(Queue);
+    if (Context)
+      olDestroyContext(Context);
+    if (Initialized)
+      ASSERT_SUCCESS(olShutDown());
+    unsetenv("OFFLOAD_DISABLE_VALIDATION");
+  }
+
+  bool Initialized = false;
+  ol_device_handle_t Device = nullptr;
+  ol_context_handle_t Context = nullptr;
+  ol_queue_handle_t Queue = nullptr;
+};
+
+TEST_F(olMemFillHostRemainderTest, RemainderDoesNotWritePastFillSize) {
+  constexpr size_t AllocSize = 16;
+  constexpr size_t FillSize = 5;
+  constexpr unsigned char Canary = 0xAA;
+  const unsigned char Pattern[4] = {1, 2, 3, 4};
+
+  void *Alloc = nullptr;
+  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED, AllocSize, &Alloc));
+  ASSERT_SUCCESS(olMemFill(Queue, Alloc, 1, &Canary, AllocSize));
+  ASSERT_SUCCESS(olMemFill(Queue, Alloc, sizeof(Pattern), Pattern, FillSize));
+  ASSERT_SUCCESS(olSyncQueue(Queue));
+
+  auto *Bytes = static_cast<unsigned char *>(Alloc);
+  for (size_t I = 0; I < FillSize; ++I)
+    ASSERT_EQ(Bytes[I], Pattern[I % sizeof(Pattern)]);
+  for (size_t I = FillSize; I < AllocSize; ++I)
+    ASSERT_EQ(Bytes[I], Canary) << "wrote past FillSize at byte " << I;
+
+  ASSERT_SUCCESS(olMemFree(Alloc));
+}
+
+TEST_F(olMemFillHostRemainderTest, PatternLargerThanFillSize) {
+  constexpr size_t AllocSize = 16;
+  constexpr size_t FillSize = 3;
+  constexpr unsigned char Canary = 0xAA;
+  const unsigned char Pattern[4] = {1, 2, 3, 4};
+
+  void *Alloc = nullptr;
+  ASSERT_SUCCESS(olMemAlloc(Device, OL_ALLOC_TYPE_MANAGED, AllocSize, &Alloc));
+  ASSERT_SUCCESS(olMemFill(Queue, Alloc, 1, &Canary, AllocSize));
+  ASSERT_SUCCESS(olMemFill(Queue, Alloc, sizeof(Pattern), Pattern, FillSize));
+  ASSERT_SUCCESS(olSyncQueue(Queue));
+
+  auto *Bytes = static_cast<unsigned char *>(Alloc);
+  for (size_t I = 0; I < FillSize; ++I)
+    ASSERT_EQ(Bytes[I], Pattern[I]);
+  for (size_t I = FillSize; I < AllocSize; ++I)
+    ASSERT_EQ(Bytes[I], Canary) << "wrote past FillSize at byte " << I;
+
+  ASSERT_SUCCESS(olMemFree(Alloc));
+}



More information about the llvm-commits mailing list