diff --git a/pyathena/converter.py b/pyathena/converter.py index e573a365..72d9d83d 100644 --- a/pyathena/converter.py +++ b/pyathena/converter.py @@ -20,6 +20,7 @@ TypeNode, TypeSignatureParser, _split_array_items, + _split_native_array_items, ) from pyathena.util import override, strtobool @@ -234,7 +235,7 @@ def _to_array(varchar_value: str | None) -> list[Any] | str | None: # If JSON parsing fails, fall back to basic parsing for simple cases pass - inner = varchar_value[1:-1].strip() + inner = varchar_value[1:-1] if not inner: return [] @@ -378,13 +379,10 @@ def _parse_array_native(inner: str) -> list[Any] | None: """ result = [] - # Smart split by comma - respect brace groupings - items = _split_array_items(inner) + # Split as Athena joins the items, respecting brace groupings + items = _split_native_array_items(inner) for item in items: - if not item: - continue - # Handle struct (ROW) values in format {a, b, c} or {key=value, ...} if item.strip().startswith("{") and item.strip().endswith("}"): # This is a struct value - parse it as a struct diff --git a/pyathena/pandas/result_set.py b/pyathena/pandas/result_set.py index 2a78aeef..54a864e0 100644 --- a/pyathena/pandas/result_set.py +++ b/pyathena/pandas/result_set.py @@ -157,14 +157,14 @@ def get_chunk(self, size: int | None = None) -> DataFrame: size: Number of rows to retrieve. If None, returns entire chunk. Returns: - DataFrame chunk. + DataFrame chunk, with date truncation applied as in iteration. """ from pandas.io.parsers import TextFileReader try: if isinstance(self._reader, TextFileReader): - return self._reader.get_chunk(size) - return next(self._reader) + return self._trunc_date(self._reader.get_chunk(size)) + return self._trunc_date(next(self._reader)) except BaseException: self.close() raise @@ -518,7 +518,10 @@ def parse_dates(self) -> list[Any | None]: def _trunc_date(self, df: DataFrame) -> DataFrame: if self._time_columns: - truncated = df.loc[:, self._time_columns].apply(lambda r: r.dt.time) + # A NULL is None, as with the GetQueryResults fallback and the other types. + truncated = df.loc[:, self._time_columns].apply( + lambda r: r.dt.time.astype(object).where(r.notna(), None) + ) for time_col in self._time_columns: df.isetitem(df.columns.get_loc(time_col), truncated[time_col]) return df diff --git a/pyathena/parser.py b/pyathena/parser.py index 120e2998..89a17584 100644 --- a/pyathena/parser.py +++ b/pyathena/parser.py @@ -49,6 +49,40 @@ def _split_array_items(inner: str) -> list[str]: return items +def _split_native_array_items(inner: str) -> list[str]: + """Split the items of an array in Athena's native format. + + Athena joins the items with ``", "``, so only a top-level comma followed by a space + separates items. Other commas, leading and trailing spaces, and empty items belong + to the items. Brace and bracket groupings are respected. + + Args: + inner: Interior content of the array without brackets, not stripped. + + Returns: + List of item strings. + """ + items: list[str] = [] + current: list[str] = [] + depth = 0 + index = 0 + while index < len(inner): + char = inner[index] + if char in "{[": + depth += 1 + elif char in "}]": + depth -= 1 + elif char == "," and depth == 0 and inner.startswith(" ", index + 1): + items.append("".join(current)) + current = [] + index += 2 + continue + current.append(char) + index += 1 + items.append("".join(current)) + return items + + @dataclass class TypeNode: """Parsed representation of an Athena DDL type signature. @@ -321,9 +355,14 @@ def _convert_typed_array(self, value: str, type_node: TypeNode) -> list[Any] | N element_type = type_node.children[0] if type_node.children else TypeNode("varchar") - # Try JSON first (only if content looks like JSON) + # Try JSON first if the elements are JSON, whose values Athena renders as JSON + # text, or if the content looks like JSON inner_preview = value[1:10] if len(value) > 10 else value[1:-1] - if '"' in inner_preview or value.startswith(("[{", "[null", "[[")): + if ( + element_type.type_name == "json" + or '"' in inner_preview + or value.startswith(("[{", "[null", "[[")) + ): try: parsed = json.loads(value) if isinstance(parsed, list): @@ -337,19 +376,16 @@ def _convert_typed_array(self, value: str, type_node: TypeNode) -> list[Any] | N pass # Native format - inner = value[1:-1].strip() + inner = value[1:-1] if not inner: return [] if "[" in inner: return None # Nested arrays not supported in native format - items = _split_array_items(inner) + items = _split_native_array_items(inner) result: list[Any] = [] for item in items: - item = item.strip() - if not item: - continue if item.startswith("{") and item.endswith("}"): if element_type.type_name in ("row", "struct"): result.append(self._convert_typed_struct(item, element_type)) diff --git a/tests/pyathena/pandas/test_cursor.py b/tests/pyathena/pandas/test_cursor.py index a6916536..d4d78c9c 100644 --- a/tests/pyathena/pandas/test_cursor.py +++ b/tests/pyathena/pandas/test_cursor.py @@ -1589,6 +1589,17 @@ def test_pandas_cursor_iter_chunks_consistency(self, pandas_cursor): for chunk1, chunk2 in zip(chunks_via_method, chunks_via_direct, strict=False): pd.testing.assert_frame_equal(chunk1, chunk2) + @pytest.mark.parametrize( + "pandas_cursor", [{"cursor_kwargs": {"chunksize": 2}}], indirect=["pandas_cursor"] + ) + def test_get_chunk_time(self, pandas_cursor): + """get_chunk() converts time columns as iteration does.""" + pandas_cursor.execute( + "SELECT * FROM (VALUES (1, CAST('12:34:56' AS TIME)), (2, NULL)) AS t(i, v) ORDER BY i" + ) + chunk = pandas_cursor.as_pandas().get_chunk() + assert chunk["v"].tolist() == [datetime(2017, 1, 1, 12, 34, 56).time(), None] + @pytest.mark.parametrize( "pandas_cursor", [ diff --git a/tests/pyathena/test_converter.py b/tests/pyathena/test_converter.py index 08092519..9293047f 100644 --- a/tests/pyathena/test_converter.py +++ b/tests/pyathena/test_converter.py @@ -726,3 +726,33 @@ def test_to_datetime_with_tz_offsets_and_zone_names(input_value, expected): if expected is not None: assert result.utcoffset() == expected.utcoffset() assert result.tzinfo is not None + + +@pytest.mark.parametrize( + ("input_value", "expected"), + [ + ("[x, , y]", ["x", "", "y"]), + ("[x, ]", ["x", ""]), + ("[x, , y]", ["x", " ", "y"]), + ("[ , x]", [" ", "x"]), + ("[a,b, c]", ["a,b", "c"]), + ("[x, null]", ["x", None]), + ("[{a=1, b=}, {a=, b=2}]", [{"a": "1", "b": ""}, {"a": "", "b": "2"}]), + ], +) +def test_native_array_items(input_value, expected): + """Native arrays are split as Athena joins them, keeping empty items.""" + assert _to_array(input_value) == expected + assert ( + DefaultTypeConverter().convert("array", input_value, type_hint="array(varchar)") == expected + ) + + +def test_typed_json_array_starting_with_non_string(): + """A typed array(json) value is parsed as JSON whatever its first element is.""" + converter = DefaultTypeConverter() + assert converter.convert("array", '[1234567890, "a,b"]', type_hint="array(json)") == [ + 1234567890, + "a,b", + ] + assert converter.convert("array", "[1, 2.5, true]", type_hint="array(json)") == [1, 2.5, True] diff --git a/tests/pyathena/test_cursor.py b/tests/pyathena/test_cursor.py index 6b012f7f..3f248587 100644 --- a/tests/pyathena/test_cursor.py +++ b/tests/pyathena/test_cursor.py @@ -1693,6 +1693,42 @@ def test_fetch_all_rows(self, cursor): cursor.execute(CONVERTED_VALUES_QUERY) assert cursor.fetchall() == [CONVERTED_VALUES_ROW] + @pytest.mark.parametrize( + "cursor", + [ + pytest.param({}, id="default"), + pytest.param( + {"work_group": ENV.managed_work_group, "s3_staging_dir": ""}, + id="managed", + marks=pytest.mark.skipif( + not ENV.managed_work_group, + reason="AWS_ATHENA_MANAGED_WORKGROUP not set", + ), + ), + ], + indirect=["cursor"], + ) + def test_fetch_complex_values(self, cursor): + """Native arrays keep empty, space-only, and comma items; array(json) keeps JSON.""" + query = """ + SELECT + ARRAY['x', '', 'y'] AS col_empty + ,ARRAY['x', ' ', 'a,b'] AS col_space_comma + ,ARRAY[json_parse('1234567890'), json_parse('"a,b"')] AS col_json + """ + cursor.execute(query) + expected = (["x", "", "y"], ["x", " ", "a,b"], [1234567890, "a,b"]) + assert cursor.fetchall() == [expected] + cursor.execute( + query, + result_set_type_hints={ + "col_empty": "array(varchar)", + "col_space_comma": "array(varchar)", + "col_json": "array(json)", + }, + ) + assert cursor.fetchall() == [expected] + @staticmethod def _metadata_view(metadata): return ( diff --git a/tests/pyathena/util.py b/tests/pyathena/util.py index b8479d36..8fd8d93b 100644 --- a/tests/pyathena/util.py +++ b/tests/pyathena/util.py @@ -110,8 +110,8 @@ def unreachable_glue(connection): # TIME values of several precisions, with and without a time zone, TIMESTAMP WITH TIME -# ZONE values with UTC offsets and a zone name, a NULL JSON value, and the row that -# every cursor should fetch for them. +# ZONE values with UTC offsets and a zone name, NULL JSON and TIME values, and the row +# that every cursor should fetch for them. CONVERTED_VALUES_QUERY = """ SELECT 1 AS col @@ -125,6 +125,7 @@ def unreachable_glue(connection): ,TIMESTAMP '2024-02-29 23:59:58.123 -08:00' AS col_timestamp_tz_negative ,TIMESTAMP '2024-02-29 23:59:58.123 America/New_York' AS col_timestamp_tz_name ,CAST(NULL AS TIMESTAMP WITH TIME ZONE) AS col_timestamp_tz_null + ,CAST(NULL AS TIME) AS col_time_null """ CONVERTED_VALUES_ROW = ( 1, @@ -138,6 +139,7 @@ def unreachable_glue(connection): datetime(2024, 2, 29, 23, 59, 58, 123000, tzinfo=timezone(-timedelta(hours=8))), datetime(2024, 2, 29, 23, 59, 58, 123000, tzinfo=gettz("America/New_York")), None, + None, )