|
1 | | -// Copyright (C) 2024-2025 Intel Corporation |
| 1 | +// Copyright (C) 2024-2026 Intel Corporation |
2 | 2 | // Under the Apache License v2.0 with LLVM Exceptions. See LICENSE.TXT. |
3 | 3 | // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception |
4 | 4 |
|
|
9 | 9 | #include <umf/experimental/memspace.h> |
10 | 10 | #include <umf/experimental/memtarget.h> |
11 | 11 |
|
| 12 | +#include "memtarget_internal.h" |
| 13 | +#include "memtarget_ops.h" |
| 14 | + |
12 | 15 | using umf_test::test; |
13 | 16 |
|
| 17 | +namespace { |
| 18 | + |
| 19 | +struct fake_target_t { |
| 20 | + unsigned id; |
| 21 | +}; |
| 22 | + |
| 23 | +int finalizeCount; |
| 24 | + |
| 25 | +umf_result_t fakeInitialize(void *params, void **memoryTarget) { |
| 26 | + if (!params || !memoryTarget) { |
| 27 | + return UMF_RESULT_ERROR_INVALID_ARGUMENT; |
| 28 | + } |
| 29 | + |
| 30 | + auto target = static_cast<fake_target_t *>(malloc(sizeof(fake_target_t))); |
| 31 | + if (!target) { |
| 32 | + return UMF_RESULT_ERROR_OUT_OF_HOST_MEMORY; |
| 33 | + } |
| 34 | + |
| 35 | + target->id = *static_cast<unsigned *>(params); |
| 36 | + *memoryTarget = target; |
| 37 | + return UMF_RESULT_SUCCESS; |
| 38 | +} |
| 39 | + |
| 40 | +void fakeFinalize(void *memoryTarget) { |
| 41 | + finalizeCount++; |
| 42 | + free(memoryTarget); |
| 43 | +} |
| 44 | + |
| 45 | +umf_result_t fakeClone(void *memoryTarget, void **outMemoryTarget) { |
| 46 | + if (!memoryTarget || !outMemoryTarget) { |
| 47 | + return UMF_RESULT_ERROR_INVALID_ARGUMENT; |
| 48 | + } |
| 49 | + |
| 50 | + auto source = static_cast<fake_target_t *>(memoryTarget); |
| 51 | + auto target = static_cast<fake_target_t *>(malloc(sizeof(fake_target_t))); |
| 52 | + if (!target) { |
| 53 | + return UMF_RESULT_ERROR_OUT_OF_HOST_MEMORY; |
| 54 | + } |
| 55 | + |
| 56 | + target->id = source->id; |
| 57 | + *outMemoryTarget = target; |
| 58 | + return UMF_RESULT_SUCCESS; |
| 59 | +} |
| 60 | + |
| 61 | +umf_result_t fakePoolCreateFromMemspace(umf_const_memspace_handle_t memspace, |
| 62 | + void **memoryTargets, size_t numTargets, |
| 63 | + umf_const_mempolicy_handle_t policy, |
| 64 | + umf_memory_pool_handle_t *pool) { |
| 65 | + (void)memspace; |
| 66 | + (void)memoryTargets; |
| 67 | + (void)numTargets; |
| 68 | + (void)policy; |
| 69 | + (void)pool; |
| 70 | + return UMF_RESULT_ERROR_NOT_SUPPORTED; |
| 71 | +} |
| 72 | + |
| 73 | +umf_result_t |
| 74 | +fakeMemoryProviderCreateFromMemspace(umf_const_memspace_handle_t memspace, |
| 75 | + void **memoryTargets, size_t numTargets, |
| 76 | + umf_const_mempolicy_handle_t policy, |
| 77 | + umf_memory_provider_handle_t *provider) { |
| 78 | + (void)memspace; |
| 79 | + (void)memoryTargets; |
| 80 | + (void)numTargets; |
| 81 | + (void)policy; |
| 82 | + (void)provider; |
| 83 | + return UMF_RESULT_ERROR_NOT_SUPPORTED; |
| 84 | +} |
| 85 | + |
| 86 | +umf_result_t fakeGetCapacity(void *memoryTarget, size_t *capacity) { |
| 87 | + (void)memoryTarget; |
| 88 | + *capacity = 4096; |
| 89 | + return UMF_RESULT_SUCCESS; |
| 90 | +} |
| 91 | + |
| 92 | +umf_result_t fakeGetBandwidth(void *srcMemoryTarget, void *dstMemoryTarget, |
| 93 | + size_t *bandwidth) { |
| 94 | + (void)srcMemoryTarget; |
| 95 | + (void)dstMemoryTarget; |
| 96 | + *bandwidth = 1; |
| 97 | + return UMF_RESULT_SUCCESS; |
| 98 | +} |
| 99 | + |
| 100 | +umf_result_t fakeGetLatency(void *srcMemoryTarget, void *dstMemoryTarget, |
| 101 | + size_t *latency) { |
| 102 | + (void)srcMemoryTarget; |
| 103 | + (void)dstMemoryTarget; |
| 104 | + *latency = 1; |
| 105 | + return UMF_RESULT_SUCCESS; |
| 106 | +} |
| 107 | + |
| 108 | +umf_result_t fakeGetType(void *memoryTarget, umf_memtarget_type_t *type) { |
| 109 | + (void)memoryTarget; |
| 110 | + *type = UMF_MEMTARGET_TYPE_NUMA; |
| 111 | + return UMF_RESULT_SUCCESS; |
| 112 | +} |
| 113 | + |
| 114 | +umf_result_t fakeGetId(void *memoryTarget, unsigned *id) { |
| 115 | + *id = static_cast<fake_target_t *>(memoryTarget)->id; |
| 116 | + return UMF_RESULT_SUCCESS; |
| 117 | +} |
| 118 | + |
| 119 | +umf_result_t fakeCompare(void *memTarget, void *otherMemTarget, int *result) { |
| 120 | + auto left = static_cast<fake_target_t *>(memTarget); |
| 121 | + auto right = static_cast<fake_target_t *>(otherMemTarget); |
| 122 | + |
| 123 | + if (finalizeCount > 0 && left->id == 2 && right->id == 2) { |
| 124 | + return UMF_RESULT_ERROR_UNKNOWN; |
| 125 | + } |
| 126 | + |
| 127 | + *result = left->id == right->id ? 0 : 1; |
| 128 | + return UMF_RESULT_SUCCESS; |
| 129 | +} |
| 130 | + |
| 131 | +const umf_memtarget_ops_t FAKE_OPS = { |
| 132 | + .version = UMF_MEMTARGET_OPS_VERSION_CURRENT, |
| 133 | + .initialize = fakeInitialize, |
| 134 | + .finalize = fakeFinalize, |
| 135 | + .clone = fakeClone, |
| 136 | + .pool_create_from_memspace = fakePoolCreateFromMemspace, |
| 137 | + .memory_provider_create_from_memspace = |
| 138 | + fakeMemoryProviderCreateFromMemspace, |
| 139 | + .get_capacity = fakeGetCapacity, |
| 140 | + .get_bandwidth = fakeGetBandwidth, |
| 141 | + .get_latency = fakeGetLatency, |
| 142 | + .get_type = fakeGetType, |
| 143 | + .get_id = fakeGetId, |
| 144 | + .compare = fakeCompare, |
| 145 | +}; |
| 146 | + |
| 147 | +int removeEverything(umf_const_memspace_handle_t memspace, |
| 148 | + umf_const_memtarget_handle_t target, void *args) { |
| 149 | + (void)memspace; |
| 150 | + (void)target; |
| 151 | + (void)args; |
| 152 | + return 0; |
| 153 | +} |
| 154 | + |
| 155 | +} // namespace |
| 156 | + |
14 | 157 | TEST_F(test, memTargetNuma) { |
15 | 158 | auto memspace = umfMemspaceHostAllGet(); |
16 | 159 | ASSERT_NE(memspace, nullptr); |
@@ -172,3 +315,58 @@ TEST_F(test, memTargetRemoveAll) { |
172 | 315 | ret = umfMemspaceDestroy(memspace); |
173 | 316 | EXPECT_EQ(ret, UMF_RESULT_SUCCESS); |
174 | 317 | } |
| 318 | + |
| 319 | +TEST_F(test, memTargetFilterRollback) { |
| 320 | + finalizeCount = 0; |
| 321 | + |
| 322 | + umf_memspace_handle_t memspace = nullptr; |
| 323 | + umf_result_t ret = umfMemspaceNew(&memspace); |
| 324 | + ASSERT_EQ(ret, UMF_RESULT_SUCCESS); |
| 325 | + ASSERT_NE(memspace, nullptr); |
| 326 | + |
| 327 | + umf_memtarget_handle_t target1 = nullptr; |
| 328 | + umf_memtarget_handle_t target2 = nullptr; |
| 329 | + umf_memtarget_handle_t target3 = nullptr; |
| 330 | + unsigned id1 = 1; |
| 331 | + unsigned id2 = 2; |
| 332 | + unsigned id3 = 3; |
| 333 | + |
| 334 | + ret = umfMemtargetCreate(&FAKE_OPS, &id1, &target1); |
| 335 | + ASSERT_EQ(ret, UMF_RESULT_SUCCESS); |
| 336 | + ret = umfMemtargetCreate(&FAKE_OPS, &id2, &target2); |
| 337 | + ASSERT_EQ(ret, UMF_RESULT_SUCCESS); |
| 338 | + ret = umfMemtargetCreate(&FAKE_OPS, &id3, &target3); |
| 339 | + ASSERT_EQ(ret, UMF_RESULT_SUCCESS); |
| 340 | + |
| 341 | + ret = umfMemspaceMemtargetAdd(memspace, target1); |
| 342 | + ASSERT_EQ(ret, UMF_RESULT_SUCCESS); |
| 343 | + ret = umfMemspaceMemtargetAdd(memspace, target2); |
| 344 | + ASSERT_EQ(ret, UMF_RESULT_SUCCESS); |
| 345 | + ret = umfMemspaceMemtargetAdd(memspace, target3); |
| 346 | + ASSERT_EQ(ret, UMF_RESULT_SUCCESS); |
| 347 | + |
| 348 | + ret = umfMemspaceUserFilter(memspace, removeEverything, nullptr); |
| 349 | + EXPECT_EQ(ret, UMF_RESULT_ERROR_UNKNOWN); |
| 350 | + ASSERT_EQ(umfMemspaceMemtargetNum(memspace), 3u); |
| 351 | + |
| 352 | + std::vector<unsigned> ids; |
| 353 | + for (size_t targetIdx = 0; targetIdx < umfMemspaceMemtargetNum(memspace); |
| 354 | + targetIdx++) { |
| 355 | + auto target = umfMemspaceMemtargetGet(memspace, targetIdx); |
| 356 | + ASSERT_NE(target, nullptr); |
| 357 | + |
| 358 | + unsigned id = 0; |
| 359 | + ret = umfMemtargetGetId(target, &id); |
| 360 | + ASSERT_EQ(ret, UMF_RESULT_SUCCESS); |
| 361 | + ids.push_back(id); |
| 362 | + } |
| 363 | + |
| 364 | + std::sort(ids.begin(), ids.end()); |
| 365 | + EXPECT_EQ(ids, (std::vector<unsigned>{1, 2, 3})); |
| 366 | + |
| 367 | + ret = umfMemspaceDestroy(memspace); |
| 368 | + EXPECT_EQ(ret, UMF_RESULT_SUCCESS); |
| 369 | + umfMemtargetDestroy(target1); |
| 370 | + umfMemtargetDestroy(target2); |
| 371 | + umfMemtargetDestroy(target3); |
| 372 | +} |
0 commit comments