From b9bbf71d3485d5e840a46a4896e65161815b82ff Mon Sep 17 00:00:00 2001 From: Abtin Keshavarzian Date: Wed, 17 Dec 2025 13:01:36 -0800 Subject: [PATCH] [num-utils] add `SafeMultiply()` for overflow-safe multiplication (#12220) This commit introduces `SafeMultiply()` in `num_utils.hpp` as a centralized and safe way to multiply two unsigned integers while checking for overflow. It updates `Coap::TxParameters::IsValid()` to use this new helper for validating `TxParameters`, replacing a less robust local `Multiply` implementation. It also updates `Heap::CAlloc()` to use this function for safely calculating the total allocation size. Unit tests are updated to verify `SafeMultiply()` implementation. --- src/core/coap/coap.cpp | 58 +++++++++++++++---------------- src/core/coap/coap.hpp | 2 ++ src/core/common/num_utils.hpp | 33 ++++++++++++++++++ src/core/utils/heap.cpp | 10 +++--- tests/unit/test_serial_number.cpp | 35 +++++++++++++++++++ 5 files changed, 103 insertions(+), 35 deletions(-) diff --git a/src/core/coap/coap.cpp b/src/core/coap/coap.cpp index 532c04ad2..6816e8bc8 100644 --- a/src/core/coap/coap.cpp +++ b/src/core/coap/coap.cpp @@ -1631,42 +1631,42 @@ void ResponsesQueue::HandleTimer(void) mTimer.FireAt(nextDequeueTime); } -/// Return product of @p aValueA and @p aValueB if no overflow otherwise 0. -static uint32_t Multiply(uint32_t aValueA, uint32_t aValueB) -{ - uint32_t result = 0; - - VerifyOrExit(aValueA); - - result = aValueA * aValueB; - result = (result / aValueA == aValueB) ? result : 0; - -exit: - return result; -} - bool TxParameters::IsValid(void) const { - bool rval = false; + bool isValid = false; + uint32_t duration; + uint32_t retryFactor; - // support fire and forget requests if (mAckTimeout == 0) { - rval = true; - } - else if ((mAckRandomFactorDenominator > 0) && (mAckRandomFactorNumerator >= mAckRandomFactorDenominator) && - (mAckTimeout >= OT_COAP_MIN_ACK_TIMEOUT) && (mMaxRetransmit <= OT_COAP_MAX_RETRANSMIT)) - { - // Calculate exchange lifetime step by step and verify no overflow. - uint32_t tmp = Multiply(mAckTimeout, (1U << (mMaxRetransmit + 1)) - 1); - - tmp = Multiply(tmp, mAckRandomFactorNumerator); - tmp /= mAckRandomFactorDenominator; - - rval = (tmp != 0 && (tmp + mAckTimeout + 2 * kDefaultMaxLatency) > tmp); + // Support fire and forget requests + isValid = true; + ExitNow(); } - return rval; + VerifyOrExit(mAckRandomFactorDenominator > 0); + VerifyOrExit(mAckRandomFactorNumerator >= mAckRandomFactorDenominator); + VerifyOrExit(mAckTimeout >= kMinAckTimeout); + VerifyOrExit(mMaxRetransmit <= kMaxRetransmit); + + // Calculate exchange lifetime max duration step by step and verify no overflow. + + static_assert(kMaxRetransmit < 31, "kMaxRetransmit is not valid"); + + retryFactor = static_cast((1U << (mMaxRetransmit + 1)) - 1); + SuccessOrExit(SafeMultiply(mAckTimeout, retryFactor, duration)); + + SuccessOrExit(SafeMultiply(duration, mAckRandomFactorNumerator, duration)); + duration /= mAckRandomFactorDenominator; + + VerifyOrExit(duration > 0); + VerifyOrExit(CanAddSafely(mAckTimeout, 2 * kDefaultMaxLatency)); + VerifyOrExit(CanAddSafely(duration, mAckTimeout + 2 * kDefaultMaxLatency)); + + isValid = true; + +exit: + return isValid; } uint32_t TxParameters::CalculateInitialRetransmissionTimeout(void) const diff --git a/src/core/coap/coap.hpp b/src/core/coap/coap.hpp index 1ee0a3ff5..f8ff74caf 100644 --- a/src/core/coap/coap.hpp +++ b/src/core/coap/coap.hpp @@ -128,6 +128,8 @@ private: static constexpr uint8_t kDefaultAckRandomFactorDenominator = 2; static constexpr uint8_t kDefaultMaxRetransmit = 4; static constexpr uint32_t kDefaultMaxLatency = 100000; // in msec + static constexpr uint8_t kMaxRetransmit = OT_COAP_MAX_RETRANSMIT; + static constexpr uint32_t kMinAckTimeout = OT_COAP_MIN_ACK_TIMEOUT; uint32_t CalculateInitialRetransmissionTimeout(void) const; uint32_t CalculateExchangeLifetime(void) const; diff --git a/src/core/common/num_utils.hpp b/src/core/common/num_utils.hpp index 99104388e..4d42a5b9d 100644 --- a/src/core/common/num_utils.hpp +++ b/src/core/common/num_utils.hpp @@ -34,6 +34,8 @@ #ifndef NUM_UTILS_HPP_ #define NUM_UTILS_HPP_ +#include "common/code_utils.hpp" +#include "common/error.hpp" #include "common/numeric_limits.hpp" #include "common/type_traits.hpp" @@ -229,6 +231,37 @@ template <> inline int ThreeWayCompare(bool aFirst, bool aSecond) return (aFirst == aSecond) ? 0 : (aFirst ? 1 : -1); } +/** + * Safely multiplies two unsigned integers and checks for overflow. + * + * @tparam UintType The value type (MUST be `uint8_t`, `uint16_t`, `uint32_t`, or `uint64_t`). + * + * @param[in] aFirstValue The first operand in the multiplication. + * @param[in] aSecondValue The second operand in the multiplication. + * @param[out] aResult A reference to return the multiplication result. + * + * @retval kErrorNone If the multiplication was performed safely without overflow. @p aResult is updated. + * @retval kErrorInvalidArgs If the multiplication would result in an overflow. + */ +template inline Error SafeMultiply(UintType aFirstValue, UintType aSecondValue, UintType &aResult) +{ + static_assert(TypeTraits::IsUint::kValue, "UintType must be an unsigned int (8, 16, 32, or 64 bit len)"); + + Error error = kErrorNone; + + if (aFirstValue == 0 || aSecondValue == 0) + { + aResult = 0; + ExitNow(); + } + + aResult = aFirstValue * aSecondValue; + VerifyOrExit(aResult / aFirstValue == aSecondValue, error = kErrorInvalidArgs); + +exit: + return error; +} + /** * This template function divides two numbers and rounds the result to the closest integer. * diff --git a/src/core/utils/heap.cpp b/src/core/utils/heap.cpp index 6489783dd..a865fbbce 100644 --- a/src/core/utils/heap.cpp +++ b/src/core/utils/heap.cpp @@ -39,6 +39,7 @@ #include "common/code_utils.hpp" #include "common/debug.hpp" +#include "common/num_utils.hpp" #include "common/numeric_limits.hpp" namespace ot { @@ -66,7 +67,6 @@ void *Heap::CAlloc(size_t aCount, size_t aSize) void *ret = nullptr; Block *prev = nullptr; Block *curr = nullptr; - size_t totalSize; uint16_t size; // Verify that the requested allocation size will not cause an overflow. @@ -79,12 +79,10 @@ void *Heap::CAlloc(size_t aCount, size_t aSize) VerifyOrExit(aCount <= NumericLimits::kMax); VerifyOrExit(aSize <= NumericLimits::kMax); - totalSize = aCount * aSize; - VerifyOrExit(totalSize <= NumericLimits::kMax - kTotalSizeGuard); + SuccessOrExit(SafeMultiply(static_cast(aCount), static_cast(aSize), size)); - size = static_cast(totalSize); - - VerifyOrExit(size); + VerifyOrExit(size > 0); + VerifyOrExit(size <= NumericLimits::kMax - kTotalSizeGuard); size += kAlignSize - 1 - kBlockRemainderSize; size &= ~(kAlignSize - 1); diff --git a/tests/unit/test_serial_number.cpp b/tests/unit/test_serial_number.cpp index bba347fdb..f894c5b5a 100644 --- a/tests/unit/test_serial_number.cpp +++ b/tests/unit/test_serial_number.cpp @@ -151,6 +151,41 @@ void TestNumUtils(void) VerifyOrQuit(ThreeWayCompare(true, false) > 0); VerifyOrQuit(ThreeWayCompare(false, true) < 0); + SuccessOrQuit(SafeMultiply(0, 0, u16)); + VerifyOrQuit(u16 == 0); + SuccessOrQuit(SafeMultiply(0, 0xffff, u16)); + VerifyOrQuit(u16 == 0); + SuccessOrQuit(SafeMultiply(0xffff, 0, u16)); + VerifyOrQuit(u16 == 0); + + SuccessOrQuit(SafeMultiply(1, 0xffff, u16)); + VerifyOrQuit(u16 == 0xffff); + SuccessOrQuit(SafeMultiply(0xffff, 1, u16)); + VerifyOrQuit(u16 == 0xffff); + + SuccessOrQuit(SafeMultiply(256, 255, u16)); + VerifyOrQuit(u16 == 65280); + SuccessOrQuit(SafeMultiply(255, 256, u16)); + VerifyOrQuit(u16 == 65280); + + VerifyOrQuit(SafeMultiply(256, 256, u16) == kErrorInvalidArgs); + + for (uint16_t num = 2; num < 255; num++) + { + uint16_t div = 0xffff / num; + + SuccessOrQuit(SafeMultiply(num, div, u16)); + VerifyOrQuit(u16 == num * div); + SuccessOrQuit(SafeMultiply(div, num, u16)); + VerifyOrQuit(u16 == num * div); + + VerifyOrQuit(SafeMultiply(num, div + 1, u16) == kErrorInvalidArgs); + VerifyOrQuit(SafeMultiply(div + 1, num, u16) == kErrorInvalidArgs); + + VerifyOrQuit(SafeMultiply(num + 1, div, u16) == kErrorInvalidArgs); + VerifyOrQuit(SafeMultiply(div, num + 1, u16) == kErrorInvalidArgs); + } + VerifyOrQuit(DivideAndRoundToClosest(2, 1) == 2); VerifyOrQuit(DivideAndRoundToClosest(1, 3) == 0); VerifyOrQuit(DivideAndRoundToClosest(1, 2) == 1);