[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