diff --git a/cpp/src/arrow/json/converter_test.cc b/cpp/src/arrow/json/converter_test.cc index fa85e704bc5e..90828639849a 100644 --- a/cpp/src/arrow/json/converter_test.cc +++ b/cpp/src/arrow/json/converter_test.cc @@ -254,9 +254,15 @@ TEST(ConverterTest, Decimal128And256PrecisionError) { std::shared_ptr parse_array; ASSERT_OK(ParseFromString(options, json_source, &parse_array)); - std::string error_msg = - "Invalid: Failed to convert JSON to " + decimal_type->ToString() + - ": 123456789012345678901234567890.0123456789 requires precision 40"; + std::string error_msg; + if (decimal_type->id() == Type::DECIMAL128) { + error_msg = + "Invalid: The string '123456789012345678901234567890.0123456789' " + "cannot be represented as decimal128"; + } else { + error_msg = "Invalid: Failed to convert JSON to " + decimal_type->ToString() + + ": 123456789012345678901234567890.0123456789 requires precision 40"; + } EXPECT_RAISES_WITH_MESSAGE_THAT( Invalid, ::testing::HasSubstr(error_msg), Convert(decimal_type, parse_array->GetFieldByName(""))); diff --git a/cpp/src/arrow/util/decimal.cc b/cpp/src/arrow/util/decimal.cc index a9d2fcb02d94..05daf0226da6 100644 --- a/cpp/src/arrow/util/decimal.cc +++ b/cpp/src/arrow/util/decimal.cc @@ -768,7 +768,9 @@ 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, bool negative) { + constexpr uint64_t kSignBit = uint64_t{1} << 63; 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,20 +785,22 @@ 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; + } + const uint64_t high = out[out_size - 1]; + if ((high & kSignBit) != 0 && + (!negative || high != kSignBit || + std::any_of(out, out + out_size - 1, [](uint64_t word) { return word != 0; }))) { + return true; + } posn += group_size; } + return false; } namespace { -struct DecimalComponents { - std::string_view whole_digits; - std::string_view fractional_digits; - int32_t exponent = 0; - char sign = 0; - bool has_exponent = false; -}; - inline bool IsSign(char c) { return c == '-' || c == '+'; } inline bool IsDot(char c) { return c == '.'; } @@ -817,7 +821,10 @@ inline size_t ParseDigitsRun(const char* s, size_t start, size_t size, return pos; } -bool ParseDecimalComponents(const char* s, size_t size, DecimalComponents* out) { +} // namespace + +bool internal::ParseDecimalComponents(const char* s, size_t size, + DecimalComponents* out) { size_t pos = 0; if (size == 0) { @@ -859,6 +866,8 @@ bool ParseDecimalComponents(const char* s, size_t size, DecimalComponents* out) return pos == size; } +namespace { + template Status DecimalFromString(const char* type_name, std::string_view s, Decimal* out, int32_t* precision, int32_t* scale) { @@ -866,8 +875,8 @@ Status DecimalFromString(const char* type_name, std::string_view s, Decimal* out return Status::Invalid("Empty string cannot be converted to ", type_name); } - DecimalComponents dec; - if (!ParseDecimalComponents(s.data(), s.size(), &dec)) { + internal::DecimalComponents dec; + if (!internal::ParseDecimalComponents(s.data(), s.size(), &dec)) { return Status::Invalid("The string '", s, "' is not a valid ", type_name, " number"); } @@ -892,16 +901,17 @@ Status DecimalFromString(const char* type_name, std::string_view s, Decimal* out parsed_scale = static_cast(dec.fractional_digits.size()); } - 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()); - *out = Decimal(bit_util::little_endian::ToNative(little_endian_array)); - if (dec.sign == '-') { - out->Negate(); - } + static_assert(Decimal::kBitWidth % 64 == 0, "decimal bit-width not a multiple of 64"); + std::array little_endian_array{}; + if (ShiftAndAddWithOverflow(dec.whole_digits, little_endian_array.data(), + little_endian_array.size(), dec.sign == '-') || + ShiftAndAddWithOverflow(dec.fractional_digits, little_endian_array.data(), + little_endian_array.size(), dec.sign == '-')) { + return Status::Invalid("The string '", s, "' cannot be represented as ", type_name); + } + Decimal parsed_value(bit_util::little_endian::ToNative(little_endian_array)); + if (dec.sign == '-') { + parsed_value.Negate(); } if (parsed_scale < 0) { @@ -910,13 +920,19 @@ Status DecimalFromString(const char* type_name, std::string_view s, Decimal* out if (-parsed_scale > Decimal::kMaxScale) { return Status::Invalid("The string '", s, "' cannot be represented as ", type_name); } - if (out != nullptr) { - *out *= Decimal::GetScaleMultiplier(-parsed_scale); + const auto& multiplier = Decimal::GetScaleMultiplier(-parsed_scale); + if (parsed_value > Decimal::GetMaxSentinel() / multiplier || + parsed_value < Decimal::GetMinSentinel() / multiplier) { + return Status::Invalid("The string '", s, "' cannot be represented as ", type_name); } + parsed_value *= multiplier; parsed_precision -= parsed_scale; parsed_scale = 0; } + if (out != nullptr) { + *out = parsed_value; + } if (precision != nullptr) { *precision = parsed_precision; } @@ -934,8 +950,8 @@ Status SimpleDecimalFromString(const char* type_name, std::string_view s, return Status::Invalid("Empty string cannot be converted to ", type_name); } - DecimalComponents dec; - if (!ParseDecimalComponents(s.data(), s.size(), &dec)) { + internal::DecimalComponents dec; + if (!internal::ParseDecimalComponents(s.data(), s.size(), &dec)) { return Status::Invalid("The string '", s, "' is not a valid ", type_name, " number"); } @@ -960,19 +976,17 @@ Status SimpleDecimalFromString(const char* type_name, std::string_view s, parsed_scale = static_cast(dec.fractional_digits.size()); } - if (out != nullptr) { - uint64_t value{0}; - ShiftAndAdd(dec.whole_digits, &value, 1); - ShiftAndAdd(dec.fractional_digits, &value, 1); - if (value > static_cast( - std::numeric_limits::max())) { - return Status::Invalid("The string '", s, "' cannot be represented as ", type_name); - } - - *out = DecimalClass(value); - if (dec.sign == '-') { - out->Negate(); - } + uint64_t value{0}; + if (ShiftAndAddWithOverflow(dec.whole_digits, &value, 1, dec.sign == '-') || + ShiftAndAddWithOverflow(dec.fractional_digits, &value, 1, dec.sign == '-') || + value > static_cast( + std::numeric_limits::max()) + + static_cast(dec.sign == '-')) { + return Status::Invalid("The string '", s, "' cannot be represented as ", type_name); + } + DecimalClass parsed_value(value); + if (dec.sign == '-') { + parsed_value.Negate(); } if (parsed_scale < 0) { @@ -981,13 +995,20 @@ Status SimpleDecimalFromString(const char* type_name, std::string_view s, if (-parsed_scale > DecimalClass::kMaxScale) { return Status::Invalid("The string '", s, "' cannot be represented as ", type_name); } - if (out != nullptr) { - *out *= DecimalClass::GetScaleMultiplier(-parsed_scale); + typename DecimalClass::ValueType scaled_value; + if (internal::MultiplyWithOverflow( + parsed_value.value(), DecimalClass::GetScaleMultiplier(-parsed_scale).value(), + &scaled_value)) { + return Status::Invalid("The string '", s, "' cannot be represented as ", type_name); } + parsed_value = DecimalClass(scaled_value); parsed_precision -= parsed_scale; parsed_scale = 0; } + if (out != nullptr) { + *out = parsed_value; + } if (precision != nullptr) { *precision = parsed_precision; } diff --git a/cpp/src/arrow/util/decimal_internal.h b/cpp/src/arrow/util/decimal_internal.h index 3845a544cff3..094379f657a0 100644 --- a/cpp/src/arrow/util/decimal_internal.h +++ b/cpp/src/arrow/util/decimal_internal.h @@ -20,6 +20,7 @@ #include #include #include +#include #include #include "arrow/type_fwd.h" @@ -30,6 +31,21 @@ namespace arrow { +namespace internal { + +struct DecimalComponents { + std::string_view whole_digits; + std::string_view fractional_digits; + int32_t exponent = 0; + char sign = 0; + bool has_exponent = false; +}; + +ARROW_EXPORT bool ParseDecimalComponents(const char* s, size_t size, + DecimalComponents* out); + +} // namespace internal + constexpr auto kInt32DecimalDigits = static_cast(std::numeric_limits::digits10); diff --git a/cpp/src/arrow/util/decimal_test.cc b/cpp/src/arrow/util/decimal_test.cc index 7022c8117802..a5f168df80a6 100644 --- a/cpp/src/arrow/util/decimal_test.cc +++ b/cpp/src/arrow/util/decimal_test.cc @@ -20,6 +20,7 @@ #include #include #include +#include #include #include #include @@ -68,6 +69,10 @@ void AssertDecimalFromString(const std::string& s, const DecimalType& expected, EXPECT_EQ(expected, d); EXPECT_EQ(expected_precision, precision); EXPECT_EQ(expected_scale, scale); + precision = scale = -1; + ASSERT_OK(DecimalType::FromString(s, nullptr, &precision, &scale)); + EXPECT_EQ(expected_precision, precision); + EXPECT_EQ(expected_scale, scale); } // Assert that the low bits of an array of integers are equal to `expected_low`, @@ -163,6 +168,53 @@ class DecimalFromStringTest : public ::testing::Test { ASSERT_OK_AND_EQ(expected_value, DecimalType::FromString("1.23E-8")); } + void TestPositiveExponentOverflow() { + const std::string exponent = "e" + std::to_string(DecimalType::kMaxScale); + for (const std::string& s : + {"99" + exponent, "-99" + exponent, "9.9" + exponent, "-9.9" + exponent}) { + ARROW_SCOPED_TRACE("s = '", s, "'"); + ASSERT_RAISES(Invalid, DecimalType::FromString(s)); + int32_t precision, scale; + ASSERT_RAISES(Invalid, DecimalType::FromString(s, nullptr, &precision, &scale)); + } + AssertDecimalFromString("0" + exponent, DecimalType(0), DecimalType::kMaxScale, 0); + AssertDecimalFromString("-0" + exponent, DecimalType(0), DecimalType::kMaxScale, 0); + } + + void TestPositiveExponentLimits() { + DecimalType maximum; + if constexpr (DecimalType::kBitWidth <= 64) { + maximum = DecimalType(std::numeric_limits::max()); + } else { + maximum = DecimalType(DecimalType::GetMaxSentinel()); + } + DecimalType minimum = maximum; + minimum.Negate(); + minimum -= DecimalType(1); + for (const DecimalType& limit : {maximum, minimum}) { + const int32_t sign_size = limit.IsNegative() ? 1 : 0; + const std::string limit_string = limit.ToIntegerString(); + AssertDecimalFromString(limit_string + "e0", limit, + static_cast(limit_string.size()) - sign_size, 0); + for (int32_t exponent = 1; exponent <= DecimalType::kMaxScale; ++exponent) { + const auto multiplier = DecimalType::GetScaleMultiplier(exponent); + const DecimalType coefficient(limit / multiplier); + const std::string digits = coefficient.ToIntegerString(); + const std::string suffix = "e" + std::to_string(exponent); + AssertDecimalFromString( + digits + suffix, DecimalType(coefficient * multiplier), + static_cast(digits.size()) - sign_size + exponent, 0); + const DecimalType overflow(coefficient + + DecimalType(limit.IsNegative() ? -1 : 1)); + const std::string s = overflow.ToIntegerString() + suffix; + ARROW_SCOPED_TRACE("s = '", s, "'"); + ASSERT_RAISES(Invalid, DecimalType::FromString(s)); + int32_t precision, scale; + ASSERT_RAISES(Invalid, DecimalType::FromString(s, nullptr, &precision, &scale)); + } + } + } + void TestSmallValues() { struct TestValue { std::string s; @@ -230,6 +282,14 @@ TEST(Decimal32Test, TestIntMinFitsPrecision) { ASSERT_FALSE(d.FitsInPrecision(9)); } +TEST(Decimal32Test, FromStringLimits) { + AssertDecimalFromString("-2147483648", Decimal32(INT32_MIN), 10, 0); + ASSERT_RAISES(Invalid, Decimal32::FromString("-2147483649")); + int32_t precision, scale; + ASSERT_RAISES(Invalid, + Decimal32::FromString("-2147483649", nullptr, &precision, &scale)); +} + TEST(Decimal64Test, TestIntMinNegate) { Decimal64 d(INT64_MIN); auto neg = d.Negate(); @@ -241,6 +301,14 @@ TEST(Decimal64Test, TestIntMinFitsPrecision) { ASSERT_FALSE(d.FitsInPrecision(18)); } +TEST(Decimal64Test, FromStringLimits) { + AssertDecimalFromString("-9223372036854775808", Decimal64(INT64_MIN), 19, 0); + ASSERT_RAISES(Invalid, Decimal64::FromString("-9223372036854775809")); + int32_t precision, scale; + ASSERT_RAISES(Invalid, Decimal64::FromString("-9223372036854775809", nullptr, + &precision, &scale)); +} + TYPED_TEST_SUITE(DecimalFromStringTest, DecimalTypes); TYPED_TEST(DecimalFromStringTest, Basics) { this->TestBasics(); } @@ -271,6 +339,14 @@ TYPED_TEST(DecimalFromStringTest, WithExponentAndNullptrScale) { this->TestWithExponentAndNullptrScale(); } +TYPED_TEST(DecimalFromStringTest, PositiveExponentOverflow) { + this->TestPositiveExponentOverflow(); +} + +TYPED_TEST(DecimalFromStringTest, PositiveExponentLimits) { + this->TestPositiveExponentLimits(); +} + TYPED_TEST(DecimalFromStringTest, SmallValues) { this->TestSmallValues(); } TYPED_TEST(DecimalFromStringTest, RandomSmallValuesRoundTrip) { @@ -435,15 +511,18 @@ 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 - // 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")); + AssertDecimalFromString("-170141183460469231731687303715884105728", + Decimal128FromLE({0, uint64_t{1} << 63}), 39, 0); + ASSERT_RAISES(Invalid, + Decimal128::FromString("170141183460469231731687303715884105728")); + ASSERT_RAISES(Invalid, + Decimal128::FromString("-170141183460469231731687303715884105729")); + int32_t precision, scale; + ASSERT_RAISES( + Invalid, Decimal128::FromString("-170141183460469231731687303715884105729", nullptr, + &precision, &scale)); // No exponent, many fractional digits AssertDecimalFromString("9.9999999999999999999999999999999999999", dec38times9pos, 38, @@ -541,15 +620,20 @@ TEST(Decimal256Test, FromStringLimits) { ASSERT_RAISES(Invalid, Decimal256::FromString("9.9e78")); ASSERT_RAISES(Invalid, Decimal256::FromString("-9.9e78")); - // XXX conversion overflows are currently not detected - // ASSERT_RAISES(Invalid, Decimal256::FromString("99e76")); - // ASSERT_RAISES(Invalid, Decimal256::FromString("-99e76")); - // ASSERT_RAISES(Invalid, - // Decimal256::FromString("9999999999999999999999999999999999999999999999999999999999999999999999999999e1")); - // ASSERT_RAISES(Invalid, - // Decimal256::FromString("-9999999999999999999999999999999999999999999999999999999999999999999999999999e1")); - // ASSERT_RAISES(Invalid, - // Decimal256::FromString("99999999999999999999999999999999999999999999999999999999999999999999999999999")); + ASSERT_RAISES(Invalid, Decimal256::FromString(std::string(78, '9'))); + AssertDecimalFromString( + "-57896044618658097711785492504343953926634992332820282019728792003956564819968", + Decimal256FromLE({0, 0, 0, uint64_t{1} << 63}), 77, 0); + ASSERT_RAISES(Invalid, Decimal256::FromString("5789604461865809771178549250434395392663" + "4992332820282019728792003956564819968")); + ASSERT_RAISES(Invalid, + Decimal256::FromString("-5789604461865809771178549250434395392663" + "4992332820282019728792003956564819969")); + int32_t precision, scale; + ASSERT_RAISES(Invalid, + Decimal256::FromString("-5789604461865809771178549250434395392663" + "4992332820282019728792003956564819969", + nullptr, &precision, &scale)); // No exponent, many fractional digits AssertDecimalFromString( diff --git a/cpp/src/gandiva/gdv_function_stubs.cc b/cpp/src/gandiva/gdv_function_stubs.cc index 6b3e9935b017..cbf27c9af6ee 100644 --- a/cpp/src/gandiva/gdv_function_stubs.cc +++ b/cpp/src/gandiva/gdv_function_stubs.cc @@ -23,11 +23,14 @@ #include #include #include +#include #include #include "arrow/util/base64.h" #include "arrow/util/bit_util.h" +#include "arrow/util/decimal_internal.h" #include "arrow/util/double_conversion_internal.h" +#include "arrow/util/int_util_overflow.h" #include "arrow/util/value_parsing.h" #include "gandiva/encrypt_utils.h" @@ -204,11 +207,59 @@ CRC_FUNCTION(utf8) CRC_FUNCTION(binary) int32_t gdv_fn_dec_from_string(int64_t context, const char* in, int32_t in_length, - int32_t* precision_from_str, int32_t* scale_from_str, - int64_t* dec_high_from_str, uint64_t* dec_low_from_str) { + int32_t out_scale, int32_t* precision_from_str, + int32_t* scale_from_str, int64_t* dec_high_from_str, + uint64_t* dec_low_from_str) { arrow::Decimal128 dec; - auto status = arrow::Decimal128::FromString(std::string(in, in_length), &dec, - precision_from_str, scale_from_str); + const std::string_view input(in, in_length); + auto status = + arrow::Decimal128::FromString(input, &dec, precision_from_str, scale_from_str); + if (!status.ok() || + static_cast(*scale_from_str) - out_scale > arrow::Decimal128::kMaxScale) { + arrow::internal::DecimalComponents components; + int32_t input_scale = 0; + if (arrow::internal::ParseDecimalComponents(input.data(), input.size(), + &components) && + !arrow::internal::SubtractWithOverflow( + static_cast(components.fractional_digits.size()), + components.exponent, &input_scale) && + input_scale > out_scale) { + std::string digits(components.whole_digits); + digits.append(components.fractional_digits); + digits.erase(0, digits.find_first_not_of('0')); + const int64_t num_digits = + static_cast(digits.size()) + out_scale - input_scale; + if (num_digits > arrow::Decimal128::kMaxPrecision) { + gdv_fn_context_set_error_msg(context, status.message().data()); + return -1; + } + bool round_up = false; + if (num_digits < 0) { + digits = "0"; + } else { + round_up = num_digits < static_cast(digits.size()) && + digits[static_cast(num_digits)] >= '5'; + digits.resize(static_cast(num_digits), '0'); + if (digits.empty()) { + digits = "0"; + } + } + status = + arrow::Decimal128::FromString(digits, &dec, precision_from_str, scale_from_str); + if (status.ok()) { + if (round_up) { + dec += arrow::Decimal128(1); + if (dec == arrow::Decimal128::GetScaleMultiplier(*precision_from_str)) { + ++(*precision_from_str); + } + } + if (components.sign == '-') { + dec.Negate(); + } + *scale_from_str = out_scale; + } + } + } if (!status.ok()) { gdv_fn_context_set_error_msg(context, status.message().data()); return -1; @@ -948,6 +999,7 @@ arrow::Status ExportedStubFunctions::AddMappings(Engine* engine) const { types->i64_type(), // context types->i8_ptr_type(), // const char* in types->i32_type(), // int32_t in_length + types->i32_type(), // int32_t out_scale types->i32_ptr_type(), // int32_t* precision_from_str types->i32_ptr_type(), // int32_t* scale_from_str types->i64_ptr_type(), // int64_t* dec_high_from_str diff --git a/cpp/src/gandiva/gdv_function_stubs.h b/cpp/src/gandiva/gdv_function_stubs.h index 4113f261ad76..9c9ad240c278 100644 --- a/cpp/src/gandiva/gdv_function_stubs.h +++ b/cpp/src/gandiva/gdv_function_stubs.h @@ -128,8 +128,9 @@ const char* gdv_fn_sha1_decimal128(int64_t context, int64_t x_high, uint64_t x_l gdv_boolean x_isvalid, int32_t* out_length); int32_t gdv_fn_dec_from_string(int64_t context, const char* in, int32_t in_length, - int32_t* precision_from_str, int32_t* scale_from_str, - int64_t* dec_high_from_str, uint64_t* dec_low_from_str); + int32_t out_scale, int32_t* precision_from_str, + int32_t* scale_from_str, int64_t* dec_high_from_str, + uint64_t* dec_low_from_str); char* gdv_fn_dec_to_string(int64_t context, int64_t x_high, uint64_t x_low, int32_t x_scale, int32_t* dec_str_len); diff --git a/cpp/src/gandiva/precompiled/decimal_wrapper.cc b/cpp/src/gandiva/precompiled/decimal_wrapper.cc index f232f35e4e29..355d2509b53c 100644 --- a/cpp/src/gandiva/precompiled/decimal_wrapper.cc +++ b/cpp/src/gandiva/precompiled/decimal_wrapper.cc @@ -403,8 +403,8 @@ void castDECIMAL_utf8(int64_t context, const char* in, int32_t in_length, int32_t precision_from_str; int32_t scale_from_str; int32_t status = - gdv_fn_dec_from_string(context, in, in_length, &precision_from_str, &scale_from_str, - &dec_high_from_str, &dec_low_from_str); + gdv_fn_dec_from_string(context, in, in_length, out_scale, &precision_from_str, + &scale_from_str, &dec_high_from_str, &dec_low_from_str); if (status != 0) { *out_high = 0; *out_low = 0; diff --git a/cpp/src/gandiva/tests/decimal_test.cc b/cpp/src/gandiva/tests/decimal_test.cc index 043bdc4605a7..138463a79c98 100644 --- a/cpp/src/gandiva/tests/decimal_test.cc +++ b/cpp/src/gandiva/tests/decimal_test.cc @@ -16,6 +16,7 @@ // under the License. #include +#include #include #include @@ -1083,18 +1084,74 @@ TEST_F(TestDecimal, TestCastDecimalVarCharInvalidInput) { // Create a row-batch with some sample data int num_records = 5; - // invalid input - auto invalid_in = MakeArrowArrayUtf8({"a10.5134", "-0.0", "-0.1", "10.516", "-1000"}, - {true, false, true, true, true}); - - // prepare input record batch - auto in_batch_1 = arrow::RecordBatch::Make(schema, num_records, {invalid_in}); + const std::string invalid_number = "not a valid decimal128 number"; + const std::string out_of_range = "cannot be represented as decimal128"; + for (const auto& [invalid, expected_error] : + std::vector>{ + {"a10.5134", invalid_number}, + {"1." + std::string(100, '5') + "x", invalid_number}, + {"1." + std::string(100, '5') + "e2147483648", invalid_number}, + {std::string(40, '9'), out_of_range}, + {std::string(40, '9') + ".1", out_of_range}, + {std::string(50, '9') + "." + std::string(40, '5'), out_of_range}, + {"1e39", out_of_range}, + {"1e2147483647", out_of_range}, + {"1e-2147483648", out_of_range}, + {"1.0e-2147483647", out_of_range}, + {"99e38", out_of_range}, + {std::string(50, '9'), out_of_range}}) { + SCOPED_TRACE(invalid); + auto invalid_in = MakeArrowArrayUtf8({invalid, "-0.0", "-0.1", "10.516", "-1000"}, + {true, false, true, true, true}); + auto in_batch = arrow::RecordBatch::Make(schema, num_records, {invalid_in}); + arrow::ArrayVector outputs; + status = projector->Evaluate(*in_batch, pool_, &outputs); + EXPECT_FALSE(status.ok()) << status.message(); + EXPECT_NE(status.message().find(expected_error), std::string::npos); + } +} - // Evaluate expression - arrow::ArrayVector outputs_1; - status = projector->Evaluate(*in_batch_1, pool_, &outputs_1); - EXPECT_FALSE(status.ok()) << status.message(); - EXPECT_NE(status.message().find("not a valid decimal128 number"), std::string::npos); +TEST_F(TestDecimal, TestCastDecimalVarCharLongInputs) { + struct TestCase { + std::string input; + int32_t precision; + int32_t scale; + std::string expected; + }; + for (const auto& test_case : std::vector{ + {"1." + std::string(50, '5'), 38, 37, "1." + std::string(36, '5') + "6"}, + {"-1." + std::string(50, '5'), 38, 37, "-1." + std::string(36, '5') + "6"}, + {"1." + std::string(100, '1'), 38, 37, "1." + std::string(37, '1')}, + {"15." + std::string(50, '5') + "e-1", 38, 37, + "1." + std::string(36, '5') + "6"}, + {"0." + std::string(37, '0') + "5" + std::string(70, '0'), 38, 37, + "0." + std::string(36, '0') + "1"}, + {"0." + std::string(100, '0') + "5", 38, 37, "0"}, + {"1e-30", 38, 0, "0"}, + {"9." + std::string(80, '9'), 4, 2, "10.00"}, + {"-9." + std::string(80, '9'), 4, 2, "-10.00"}, + {"99." + std::string(80, '9'), 4, 2, "0.00"}, + {"0001." + std::string(80, '0'), 4, 2, "1.00"}, + {"1." + std::string(100, '0') + "e+2", 6, 2, "100.00"}, + {"1234" + std::string(100, '0') + "e-102", 4, 2, "12.34"}}) { + SCOPED_TRACE(test_case.input); + auto decimal_type = arrow::decimal128(test_case.precision, test_case.scale); + auto field_str = field("in_str", utf8()); + auto schema = arrow::schema({field_str}); + auto expr = TreeExprBuilder::MakeExpression("castDECIMAL", {field_str}, + field("out", decimal_type)); + std::shared_ptr projector; + ASSERT_OK(Projector::Make(schema, {expr}, TestConfiguration(), &projector)); + + auto input = MakeArrowArrayUtf8({test_case.input, ""}, {true, false}); + auto batch = arrow::RecordBatch::Make(schema, input->length(), {input}); + arrow::ArrayVector outputs; + ASSERT_OK(projector->Evaluate(*batch, pool_, &outputs)); + auto expected = MakeArrowArrayDecimal( + decimal_type, MakeDecimalVector({test_case.expected, "0"}, test_case.scale), + {true, false}); + EXPECT_ARROW_ARRAY_EQUALS(expected, outputs[0]); + } } TEST_F(TestDecimal, TestVarCharDecimalNestedCast) {