diff --git a/core/wren/src/wren/connector/trino.py b/core/wren/src/wren/connector/trino.py index 04aa8870b..2ec607205 100644 --- a/core/wren/src/wren/connector/trino.py +++ b/core/wren/src/wren/connector/trino.py @@ -101,7 +101,12 @@ def _trino_data_type_to_arrow(node) -> pa.DataType: return _TRINO_DATA_TYPE_TO_ARROW[kind] if kind == T.DECIMAL: - precision, scale = 38, 9 + # Trino DECIMAL semantics: DECIMAL(p) is precision p with scale 0, and + # bare DECIMAL is DECIMAL(38, 0). Defaulting the scale to a non-zero + # value mistyped precision-only DECIMAL columns and, for small + # precisions (e.g. DECIMAL(5)), produced scale > precision, which + # PyArrow rejects with an ArrowInvalid. + precision, scale = 38, 0 params = node.expressions if len(params) >= 1: with contextlib.suppress(AttributeError, ValueError): diff --git a/core/wren/tests/connectors/test_trino.py b/core/wren/tests/connectors/test_trino.py index a917c3741..ce1714416 100644 --- a/core/wren/tests/connectors/test_trino.py +++ b/core/wren/tests/connectors/test_trino.py @@ -64,7 +64,10 @@ _SCHEMA = "default" # Decimal ("decimal(10,2)", pa.decimal128(10, 2)), ("decimal(38,9)", pa.decimal128(38, 9)), - ("decimal", pa.decimal128(38, 9)), + # DECIMAL(p) is scale 0 and bare DECIMAL is DECIMAL(38, 0) in Trino. + ("decimal", pa.decimal128(38, 0)), + ("decimal(10)", pa.decimal128(10, 0)), + ("decimal(5)", pa.decimal128(5, 0)), # Time / timestamp ("time", pa.time64("us")), ("time(3)", pa.time64("us")), diff --git a/core/wren/tests/unit/test_athena_connector.py b/core/wren/tests/unit/test_athena_connector.py index d67c0f2a0..6153397d5 100644 --- a/core/wren/tests/unit/test_athena_connector.py +++ b/core/wren/tests/unit/test_athena_connector.py @@ -74,8 +74,20 @@ def test_parse_athena_type_primitives(type_str, expected): assert _parse_athena_type(type_str) == expected -def test_parse_athena_type_decimal(): - assert _parse_athena_type("decimal(12,4)") == pa.decimal128(12, 4) +@pytest.mark.parametrize( + ("type_str", "expected"), + [ + ("decimal(12,4)", pa.decimal128(12, 4)), + # bare DECIMAL and precision-only DECIMAL(p) default to scale 0 + # (kept aligned with the Trino parser). + ("decimal", pa.decimal128(38, 0)), + ("decimal(10)", pa.decimal128(10, 0)), + # small precision must not produce scale > precision (ArrowInvalid). + ("decimal(5)", pa.decimal128(5, 0)), + ], +) +def test_parse_athena_type_decimal(type_str, expected): + assert _parse_athena_type(type_str) == expected def test_parse_athena_type_decimal_precision_only_is_scale_zero(): @@ -94,6 +106,21 @@ def test_parse_athena_type_decimal_small_precision_only(): assert _parse_athena_type("decimal(5)") == pa.decimal128(5, 0) +def test_parse_athena_type_decimal_bare_defaults_to_scale_0(): + # Bare DECIMAL is DECIMAL(38, 0) in Athena/Trino semantics. + assert _parse_athena_type("decimal") == pa.decimal128(38, 0) + + +def test_parse_athena_type_decimal_precision_only_has_scale_0(): + # DECIMAL(p) is precision p with scale 0, not a default non-zero scale. + assert _parse_athena_type("decimal(10)") == pa.decimal128(10, 0) + + +def test_parse_athena_type_decimal_small_precision_only(): + # Small precision-only case must not produce scale > precision (ArrowInvalid). + assert _parse_athena_type("decimal(5)") == pa.decimal128(5, 0) + + def test_parse_athena_type_array(): assert _parse_athena_type("array(varchar)") == pa.list_(pa.string())