diff --git a/awswrangler/_data_types.py b/awswrangler/_data_types.py index 432b3f72e..3b54640c9 100644 --- a/awswrangler/_data_types.py +++ b/awswrangler/_data_types.py @@ -344,7 +344,7 @@ def athena2pyarrow(dtype: str, df_type: str | None = None) -> pa.DataType: # no return pa.timestamp(unit="ns") if dtype == "date": return pa.date32() - if dtype in ("binary" or "varbinary"): + if dtype in ("binary", "varbinary"): return pa.binary() if dtype.startswith("decimal") is True: precision, scale = dtype.replace("decimal(", "").replace(")", "").split(sep=",") diff --git a/tests/unit/test_data_types.py b/tests/unit/test_data_types.py new file mode 100644 index 000000000..233ad39fc --- /dev/null +++ b/tests/unit/test_data_types.py @@ -0,0 +1,29 @@ +import pyarrow as pa +import pytest + +from awswrangler._data_types import athena2pandas, athena2pyarrow +from awswrangler.exceptions import UnsupportedType + + +@pytest.mark.parametrize( + "dtype,expected", + [ + ("binary", pa.binary()), + ("varbinary", pa.binary()), + ("BINARY", pa.binary()), + ("VARBINARY", pa.binary()), + ], +) +def test_athena2pyarrow_binary_types(dtype, expected): + assert athena2pyarrow(dtype) == expected + + +@pytest.mark.parametrize("dtype", ["i", "n", "ary", "bin"]) +def test_athena2pyarrow_rejects_binary_substrings(dtype): + with pytest.raises(UnsupportedType, match=f"Unsupported Athena type: {dtype}"): + athena2pyarrow(dtype) + + +@pytest.mark.parametrize("dtype", ["binary", "varbinary"]) +def test_athena2pandas_binary_types(dtype): + assert athena2pandas(dtype) == "bytes"