mirror of
https://github.com/Canner/WrenAI.git
synced 2026-09-24 23:29:49 +08:00
fix(trino): treat DECIMAL(p) as scale 0, not a default non-zero scale (#2404)
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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")),
|
||||
|
||||
@@ -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())
|
||||
|
||||
|
||||
Reference in New Issue
Block a user