diff --git a/cpp/src/arrow/util/decimal.cc b/cpp/src/arrow/util/decimal.cc index a9d2fcb02d9..39ea6755b97 100644 --- a/cpp/src/arrow/util/decimal.cc +++ b/cpp/src/arrow/util/decimal.cc @@ -768,7 +768,8 @@ std::string Decimal128::ToString(int32_t scale) const { // Iterates over input and for each group of kInt64DecimalDigits multiple out by // the appropriate power of 10 necessary to add source parsed as uint64 and // then adds the parsed value of source. -static inline void ShiftAndAdd(std::string_view input, uint64_t out[], size_t out_size) { +static inline bool ShiftAndAddWithOverflow(std::string_view input, uint64_t out[], + size_t out_size) { for (size_t posn = 0; posn < input.size();) { const size_t group_size = std::min(kInt64DecimalDigits, input.size() - posn); const uint64_t multiple = kUInt64PowersOfTen[group_size]; @@ -783,8 +784,25 @@ static inline void ShiftAndAdd(std::string_view input, uint64_t out[], size_t ou out[i] = static_cast(tmp & 0xFFFFFFFFFFFFFFFFULL); chunk = static_cast(tmp >> 64); } + if (chunk != 0) { + return true; + } posn += group_size; } + return false; +} + +static inline bool MagnitudeOverflowsSignedDecimal(const uint64_t out[], size_t out_size, + bool negative) { + constexpr uint64_t kSignBit = uint64_t{1} << 63; + const uint64_t high = out[out_size - 1]; + if (high < kSignBit) { + return false; + } + if (!negative || high > kSignBit) { + return true; + } + return std::any_of(out, out + out_size - 1, [](uint64_t word) { return word != 0; }); } namespace { @@ -895,9 +913,14 @@ Status DecimalFromString(const char* type_name, std::string_view s, Decimal* out if (out != nullptr) { static_assert(Decimal::kBitWidth % 64 == 0, "decimal bit-width not a multiple of 64"); std::array little_endian_array{}; - ShiftAndAdd(dec.whole_digits, little_endian_array.data(), little_endian_array.size()); - ShiftAndAdd(dec.fractional_digits, little_endian_array.data(), - little_endian_array.size()); + if (ShiftAndAddWithOverflow(dec.whole_digits, little_endian_array.data(), + little_endian_array.size()) || + ShiftAndAddWithOverflow(dec.fractional_digits, little_endian_array.data(), + little_endian_array.size()) || + MagnitudeOverflowsSignedDecimal(little_endian_array.data(), + little_endian_array.size(), dec.sign == '-')) { + return Status::Invalid("The string '", s, "' cannot be represented as ", type_name); + } *out = Decimal(bit_util::little_endian::ToNative(little_endian_array)); if (dec.sign == '-') { out->Negate(); @@ -962,9 +985,9 @@ Status SimpleDecimalFromString(const char* type_name, std::string_view s, if (out != nullptr) { uint64_t value{0}; - ShiftAndAdd(dec.whole_digits, &value, 1); - ShiftAndAdd(dec.fractional_digits, &value, 1); - if (value > static_cast( + if (ShiftAndAddWithOverflow(dec.whole_digits, &value, 1) || + ShiftAndAddWithOverflow(dec.fractional_digits, &value, 1) || + value > static_cast( std::numeric_limits::max())) { return Status::Invalid("The string '", s, "' cannot be represented as ", type_name); } diff --git a/cpp/src/arrow/util/decimal_test.cc b/cpp/src/arrow/util/decimal_test.cc index 7022c811780..c667daf9759 100644 --- a/cpp/src/arrow/util/decimal_test.cc +++ b/cpp/src/arrow/util/decimal_test.cc @@ -435,15 +435,19 @@ TEST(Decimal128Test, FromStringLimits) { ASSERT_RAISES(Invalid, Decimal128::FromString("-9e39")); ASSERT_RAISES(Invalid, Decimal128::FromString("9.9e40")); ASSERT_RAISES(Invalid, Decimal128::FromString("-9.9e40")); - // XXX conversion overflows are currently not detected + // XXX conversion overflows after parsing are currently not detected // ASSERT_RAISES(Invalid, Decimal128::FromString("99e38")); // ASSERT_RAISES(Invalid, Decimal128::FromString("-99e38")); // ASSERT_RAISES(Invalid, // Decimal128::FromString("999999999999999999999999999999999999999e1")); // ASSERT_RAISES(Invalid, // Decimal128::FromString("-999999999999999999999999999999999999999e1")); - // ASSERT_RAISES(Invalid, - // Decimal128::FromString("999999999999999999999999999999999999999")); + ASSERT_RAISES(Invalid, Decimal128::FromString( + "1.55555555555555555555555555555555555555555555555555")); + ASSERT_RAISES(Invalid, + Decimal128::FromString("170141183460469231731687303715884105728")); + ASSERT_RAISES(Invalid, + Decimal128::FromString("-170141183460469231731687303715884105729")); // No exponent, many fractional digits AssertDecimalFromString("9.9999999999999999999999999999999999999", dec38times9pos, 38, @@ -541,7 +545,8 @@ TEST(Decimal256Test, FromStringLimits) { ASSERT_RAISES(Invalid, Decimal256::FromString("9.9e78")); ASSERT_RAISES(Invalid, Decimal256::FromString("-9.9e78")); - // XXX conversion overflows are currently not detected + // XXX precision limits and conversion overflows after parsing are currently not + // detected // ASSERT_RAISES(Invalid, Decimal256::FromString("99e76")); // ASSERT_RAISES(Invalid, Decimal256::FromString("-99e76")); // ASSERT_RAISES(Invalid, @@ -550,6 +555,9 @@ TEST(Decimal256Test, FromStringLimits) { // Decimal256::FromString("-9999999999999999999999999999999999999999999999999999999999999999999999999999e1")); // ASSERT_RAISES(Invalid, // Decimal256::FromString("99999999999999999999999999999999999999999999999999999999999999999999999999999")); + ASSERT_RAISES(Invalid, Decimal256::FromString(std::string(78, '9'))); + ASSERT_RAISES(Invalid, Decimal256::FromString("5789604461865809771178549250434395392663" + "4992332820282019728792003956564819968")); // No exponent, many fractional digits AssertDecimalFromString(