Skip to content

Commit 2cd0bdf

Browse files
Merge pull request #770 from pyathena-dev/fix/759-binary-compliance
Restore SQLAlchemy binary compliance and preserve CSV binary NULLs
2 parents f7b2369 + cf0514e commit 2cd0bdf

26 files changed

Lines changed: 866 additions & 61 deletions

‎docs/null_handling.md‎

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -60,6 +60,17 @@ which correctly interprets unquoted empty values as NULL, while `S3FSCursor` use
6060
`AthenaCSVReader` that respects CSV quoting rules.
6161
```
6262

63+
## Binary Values
64+
65+
The string comparison above does not apply to `VARBINARY` columns.
66+
With the default CSV settings and converters, pandas and Arrow cursors distinguish SQL NULL from empty binary
67+
values when reading CSV results: `fetchone()`, `fetchmany()`, and `fetchall()` return `None`
68+
for NULL and `b''` for an empty binary value.
69+
This also applies to their asynchronous variants and pandas chunked reads.
70+
Pandas DataFrames preserve the same values.
71+
Arrow Tables retain the CSV hexadecimal strings, with NULL represented as an Arrow null;
72+
fetch methods convert the hexadecimal strings to Python bytes.
73+
6374
## Default Cursor (API-based)
6475

6576
The default `Cursor` and `DictCursor` fetch results directly from the Athena API,

‎pyathena/__init__.py‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -63,6 +63,7 @@ def __hash__(self):
6363
Date: type[datetime.date] = datetime.date
6464
Time: type[datetime.time] = datetime.time
6565
Timestamp: type[datetime.datetime] = datetime.datetime
66+
Binary: type[bytes] = bytes
6667

6768

6869
@overload

‎pyathena/aio/sqlalchemy/base.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -171,6 +171,8 @@ class AsyncAdapt_pyathena_dbapi:
171171
"""
172172

173173
paramstyle = "pyformat"
174+
Binary = pyathena.Binary
175+
BINARY = pyathena.BINARY
174176

175177
# DBAPI exception hierarchy
176178
Error = Error

‎pyathena/arrow/result_set.py‎

Lines changed: 12 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -276,6 +276,7 @@ def _read_csv(self) -> Table:
276276
):
277277
return pa.Table.from_pydict({})
278278
length = self._get_content_length()
279+
binary_columns = {d[0] for d in self.description or [] if d[1] == "varbinary"}
279280
if length and self.output_location.endswith(".txt"):
280281
description = self.description if self.description else []
281282
column_names = [d[0] for d in description]
@@ -296,6 +297,7 @@ def _read_csv(self) -> Table:
296297
parse_opts = csv.ParseOptions(
297298
delimiter=",",
298299
quote_char='"',
300+
ignore_empty_lines=not binary_columns,
299301
double_quote=True,
300302
escape_char=False,
301303
)
@@ -304,16 +306,25 @@ def _read_csv(self) -> Table:
304306

305307
bucket, key = parse_output_location(self.output_location)
306308
try:
307-
return csv.read_csv(
309+
table = csv.read_csv(
308310
self._fs.open_input_stream(f"{bucket}/{key}"),
309311
read_options=read_opts,
310312
parse_options=parse_opts,
311313
convert_options=csv.ConvertOptions(
314+
strings_can_be_null=bool(binary_columns),
312315
quoted_strings_can_be_null=False,
313316
timestamp_parsers=self.timestamp_parsers,
314317
column_types=self.column_types,
315318
),
316319
)
320+
if binary_columns:
321+
for index, field in enumerate(table.schema):
322+
if field.name not in binary_columns and (
323+
pa.types.is_string(field.type) or pa.types.is_binary(field.type)
324+
):
325+
# Preserve the existing CSV behavior for non-binary Athena columns.
326+
table = table.set_column(index, field, table.column(index).fill_null(""))
327+
return table
317328
except Exception as e:
318329
_logger.exception(f"Failed to read {bucket}/{key}.")
319330
raise OperationalError(*e.args) from e

‎pyathena/formatter.py‎

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -251,6 +251,12 @@ def _format_str(formatter: Formatter, escaper: Callable[[str], str], val: Any) -
251251
return escaper(val)
252252

253253

254+
def _format_binary(
255+
formatter: Formatter, escaper: Callable[[str], str], val: bytes | bytearray | memoryview
256+
) -> str:
257+
return f"X'{val.hex()}'"
258+
259+
254260
def _format_seq(formatter: Formatter, escaper: Callable[[str], str], val: Any) -> Any:
255261
results = []
256262
for v in val:
@@ -291,6 +297,9 @@ def _format_decimal(formatter: Formatter, escaper: Callable[[str], str], val: An
291297
Decimal: _format_decimal,
292298
bool: _format_bool,
293299
str: _format_str,
300+
bytes: _format_binary,
301+
bytearray: _format_binary,
302+
memoryview: _format_binary,
294303
list: _format_seq,
295304
set: _format_seq,
296305
tuple: _format_seq,
@@ -307,6 +316,7 @@ class DefaultParameterFormatter(Formatter):
307316
Supported types:
308317
- None: Converts to SQL NULL
309318
- Strings: Properly escaped and quoted
319+
- Binary data: bytes, bytearray, memoryview as hexadecimal literals
310320
- Numbers: int, float, Decimal
311321
- Dates and times: date, datetime, time
312322
- Booleans: Converted to SQL boolean literals

‎pyathena/pandas/cursor.py‎

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -294,9 +294,10 @@ def iter_chunks(self) -> Generator[DataFrame, None, None]:
294294

295295
import gc
296296

297-
for chunk_count, chunk in enumerate(result_set.iter_chunks(), 1):
298-
yield chunk
297+
with result_set.iter_chunks() as chunks:
298+
for chunk_count, chunk in enumerate(chunks, 1):
299+
yield chunk
299300

300-
# Suggest garbage collection every 10 chunks for large datasets
301-
if chunk_count % 10 == 0:
302-
gc.collect()
301+
# Suggest garbage collection every 10 chunks for large datasets
302+
if chunk_count % 10 == 0:
303+
gc.collect()

‎pyathena/pandas/reader.py‎

Lines changed: 64 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,64 @@
1+
from __future__ import annotations
2+
3+
import re
4+
from io import RawIOBase
5+
from typing import Any
6+
7+
from pyathena.s3fs.reader import AthenaCSVReader
8+
9+
_BINARY_NULL = "__PYATHENA_BINARY_NULL__"
10+
_CSV_FIELD = re.compile(r'(?:^|,)(?P<value>"[^"]*(?:""[^"]*)*"|[^,]*)')
11+
12+
13+
class BinaryCSVReader(RawIOBase):
14+
"""Preserve binary NULL fields before pandas discards CSV quoting information.
15+
16+
The marker cannot occur in Athena's hexadecimal encoding of binary values.
17+
Only unquoted empty binary fields are rewritten; all other CSV text is
18+
passed through unchanged. Records are streamed to support chunked reads.
19+
"""
20+
21+
def __init__(self, stream: Any, binary_columns: set[int]) -> None:
22+
super().__init__()
23+
self._reader = AthenaCSVReader(stream)
24+
self._binary_columns = binary_columns
25+
self._header = True
26+
self._buffer = b""
27+
28+
def readable(self) -> bool:
29+
return True
30+
31+
def readinto(self, buffer: Any) -> int:
32+
if self.closed:
33+
raise ValueError("I/O operation on closed file.")
34+
if not len(buffer):
35+
return 0
36+
if not self._buffer:
37+
try:
38+
record = self._reader._read_record()
39+
except StopIteration:
40+
self._reader.close()
41+
return 0
42+
if self._header:
43+
self._header = False
44+
else:
45+
parts: list[str] = []
46+
start = 0
47+
for index, field in enumerate(_CSV_FIELD.finditer(record.rstrip("\r\n"))):
48+
if index in self._binary_columns and not field.group("value"):
49+
pos = field.start("value")
50+
parts.extend((record[start:pos], _BINARY_NULL))
51+
start = pos
52+
parts.append(record[start:])
53+
record = "".join(parts)
54+
self._buffer = record.encode("utf-8")
55+
size = min(len(buffer), len(self._buffer))
56+
buffer[:size] = self._buffer[:size]
57+
self._buffer = self._buffer[size:]
58+
return size
59+
60+
def close(self) -> None:
61+
try:
62+
self._reader.close()
63+
finally:
64+
super().close()

0 commit comments

Comments
 (0)