C++17中如何安全检查有符号整数乘法是否溢出?
检查有符号整数乘法溢出的高效便捷方案(C++17)
最优方案:编译器内置溢出检查函数
GCC、Clang、MSVC(2019及以后版本)均提供了原生的有符号乘法溢出检查函数,这类函数直接利用硬件指令的溢出标志,无除法开销,能被编译器完全优化,且语义明确不会被误优化,是最便捷高效的选择。
针对long long(64位有符号整数)
使用__builtin_smulll_overflow,函数返回bool值表示是否溢出,同时可将乘积存入第三个参数:
#include <cstdint> // 检查long long乘法是否溢出,可选返回计算结果 bool is_ll_multiply_overflow(long long a, long long b, long long* product = nullptr) { long long res; const bool overflow = __builtin_smulll_overflow(a, b, &res); if (product != nullptr) { *product = res; } return overflow; }
针对long(32/64位,取决于平台)
使用__builtin_smul_overflow,用法与上述一致:
#include <cstdint> // 检查long乘法是否溢出,可选返回计算结果 bool is_l_multiply_overflow(long a, long b, long* product = nullptr) { long res; const bool overflow = __builtin_smul_overflow(a, b, &res); if (product != nullptr) { *product = res; } return overflow; }
跨编译器通用包装(可选)
如果需要兼容多编译器,可以用模板+条件编译统一接口:
#include <cstdint> #include <type_traits> template <typename T, std::enable_if_t<std::is_signed_v<T>, bool> = true> bool multiply_overflow(T a, T b, T* product = nullptr) { static_assert(sizeof(T) == 4 || sizeof(T) == 8, "仅支持32/64位有符号整数"); if constexpr (sizeof(T) == 8) { long long res; const bool overflow = __builtin_smulll_overflow(static_cast<long long>(a), static_cast<long long>(b), &res); if (product) *product = static_cast<T>(res); return overflow; } else { long res; const bool overflow = __builtin_smul_overflow(static_cast<long>(a), static_cast<long>(b), &res); if (product) *product = static_cast<T>(res); return overflow; } }
需要规避的错误方案
先乘后查:这类写法因有符号算术溢出属于未定义行为,编译器可能基于“不会溢出”的假设移除检查,存在严重安全隐患:
// 错误示例:溢出判断会被编译器优化移除 long long a, b; long long product = a * b; if (product / a != b) { /* 无效的溢出检查 */ }除法判断:如
b <= std::numeric_limits<long long>::max() / a这类写法,需要处理正负号的多种分支,且除法运算开销远高于内置函数,不符合性能要求。
替代方案:利用__int128(无编译器内置时)
如果编译器支持__int128(GCC、Clang、MSVC均支持),可以通过扩展精度乘法判断溢出,同样无除法开销:
#include <cstdint> #include <limits> #include <type_traits> template <typename T, std::enable_if_t<std::is_signed_v<T> && sizeof(T) == 8, bool> = true> bool ll_multiply_overflow(T a, T b) { using U = std::make_unsigned_t<T>; const U ua = static_cast<U>(a); const U ub = static_cast<U>(b); // 用__int128计算完整乘积,拆分高低位 const __int128 full_product = static_cast<__int128>(ua) * ub; const U product_high = static_cast<U>(full_product >> 64); const T sign_a = a >> 63; // 符号位扩展,正数为0,负数为-1 if (sign_a == (b >> 63)) { // 同号相乘:高位必须等于符号位扩展值,否则溢出 return static_cast<T>(product_high) != sign_a; } else { // 异号相乘:仅当其中一个是LLONG_MIN且另一个是-1时溢出 return (a == std::numeric_limits<T>::min() && b == -1) || (b == std::numeric_limits<T>::min() && a == -1); } }
内容的提问来源于stack exchange,提问作者user3188445
相关产品推荐
相关产品推荐

