diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index 6e285f19..b6a287a2 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -1365,60 +1365,56 @@ events = Table('events', metadata, #### Querying JSON data -When querying JSON data, PyAthena automatically parses JSON strings into Python dictionaries: +PyAthena's default converters decode results of Athena's `json` type, and the `JSON` type returns those values unchanged. +A JSON object becomes a `dict`, an array a `list`, and a string scalar a `str`. +If the cursor returns the text of a `json` result instead, for example through a custom converter that does not decode it, the `JSON` type returns that text unchanged. ```python from sqlalchemy import select, literal_column from sqlalchemy.sql import type_coerce from sqlalchemy.types import JSON -# Query with explicit type coercion result = connection.execute( select( type_coerce( - literal_column('CAST(\'{"name": "test", "value": 123}\' AS JSON)'), + literal_column('json_parse(\'{"name": "test", "items": [1, 2, 3]}\')'), JSON ).label("json_col") ) ).fetchone() -# Result is automatically parsed as a dictionary -print(result.json_col) # {"name": "test", "value": 123} +print(result.json_col) # {'items': [1, 2, 3], 'name': 'test'} print(type(result.json_col)) # ``` -#### Important limitations - -Athena's JSON type support has specific limitations: - -- **JSON objects and arrays are supported** - `CAST('...' AS JSON)` accepts an object or a top-level array such as `[1, 2, 3]` -- **Arrays within objects are supported** - JSON objects can contain arrays as property values -- **DML only** - JSON type is supported for SELECT queries but not in CREATE TABLE statements; compiling `CREATE TABLE` with a `JSON` column raises `CompileError` +`json_parse()` parses text into a JSON value, while `CAST('...' AS JSON)` returns the text as a JSON string scalar: ```python -# Supported: JSON object with nested array result = connection.execute( select( - type_coerce( - literal_column('CAST(\'{"items": [1, 2, 3]}\' AS JSON)'), - JSON - ).label("json_col") + type_coerce(literal_column("json_parse('[1, 2, 3]')"), JSON).label("parsed"), + type_coerce(literal_column("CAST('[1, 2, 3]' AS JSON)"), JSON).label("cast"), ) ).fetchone() -print(result.json_col) # {"items": [1, 2, 3]} -# Supported: Top-level array +print(result.parsed) # [1, 2, 3] +print(result.cast) # '[1, 2, 3]' +``` + +The `JSON` type decodes JSON text in columns of other types, such as a `varchar` column, with `json.loads()` or the dialect's `json_deserializer`: + +```python result = connection.execute( - select( - type_coerce( - literal_column("CAST('[1, 2, 3]' AS JSON)"), - JSON - ).label("json_col") - ) + select(type_coerce(literal_column("'{\"a\": 1}'"), JSON).label("json_col")) ).fetchone() -print(result.json_col) # [1, 2, 3] + +print(result.json_col) # {'a': 1} ``` +#### Important limitations + +- **DML only** - JSON type is supported for SELECT queries but not in CREATE TABLE statements; compiling `CREATE TABLE` with a `JSON` column raises `CompileError` + #### Best practices 1. **Use with SELECT queries** - JSON type works best for querying existing data diff --git a/pyathena/arrow/converter.py b/pyathena/arrow/converter.py index ea648255..88943efb 100644 --- a/pyathena/arrow/converter.py +++ b/pyathena/arrow/converter.py @@ -9,11 +9,11 @@ from pyathena.converter import ( Converter, + _csv_to_json, _to_binary, _to_date, _to_decimal, _to_default, - _to_json, _to_time, ) from pyathena.util import override @@ -26,7 +26,7 @@ "time": _to_time, "decimal": _to_decimal, "varbinary": _to_binary, - "json": _to_json, + "json": _csv_to_json, } diff --git a/pyathena/arrow/result_set.py b/pyathena/arrow/result_set.py index 77978264..abdc66ab 100644 --- a/pyathena/arrow/result_set.py +++ b/pyathena/arrow/result_set.py @@ -12,7 +12,7 @@ from pyathena import OperationalError from pyathena.arrow.util import to_column_info -from pyathena.converter import Converter +from pyathena.converter import Converter, _json_text_converter, _to_default from pyathena.error import ProgrammingError from pyathena.model import AthenaQueryExecution from pyathena.result_set import AthenaResultSet @@ -146,8 +146,8 @@ def __init__( import pyarrow as pa self._table = pa.Table.from_pydict({}) - # The fetch methods convert only the values read from a result file. - # GetQueryResults values are already converted. + # The fetch methods convert the values read from a result file. GetQueryResults + # values are already converted, except json values, which stay text. self._convert_rows = bool(self.output_location) self._batches = iter(self._table.to_batches(arraysize)) @@ -254,11 +254,16 @@ def _fetch(self) -> None: return else: dict_rows = rows.to_pydict() - if self._convert_rows: - converters = self.converters + converters = ( + self.converters if self._convert_rows else self._json_converters(self.converters) + ) + if converters: column_names = dict_rows.keys() processed_rows = [ - tuple(converters[k](v) for k, v in zip(column_names, row, strict=False)) + tuple( + converters.get(k, _to_default)(v) + for k, v in zip(column_names, row, strict=False) + ) for row in zip(*dict_rows.values(), strict=False) ] else: @@ -381,11 +386,12 @@ def _as_arrow_from_api(self, converter: Converter | None = None) -> Table: Args: converter: Type converter for result values. Defaults to - ``DefaultTypeConverter`` if not specified. + ``DefaultTypeConverter`` with json values kept as text, as in + the CSV result file. """ import pyarrow as pa - rows = self._fetch_all_rows(converter) + rows = self._fetch_all_rows(converter or _json_text_converter()) if not rows: return pa.Table.from_pydict({}) description = self.description if self.description else [] diff --git a/pyathena/converter.py b/pyathena/converter.py index 3ea1577a..eff98742 100644 --- a/pyathena/converter.py +++ b/pyathena/converter.py @@ -124,6 +124,21 @@ def _to_json(varchar_value: str | None) -> Any | None: return json.loads(varchar_value) +def _csv_to_json(varchar_value: str | None) -> Any | None: + """Convert an Athena JSON value read from a CSV result file. + + Args: + varchar_value: The value as JSON text, or None. An empty string, which + CSV results use for SQL NULL, is also treated as NULL. + + Returns: + The decoded value, or None for SQL NULL. + """ + if not varchar_value: + return None + return json.loads(varchar_value) + + def _to_array(varchar_value: str | None) -> list[Any] | str | None: """Convert array data to Python list. @@ -727,3 +742,14 @@ def _parse_type_hint(self, type_hint: str) -> TypeNode: if normalized not in self._parsed_hints: self._parsed_hints[normalized] = self._parser.parse(normalized) return self._parsed_hints[normalized] + + +def _json_text_converter() -> DefaultTypeConverter: + """Return a ``DefaultTypeConverter`` that keeps json values as text. + + Returns: + The converter. + """ + converter = DefaultTypeConverter() + converter.set("json", _to_default) + return converter diff --git a/pyathena/pandas/converter.py b/pyathena/pandas/converter.py index a8097587..a3b207e9 100644 --- a/pyathena/pandas/converter.py +++ b/pyathena/pandas/converter.py @@ -9,11 +9,11 @@ from pyathena.converter import ( Converter, + _csv_to_json, _to_binary, _to_boolean, _to_decimal, _to_default, - _to_json, ) from pyathena.util import override @@ -24,7 +24,7 @@ "boolean": _to_boolean, "decimal": _to_decimal, "varbinary": _to_binary, - "json": _to_json, + "json": _csv_to_json, } diff --git a/pyathena/polars/result_set.py b/pyathena/polars/result_set.py index 03985075..34dbebe4 100644 --- a/pyathena/polars/result_set.py +++ b/pyathena/polars/result_set.py @@ -20,7 +20,7 @@ ) from pyathena import OperationalError -from pyathena.converter import Converter +from pyathena.converter import Converter, _json_text_converter from pyathena.error import ProgrammingError from pyathena.model import AthenaQueryExecution from pyathena.polars.util import to_column_info @@ -262,7 +262,7 @@ def __init__( # Note: _as_polars() and _create_dataframe_iterator() update _metadata for unload # queries, so the converters and column names must be read AFTER them. self._df: pl.DataFrame | None = None - # Converters for the rows of self._df. GetQueryResults values are already converted. + # Converters for the rows of self._df. self._df_converters: dict[str, Callable[[str | None], Any | None]] = {} if self.state == AthenaQueryExecution.STATE_SUCCEEDED and self.output_location: if self._chunksize is None: @@ -272,6 +272,8 @@ def __init__( self._df_iter = self._create_dataframe_iterator() elif self.state == AthenaQueryExecution.STATE_SUCCEEDED: self._df = self._as_polars_from_api() + # GetQueryResults values are already converted, except json values kept as text. + self._df_converters = self._json_converters(self.converters) else: self._df = pl.DataFrame() if self._df is not None: @@ -539,11 +541,12 @@ def _as_polars_from_api(self, converter: Converter | None = None) -> pl.DataFram Args: converter: Type converter for result values. Defaults to - ``DefaultTypeConverter`` if not specified. + ``DefaultTypeConverter`` with json values kept as text, as in + the CSV result file. """ import polars as pl - rows = self._fetch_all_rows(converter) + rows = self._fetch_all_rows(converter or _json_text_converter()) if not rows: return pl.DataFrame() description = self.description if self.description else [] diff --git a/pyathena/result_set.py b/pyathena/result_set.py index 410c471b..65842761 100644 --- a/pyathena/result_set.py +++ b/pyathena/result_set.py @@ -18,6 +18,8 @@ from pyathena.util import RetryConfig, override, parse_output_location, retry_api_call if TYPE_CHECKING: + from collections.abc import Callable + from pyathena.connection import Connection _logger = logging.getLogger(__name__) @@ -679,6 +681,20 @@ def _is_first_row_column_labels(self, rows: list[dict[str, Any]]) -> bool: return False return True + def _json_converters( + self, converters: dict[str, Callable[[str | None], Any | None]] + ) -> dict[str, Callable[[str | None], Any | None]]: + """Select the converters of the json columns. + + Args: + converters: The converters keyed by column name. + + Returns: + The converters of the columns whose Athena type is json. + """ + description = self.description if self.description else [] + return {d[0]: converters[d[0]] for d in description if d[1] == "json"} + def _fetch_all_rows( self, converter: Converter | None = None, diff --git a/pyathena/s3fs/result_set.py b/pyathena/s3fs/result_set.py index 74dff6c7..0bc268fc 100644 --- a/pyathena/s3fs/result_set.py +++ b/pyathena/s3fs/result_set.py @@ -120,8 +120,9 @@ def __init__( if self.state == AthenaQueryExecution.STATE_SUCCEEDED and self.output_location: self._init_csv_reader() elif self.state == AthenaQueryExecution.STATE_SUCCEEDED: - # Managed query result storage: no output_location, use API - rows = self._fetch_all_rows() + # Managed query result storage: no output_location, use API. + # The converter reads text values, as from the CSV result file. + rows = self._fetch_all_rows(self._converter) self._rows.extend(rows) # If CSV reader was not initialized (e.g., CTAS, DDL), diff --git a/pyathena/sqlalchemy/base.py b/pyathena/sqlalchemy/base.py index d620fce6..f408fdf0 100644 --- a/pyathena/sqlalchemy/base.py +++ b/pyathena/sqlalchemy/base.py @@ -38,6 +38,7 @@ AthenaArray, AthenaBinary, AthenaDate, + AthenaJSON, AthenaMap, AthenaStruct, AthenaTimestamp, @@ -212,6 +213,7 @@ class AthenaDialect(DefaultDialect): types.ARRAY: AthenaArray, types.Date: AthenaDate, types.DateTime: AthenaTimestamp, + types.JSON: AthenaJSON, } ischema_names: dict[str, type[Any]] = ischema_names diff --git a/pyathena/sqlalchemy/types.py b/pyathena/sqlalchemy/types.py index 07c8e69d..74872d11 100644 --- a/pyathena/sqlalchemy/types.py +++ b/pyathena/sqlalchemy/types.py @@ -6,7 +6,7 @@ from __future__ import annotations -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any from sqlalchemy import types from sqlalchemy.sql import sqltypes @@ -19,7 +19,7 @@ if TYPE_CHECKING: from sqlalchemy import Dialect - from sqlalchemy.sql.type_api import _LiteralProcessorType + from sqlalchemy.sql.type_api import _LiteralProcessorType, _ResultProcessorType __all__ = [ "ARRAY", @@ -29,6 +29,7 @@ "AthenaArray", "AthenaBinary", "AthenaDate", + "AthenaJSON", "AthenaMap", "AthenaStruct", "AthenaTimestamp", @@ -47,6 +48,37 @@ def process(value: bytes) -> str: return process +class AthenaJSON(types.JSON): + """SQLAlchemy JSON type that keeps the values PyAthena has decoded. + + PyAthena's default converters decode results of the Athena ``json`` type, + so this type returns them unchanged, and a JSON string scalar stays a + ``str``. If the cursor returns the text of a ``json`` result instead, for + example through a custom converter that does not decode it, this type + returns that text unchanged. Results of other Athena types, such as JSON + text in a ``varchar`` column, are decoded with the dialect's JSON + deserializer. + """ + + @override + def result_processor( + self, dialect: Dialect, coltype: object + ) -> _ResultProcessorType[Any] | None: + """Return a processor decoding JSON text of Athena types other than ``json``. + + Args: + dialect: The dialect fetching the value. + coltype: The Athena type name from the cursor description. + + Returns: + The processor, or None for the Athena ``json`` type. + """ + if coltype == "json": + return None + processor: _ResultProcessorType[Any] | None = super().result_processor(dialect, coltype) + return processor + + class Tinyint(sqltypes.Integer): """SQLAlchemy type for Athena TINYINT (8-bit signed integer). diff --git a/tests/pyathena/arrow/test_cursor.py b/tests/pyathena/arrow/test_cursor.py index f70e99f6..df58ada7 100644 --- a/tests/pyathena/arrow/test_cursor.py +++ b/tests/pyathena/arrow/test_cursor.py @@ -977,8 +977,15 @@ def test_fetch_all_rows(self, arrow_cursor): ,X'0102' AS col_varbinary ,json_parse('{"a": 1}') AS col_json ,CAST('{"a": 1}' AS JSON) AS col_json_string + ,CAST(NULL AS JSON) AS col_json_null + UNION ALL + SELECT + 2, CAST('12:34:56' AS TIME), X'0102', json_parse('[1, "x"]'), json_parse('"s"'), NULL + ORDER BY col """ ) assert arrow_cursor.fetchall() == [ - (1, datetime(2017, 1, 1, 12, 34, 56).time(), b"\x01\x02", {"a": 1}, '{"a": 1}') + (1, datetime(2017, 1, 1, 12, 34, 56).time(), b"\x01\x02", {"a": 1}, '{"a": 1}', None), + (2, datetime(2017, 1, 1, 12, 34, 56).time(), b"\x01\x02", [1, "x"], "s", None), ] + assert arrow_cursor.as_arrow().schema.field("col_json").type == pa.string() diff --git a/tests/pyathena/pandas/test_cursor.py b/tests/pyathena/pandas/test_cursor.py index 3d0f50e5..db5c55c7 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -1590,5 +1590,14 @@ def test_pandas_cursor_iter_chunks_consistency(self, pandas_cursor): indirect=["pandas_cursor"], ) def test_fetch_all_rows(self, pandas_cursor): - pandas_cursor.execute("SELECT 1 AS col, CAST('12:34:56' AS TIME) AS col_time") - assert pandas_cursor.fetchall() == [(1, datetime(2017, 1, 1, 12, 34, 56).time())] + pandas_cursor.execute( + "SELECT 1 AS col, CAST('12:34:56' AS TIME) AS col_time, X'0001' AS col_binary, " + "json_parse('{\"a\": 1}') AS col_json, CAST('[1, 2]' AS JSON) AS col_json_string, " + "CAST(NULL AS JSON) AS col_json_null " + "UNION ALL SELECT 2, CAST('12:34:56' AS TIME), X'0001', json_parse('[1, \"x\"]'), " + "json_parse('\"s\"'), NULL ORDER BY col" + ) + assert pandas_cursor.fetchall() == [ + (1, datetime(2017, 1, 1, 12, 34, 56).time(), b"\x00\x01", {"a": 1}, "[1, 2]", None), + (2, datetime(2017, 1, 1, 12, 34, 56).time(), b"\x00\x01", [1, "x"], "s", None), + ] diff --git a/tests/pyathena/polars/test_cursor.py b/tests/pyathena/polars/test_cursor.py index c9e30da0..c8727f03 100644 --- a/tests/pyathena/polars/test_cursor.py +++ b/tests/pyathena/polars/test_cursor.py @@ -751,5 +751,15 @@ def test_null_vs_empty_string(self, polars_cursor): indirect=["polars_cursor"], ) def test_fetch_all_rows(self, polars_cursor): - polars_cursor.execute("SELECT 1 AS col, CAST('12:34:56' AS TIME) AS col_time") - assert polars_cursor.fetchall() == [(1, datetime(2017, 1, 1, 12, 34, 56).time())] + polars_cursor.execute( + "SELECT 1 AS col, CAST('12:34:56' AS TIME) AS col_time, X'0001' AS col_binary, " + "json_parse('{\"a\": 1}') AS col_json, CAST('[1, 2]' AS JSON) AS col_json_string, " + "CAST(NULL AS JSON) AS col_json_null " + "UNION ALL SELECT 2, CAST('12:34:56' AS TIME), X'0001', json_parse('[1, \"x\"]'), " + "json_parse('\"s\"'), NULL ORDER BY col" + ) + assert polars_cursor.fetchall() == [ + (1, datetime(2017, 1, 1, 12, 34, 56).time(), b"\x00\x01", {"a": 1}, "[1, 2]", None), + (2, datetime(2017, 1, 1, 12, 34, 56).time(), b"\x00\x01", [1, "x"], "s", None), + ] + assert polars_cursor.as_polars()["col_json"].dtype == pl.String diff --git a/tests/pyathena/s3fs/test_cursor.py b/tests/pyathena/s3fs/test_cursor.py index 777303f4..1d12d704 100644 --- a/tests/pyathena/s3fs/test_cursor.py +++ b/tests/pyathena/s3fs/test_cursor.py @@ -8,7 +8,9 @@ import pytest +from pyathena.converter import _to_default from pyathena.error import DatabaseError, ProgrammingError +from pyathena.s3fs.converter import DefaultS3FSTypeConverter from pyathena.s3fs.cursor import S3FSCursor from pyathena.s3fs.reader import AthenaCSVReader, DefaultCSVReader from pyathena.s3fs.result_set import AthenaS3FSResultSet @@ -17,6 +19,12 @@ from tests.pyathena.util import cached_file_systems +def _s3fs_converter_with_json_text(): + converter = DefaultS3FSTypeConverter() + converter.set("json", _to_default) + return converter + + class TestS3FSCursor: def test_fetchone(self, s3fs_cursor): s3fs_cursor.execute("SELECT * FROM one_row") @@ -528,3 +536,26 @@ def test_quoted_string_with_comma(self, csv_reader_class): def test_fetch_all_rows(self, s3fs_cursor): s3fs_cursor.execute("SELECT 1 AS col") assert s3fs_cursor.fetchall() == [(1,)] + + @pytest.mark.parametrize( + "s3fs_cursor", + [ + pytest.param({"converter": _s3fs_converter_with_json_text()}, id="default"), + pytest.param( + { + "work_group": ENV.managed_work_group, + "s3_staging_dir": "", + "converter": _s3fs_converter_with_json_text(), + }, + id="managed", + marks=pytest.mark.skipif( + not ENV.managed_work_group, + reason="AWS_ATHENA_MANAGED_WORKGROUP not set", + ), + ), + ], + indirect=["s3fs_cursor"], + ) + def test_fetch_all_rows_custom_converter(self, s3fs_cursor): + s3fs_cursor.execute("SELECT 1 AS col, json_parse('{\"a\": 1}') AS col_json") + assert s3fs_cursor.fetchall() == [(1, '{"a":1}')] diff --git a/tests/pyathena/sqlalchemy/test_base.py b/tests/pyathena/sqlalchemy/test_base.py index 9e0536bc..eb6055d4 100644 --- a/tests/pyathena/sqlalchemy/test_base.py +++ b/tests/pyathena/sqlalchemy/test_base.py @@ -891,67 +891,36 @@ def test_basic_query(self, engine): assert rows[0].number_of_rows == 1 assert len(rows[0]) == 1 - def test_json_type_with_cast(self, engine): - """Test JSON type support with CAST operation in SELECT query.""" - engine, conn = engine - # Note: Athena JSON type support has limitations - # - JSON objects are supported - # - Direct CAST of JSON arrays is not supported - # - JSON is primarily used with DML operations, not DDL - - # Test 1: Simple JSON object with type_coerce for proper type handling - result = conn.execute( - select( - type_coerce( - literal_column('CAST(\'{"name": "test", "value": 123}\' AS JSON)'), - types.JSON, - ).label("json_col") - ) - ).fetchone() - assert result.json_col == {"name": "test", "value": 123} - assert isinstance(result.json_col, dict) - - # Test 2: Nested JSON object with arrays inside - # (Arrays are supported as part of JSON objects, just not as top-level CAST) - nested_json_str = '{"user": {"id": 1, "name": "Alice"}, "scores": [95, 87, 92]}' - result = conn.execute( - select( - type_coerce(literal_column(f"CAST('{nested_json_str}' AS JSON)"), types.JSON).label( - "nested_json" - ) - ) - ).fetchone() - assert result.nested_json == { - "user": {"id": 1, "name": "Alice"}, - "scores": [95, 87, 92], + def test_json_type(self, engine): + engine, conn = engine + nested = '{"user": {"id": 1, "name": "Alice"}, "scores": [95, 87, 92]}' + scalars = '{"str": "value", "num": 42, "bool": true, "nil": null}' + columns = { + "obj": f"json_parse('{nested}')", + "arr": "json_parse('[1, 2, 3]')", + "scalars": f"json_parse('{scalars}')", + # Athena returns a cast of text as a JSON string scalar. + "cast_text": "CAST('{\"a\": 1}' AS JSON)", + "missing": "CAST(NULL AS JSON)", + # JSON text in a varchar column is decoded. + "text_obj": "'{\"a\": 1}'", } - assert result.nested_json["user"]["name"] == "Alice" - assert result.nested_json["scores"][0] == 95 - assert isinstance(result.nested_json["scores"], list) - - # Test 3: JSON with null value result = conn.execute( select( - type_coerce(literal_column("CAST('{\"key\": null}' AS JSON)"), types.JSON).label( - "json_with_null" + *( + type_coerce(literal_column(expr), types.JSON).label(name) + for name, expr in columns.items() ) ) - ).fetchone() - assert result.json_with_null == {"key": None} - assert result.json_with_null["key"] is None - - # Test 4: JSON with various types - result = conn.execute( - select( - type_coerce( - literal_column( - 'CAST(\'{"str": "value", "num": 42, "bool": true, "nil": null}\' AS JSON)' - ), - types.JSON, - ).label("json_types") - ) - ).fetchone() - assert result.json_types == {"str": "value", "num": 42, "bool": True, "nil": None} + ).one() + assert result._asdict() == { + "obj": {"user": {"id": 1, "name": "Alice"}, "scores": [95, 87, 92]}, + "arr": [1, 2, 3], + "scalars": {"str": "value", "num": 42, "bool": True, "nil": None}, + "cast_text": '{"a": 1}', + "missing": None, + "text_obj": {"a": 1}, + } def test_select_nested_struct_query(self, engine): """Test SELECT query with nested STRUCT (ROW) types (Issue #627).""" diff --git a/tests/pyathena/sqlalchemy/test_types.py b/tests/pyathena/sqlalchemy/test_types.py index 665a16cb..a5be6b0d 100644 --- a/tests/pyathena/sqlalchemy/test_types.py +++ b/tests/pyathena/sqlalchemy/test_types.py @@ -1,7 +1,48 @@ +import json + +import pytest from sqlalchemy import types from pyathena.sqlalchemy.base import ischema_names +from pyathena.sqlalchemy.rest import AthenaRestDialect def test_double_column_type(): assert ischema_names["double"] is types.DOUBLE + + +class _JSONText(types.TypeDecorator): + impl = types.JSON + cache_ok = True + + +class TestAthenaJSON: + @staticmethod + def _process(type_, value, coltype, dialect=None): + dialect = dialect or AthenaRestDialect() + processor = type_.dialect_impl(dialect).result_processor(dialect, coltype) + return processor(value) if processor else value + + @pytest.mark.parametrize("type_", [types.JSON(), _JSONText()]) + @pytest.mark.parametrize( + "value", + [{"a": 1}, [1, 2, 3], '{"a": 1}', "text", 1, True, None], + ) + def test_keeps_converted_json_results(self, type_, value): + assert self._process(type_, value, "json") == value + + @pytest.mark.parametrize("type_", [types.JSON(), _JSONText()]) + @pytest.mark.parametrize( + ("value", "expected"), + [('{"a": 1}', {"a": 1}), ("[1, 2, 3]", [1, 2, 3]), ('"text"', "text"), (None, None)], + ) + def test_decodes_json_text_of_other_types(self, type_, value, expected): + assert self._process(type_, value, "varchar") == expected + + def test_decodes_json_text_with_dialect_deserializer(self): + dialect = AthenaRestDialect(json_deserializer=lambda v: ("custom", json.loads(v))) + assert self._process(types.JSON(), '{"a": 1}', "varchar", dialect) == ( + "custom", + {"a": 1}, + ) + assert self._process(types.JSON(), {"a": 1}, "json", dialect) == {"a": 1} diff --git a/tests/pyathena/test_converter.py b/tests/pyathena/test_converter.py index 3dd10249..5d7ad401 100644 --- a/tests/pyathena/test_converter.py +++ b/tests/pyathena/test_converter.py @@ -5,6 +5,7 @@ from pyathena.converter import ( DefaultTypeConverter, + _csv_to_json, _to_array, _to_datetime, _to_datetime_with_tz, @@ -325,6 +326,33 @@ def test_to_array_invalid_formats(input_value): assert _to_array(input_value) is None +@pytest.mark.parametrize( + ("input_value", "expected"), + [ + (None, None), + ("", None), + ('{"a":1}', {"a": 1}), + ("[1,2]", [1, 2]), + ('""', ""), + ('"[1, 2]"', "[1, 2]"), + ("null", None), + ], +) +def test_csv_to_json(input_value, expected): + assert _csv_to_json(input_value) == expected + + +@pytest.mark.parametrize( + ("value", "type_hint", "expected"), + [ + ('[""]', "array(json)", [""]), + ('{"k": ""}', "map(varchar,json)", {"k": ""}), + ], +) +def test_nested_json_empty_string(value, type_hint, expected): + assert DefaultTypeConverter().convert(type_hint.split("(")[0], value, type_hint) == expected + + class TestDefaultTypeConverter: @pytest.mark.parametrize( ("input_value", "expected"),