Skip to content

Commit db4c71e

Browse files
Support ORM ARRAY updates and preserve typed assignment expressions
1 parent 5f1aed0 commit db4c71e

4 files changed

Lines changed: 111 additions & 11 deletions

File tree

‎docs/sqlalchemy.md‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1094,6 +1094,7 @@ with engine.begin() as conn:
10941094
```
10951095

10961096
Element assignment beyond the end extends the array and fills intervening positions with NULL.
1097+
NULL padding uses Athena's `repeat()` function and follows its size limits.
10971098
A NULL destination array is treated as empty for partial updates.
10981099
Assigning `None` to an element stores NULL.
10991100
Nested indices rebuild the corresponding inner arrays, including missing inner arrays.

‎pyathena/sqlalchemy/_array_update.py‎

Lines changed: 25 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,10 @@ def process(value):
3232
def literal_processor(self, dialect):
3333
return lambda value: _literal_complex(value, self.item_type, dialect)
3434

35+
def bind_expression(self, bindvalue):
36+
expression = self.item_type.bind_expression(bindvalue)
37+
return bindvalue if expression is None else expression
38+
3539

3640
class _IndexType(types.TypeDecorator[int]):
3741
impl = types.Integer
@@ -62,7 +66,9 @@ def __init__(self, column, path, value, value_type):
6266
self.type = column.type
6367
self.value_type = value_type
6468
self.value = (
65-
value._with_binary_element_type(_AssignmentType(value_type))
69+
value._with_binary_element_type(
70+
_AssignmentType(value_type if value.type._isnull else value.type)
71+
)
6672
if isinstance(value, BindParameter)
6773
else value
6874
)
@@ -89,7 +95,10 @@ def rewrite_array_update(statement):
8995
base = base.left
9096
name = base if isinstance(base, str) else getattr(base, "key", None)
9197
if path:
92-
if not isinstance(base, Column) or base.table is not statement.table:
98+
if (
99+
not isinstance(base, Column)
100+
or base.table._deannotate() is not statement.table._deannotate()
101+
):
93102
raise exc.CompileError("ARRAY updates require a column of the target table")
94103
if name in seen:
95104
raise exc.CompileError("Only one assignment per ARRAY column is supported")
@@ -119,6 +128,7 @@ def _index_sql(compiler, index: ColumnElement[Any], **kw):
119128
if (
120129
isinstance(index, BindParameter)
121130
and not index.required
131+
and index.callable is None
122132
and (type(index.value) is not int or index.value <= 0)
123133
):
124134
raise exc.CompileError("ARRAY write indices must be positive integers after normalization")
@@ -139,7 +149,12 @@ def compile_array_update(compiler, expression, **kw):
139149
final_slice = isinstance(expression.path[-1], Slice)
140150
if final_slice and (
141151
isinstance(value, Null)
142-
or (isinstance(value, BindParameter) and not value.required and value.value is None)
152+
or (
153+
isinstance(value, BindParameter)
154+
and not value.required
155+
and value.callable is None
156+
and value.value is None
157+
)
143158
):
144159
raise exc.CompileError("An ARRAY slice assignment requires a non-NULL array")
145160
rhs = compiler.process(value, **kw)
@@ -189,15 +204,17 @@ def rebuild(array, array_type, path):
189204
f"concat({prefix}, {padding}, {rhs}, {suffix})", bound.step, array_type, **kw
190205
)
191206
index = _index_sql(compiler, bound, **kw)
192-
variable = compiler._array_lambda_name()
193207
previous = f"element_at({array}, {index})"
194208
replacement = (
195209
rebuild(previous, _array_item_type(array_type), path[1:]) if len(path) > 1 else rhs
196210
)
197-
return (
198-
f"transform(sequence(1, greatest(cardinality({array}), {index})), "
199-
f"{variable} -> IF({variable} = {index}, {replacement}, "
200-
f"element_at({array}, {variable})))"
211+
prefix = f"slice({array}, 1, least({index} - 1, cardinality({array})))"
212+
element_type = compiler._complex_dml_type(_array_item_type(array_type))
213+
padding = (
214+
f"repeat(CAST(NULL AS {element_type}), "
215+
f"CAST(greatest({index} - 1 - cardinality({array}), 0) AS INTEGER))"
201216
)
217+
suffix = f"slice({array}, {index} + 1, greatest(cardinality({array}) - {index}, 0))"
218+
return f"concat({prefix}, {padding}, ARRAY[{replacement}], {suffix})"
202219

203220
return rebuild(compiler.process(expression.column, **kw), expression.type, expression.path)

‎tests/pyathena/sqlalchemy/test_array_update.py‎

Lines changed: 50 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
import pytest
2-
from sqlalchemy import Column, Integer, MetaData, Table, bindparam, func, types
2+
from sqlalchemy import Column, Integer, MetaData, Table, bindparam, func, types, update
33
from sqlalchemy import exc as sa_exc
4+
from sqlalchemy.orm import declarative_base
45

56
from pyathena.formatter import DefaultParameterFormatter
67
from pyathena.sqlalchemy.base import AthenaDialect
@@ -79,15 +80,16 @@ def test_nested_and_zero_indexed_update():
7980
table = array_table(AthenaArray(Integer, dimensions=2, zero_indexes=True))
8081
statement = table.update().values({table.c["items"][0][2]: 7})
8182
sql = str(statement.compile(dialect=AthenaDialect(), compile_kwargs={"literal_binds": True}))
82-
assert sql.count("transform(sequence(") == 2
83+
assert "ARRAY[concat(" in sql
84+
assert "sequence(" not in sql
8385
assert "IF(1 > 0, 1," in sql
8486
assert "IF(3 > 0, 3," in sql
8587

8688

8789
def test_generic_array_partial_update():
8890
table = array_table(types.ARRAY(Integer))
8991
sql = str(table.update().values({table.c["items"][1]: 2}).compile(dialect=AthenaDialect()))
90-
assert "SET items=transform(" in sql
92+
assert "SET items=concat(" in sql
9193

9294

9395
def test_write_index_expression_keeps_its_argument_types():
@@ -121,3 +123,48 @@ def test_binary_element_assignment_uses_native_hex_parameter():
121123
for name, value in compiled.params.items()
122124
}
123125
assert "FROM_HEX('00ff')" in DefaultParameterFormatter().format(str(compiled), params)
126+
127+
128+
class PrefixString(types.TypeDecorator):
129+
impl = types.String
130+
cache_ok = True
131+
132+
def process_bind_param(self, value, dialect):
133+
return f"prefix:{value}"
134+
135+
def bind_expression(self, bindvalue):
136+
return func.upper(bindvalue)
137+
138+
139+
def test_explicit_assignment_type_and_callable_bindings():
140+
table = array_table(AthenaArray(types.String))
141+
stmt = table.update().values(
142+
{
143+
table.c["items"][bindparam("index", callable_=lambda: 1)]: bindparam(
144+
"value", type_=PrefixString(), callable_=lambda: "a"
145+
)
146+
}
147+
)
148+
compiled = stmt.compile(dialect=AthenaDialect())
149+
assert compiled._bind_processors["value"]("a") == "prefix:a"
150+
assert "upper(%(value)s)" in str(compiled)
151+
assert compiled.params["index"] == 1
152+
assert compiled.params["value"] == "a"
153+
table.update().values(
154+
{table.c["items"][1:2]: bindparam("values", callable_=lambda: ["a"])}
155+
).compile(dialect=AthenaDialect())
156+
157+
158+
def test_orm_partial_update_and_renamed_attribute_conflicts():
159+
base = declarative_base()
160+
161+
class Model(base):
162+
__tablename__ = "arrays"
163+
id = Column(Integer, primary_key=True)
164+
values = Column("stored", AthenaArray(Integer), key="db_key")
165+
166+
sql = str(update(Model).values({Model.values[1]: 2}).compile(dialect=AthenaDialect()))
167+
assert "UPDATE arrays SET stored=concat(" in sql
168+
for whole in (Model.values, "values"):
169+
with pytest.raises(sa_exc.CompileError, match="one assignment"):
170+
update(Model).values({Model.values[1]: 2, whole: []}).compile(dialect=AthenaDialect())

‎tests/sqlalchemy/test_suite.py‎

Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,9 +20,11 @@
2020
select,
2121
text,
2222
types,
23+
update,
2324
)
2425
from sqlalchemy import exc as sa_exc
2526
from sqlalchemy import testing as sa_testing
27+
from sqlalchemy.orm import Session, registry
2628
from sqlalchemy.sql.elements import quoted_name
2729
from sqlalchemy.testing import eq_, fixtures
2830
from sqlalchemy.testing.schema import Column, Table
@@ -189,6 +191,39 @@ def test_expression_values_and_indices(self, connection, metadata):
189191
connection.execute(table.update().values({items[1:2]: items[2:3].concat([4])}))
190192
eq_(connection.execute(select(table)).one(), (2, [2, 9, 4, 9], [b"\x00\xff"]))
191193

194+
def test_orm_and_long_array_update(self, connection, metadata):
195+
table = Table(
196+
"array_orm_updates",
197+
metadata,
198+
Column("id", Integer, primary_key=True),
199+
Column("items", AthenaArray(Integer)),
200+
)
201+
table.create(connection)
202+
connection.execute(
203+
table.insert().values(
204+
id=1,
205+
items=func.concat(func.sequence(1, 10000), literal([10001], AthenaArray(Integer))),
206+
)
207+
)
208+
mapping = registry()
209+
210+
class Record:
211+
pass
212+
213+
mapping.map_imperatively(Record, table)
214+
try:
215+
with Session(bind=connection) as session:
216+
session.execute(update(Record).where(Record.id == 1).values({Record.items[1]: 99}))
217+
session.flush()
218+
row = connection.execute(select(table.c["items"])).scalar_one()
219+
eq_((len(row), row[0], row[-1]), (10001, 99, 10001))
220+
connection.execute(
221+
table.update().values({table.c["items"][2]: select(literal(77)).scalar_subquery()})
222+
)
223+
eq_(connection.execute(select(table.c["items"])).scalar_one()[:3], [99, 77, 3])
224+
finally:
225+
mapping.dispose()
226+
192227
def test_null_slice_binding_rejected(self, connection, metadata):
193228
table = Table("array_null_slice", metadata, Column("items", types.ARRAY(Integer)))
194229
table.create(connection)

0 commit comments

Comments
 (0)