diff --git a/aws_lambda_powertools/shared/dynamodb_deserializer.py b/aws_lambda_powertools/shared/dynamodb_deserializer.py index d90e0a47554..6ff6e67f545 100644 --- a/aws_lambda_powertools/shared/dynamodb_deserializer.py +++ b/aws_lambda_powertools/shared/dynamodb_deserializer.py @@ -15,6 +15,10 @@ traps=[Clamped, Overflow, Inexact, Rounded, Underflow], ) +# Allow removing trailing zeros, but reject any loss of precision. +_NORMALIZE_CONTEXT = Context(prec=38) +_NORMALIZE_CONTEXT.traps[Inexact] = True + class TypeDeserializer: """ @@ -75,20 +79,18 @@ def _deserialize_bool(self, value: bool) -> bool: return value def _deserialize_n(self, value: str) -> Decimal: - # value is None or "."? It's zero - # then return early - value = value.lstrip("0") - if not value or value == ".": + # Preserve the existing fallback for empty strings and a lone decimal point. + if value in ("", "."): return DYNAMODB_CONTEXT.create_decimal(0) - if len(value) > 38: + number = Decimal(value) + if len(number.as_tuple().digits) > 38: # See: https://docs.aws.amazon.com/amazondynamodb/latest/developerguide/HowItWorks.NamingRulesDataTypes.html#HowItWorks.DataTypes.Number - # Calculate the number of trailing zeros after the 38th character - tail = len(value[38:]) - len(value[38:].rstrip("0")) - # Trim the value: remove trailing zeros if any, or just take the first 38 characters - value = value[:-tail] if tail > 0 else value[:38] + # Trailing zeros don't count towards the 38 digits of precision, so move them + # into the exponent. Cutting characters off the string would change the value. + number = number.normalize(_NORMALIZE_CONTEXT) - return DYNAMODB_CONTEXT.create_decimal(value) + return DYNAMODB_CONTEXT.create_decimal(number) def _deserialize_s(self, value: str) -> str: return value diff --git a/tests/unit/data_classes/required_dependencies/test_dynamo_db_stream_event.py b/tests/unit/data_classes/required_dependencies/test_dynamo_db_stream_event.py index 02fdab9582e..ebf370c3f25 100644 --- a/tests/unit/data_classes/required_dependencies/test_dynamo_db_stream_event.py +++ b/tests/unit/data_classes/required_dependencies/test_dynamo_db_stream_event.py @@ -2,6 +2,8 @@ from decimal import Clamped, Context, Inexact, Overflow, Rounded, Underflow +import pytest + from aws_lambda_powertools.utilities.data_classes.dynamo_db_stream_event import ( DynamoDBRecordEventName, DynamoDBStreamEvent, @@ -93,23 +95,23 @@ def test_dynamodb_stream_record_deserialization_large_int(): record = StreamRecord(data) assert record.new_image == { "Name": "Joe", - "Age": DECIMAL_CONTEXT.create_decimal("11011111111111111000000000000000000000"), + "Age": DECIMAL_CONTEXT.create_decimal("1.1011111111111111E+46"), } -def test_dynamodb_stream_record_deserialization_large_int_without_trailing_zeros(): +@pytest.mark.parametrize("image", ["Keys", "NewImage", "OldImage"]) +def test_dynamodb_stream_record_deserialization_rejects_inexact_numbers(image): data = { - "Keys": {"key1": {"attr1": "value1"}}, - "NewImage": { + image: { "Name": {"S": "Joe"}, "Age": {"N": "000000011011111111111112222222222221111111111111111111111"}, }, } record = StreamRecord(data) - assert record.new_image == { - "Name": "Joe", - "Age": DECIMAL_CONTEXT.create_decimal("11011111111111112222222222221111111111"), - } + attribute = {"Keys": "keys", "NewImage": "new_image", "OldImage": "old_image"}[image] + + with pytest.raises(Inexact): + getattr(record, attribute) def test_dynamodb_stream_record_deserialization_zero_value(): diff --git a/tests/unit/shared/test_dynamodb_deserializer.py b/tests/unit/shared/test_dynamodb_deserializer.py index 7e6c2cc3885..6ab7eb82d56 100644 --- a/tests/unit/shared/test_dynamodb_deserializer.py +++ b/tests/unit/shared/test_dynamodb_deserializer.py @@ -1,5 +1,6 @@ from __future__ import annotations +from decimal import ROUND_UP, Decimal, Inexact, Overflow, Rounded, localcontext from typing import Any import pytest @@ -51,3 +52,89 @@ def test_deserializer_error(): with pytest.raises(TypeError): model.data.get("Id") + + +@pytest.mark.parametrize( + "value", + [ + "-12345678901234567890123456789012345678", + "1.2345678901234567890123456789012345678", + "110111111111111110000000000000000000000", + "12345678901234567890123456789012345678000", + ], +) +def test_deserializer_keeps_value_of_numbers_with_38_digits_of_precision(value): + assert TypeDeserializer().deserialize({"N": value}) == Decimal(value) + + +@pytest.mark.parametrize( + "value", + [ + "9" * 39, + "-" + "9" * 39, + "1" * 37 + "25", + "1" * 37 + "35", + "1." + "2" * 38, + "123456789012345678901234567890123456789000", + "123456789012345678901234567890123456789E-100", + "1" * 39 + "E-9999999", + ], +) +def test_deserializer_rejects_inexact_numbers(value): + with pytest.raises(Inexact): + TypeDeserializer().deserialize({"N": value}) + + +@pytest.mark.parametrize("value", ["0E+0", "0E-100", "0e-130", "000E+5", "+0E-10", "-0E-100"]) +def test_deserializer_preserves_scientific_zero(value): + assert TypeDeserializer().deserialize({"N": value}) == Decimal(0) + + +@pytest.mark.parametrize("value", ["", ".", "0", "000", "000.", "0.000", "-0", "+0"]) +def test_deserializer_preserves_existing_zero_handling(value): + assert TypeDeserializer().deserialize({"N": value}) == Decimal(0) + + +@pytest.mark.parametrize( + "value", + [ + "00012345678901234567890123456789012345678000", + "-12345678901234567890123456789012345678000", + "1.2345678901234567890123456789012345678000", + "12345678901234567890123456789012345678000E-40", + "1000000000000000000000000000000000000000E-167", + "99999999999999999999999999999999999999000E+85", + ], +) +def test_deserializer_removes_only_exact_trailing_zeros(value): + # Application Decimal settings must not affect deserialization. + with localcontext() as context: + context.prec = 6 + context.rounding = ROUND_UP + context.traps[Inexact] = True + context.traps[Rounded] = True + assert TypeDeserializer().deserialize({"N": value}) == Decimal(value) + + +@pytest.mark.parametrize( + "value", + [ + {"NS": ["1", "9" * 39]}, + {"L": [{"N": "9" * 39}]}, + {"M": {"amount": {"N": "9" * 39}}}, + ], +) +def test_deserializer_rejects_inexact_numbers_in_collections(value): + with pytest.raises(Inexact): + TypeDeserializer().deserialize(value) + + +def test_deserializer_keeps_distinct_number_set_values(): + values = ["1.2345678901234567890123456789012345677", "1.2345678901234567890123456789012345678"] + + assert TypeDeserializer().deserialize({"NS": values}) == {Decimal(value) for value in values} + + +def test_deserializer_normalization_preserves_overflow_errors(): + with pytest.raises(Overflow): + TypeDeserializer().deserialize({"N": "1" * 38 + "00E+9999999"})