已合并
Fix: 修复8个issue(#41/#42/#50/#52/#53/#56/#59/#60) 构建脚本清理/空指针防护/尺寸截断/UAF #38
Fix: 修复8个issue(#41/#42/#50/#52/#53/#56/#59/#60) 构建脚本清理/空指针防护/尺寸截断/UAF #38
已合并
syy_3597创建于 14 天前
7 个文件变更+42-30
Mbuild.sh+15-13
@@ -20,7 +20,8 @@
20# --ops=OP_LIST 指定要编译的算子列表 (逗号分隔)20# --ops=OP_LIST 指定要编译的算子列表 (逗号分隔)
21# --run 编译后执行测试21# --run 编译后执行测试
22# --pkg 编译并打包成 .run 文件22# --pkg 编译并打包成 .run 文件
23-# --soc=SOC 指定目标 SoC 型号 (支持: Ascend950、Ascend910B, 支持小写输入)23+# --soc=SOC 指定目标 SoC 型号 (支持: Ascend950、Ascend910B
24+# Ascend910_93、Ascend910、Ascend310P,支持小写输入)
24# -j[N] 编译线程数,默认为 8,例如: -j1625# -j[N] 编译线程数,默认为 8,例如: -j16
25# --test-timeout=N 测试超时时间(秒),默认为 30026# --test-timeout=N 测试超时时间(秒),默认为 300
26# --cann_3rd_lib_path=PATH27# --cann_3rd_lib_path=PATH
@@ -51,8 +52,12 @@
51# ./build.sh --test-timeout=600 --run # 运行测试 with 600s timeout52# ./build.sh --test-timeout=600 --run # 运行测试 with 600s timeout
52#53#
53# 支持的 SoC 型号:54# 支持的 SoC 型号:
54-# Ascend950 (dav-3510, 默认)55+# Ascend950 (dav-3510, 默认)
55-# Ascend910B (dav-2201)56+# Ascend910B (dav-2201)
57+# Ascend910_93(dav-2201)
58+# Ascend910 (dav-2101)
59+# Ascend310P (dav-2101)
60+# 注: Ascend310B 可识别但当前版本暂不支持构建
56##############################################################################61##############################################################################
57 62 
58set -e63set -e
@@ -91,13 +96,6 @@ fi
91# 统一导出 ASCEND_HOME_PATH,使后续 check_ascend_env 与 CMake 均可使用96# 统一导出 ASCEND_HOME_PATH,使后续 check_ascend_env 与 CMake 均可使用
92export ASCEND_HOME_PATH="${_ASCEND_INSTALL_PATH}"97export ASCEND_HOME_PATH="${_ASCEND_INSTALL_PATH}"
93 98 
94-# SoC 名称标准化函数(首字母大写,其余小写)
95-normalize_soc_name() {
96- local soc="$1"
97- # 转换为小写,然后首字母大写,最后字符大写
98- echo "${soc}" | sed 's/.*/\L&/; s/^./\U&/; s/.$/\U&/'
99-}
100- 
101# SoC 名称映射到完整的 SOC_VERSION(CANN ASC 编译器要求的格式)99# SoC 名称映射到完整的 SOC_VERSION(CANN ASC 编译器要求的格式)
102get_soc_version() {100get_soc_version() {
103 local soc_name="$1"101 local soc_name="$1"
@@ -253,10 +251,14 @@ Options:
253 --make_clean Clean build artifacts"251 --make_clean Clean build artifacts"
254 252 
255Supported SoC models:253Supported SoC models:
256- Ascend950 (dav-3510, default)254+ Ascend950 (dav-3510, default)
257- Ascend910B (dav-2201)255+ Ascend910B (dav-2201)
256+ Ascend910_93 (dav-2201)
257+ Ascend910 (dav-2101)
258+ Ascend310P (dav-2101)
258 259 
259-Note: Other SoC models are not officially supported in the current version.260+Note: Ascend310B is recognized but not supported in the current version.
261+ Other SoC models are not supported.
260 262 
261Examples:263Examples:
262 $(basename "$0") # Build all operators (default 8 threads)264 $(basename "$0") # Build all operators (default 8 threads)
@@ -28,15 +28,10 @@ aclfftResult aclfftDestroy(aclfftHandle plan) {
28 // 注意:允许销毁未初始化的 Plan28 // 注意:允许销毁未初始化的 Plan
29 ACLFFT_CHECK_NULL(impl);29 ACLFFT_CHECK_NULL(impl);
30 30 
31- // 防重复销毁31+ // 说明:原 is_destroyed 防重复销毁检查本身构成 use-after-free——对象在首次
32- if (impl->is_destroyed) {32+ // 销毁时即被 delete,标志随对象一同释放,二次调用读取的是已释放内存。
33- return ACLFFT_INVALID_PLAN;33+ // 重复销毁应由调用方保证(内部调用方在销毁后均置 *plan=nullptr),
34- }34+ // 对已销毁句柄的再次传入属未定义行为,此处不再解引用已释放对象(issue #60)。
35- 
36- if (impl->has_operator_state && impl->operator_state != nullptr) {}
37- 
38- // 标记为已销毁
39- impl->is_destroyed = true;
40 35 
41 // 释放 Plan 对象36 // 释放 Plan 对象
42 delete impl;37 delete impl;
@@ -55,6 +55,8 @@ aclfftResult aclfftExecC2C_1D(aclfftHandle plan,
55 aclfftComplex* odata,55 aclfftComplex* odata,
56 int direction) {56 int direction) {
57 aclfftHandle_t* impl = plan;57 aclfftHandle_t* impl = plan;
58+ // 防护: 本函数为 weak 符号可被外部直接调用,plan/idata/odata 可能为 NULL(issue #50)
59+ ACLFFT_CHECK_PARAM(impl != nullptr && idata != nullptr && odata != nullptr, ACLFFT_INVALID_VALUE);
58 ACLFFT_CHECK_PARAM(impl->rank == 1, ACLFFT_INVALID_VALUE);60 ACLFFT_CHECK_PARAM(impl->rank == 1, ACLFFT_INVALID_VALUE);
59 61 
60 const uint32_t n = impl->lengths[0];62 const uint32_t n = impl->lengths[0];
@@ -112,6 +112,11 @@ static std::vector<float> GenerateMixTwiddle(const std::vector<int32_t> &radixLi
112extern "C" aclError aclfftFft1DC2CMix(float *x, float *y, uint32_t n, int32_t norm,112extern "C" aclError aclfftFft1DC2CMix(float *x, float *y, uint32_t n, int32_t norm,
113 uint32_t batches, int isForward, void *stream)113 uint32_t batches, int isForward, void *stream)
114{114{
115+ // 防护: 本函数经 ACLFFT_API 导出可被外部直接调用,需校验输入输出指针(issue #52)
116+ if (x == nullptr || y == nullptr) {
117+ std::cerr << "[ops-fft] aclfftFft1DC2CMix: input/output pointer is null" << std::endl;
118+ return ACL_ERROR_INVALID_PARAM;
119+ }
115 (void)norm;120 (void)norm;
116 auto ascendcPlatform = platform_ascendc::PlatformAscendCManager::GetInstance();121 auto ascendcPlatform = platform_ascendc::PlatformAscendCManager::GetInstance();
117 uint32_t coreNum = ascendcPlatform->GetCoreNumAiv();122 uint32_t coreNum = ascendcPlatform->GetCoreNumAiv();
@@ -62,6 +62,11 @@ static std::vector<float> InitPQMatrix(int64_t fftN, bool forward, bool isP)
62 62 
63extern "C" aclError aclfftFft2DDd(float *x, float *y, uint32_t fftX, uint32_t fftY,63extern "C" aclError aclfftFft2DDd(float *x, float *y, uint32_t fftX, uint32_t fftY,
64 uint32_t batches, int isForward, void *stream) {64 uint32_t batches, int isForward, void *stream) {
65+ // 防护: 本函数经 ACLFFT_API 导出可被外部直接调用,需校验输入输出指针(issue #53)
66+ if (x == nullptr || y == nullptr) {
67+ std::cerr << "[ops-fft] aclfftFft2DDd: input/output pointer is null" << std::endl;
68+ return ACL_ERROR_INVALID_PARAM;
69+ }
65 auto ascendcPlatform = platform_ascendc::PlatformAscendCManager::GetInstance();70 auto ascendcPlatform = platform_ascendc::PlatformAscendCManager::GetInstance();
66 uint32_t coreNum = ascendcPlatform->GetCoreNumAic();71 uint32_t coreNum = ascendcPlatform->GetCoreNumAic();
67 if (coreNum == 0) {72 if (coreNum == 0) {
@@ -68,11 +68,12 @@ extern "C" aclError aclfftRfft1DDft(float *x, float *y, uint32_t n, int32_t norm
68 68 
69 std::vector<float> dftMatrix = GenerateDftMatrixR2C(fftN);69 std::vector<float> dftMatrix = GenerateDftMatrixR2C(fftN);
70 70 
71- uint32_t inputSize = batches * fftN * sizeof(float);71+ // size_t 计算,避免 batches*fftN 等中间结果在 uint32_t 域溢出截断(issue #56)
72- uint32_t dftMatrixSize = dftMatrix.size() * sizeof(float);72+ size_t inputSize = static_cast<size_t>(batches) * static_cast<size_t>(fftN) * sizeof(float);
73- uint32_t outputSize = batches * (fftN / 2 + 1) * sizeof(float) * 2;73+ size_t dftMatrixSize = dftMatrix.size() * sizeof(float);
74- uint32_t sysWorkspaceSize = ascendcPlatform->GetLibApiWorkSpaceSize();74+ size_t outputSize = static_cast<size_t>(batches) * (fftN / 2 + 1) * sizeof(float) * 2;
75- uint32_t tilingSize = sizeof(Rfft1DDftTilingData);75+ size_t sysWorkspaceSize = static_cast<size_t>(ascendcPlatform->GetLibApiWorkSpaceSize());
76+ size_t tilingSize = sizeof(Rfft1DDftTilingData);
76 77 
77 void *dev_input = nullptr;78 void *dev_input = nullptr;
78 void *dev_dft = nullptr;79 void *dev_dft = nullptr;
@@ -317,9 +317,11 @@ extern "C" aclError aclfftRfft1D(float *x, float *y, uint32_t n, int32_t norm, u
317 317 
318 uint32_t sysWorkspaceSize = ascendcPlatform->GetLibApiWorkSpaceSize();318 uint32_t sysWorkspaceSize = ascendcPlatform->GetLibApiWorkSpaceSize();
319 319 
320- const uint32_t inputSize = n * batches * sizeof(float);320+ // size_t 计算,避免 ((n/2)+1)*2*batches 等中间结果在 uint32_t 域溢出截断(issue #59)
321- const uint32_t dftSize = dft.size() * sizeof(float);321+ const size_t inputSize = static_cast<size_t>(n) * static_cast<size_t>(batches) * sizeof(float);
322- const uint32_t outputSize = ((n / RFFT_SYMMETRY_DIVISOR) + 1) * COMPLEX_PART * batches * sizeof(float);322+ const size_t dftSize = dft.size() * sizeof(float);
323+ const size_t outputSize = (static_cast<size_t>(n / RFFT_SYMMETRY_DIVISOR) + 1) * COMPLEX_PART *
324+ static_cast<size_t>(batches) * sizeof(float);
323 325 
324 void *dev_x = nullptr;326 void *dev_x = nullptr;
325 void *dev_dft = nullptr;327 void *dev_dft = nullptr;