已合并
[feat] support environment variable LD_PRELOAD #34371
liujunzhu创建于 4月25日
[feat] support environment variable LD_PRELOAD #34371
已合并
共 2 个文件变更+226-9
| @@ -0,0 +1,193 @@ | |||
| 1 | +# Integration test: verify that LD_PRELOAD-injected overrides of ACL symbols | ||
| 2 | +# loaded via FunctionLoader (e.g. aclrtMallocAlign32) are picked up. | ||
| 3 | +# | ||
| 4 | +# Requires: an NPU environment with libascendcl available, a C compiler. | ||
| 5 | +# Run: python test/npu/test_ld_preload_acl_hook.py | ||
| 6 | + | ||
| 7 | +import os | ||
| 8 | +import shutil | ||
| 9 | +import subprocess | ||
| 10 | +import sys | ||
| 11 | +import tempfile | ||
| 12 | +import textwrap | ||
| 13 | +import unittest | ||
| 14 | + | ||
| 15 | +import torch_npu | ||
| 16 | +from torch_npu.testing.testcase import TestCase, run_tests | ||
| 17 | + | ||
| 18 | + | ||
| 19 | +HOOK_SRC = r""" | ||
| 20 | +#define _GNU_SOURCE | ||
| 21 | +#include <dlfcn.h> | ||
| 22 | +#include <stdio.h> | ||
| 23 | +#include <stdlib.h> | ||
| 24 | + | ||
| 25 | +typedef int aclError; | ||
| 26 | +typedef int aclrtMemMallocPolicy; | ||
| 27 | + | ||
| 28 | +__attribute__((visibility("default"))) | ||
| 29 | +aclError aclrtMallocAlign32(void **devPtr, size_t size, aclrtMemMallocPolicy policy) | ||
| 30 | +{ | ||
| 31 | + const char *flag = getenv("ACL_HOOK_FLAG_FILE"); | ||
| 32 | + if (flag != NULL) { | ||
| 33 | + FILE *fp = fopen(flag, "w"); | ||
| 34 | + if (fp != NULL) { | ||
| 35 | + fprintf(fp, "hit\n"); | ||
| 36 | + fclose(fp); | ||
| 37 | + } | ||
| 38 | + } | ||
| 39 | + typedef aclError (*real_fn_t)(void **, size_t, aclrtMemMallocPolicy); | ||
| 40 | + static real_fn_t real = NULL; | ||
| 41 | + if (real == NULL) { | ||
| 42 | + real = (real_fn_t)dlsym(RTLD_NEXT, "aclrtMallocAlign32"); | ||
| 43 | + } | ||
| 44 | + if (real == NULL) { | ||
| 45 | + return -1; | ||
| 46 | + } | ||
| 47 | + return real(devPtr, size, policy); | ||
| 48 | +} | ||
| 49 | +""" | ||
| 50 | + | ||
| 51 | + | ||
| 52 | +# Minimal Python payload: allocate a tensor on NPU, which eventually exercises | ||
| 53 | +# aclrtMallocAlign32 via the caching allocator's large-block path. | ||
| 54 | +PAYLOAD_SRC = textwrap.dedent( | ||
| 55 | + """ | ||
| 56 | + import sys | ||
| 57 | + import torch | ||
| 58 | + import torch_npu | ||
| 59 | + torch.npu.set_device(0) | ||
| 60 | + t = torch.empty(1024 * 1024, dtype=torch.float32, device='npu') | ||
| 61 | + del t | ||
| 62 | + torch.npu.empty_cache() | ||
| 63 | + sys.exit(0) | ||
| 64 | + """ | ||
| 65 | +) | ||
| 66 | + | ||
| 67 | + | ||
| 68 | +def _compile_hook(tmpdir, cc): | ||
| 69 | + src_path = os.path.join(tmpdir, "hook.c") | ||
| 70 | + so_path = os.path.join(tmpdir, "libacl_hook_test.so") | ||
| 71 | + with open(src_path, "w") as f: | ||
| 72 | + f.write(HOOK_SRC) | ||
| 73 | + subprocess.check_call( | ||
| 74 | + [cc, "-shared", "-fPIC", "-O2", "-o", so_path, src_path, "-ldl"] | ||
| 75 | + ) | ||
| 76 | + return so_path | ||
| 77 | + | ||
| 78 | + | ||
| 79 | +def _compile_unrelated(tmpdir, cc): | ||
| 80 | + src_path = os.path.join(tmpdir, "other.c") | ||
| 81 | + so_path = os.path.join(tmpdir, "libunrelated_hook.so") | ||
| 82 | + with open(src_path, "w") as f: | ||
| 83 | + f.write("int unrelated_noop(void) { return 0; }\n") | ||
| 84 | + subprocess.check_call( | ||
| 85 | + [cc, "-shared", "-fPIC", "-O2", "-o", so_path, src_path] | ||
| 86 | + ) | ||
| 87 | + return so_path | ||
| 88 | + | ||
| 89 | + | ||
| 90 | +def _run_payload(env): | ||
| 91 | + with tempfile.NamedTemporaryFile("w", suffix=".py", delete=False) as f: | ||
| 92 | + f.write(PAYLOAD_SRC) | ||
| 93 | + script = f.name | ||
| 94 | + try: | ||
| 95 | + return subprocess.run( | ||
| 96 | + [sys.executable, script], | ||
| 97 | + env=env, | ||
| 98 | + capture_output=True, | ||
| 99 | + text=True, | ||
| 100 | + timeout=120, | ||
| 101 | + ) | ||
| 102 | + finally: | ||
| 103 | + os.unlink(script) | ||
| 104 | + | ||
| 105 | + | ||
| 106 | + | ||
| 107 | + not torch_npu.npu.is_available(), "NPU not available; skipping LD_PRELOAD test" | ||
| 108 | +) | ||
| 109 | +class TestLdPreloadAclHook(TestCase): | ||
| 110 | + | ||
| 111 | + def setUpClass(cls): | ||
| 112 | + cc = os.environ.get("CC", "gcc") | ||
| 113 | + if shutil.which(cc) is None: | ||
| 114 | + raise unittest.SkipTest( | ||
| 115 | + "C compiler '{}' not found; skipping LD_PRELOAD test".format(cc) | ||
| 116 | + ) | ||
| 117 | + cls._cc = cc | ||
| 118 | + cls._tmpdir = tempfile.mkdtemp(prefix="acl_hook_") | ||
| 119 | + try: | ||
| 120 | + cls._hook_so = _compile_hook(cls._tmpdir, cc) | ||
| 121 | + cls._unrelated_so = _compile_unrelated(cls._tmpdir, cc) | ||
| 122 | + except Exception: | ||
| 123 | + shutil.rmtree(cls._tmpdir, ignore_errors=True) | ||
| 124 | + raise | ||
| 125 | + | ||
| 126 | + | ||
| 127 | + def tearDownClass(cls): | ||
| 128 | + shutil.rmtree(cls._tmpdir, ignore_errors=True) | ||
| 129 | + | ||
| 130 | + def _base_env(self): | ||
| 131 | + env = os.environ.copy() | ||
| 132 | + env.pop("LD_PRELOAD", None) | ||
| 133 | + return env | ||
| 134 | + | ||
| 135 | + def test_without_preload_behaves_unchanged(self): | ||
| 136 | + """No LD_PRELOAD: allocation succeeds, hook never invoked.""" | ||
| 137 | + flag = os.path.join(self._tmpdir, "flag_no_preload") | ||
| 138 | + env = self._base_env() | ||
| 139 | + env["ACL_HOOK_FLAG_FILE"] = flag | ||
| 140 | + result = _run_payload(env) | ||
| 141 | + if result.returncode != 0: | ||
| 142 | + sys.stderr.write(result.stderr) | ||
| 143 | + self.assertEqual(result.returncode, 0) | ||
| 144 | + self.assertFalse( | ||
| 145 | + os.path.exists(flag), "hook must not be called when LD_PRELOAD is unset" | ||
| 146 | + ) | ||
| 147 | + | ||
| 148 | + def test_with_preload_hook_is_invoked(self): | ||
| 149 | + """LD_PRELOAD set to hook .so: hook is called for aclrtMallocAlign32.""" | ||
| 150 | + flag = os.path.join(self._tmpdir, "flag_preload") | ||
| 151 | + env = self._base_env() | ||
| 152 | + env["LD_PRELOAD"] = self._hook_so | ||
| 153 | + env["ACL_HOOK_FLAG_FILE"] = flag | ||
| 154 | + result = _run_payload(env) | ||
| 155 | + if result.returncode != 0: | ||
| 156 | + sys.stderr.write(result.stderr) | ||
| 157 | + self.assertEqual(result.returncode, 0) | ||
| 158 | + self.assertTrue( | ||
| 159 | + os.path.exists(flag), | ||
| 160 | + "hook must be called when LD_PRELOAD overrides aclrtMallocAlign32", | ||
| 161 | + ) | ||
| 162 | + | ||
| 163 | + def test_preload_without_symbol_falls_back(self): | ||
| 164 | + """LD_PRELOAD set to a .so that does NOT define ACL symbols: | ||
| 165 | + allocation must still succeed via fallback to libascendcl.""" | ||
| 166 | + flag = os.path.join(self._tmpdir, "flag_unrelated") | ||
| 167 | + env = self._base_env() | ||
| 168 | + env["LD_PRELOAD"] = self._unrelated_so | ||
| 169 | + env["ACL_HOOK_FLAG_FILE"] = flag | ||
| 170 | + result = _run_payload(env) | ||
| 171 | + if result.returncode != 0: | ||
| 172 | + sys.stderr.write(result.stderr) | ||
| 173 | + self.assertEqual(result.returncode, 0) | ||
| 174 | + self.assertFalse( | ||
| 175 | + os.path.exists(flag), | ||
| 176 | + "unrelated preload must not cause the hook flag to be set", | ||
| 177 | + ) | ||
| 178 | + | ||
| 179 | + def test_multiple_preload_sos(self): | ||
| 180 | + """LD_PRELOAD with multiple .so (hook first): hook still wins.""" | ||
| 181 | + flag = os.path.join(self._tmpdir, "flag_multi") | ||
| 182 | + env = self._base_env() | ||
| 183 | + env["LD_PRELOAD"] = "{}:{}".format(self._hook_so, self._unrelated_so) | ||
| 184 | + env["ACL_HOOK_FLAG_FILE"] = flag | ||
| 185 | + result = _run_payload(env) | ||
| 186 | + if result.returncode != 0: | ||
| 187 | + sys.stderr.write(result.stderr) | ||
| 188 | + self.assertEqual(result.returncode, 0) | ||
| 189 | + self.assertTrue(os.path.exists(flag), "first-loaded preload must win") | ||
| 190 | + | ||
| 191 | + | ||
| 192 | +if __name__ == "__main__": | ||
| 193 | + run_tests() | ||
| @@ -1,3 +1,4 @@ | |||
| 1 | + | ||
| 1 | 2 | ||
| 2 | 3 | ||
| 3 | 4 | ||
| @@ -24,14 +25,7 @@ void FunctionLoader::Set(const std::string &name) | |||
| 24 | 25 | ||
| 25 | void *FunctionLoader::Get(const std::string &name) | 26 | void *FunctionLoader::Get(const std::string &name) |
| 26 | { | 27 | { |
| 27 | - if (this->handle == nullptr) { | 28 | + std::lock_guard<std::mutex> lock(this->mu_); |
| 28 | - auto handle = dlopen(this->fileName.c_str(), this->flags); | ||
| 29 | - if (handle == nullptr) { | ||
| 30 | - AT_ERROR(dlerror()); | ||
| 31 | - return nullptr; | ||
| 32 | - } | ||
| 33 | - this->handle = handle; | ||
| 34 | - } | ||
| 35 | 29 | ||
| 36 | auto itr = registry.find(name); | 30 | auto itr = registry.find(name); |
| 37 | if (itr == registry.end()) { | 31 | if (itr == registry.end()) { |
| @@ -43,7 +37,37 @@ void *FunctionLoader::Get(const std::string &name) | |||
| 43 | return itr->second; | 37 | return itr->second; |
| 44 | } | 38 | } |
| 45 | 39 | ||
| 46 | - auto func = dlsym(this->handle, name.c_str()); | 40 | + // When LD_PRELOAD is set, prefer RTLD_DEFAULT so that symbols overridden |
| 41 | + // by the user's preloaded .so take precedence over libascendcl.so's own | ||
| 42 | + // implementation. Opt-in by the presence of LD_PRELOAD itself; without | ||
| 43 | + // LD_PRELOAD the behavior is identical to the original path. | ||
| 44 | + static const bool preload_enabled = []() { | ||
| 45 | + const char *env = std::getenv("LD_PRELOAD"); | ||
| 46 | + bool enabled = (env != nullptr && env[0] != '\0'); | ||
| 47 | + if (enabled) { | ||
| 48 | + TORCH_NPU_WARN_ONCE("LD_PRELOAD detected, FunctionLoader prefers " | ||
| 49 | + "RTLD_DEFAULT for symbol resolution."); | ||
| 50 | + } | ||
| 51 | + return enabled; | ||
| 52 | + }(); | ||
| 53 | + | ||
| 54 | + void *func = nullptr; | ||
| 55 | + if (preload_enabled) { | ||
| 56 | + func = dlsym(RTLD_DEFAULT, name.c_str()); | ||
| 57 | + } | ||
| 58 | + | ||
| 59 | + if (func == nullptr) { | ||
| 60 | + if (this->handle == nullptr) { | ||
| 61 | + auto handle = dlopen(this->fileName.c_str(), this->flags); | ||
| 62 | + if (handle == nullptr) { | ||
| 63 | + AT_ERROR(dlerror()); | ||
| 64 | + return nullptr; | ||
| 65 | + } | ||
| 66 | + this->handle = handle; | ||
| 67 | + } | ||
| 68 | + func = dlsym(this->handle, name.c_str()); | ||
| 69 | + } | ||
| 70 | + | ||
| 47 | if (func == nullptr) { | 71 | if (func == nullptr) { |
| 48 | return nullptr; | 72 | return nullptr; |
| 49 | } | 73 | } |