Skip to content

Commit 165f37c

Browse files
Preserve ARRAY update expression arguments and binary values
1 parent 777cfd2 commit 165f37c

4 files changed

Lines changed: 67 additions & 10 deletions

File tree

‎pyathena/formatter.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@
2222
class _ComplexParameter:
2323
"""Typed complex value supplied by the SQLAlchemy dialect."""
2424

25-
constructor: Literal["ARRAY", "MAP", "ROW", "JSON_PARSE"]
25+
constructor: Literal["ARRAY", "MAP", "ROW", "JSON_PARSE", "FROM_HEX"]
2626
values: tuple[Any, ...]
2727

2828

‎pyathena/sqlalchemy/_array_update.py‎

Lines changed: 11 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -3,11 +3,12 @@
33
from typing import Any
44

55
from sqlalchemy import exc, types, util
6-
from sqlalchemy.sql import operators, visitors
6+
from sqlalchemy.sql import operators
77
from sqlalchemy.sql.elements import BinaryExpression, BindParameter, ColumnElement, Null, Slice
88
from sqlalchemy.sql.schema import Column
99
from sqlalchemy.sql.visitors import InternalTraversal
1010

11+
from pyathena.formatter import _ComplexParameter
1112
from pyathena.sqlalchemy.types import _array_item_type, _bind_complex, _literal_complex
1213

1314

@@ -20,7 +21,13 @@ def __init__(self, item_type):
2021
self.item_type = item_type
2122

2223
def bind_processor(self, dialect):
23-
return lambda value: _bind_complex(value, self.item_type, dialect)
24+
def process(value):
25+
value = _bind_complex(value, self.item_type, dialect)
26+
if isinstance(value, (bytes, bytearray)):
27+
return _ComplexParameter("FROM_HEX", (value.hex(),))
28+
return value
29+
30+
return process
2431

2532
def literal_processor(self, dialect):
2633
return lambda value: _literal_complex(value, self.item_type, dialect)
@@ -118,12 +125,8 @@ def _index_sql(compiler, index: ColumnElement[Any], **kw):
118125
if not isinstance(index.type, (types.Integer, types.NullType, types.ARRAY)):
119126
raise exc.CompileError("ARRAY write indices must be integers")
120127

121-
def type_bind(element: Any, **kwargs: Any) -> Any:
122-
if isinstance(element, BindParameter):
123-
return element._with_binary_element_type(_IndexType())
124-
return None
125-
126-
index = visitors.replacement_traverse(index, {}, type_bind)
128+
if isinstance(index, BindParameter):
129+
index = index._with_binary_element_type(_IndexType())
127130
sql = compiler.process(index, **kw)
128131
failure = (
129132
f"CAST(concat('Invalid ARRAY index: ', coalesce(CAST({sql} AS VARCHAR), 'NULL')) AS BIGINT)"

‎tests/pyathena/sqlalchemy/test_array_update.py‎

Lines changed: 34 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11
import pytest
2-
from sqlalchemy import Column, Integer, MetaData, Table, bindparam, types
2+
from sqlalchemy import Column, Integer, MetaData, Table, bindparam, func, types
33
from sqlalchemy import exc as sa_exc
44

55
from pyathena.formatter import DefaultParameterFormatter
@@ -88,3 +88,36 @@ def test_generic_array_partial_update():
8888
table = array_table(types.ARRAY(Integer))
8989
sql = str(table.update().values({table.c["items"][1]: 2}).compile(dialect=AthenaDialect()))
9090
assert "SET items=transform(" in sql
91+
92+
93+
def test_write_index_expression_keeps_its_argument_types():
94+
table = array_table()
95+
index = func.length("abc")
96+
statement = table.update().values({table.c["items"][index]: 9})
97+
compiled = statement.compile(dialect=AthenaDialect())
98+
params = {
99+
name: compiled._bind_processors.get(name, lambda value: value)(value)
100+
for name, value in compiled.params.items()
101+
}
102+
assert "length('abc')" in DefaultParameterFormatter().format(str(compiled), params)
103+
104+
105+
def test_ordered_partial_update_with_sql_expression():
106+
table = array_table()
107+
items = table.c["items"]
108+
statement = table.update().ordered_values((items[2], items[1] + 1), (table.c.id, 2))
109+
compiled = str(statement.compile(dialect=AthenaDialect()))
110+
assert compiled.index("SET items=") < compiled.index(", id=")
111+
assert "element_at(arrays.items" in compiled
112+
113+
114+
def test_binary_element_assignment_uses_native_hex_parameter():
115+
table = array_table(AthenaArray(types.BINARY))
116+
compiled = (
117+
table.update().values({table.c["items"][1]: b"\x00\xff"}).compile(dialect=AthenaDialect())
118+
)
119+
params = {
120+
name: compiled._bind_processors.get(name, lambda value: value)(value)
121+
for name, value in compiled.params.items()
122+
}
123+
assert "FROM_HEX('00ff')" in DefaultParameterFormatter().format(str(compiled), params)

‎tests/sqlalchemy/test_suite.py‎

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -168,6 +168,27 @@ def test_nested_zero_indexed_and_cached_bindings(self, connection, metadata):
168168
with pytest.raises(sa_exc.DBAPIError):
169169
connection.execute(statement, {"row_id": 1, "outer": -1, "inner": 0, "value": 9})
170170

171+
def test_expression_values_and_indices(self, connection, metadata):
172+
table = Table(
173+
"array_expression_updates",
174+
metadata,
175+
Column("id", Integer),
176+
Column("items", AthenaArray(Integer)),
177+
Column("binary_items", AthenaArray(types.BINARY)),
178+
)
179+
table.create(connection)
180+
connection.execute(table.insert().values(id=1, items=[1, 2, 3], binary_items=[b"abc"]))
181+
items = table.c["items"]
182+
connection.execute(
183+
table.update().ordered_values(
184+
(items[func.length("abc")], items[1] + 8),
185+
(table.c.binary_items[1], b"\x00\xff"),
186+
(table.c.id, 2),
187+
)
188+
)
189+
connection.execute(table.update().values({items[1:2]: items[2:3].concat([4])}))
190+
eq_(connection.execute(select(table)).one(), (2, [2, 9, 4, 9], [b"\x00\xff"]))
191+
171192
def test_null_slice_binding_rejected(self, connection, metadata):
172193
table = Table("array_null_slice", metadata, Column("items", types.ARRAY(Integer)))
173194
table.create(connection)

0 commit comments

Comments
 (0)