From 054200de71e39e736ad2cc2d4f5a01b2d720a83c Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 19 Sep 2026 16:33:10 +0900 Subject: [PATCH 01/10] Implement single-statement SQLAlchemy ARRAY partial updates --- pyathena/sqlalchemy/_array_update.py | 198 ++++++++++++++++++ pyathena/sqlalchemy/compiler.py | 9 + .../pyathena/sqlalchemy/test_array_update.py | 90 ++++++++ 3 files changed, 297 insertions(+) create mode 100644 pyathena/sqlalchemy/_array_update.py create mode 100644 tests/pyathena/sqlalchemy/test_array_update.py diff --git a/pyathena/sqlalchemy/_array_update.py b/pyathena/sqlalchemy/_array_update.py new file mode 100644 index 00000000..09e409fd --- /dev/null +++ b/pyathena/sqlalchemy/_array_update.py @@ -0,0 +1,198 @@ +from __future__ import annotations + +from typing import Any + +from sqlalchemy import exc, types, util +from sqlalchemy.sql import operators, visitors +from sqlalchemy.sql.elements import BinaryExpression, BindParameter, ColumnElement, Null, Slice +from sqlalchemy.sql.schema import Column +from sqlalchemy.sql.visitors import InternalTraversal + +from pyathena.sqlalchemy.types import _array_item_type, _bind_complex, _literal_complex + + +class _AssignmentType(types.TypeDecorator[Any]): + impl = types.NullType + cache_ok = True + + def __init__(self, item_type): + super().__init__() + self.item_type = item_type + + def bind_processor(self, dialect): + return lambda value: _bind_complex(value, self.item_type, dialect) + + def literal_processor(self, dialect): + return lambda value: _literal_complex(value, self.item_type, dialect) + + +class _IndexType(types.TypeDecorator[int]): + impl = types.Integer + cache_ok = True + + def process_bind_param(self, value, dialect): + if type(value) is not int: + raise ValueError("ARRAY write indices must be non-NULL integers") + return value + + def process_literal_param(self, value, dialect): + return self.process_bind_param(value, dialect) + + +class _ArrayUpdate(ColumnElement[Any]): + __visit_name__ = "athena_array_update" + inherit_cache = True + _traverse_internals = [ # noqa: RUF012 + ("column", InternalTraversal.dp_clauseelement), + ("path", InternalTraversal.dp_clauseelement_list), + ("value", InternalTraversal.dp_clauseelement), + ("type", InternalTraversal.dp_type), + ] + + def __init__(self, column, path, value, value_type): + self.column = column + self.path = path + self.type = column.type + self.value_type = value_type + self.value = ( + value._with_binary_element_type(_AssignmentType(value_type)) + if isinstance(value, BindParameter) + else value + ) + + @property + def _from_objects(self): + return self.column._from_objects + self.value._from_objects + + +def rewrite_array_update(statement): + values = statement._ordered_values + if values is None: + values = list((statement._values or {}).items()) + rewritten = [] + seen = set() + partial = set() + for key, value in values: + base = key + path: list[Any] = [] + while isinstance(base, BinaryExpression) and base.operator is operators.getitem: + if not isinstance(base.left.type, types.ARRAY): + break + path.insert(0, base.right) + base = base.left + name = base if isinstance(base, str) else getattr(base, "key", None) + if path: + if not isinstance(base, Column) or base.table is not statement.table: + raise exc.CompileError("ARRAY updates require a column of the target table") + if name in seen: + raise exc.CompileError("Only one assignment per ARRAY column is supported") + if any(isinstance(index, Slice) for index in path[:-1]): + raise exc.CompileError("Only the final ARRAY update index can be a slice") + partial.add(name) + value_type = key.type + value = _ArrayUpdate(base, path, value, value_type) + key = base + elif name in partial: + raise exc.CompileError("Only one assignment per ARRAY column is supported") + seen.add(name) + rewritten.append((key, value)) + if not partial: + return statement + result = statement._clone() + if statement._ordered_values is not None: + result._ordered_values = rewritten + else: + result._values = util.immutabledict(rewritten) + return result + + +def _index_sql(compiler, index: ColumnElement[Any], **kw): + if isinstance(index, Null): + raise exc.CompileError("ARRAY write indices must be non-NULL positive integers") + if ( + isinstance(index, BindParameter) + and not index.required + and (type(index.value) is not int or index.value <= 0) + ): + raise exc.CompileError("ARRAY write indices must be positive integers after normalization") + if not isinstance(index.type, (types.Integer, types.NullType, types.ARRAY)): + raise exc.CompileError("ARRAY write indices must be integers") + + def type_bind(element: Any, **kwargs: Any) -> Any: + if isinstance(element, BindParameter): + return element._with_binary_element_type(_IndexType()) + return None + + index = visitors.replacement_traverse(index, {}, type_bind) + sql = compiler.process(index, **kw) + failure = ( + f"CAST(concat('Invalid ARRAY index: ', coalesce(CAST({sql} AS VARCHAR), 'NULL')) AS BIGINT)" + ) + return f"IF({sql} > 0, {sql}, {failure})" + + +def compile_array_update(compiler, expression, **kw): + value = expression.value + final_slice = isinstance(expression.path[-1], Slice) + if final_slice and ( + isinstance(value, Null) + or (isinstance(value, BindParameter) and not value.required and value.value is None) + ): + raise exc.CompileError("An ARRAY slice assignment requires a non-NULL array") + rhs = compiler.process(value, **kw) + rhs_type = compiler._complex_dml_type(expression.value_type) + rhs = f"CAST({rhs} AS {rhs_type})" + if final_slice: + # Reject SQL expressions that evaluate to NULL without issuing a second statement. + failure = ( + f"slice(CAST(ARRAY[] AS {rhs_type}), " + "CAST(concat('NULL ARRAY slice assignment', coalesce(CAST(cardinality(" + f"{rhs}) AS VARCHAR), '')) AS BIGINT), 0)" + ) + rhs = f"IF({rhs} IS NULL, {failure}, {rhs})" + + def rebuild(array, array_type, path): + array_sql_type = compiler._complex_dml_type(array_type) + array = f"coalesce({array}, CAST(ARRAY[] AS {array_sql_type}))" + bound = path[0] + if isinstance(bound, Slice): + if not isinstance(bound.step, Null) and not ( + isinstance(bound.step, BindParameter) + and bound.step.unique + and type(bound.step.value) is int + and bound.step.value == 1 + ): + raise exc.CompileError("Athena ARRAY slices support only step=None or step=1") + start = ( + "1" if isinstance(bound.start, Null) else _index_sql(compiler, bound.start, **kw) + ) + stop = ( + f"cardinality({array})" + if isinstance(bound.stop, Null) + else _index_sql(compiler, bound.stop, **kw) + ) + prefix = f"slice({array}, 1, least({start} - 1, cardinality({array})))" + element_type = compiler._complex_dml_type(_array_item_type(array_type)) + padding = ( + f"repeat(CAST(NULL AS {element_type}), " + f"greatest({start} - 1 - cardinality({array}), 0))" + ) + tail_start = f"greatest({start}, {stop} + 1)" + suffix = ( + f"slice({array}, {tail_start}, " + f"greatest(cardinality({array}) - {tail_start} + 1, 0))" + ) + return f"concat({prefix}, {padding}, {rhs}, {suffix})" + index = _index_sql(compiler, bound, **kw) + variable = compiler._array_lambda_name() + previous = f"element_at({array}, {index})" + replacement = ( + rebuild(previous, _array_item_type(array_type), path[1:]) if len(path) > 1 else rhs + ) + return ( + f"transform(sequence(1, greatest(cardinality({array}), {index})), " + f"{variable} -> IF({variable} = {index}, {replacement}, " + f"element_at({array}, {variable})))" + ) + + return rebuild(compiler.process(expression.column, **kw), expression.type, expression.path) diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index feb54129..4de1543a 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -34,6 +34,7 @@ AthenaRowFormatSerde, ) from pyathena.sqlalchemy.array import _ArraySliceStepType, _ArrayTypeInspector +from pyathena.sqlalchemy._array_update import compile_array_update, rewrite_array_update from pyathena.sqlalchemy.preparer import AthenaDDLIdentifierPreparer from pyathena.sqlalchemy.types import ( AthenaMap, @@ -269,6 +270,14 @@ def _original_froms(elements): while element._is_clone_of is not None: element = element._is_clone_of yield element + def visit_update(self, update_stmt, visiting_cte=None, **kw): + """Rewrite partial array assignments into one native Athena UPDATE.""" + return super().visit_update( + rewrite_array_update(update_stmt), visiting_cte=visiting_cte, **kw + ) + + def visit_athena_array_update(self, expression, **kw): + return compile_array_update(self, expression, **kw) def _array_lambda_name(self): names = { diff --git a/tests/pyathena/sqlalchemy/test_array_update.py b/tests/pyathena/sqlalchemy/test_array_update.py new file mode 100644 index 00000000..26315573 --- /dev/null +++ b/tests/pyathena/sqlalchemy/test_array_update.py @@ -0,0 +1,90 @@ +import pytest +from sqlalchemy import Column, Integer, MetaData, Table, bindparam, types +from sqlalchemy import exc as sa_exc + +from pyathena.formatter import DefaultParameterFormatter +from pyathena.sqlalchemy.base import AthenaDialect +from pyathena.sqlalchemy.types import AthenaArray + + +def array_table(type_=None): + return Table( + "arrays", MetaData(), Column("id", Integer), Column("items", type_ or AthenaArray(Integer)) + ) + + +@pytest.mark.parametrize( + ("target", "value"), + [(1, 2), (4, None), (slice(2, 3), [4]), (slice(2, 2), []), (slice(None), [])], +) +def test_partial_update_compiles_to_one_whole_column_assignment(target, value): + table = array_table() + statement = table.update().values({table.c["items"][target]: value}).where(table.c.id == 1) + original_key = statement._generate_cache_key().key + compiled = statement.compile(dialect=AthenaDialect()) + sql = str(compiled) + assert sql.startswith("UPDATE arrays SET items=") + assert "SET element_at" not in sql + assert "SELECT" not in sql + assert "WHERE arrays.id =" in sql + assert statement._generate_cache_key().key == original_key + parameters = { + name: compiled._bind_processors.get(name, lambda v: v)(value) + for name, value in compiled.params.items() + } + formatted = DefaultParameterFormatter().format(sql, parameters) + assert "ARRAY[" in formatted + + +@pytest.mark.parametrize("index", [0, -1, None, 1.5, True]) +def test_invalid_partial_update_index(index): + table = array_table() + with pytest.raises(sa_exc.CompileError, match="indices"): + table.update().values({table.c["items"][index]: 1}).compile(dialect=AthenaDialect()) + + +def test_multiple_updates_to_one_array_are_rejected(): + table = array_table() + values = table.c["items"] + for assignments in ( + {values[1]: 2, values[2]: 3}, + {values: [], values[1]: 2}, + {values[1]: 2, "items": []}, + ): + with pytest.raises(sa_exc.CompileError, match="one assignment"): + table.update().values(assignments).compile(dialect=AthenaDialect()) + + +def test_slice_null_and_nested_slice_rejected(): + table = array_table(AthenaArray(Integer, dimensions=2)) + with pytest.raises(sa_exc.CompileError, match="non-NULL array"): + table.update().values({table.c["items"][1:2]: None}).compile(dialect=AthenaDialect()) + with pytest.raises(sa_exc.CompileError, match="final"): + table.update().values({table.c["items"][1:2][1]: [2]}).compile(dialect=AthenaDialect()) + + +def test_bound_indices_and_values_are_reused_without_mutation(): + table = array_table() + expression = table.c["items"][bindparam("index")] + statement = table.update().values({expression: bindparam("value"), table.c.id: 2}) + compiled = statement.compile(dialect=AthenaDialect()) + assert set(compiled.params) == {"index", "value", "id"} + assert compiled._bind_processors["index"](2) == 2 + with pytest.raises(ValueError, match="integers"): + compiled._bind_processors["index"](1.5) + assert str(statement.compile(dialect=AthenaDialect())) == str(compiled) + + +def test_nested_and_zero_indexed_update(): + table = array_table(AthenaArray(Integer, dimensions=2, zero_indexes=True)) + statement = table.update().values({table.c["items"][0][2]: 7}) + sql = str(statement.compile(dialect=AthenaDialect(), compile_kwargs={"literal_binds": True})) + assert sql.count("transform(sequence(") == 2 + assert "IF(1 > 0, 1," in sql + assert "IF(3 > 0, 3," in sql + + +def test_generic_array_partial_update(): + table = array_table(types.ARRAY(Integer)) + sql = str(table.update().values({table.c["items"][1]: 2}).compile(dialect=AthenaDialect())) + assert "SET items=transform(" in sql From a322fa11ab07c4968d2d83dded04f8b1abd18346 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 19 Sep 2026 16:36:26 +0900 Subject: [PATCH 02/10] Cover ARRAY update resize and nested assignment behavior --- docs/sqlalchemy.md | 32 +++++++++ pyathena/sqlalchemy/_array_update.py | 4 +- tests/sqlalchemy/test_suite.py | 104 +++++++++++++++++++++++++++ 3 files changed, 139 insertions(+), 1 deletion(-) diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index edff37a7..99757c83 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -1118,6 +1118,38 @@ For example, `~(any_(flags) == True)` can become an element-wise `!= True` compa To negate the whole match, always explicitly group the comparison first: `~(any_(flags) == True).self_group()`. Quantifiers over subqueries retain their usual SQL compilation. +#### Partial ARRAY updates + +Use indexed or sliced columns as UPDATE assignment keys on Iceberg tables. +PyAthena compiles each assignment into a single server-side UPDATE of the whole array. +Other columns can be assigned in the same statement. + +```python +numbers = table.c.numbers +with engine.begin() as conn: + conn.execute(table.update().values({numbers[2]: 10})) + conn.execute(table.update().values({numbers[2:3]: [20, 30, 40]})) +``` + +Element assignment beyond the end extends the array and fills intervening positions with NULL. +A NULL destination array is treated as empty for partial updates. +Assigning `None` to an element stores NULL. +Nested indices rebuild the corresponding inner arrays, including missing inner arrays. +The same one-based default and `zero_indexes=True` translation apply to reads and writes. + +Slice assignment replaces an inclusive range with any number of elements. +An empty replacement array deletes the range; a longer or shorter replacement resizes the array. +A reversed range inserts before its start position. +A start beyond the end pads with NULL before inserting, and a stop beyond the end only removes existing elements. +Omitted boundaries mean the beginning or end. +These resize rules are PyAthena-specific and do not promise full PostgreSQL array-assignment compatibility. + +Write indices and explicit slice boundaries must be non-NULL positive integers after normalization. +A slice replacement must be a non-NULL array; use `[]` to delete elements. +Only `step=None` and `step=1` are supported, and only the final component of a nested update path may be a slice. +PyAthena rejects multiple partial assignments to the same array column, or a partial assignment combined with a whole-column assignment to that column. +Use one whole-array expression when an update needs several changes to the same array. + #### Querying ARRAY data Use `select()` with `ARRAY` or `AthenaArray` columns whose element types are known, either declared explicitly or reflected from Athena, to receive typed Python collections. diff --git a/pyathena/sqlalchemy/_array_update.py b/pyathena/sqlalchemy/_array_update.py index 09e409fd..c9fefaa3 100644 --- a/pyathena/sqlalchemy/_array_update.py +++ b/pyathena/sqlalchemy/_array_update.py @@ -182,7 +182,9 @@ def rebuild(array, array_type, path): f"slice({array}, {tail_start}, " f"greatest(cardinality({array}) - {tail_start} + 1, 0))" ) - return f"concat({prefix}, {padding}, {rhs}, {suffix})" + return compiler._array_slice_step( + f"concat({prefix}, {padding}, {rhs}, {suffix})", bound.step, **kw + ) index = _index_sql(compiler, bound, **kw) variable = compiler._array_lambda_name() previous = f"element_at({array}, {index})" diff --git a/tests/sqlalchemy/test_suite.py b/tests/sqlalchemy/test_suite.py index 64949284..a3cf5eaa 100644 --- a/tests/sqlalchemy/test_suite.py +++ b/tests/sqlalchemy/test_suite.py @@ -161,6 +161,110 @@ def process_result_value(self, value, dialect): return tuple(value) if value is not None else None +class ArrayUpdateTest(fixtures.TestBase): + __backend__ = True + __requires__ = ("array_type",) + + def test_element_resize_and_null(self, connection, metadata): + table = Table( + "array_element_updates", + metadata, + Column("id", Integer), + Column("items", AthenaArray(Integer)), + Column("marker", Integer), + ) + table.create(connection) + connection.execute( + table.insert(), + [ + {"id": 1, "items": [1, 2, 3]}, + {"id": 2, "items": []}, + {"id": 3, "items": None}, + ], + ) + items = table.c["items"] + connection.execute( + table.update().where(table.c.id == 1).values({items[5]: 9, table.c.marker: 42}) + ) + connection.execute(table.update().where(table.c.id == 1).values({items[2]: None})) + connection.execute(table.update().where(table.c.id > 1).values({items[2]: 7})) + eq_( + connection.execute(select(items, table.c.marker).order_by(table.c.id)).all(), + [([1, None, 3, None, 9], 42), ([None, 7], None), ([None, 7], None)], + ) + + def test_slice_resize(self, connection, metadata): + cases = [ + ([1, 2, 3], slice(2, 2), [8, 9], [1, 8, 9, 3]), + ([1, 2, 3], slice(2, 3), [8], [1, 8]), + ([1, 2, 3], slice(2, 3), [], [1]), + ([1, 2, 3], slice(3, 1), [8], [1, 2, 8, 3]), + ([1, 2, 3], slice(5, 9), [8], [1, 2, 3, None, 8]), + ([1, 2, 3], slice(2, 9), [8], [1, 8]), + ([1, 2, 3], slice(None), [8, 9], [8, 9]), + ([], slice(1, 2), [], []), + (None, slice(2, 2), [8], [None, 8]), + ] + table = Table( + "array_slice_updates", + metadata, + Column("id", Integer), + Column("items", AthenaArray(Integer)), + ) + table.create(connection) + connection.execute( + table.insert(), + [{"id": i, "items": before} for i, (before, _, _, _) in enumerate(cases)], + ) + for i, (_, bounds, replacement, _) in enumerate(cases): + connection.execute( + table.update() + .where(table.c.id == i) + .values({table.c["items"][bounds]: replacement}) + ) + eq_( + connection.execute(select(table.c["items"]).order_by(table.c.id)).scalars().all(), + [expected for _, _, _, expected in cases], + ) + + def test_nested_zero_indexed_and_cached_bindings(self, connection, metadata): + table = Table( + "array_nested_updates", + metadata, + Column("id", Integer), + Column("items", AthenaArray(Integer, dimensions=2, zero_indexes=True)), + ) + table.create(connection) + connection.execute( + table.insert(), [{"id": 1, "items": [[1], None]}, {"id": 2, "items": None}] + ) + items = table.c["items"] + statement = ( + table.update() + .where(table.c.id == bindparam("row_id")) + .values({items[bindparam("outer")][bindparam("inner")]: bindparam("value")}) + ) + connection.execute(statement, {"row_id": 1, "outer": 1, "inner": 1, "value": 7}) + connection.execute(statement, {"row_id": 2, "outer": 0, "inner": 0, "value": 8}) + connection.execute(table.update().where(table.c.id == 1).values({items[0][:0]: [4, 5]})) + eq_( + connection.execute(select(items).order_by(table.c.id)).scalars().all(), + [[[4, 5], [None, 7]], [[8]]], + ) + with pytest.raises(sa_exc.DBAPIError): + connection.execute(statement, {"row_id": 1, "outer": -1, "inner": 0, "value": 9}) + + def test_null_slice_binding_rejected(self, connection, metadata): + table = Table("array_null_slice", metadata, Column("items", types.ARRAY(Integer))) + table.create(connection) + connection.execute(table.insert().values(items=[1, 2])) + statement = table.update().values({table.c["items"][1:2]: bindparam("replacement")}) + connection.execute(statement, {"replacement": [3]}) + with pytest.raises(sa_exc.DBAPIError): + connection.execute(statement, {"replacement": None}) + eq_(connection.execute(select(table.c["items"])).scalar_one(), [3]) + + class ArrayExpressionTest(fixtures.TestBase): __backend__ = True __requires__ = ("array_type",) From bc37103db4131bb1e88c9a5776131f5c456f774a Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 19 Sep 2026 16:42:28 +0900 Subject: [PATCH 03/10] Use Athena's integer repeat count for ARRAY slice padding --- pyathena/sqlalchemy/_array_update.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyathena/sqlalchemy/_array_update.py b/pyathena/sqlalchemy/_array_update.py index c9fefaa3..fddb20ac 100644 --- a/pyathena/sqlalchemy/_array_update.py +++ b/pyathena/sqlalchemy/_array_update.py @@ -175,7 +175,7 @@ def rebuild(array, array_type, path): element_type = compiler._complex_dml_type(_array_item_type(array_type)) padding = ( f"repeat(CAST(NULL AS {element_type}), " - f"greatest({start} - 1 - cardinality({array}), 0))" + f"CAST(greatest({start} - 1 - cardinality({array}), 0) AS INTEGER))" ) tail_start = f"greatest({start}, {stop} + 1)" suffix = ( From b91e309b617244de324f2c3ee9b6d0b774f26ec3 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 19 Sep 2026 16:47:03 +0900 Subject: [PATCH 04/10] Preserve ARRAY update expression arguments and binary values --- pyathena/formatter.py | 2 +- pyathena/sqlalchemy/_array_update.py | 19 +++++----- .../pyathena/sqlalchemy/test_array_update.py | 35 ++++++++++++++++++- tests/sqlalchemy/test_suite.py | 21 +++++++++++ 4 files changed, 67 insertions(+), 10 deletions(-) diff --git a/pyathena/formatter.py b/pyathena/formatter.py index f4f5caea..0c6a7c13 100644 --- a/pyathena/formatter.py +++ b/pyathena/formatter.py @@ -22,7 +22,7 @@ class _ComplexParameter: """Typed complex value supplied by the SQLAlchemy dialect.""" - constructor: Literal["ARRAY", "MAP", "ROW", "JSON_PARSE"] + constructor: Literal["ARRAY", "MAP", "ROW", "JSON_PARSE", "FROM_HEX"] values: tuple[Any, ...] diff --git a/pyathena/sqlalchemy/_array_update.py b/pyathena/sqlalchemy/_array_update.py index fddb20ac..fa7fda04 100644 --- a/pyathena/sqlalchemy/_array_update.py +++ b/pyathena/sqlalchemy/_array_update.py @@ -3,11 +3,12 @@ from typing import Any from sqlalchemy import exc, types, util -from sqlalchemy.sql import operators, visitors +from sqlalchemy.sql import operators from sqlalchemy.sql.elements import BinaryExpression, BindParameter, ColumnElement, Null, Slice from sqlalchemy.sql.schema import Column from sqlalchemy.sql.visitors import InternalTraversal +from pyathena.formatter import _ComplexParameter from pyathena.sqlalchemy.types import _array_item_type, _bind_complex, _literal_complex @@ -20,7 +21,13 @@ def __init__(self, item_type): self.item_type = item_type def bind_processor(self, dialect): - return lambda value: _bind_complex(value, self.item_type, dialect) + def process(value): + value = _bind_complex(value, self.item_type, dialect) + if isinstance(value, (bytes, bytearray)): + return _ComplexParameter("FROM_HEX", (value.hex(),)) + return value + + return process def literal_processor(self, dialect): return lambda value: _literal_complex(value, self.item_type, dialect) @@ -118,12 +125,8 @@ def _index_sql(compiler, index: ColumnElement[Any], **kw): if not isinstance(index.type, (types.Integer, types.NullType, types.ARRAY)): raise exc.CompileError("ARRAY write indices must be integers") - def type_bind(element: Any, **kwargs: Any) -> Any: - if isinstance(element, BindParameter): - return element._with_binary_element_type(_IndexType()) - return None - - index = visitors.replacement_traverse(index, {}, type_bind) + if isinstance(index, BindParameter): + index = index._with_binary_element_type(_IndexType()) sql = compiler.process(index, **kw) failure = ( f"CAST(concat('Invalid ARRAY index: ', coalesce(CAST({sql} AS VARCHAR), 'NULL')) AS BIGINT)" diff --git a/tests/pyathena/sqlalchemy/test_array_update.py b/tests/pyathena/sqlalchemy/test_array_update.py index 26315573..d249b6a2 100644 --- a/tests/pyathena/sqlalchemy/test_array_update.py +++ b/tests/pyathena/sqlalchemy/test_array_update.py @@ -1,5 +1,5 @@ import pytest -from sqlalchemy import Column, Integer, MetaData, Table, bindparam, types +from sqlalchemy import Column, Integer, MetaData, Table, bindparam, func, types from sqlalchemy import exc as sa_exc from pyathena.formatter import DefaultParameterFormatter @@ -88,3 +88,36 @@ def test_generic_array_partial_update(): table = array_table(types.ARRAY(Integer)) sql = str(table.update().values({table.c["items"][1]: 2}).compile(dialect=AthenaDialect())) assert "SET items=transform(" in sql + + +def test_write_index_expression_keeps_its_argument_types(): + table = array_table() + index = func.length("abc") + statement = table.update().values({table.c["items"][index]: 9}) + compiled = statement.compile(dialect=AthenaDialect()) + params = { + name: compiled._bind_processors.get(name, lambda value: value)(value) + for name, value in compiled.params.items() + } + assert "length('abc')" in DefaultParameterFormatter().format(str(compiled), params) + + +def test_ordered_partial_update_with_sql_expression(): + table = array_table() + items = table.c["items"] + statement = table.update().ordered_values((items[2], items[1] + 1), (table.c.id, 2)) + compiled = str(statement.compile(dialect=AthenaDialect())) + assert compiled.index("SET items=") < compiled.index(", id=") + assert "element_at(arrays.items" in compiled + + +def test_binary_element_assignment_uses_native_hex_parameter(): + table = array_table(AthenaArray(types.BINARY)) + compiled = ( + table.update().values({table.c["items"][1]: b"\x00\xff"}).compile(dialect=AthenaDialect()) + ) + params = { + name: compiled._bind_processors.get(name, lambda value: value)(value) + for name, value in compiled.params.items() + } + assert "FROM_HEX('00ff')" in DefaultParameterFormatter().format(str(compiled), params) diff --git a/tests/sqlalchemy/test_suite.py b/tests/sqlalchemy/test_suite.py index a3cf5eaa..2029cb25 100644 --- a/tests/sqlalchemy/test_suite.py +++ b/tests/sqlalchemy/test_suite.py @@ -254,6 +254,27 @@ def test_nested_zero_indexed_and_cached_bindings(self, connection, metadata): with pytest.raises(sa_exc.DBAPIError): connection.execute(statement, {"row_id": 1, "outer": -1, "inner": 0, "value": 9}) + def test_expression_values_and_indices(self, connection, metadata): + table = Table( + "array_expression_updates", + metadata, + Column("id", Integer), + Column("items", AthenaArray(Integer)), + Column("binary_items", AthenaArray(types.BINARY)), + ) + table.create(connection) + connection.execute(table.insert().values(id=1, items=[1, 2, 3], binary_items=[b"abc"])) + items = table.c["items"] + connection.execute( + table.update().ordered_values( + (items[func.length("abc")], items[1] + 8), + (table.c.binary_items[1], b"\x00\xff"), + (table.c.id, 2), + ) + ) + connection.execute(table.update().values({items[1:2]: items[2:3].concat([4])})) + eq_(connection.execute(select(table)).one(), (2, [2, 9, 4, 9], [b"\x00\xff"])) + def test_null_slice_binding_rejected(self, connection, metadata): table = Table("array_null_slice", metadata, Column("items", types.ARRAY(Integer))) table.create(connection) From 13a907dc7f663fd85df4f957f0901037d5e09125 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 19 Sep 2026 16:49:23 +0900 Subject: [PATCH 05/10] Use the shared ARRAY step validation in partial updates --- pyathena/sqlalchemy/_array_update.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyathena/sqlalchemy/_array_update.py b/pyathena/sqlalchemy/_array_update.py index fa7fda04..3b14f44c 100644 --- a/pyathena/sqlalchemy/_array_update.py +++ b/pyathena/sqlalchemy/_array_update.py @@ -186,7 +186,7 @@ def rebuild(array, array_type, path): f"greatest(cardinality({array}) - {tail_start} + 1, 0))" ) return compiler._array_slice_step( - f"concat({prefix}, {padding}, {rhs}, {suffix})", bound.step, **kw + f"concat({prefix}, {padding}, {rhs}, {suffix})", bound.step, array_type, **kw ) index = _index_sql(compiler, bound, **kw) variable = compiler._array_lambda_name() From 23f7b190c520e7a42ba5eb9a48a15a3b6d4a122b Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Sat, 19 Sep 2026 17:06:05 +0900 Subject: [PATCH 06/10] Support ORM ARRAY updates and preserve typed assignment expressions --- docs/sqlalchemy.md | 1 + pyathena/sqlalchemy/_array_update.py | 33 +++++++++--- .../pyathena/sqlalchemy/test_array_update.py | 53 +++++++++++++++++-- tests/sqlalchemy/test_suite.py | 35 ++++++++++++ 4 files changed, 111 insertions(+), 11 deletions(-) diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index 99757c83..ca8a3ca2 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -1132,6 +1132,7 @@ with engine.begin() as conn: ``` Element assignment beyond the end extends the array and fills intervening positions with NULL. +NULL padding uses Athena's `repeat()` function and follows its size limits. A NULL destination array is treated as empty for partial updates. Assigning `None` to an element stores NULL. Nested indices rebuild the corresponding inner arrays, including missing inner arrays. diff --git a/pyathena/sqlalchemy/_array_update.py b/pyathena/sqlalchemy/_array_update.py index 3b14f44c..19c0da0b 100644 --- a/pyathena/sqlalchemy/_array_update.py +++ b/pyathena/sqlalchemy/_array_update.py @@ -32,6 +32,10 @@ def process(value): def literal_processor(self, dialect): return lambda value: _literal_complex(value, self.item_type, dialect) + def bind_expression(self, bindvalue): + expression = self.item_type.bind_expression(bindvalue) + return bindvalue if expression is None else expression + class _IndexType(types.TypeDecorator[int]): impl = types.Integer @@ -62,7 +66,9 @@ def __init__(self, column, path, value, value_type): self.type = column.type self.value_type = value_type self.value = ( - value._with_binary_element_type(_AssignmentType(value_type)) + value._with_binary_element_type( + _AssignmentType(value_type if value.type._isnull else value.type) + ) if isinstance(value, BindParameter) else value ) @@ -89,7 +95,10 @@ def rewrite_array_update(statement): base = base.left name = base if isinstance(base, str) else getattr(base, "key", None) if path: - if not isinstance(base, Column) or base.table is not statement.table: + if ( + not isinstance(base, Column) + or base.table._deannotate() is not statement.table._deannotate() + ): raise exc.CompileError("ARRAY updates require a column of the target table") if name in seen: raise exc.CompileError("Only one assignment per ARRAY column is supported") @@ -119,6 +128,7 @@ def _index_sql(compiler, index: ColumnElement[Any], **kw): if ( isinstance(index, BindParameter) and not index.required + and index.callable is None and (type(index.value) is not int or index.value <= 0) ): raise exc.CompileError("ARRAY write indices must be positive integers after normalization") @@ -139,7 +149,12 @@ def compile_array_update(compiler, expression, **kw): final_slice = isinstance(expression.path[-1], Slice) if final_slice and ( isinstance(value, Null) - or (isinstance(value, BindParameter) and not value.required and value.value is None) + or ( + isinstance(value, BindParameter) + and not value.required + and value.callable is None + and value.value is None + ) ): raise exc.CompileError("An ARRAY slice assignment requires a non-NULL array") rhs = compiler.process(value, **kw) @@ -189,15 +204,17 @@ def rebuild(array, array_type, path): f"concat({prefix}, {padding}, {rhs}, {suffix})", bound.step, array_type, **kw ) index = _index_sql(compiler, bound, **kw) - variable = compiler._array_lambda_name() previous = f"element_at({array}, {index})" replacement = ( rebuild(previous, _array_item_type(array_type), path[1:]) if len(path) > 1 else rhs ) - return ( - f"transform(sequence(1, greatest(cardinality({array}), {index})), " - f"{variable} -> IF({variable} = {index}, {replacement}, " - f"element_at({array}, {variable})))" + prefix = f"slice({array}, 1, least({index} - 1, cardinality({array})))" + element_type = compiler._complex_dml_type(_array_item_type(array_type)) + padding = ( + f"repeat(CAST(NULL AS {element_type}), " + f"CAST(greatest({index} - 1 - cardinality({array}), 0) AS INTEGER))" ) + suffix = f"slice({array}, {index} + 1, greatest(cardinality({array}) - {index}, 0))" + return f"concat({prefix}, {padding}, ARRAY[{replacement}], {suffix})" return rebuild(compiler.process(expression.column, **kw), expression.type, expression.path) diff --git a/tests/pyathena/sqlalchemy/test_array_update.py b/tests/pyathena/sqlalchemy/test_array_update.py index d249b6a2..ebb62d28 100644 --- a/tests/pyathena/sqlalchemy/test_array_update.py +++ b/tests/pyathena/sqlalchemy/test_array_update.py @@ -1,6 +1,7 @@ import pytest -from sqlalchemy import Column, Integer, MetaData, Table, bindparam, func, types +from sqlalchemy import Column, Integer, MetaData, Table, bindparam, func, types, update from sqlalchemy import exc as sa_exc +from sqlalchemy.orm import declarative_base from pyathena.formatter import DefaultParameterFormatter from pyathena.sqlalchemy.base import AthenaDialect @@ -79,7 +80,8 @@ def test_nested_and_zero_indexed_update(): table = array_table(AthenaArray(Integer, dimensions=2, zero_indexes=True)) statement = table.update().values({table.c["items"][0][2]: 7}) sql = str(statement.compile(dialect=AthenaDialect(), compile_kwargs={"literal_binds": True})) - assert sql.count("transform(sequence(") == 2 + assert "ARRAY[concat(" in sql + assert "sequence(" not in sql assert "IF(1 > 0, 1," in sql assert "IF(3 > 0, 3," in sql @@ -87,7 +89,7 @@ def test_nested_and_zero_indexed_update(): def test_generic_array_partial_update(): table = array_table(types.ARRAY(Integer)) sql = str(table.update().values({table.c["items"][1]: 2}).compile(dialect=AthenaDialect())) - assert "SET items=transform(" in sql + assert "SET items=concat(" in sql def test_write_index_expression_keeps_its_argument_types(): @@ -121,3 +123,48 @@ def test_binary_element_assignment_uses_native_hex_parameter(): for name, value in compiled.params.items() } assert "FROM_HEX('00ff')" in DefaultParameterFormatter().format(str(compiled), params) + + +class PrefixString(types.TypeDecorator): + impl = types.String + cache_ok = True + + def process_bind_param(self, value, dialect): + return f"prefix:{value}" + + def bind_expression(self, bindvalue): + return func.upper(bindvalue) + + +def test_explicit_assignment_type_and_callable_bindings(): + table = array_table(AthenaArray(types.String)) + stmt = table.update().values( + { + table.c["items"][bindparam("index", callable_=lambda: 1)]: bindparam( + "value", type_=PrefixString(), callable_=lambda: "a" + ) + } + ) + compiled = stmt.compile(dialect=AthenaDialect()) + assert compiled._bind_processors["value"]("a") == "prefix:a" + assert "upper(%(value)s)" in str(compiled) + assert compiled.params["index"] == 1 + assert compiled.params["value"] == "a" + table.update().values( + {table.c["items"][1:2]: bindparam("values", callable_=lambda: ["a"])} + ).compile(dialect=AthenaDialect()) + + +def test_orm_partial_update_and_renamed_attribute_conflicts(): + base = declarative_base() + + class Model(base): + __tablename__ = "arrays" + id = Column(Integer, primary_key=True) + values = Column("stored", AthenaArray(Integer), key="db_key") + + sql = str(update(Model).values({Model.values[1]: 2}).compile(dialect=AthenaDialect())) + assert "UPDATE arrays SET stored=concat(" in sql + for whole in (Model.values, "values"): + with pytest.raises(sa_exc.CompileError, match="one assignment"): + update(Model).values({Model.values[1]: 2, whole: []}).compile(dialect=AthenaDialect()) diff --git a/tests/sqlalchemy/test_suite.py b/tests/sqlalchemy/test_suite.py index 2029cb25..2a12dd4d 100644 --- a/tests/sqlalchemy/test_suite.py +++ b/tests/sqlalchemy/test_suite.py @@ -23,10 +23,12 @@ select, text, types, + update, ) from sqlalchemy import Table as SATable from sqlalchemy import exc as sa_exc from sqlalchemy import testing as sa_testing +from sqlalchemy.orm import Session, registry from sqlalchemy.sql.elements import quoted_name from sqlalchemy.testing import eq_, fixtures from sqlalchemy.testing.schema import Column, Table @@ -275,6 +277,39 @@ def test_expression_values_and_indices(self, connection, metadata): connection.execute(table.update().values({items[1:2]: items[2:3].concat([4])})) eq_(connection.execute(select(table)).one(), (2, [2, 9, 4, 9], [b"\x00\xff"])) + def test_orm_and_long_array_update(self, connection, metadata): + table = Table( + "array_orm_updates", + metadata, + Column("id", Integer, primary_key=True), + Column("items", AthenaArray(Integer)), + ) + table.create(connection) + connection.execute( + table.insert().values( + id=1, + items=func.concat(func.sequence(1, 10000), literal([10001], AthenaArray(Integer))), + ) + ) + mapping = registry() + + class Record: + pass + + mapping.map_imperatively(Record, table) + try: + with Session(bind=connection) as session: + session.execute(update(Record).where(Record.id == 1).values({Record.items[1]: 99})) + session.flush() + row = connection.execute(select(table.c["items"])).scalar_one() + eq_((len(row), row[0], row[-1]), (10001, 99, 10001)) + connection.execute( + table.update().values({table.c["items"][2]: select(literal(77)).scalar_subquery()}) + ) + eq_(connection.execute(select(table.c["items"])).scalar_one()[:3], [99, 77, 3]) + finally: + mapping.dispose() + def test_null_slice_binding_rejected(self, connection, metadata): table = Table("array_null_slice", metadata, Column("items", types.ARRAY(Integer))) table.create(connection) From 97a5b6b0e0bf4fe3e27a98bb99e2b976b26cda09 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 21 Sep 2026 03:47:16 +0900 Subject: [PATCH 07/10] Organize ARRAY updates by responsibility and preserve typed values --- docs/sqlalchemy.md | 12 +- pyathena/sqlalchemy/_array_update.py | 220 --------------- pyathena/sqlalchemy/array.py | 260 +++++++++++++++++- pyathena/sqlalchemy/compiler.py | 13 +- tests/pyathena/sqlalchemy/test_array.py | 197 +++++++++++++ .../pyathena/sqlalchemy/test_array_update.py | 170 ------------ tests/sqlalchemy/test_suite.py | 27 +- 7 files changed, 491 insertions(+), 408 deletions(-) delete mode 100644 pyathena/sqlalchemy/_array_update.py delete mode 100644 tests/pyathena/sqlalchemy/test_array_update.py diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index ca8a3ca2..7ae07a54 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -1124,11 +1124,16 @@ Use indexed or sliced columns as UPDATE assignment keys on Iceberg tables. PyAthena compiles each assignment into a single server-side UPDATE of the whole array. Other columns can be assigned in the same statement. +For an existing Iceberg table named `orders` with an ARRAY column `item_ids`: + ```python -numbers = table.c.numbers +from sqlalchemy import MetaData, Table + +orders = Table("orders", MetaData(), autoload_with=engine) +item_ids = orders.c.item_ids with engine.begin() as conn: - conn.execute(table.update().values({numbers[2]: 10})) - conn.execute(table.update().values({numbers[2:3]: [20, 30, 40]})) + conn.execute(orders.update().values({item_ids[2]: 10})) + conn.execute(orders.update().values({item_ids[2:3]: [20, 30, 40]})) ``` Element assignment beyond the end extends the array and fills intervening positions with NULL. @@ -1150,6 +1155,7 @@ A slice replacement must be a non-NULL array; use `[]` to delete elements. Only `step=None` and `step=1` are supported, and only the final component of a nested update path may be a slice. PyAthena rejects multiple partial assignments to the same array column, or a partial assignment combined with a whole-column assignment to that column. Use one whole-array expression when an update needs several changes to the same array. +Partial updates can evaluate indices, boundaries, and replacement SQL expressions more than once; use deterministic expressions. #### Querying ARRAY data diff --git a/pyathena/sqlalchemy/_array_update.py b/pyathena/sqlalchemy/_array_update.py deleted file mode 100644 index 19c0da0b..00000000 --- a/pyathena/sqlalchemy/_array_update.py +++ /dev/null @@ -1,220 +0,0 @@ -from __future__ import annotations - -from typing import Any - -from sqlalchemy import exc, types, util -from sqlalchemy.sql import operators -from sqlalchemy.sql.elements import BinaryExpression, BindParameter, ColumnElement, Null, Slice -from sqlalchemy.sql.schema import Column -from sqlalchemy.sql.visitors import InternalTraversal - -from pyathena.formatter import _ComplexParameter -from pyathena.sqlalchemy.types import _array_item_type, _bind_complex, _literal_complex - - -class _AssignmentType(types.TypeDecorator[Any]): - impl = types.NullType - cache_ok = True - - def __init__(self, item_type): - super().__init__() - self.item_type = item_type - - def bind_processor(self, dialect): - def process(value): - value = _bind_complex(value, self.item_type, dialect) - if isinstance(value, (bytes, bytearray)): - return _ComplexParameter("FROM_HEX", (value.hex(),)) - return value - - return process - - def literal_processor(self, dialect): - return lambda value: _literal_complex(value, self.item_type, dialect) - - def bind_expression(self, bindvalue): - expression = self.item_type.bind_expression(bindvalue) - return bindvalue if expression is None else expression - - -class _IndexType(types.TypeDecorator[int]): - impl = types.Integer - cache_ok = True - - def process_bind_param(self, value, dialect): - if type(value) is not int: - raise ValueError("ARRAY write indices must be non-NULL integers") - return value - - def process_literal_param(self, value, dialect): - return self.process_bind_param(value, dialect) - - -class _ArrayUpdate(ColumnElement[Any]): - __visit_name__ = "athena_array_update" - inherit_cache = True - _traverse_internals = [ # noqa: RUF012 - ("column", InternalTraversal.dp_clauseelement), - ("path", InternalTraversal.dp_clauseelement_list), - ("value", InternalTraversal.dp_clauseelement), - ("type", InternalTraversal.dp_type), - ] - - def __init__(self, column, path, value, value_type): - self.column = column - self.path = path - self.type = column.type - self.value_type = value_type - self.value = ( - value._with_binary_element_type( - _AssignmentType(value_type if value.type._isnull else value.type) - ) - if isinstance(value, BindParameter) - else value - ) - - @property - def _from_objects(self): - return self.column._from_objects + self.value._from_objects - - -def rewrite_array_update(statement): - values = statement._ordered_values - if values is None: - values = list((statement._values or {}).items()) - rewritten = [] - seen = set() - partial = set() - for key, value in values: - base = key - path: list[Any] = [] - while isinstance(base, BinaryExpression) and base.operator is operators.getitem: - if not isinstance(base.left.type, types.ARRAY): - break - path.insert(0, base.right) - base = base.left - name = base if isinstance(base, str) else getattr(base, "key", None) - if path: - if ( - not isinstance(base, Column) - or base.table._deannotate() is not statement.table._deannotate() - ): - raise exc.CompileError("ARRAY updates require a column of the target table") - if name in seen: - raise exc.CompileError("Only one assignment per ARRAY column is supported") - if any(isinstance(index, Slice) for index in path[:-1]): - raise exc.CompileError("Only the final ARRAY update index can be a slice") - partial.add(name) - value_type = key.type - value = _ArrayUpdate(base, path, value, value_type) - key = base - elif name in partial: - raise exc.CompileError("Only one assignment per ARRAY column is supported") - seen.add(name) - rewritten.append((key, value)) - if not partial: - return statement - result = statement._clone() - if statement._ordered_values is not None: - result._ordered_values = rewritten - else: - result._values = util.immutabledict(rewritten) - return result - - -def _index_sql(compiler, index: ColumnElement[Any], **kw): - if isinstance(index, Null): - raise exc.CompileError("ARRAY write indices must be non-NULL positive integers") - if ( - isinstance(index, BindParameter) - and not index.required - and index.callable is None - and (type(index.value) is not int or index.value <= 0) - ): - raise exc.CompileError("ARRAY write indices must be positive integers after normalization") - if not isinstance(index.type, (types.Integer, types.NullType, types.ARRAY)): - raise exc.CompileError("ARRAY write indices must be integers") - - if isinstance(index, BindParameter): - index = index._with_binary_element_type(_IndexType()) - sql = compiler.process(index, **kw) - failure = ( - f"CAST(concat('Invalid ARRAY index: ', coalesce(CAST({sql} AS VARCHAR), 'NULL')) AS BIGINT)" - ) - return f"IF({sql} > 0, {sql}, {failure})" - - -def compile_array_update(compiler, expression, **kw): - value = expression.value - final_slice = isinstance(expression.path[-1], Slice) - if final_slice and ( - isinstance(value, Null) - or ( - isinstance(value, BindParameter) - and not value.required - and value.callable is None - and value.value is None - ) - ): - raise exc.CompileError("An ARRAY slice assignment requires a non-NULL array") - rhs = compiler.process(value, **kw) - rhs_type = compiler._complex_dml_type(expression.value_type) - rhs = f"CAST({rhs} AS {rhs_type})" - if final_slice: - # Reject SQL expressions that evaluate to NULL without issuing a second statement. - failure = ( - f"slice(CAST(ARRAY[] AS {rhs_type}), " - "CAST(concat('NULL ARRAY slice assignment', coalesce(CAST(cardinality(" - f"{rhs}) AS VARCHAR), '')) AS BIGINT), 0)" - ) - rhs = f"IF({rhs} IS NULL, {failure}, {rhs})" - - def rebuild(array, array_type, path): - array_sql_type = compiler._complex_dml_type(array_type) - array = f"coalesce({array}, CAST(ARRAY[] AS {array_sql_type}))" - bound = path[0] - if isinstance(bound, Slice): - if not isinstance(bound.step, Null) and not ( - isinstance(bound.step, BindParameter) - and bound.step.unique - and type(bound.step.value) is int - and bound.step.value == 1 - ): - raise exc.CompileError("Athena ARRAY slices support only step=None or step=1") - start = ( - "1" if isinstance(bound.start, Null) else _index_sql(compiler, bound.start, **kw) - ) - stop = ( - f"cardinality({array})" - if isinstance(bound.stop, Null) - else _index_sql(compiler, bound.stop, **kw) - ) - prefix = f"slice({array}, 1, least({start} - 1, cardinality({array})))" - element_type = compiler._complex_dml_type(_array_item_type(array_type)) - padding = ( - f"repeat(CAST(NULL AS {element_type}), " - f"CAST(greatest({start} - 1 - cardinality({array}), 0) AS INTEGER))" - ) - tail_start = f"greatest({start}, {stop} + 1)" - suffix = ( - f"slice({array}, {tail_start}, " - f"greatest(cardinality({array}) - {tail_start} + 1, 0))" - ) - return compiler._array_slice_step( - f"concat({prefix}, {padding}, {rhs}, {suffix})", bound.step, array_type, **kw - ) - index = _index_sql(compiler, bound, **kw) - previous = f"element_at({array}, {index})" - replacement = ( - rebuild(previous, _array_item_type(array_type), path[1:]) if len(path) > 1 else rhs - ) - prefix = f"slice({array}, 1, least({index} - 1, cardinality({array})))" - element_type = compiler._complex_dml_type(_array_item_type(array_type)) - padding = ( - f"repeat(CAST(NULL AS {element_type}), " - f"CAST(greatest({index} - 1 - cardinality({array}), 0) AS INTEGER))" - ) - suffix = f"slice({array}, {index} + 1, greatest(cardinality({array}) - {index}, 0))" - return f"concat({prefix}, {padding}, ARRAY[{replacement}], {suffix})" - - return rebuild(compiler.process(expression.column, **kw), expression.type, expression.path) diff --git a/pyathena/sqlalchemy/array.py b/pyathena/sqlalchemy/array.py index cae77efb..8a78469c 100644 --- a/pyathena/sqlalchemy/array.py +++ b/pyathena/sqlalchemy/array.py @@ -8,9 +8,10 @@ from decimal import Decimal from typing import TYPE_CHECKING, Any -from sqlalchemy import cast, exc, types +from sqlalchemy import cast, exc, types, util from sqlalchemy.sql import operators, sqltypes -from sqlalchemy.sql.elements import ColumnElement, Slice +from sqlalchemy.sql.elements import BinaryExpression, BindParameter, ColumnElement, Null, Slice +from sqlalchemy.sql.schema import Column from sqlalchemy.sql.type_api import TypeEngine from sqlalchemy.sql.visitors import InternalTraversal @@ -220,23 +221,23 @@ def has_unknown_element(type_: TypeEngine[Any]) -> bool: class _ArrayValueProcessor: - """Convert one declared ARRAY type between Python values and Athena transport. + """Convert ARRAY values and their typed elements to and from Athena transport. SQLAlchemy constructs processors per type and dialect. Keep that context here and share the recursive ARRAY/MAP/ROW traversal across bind parameters, SQL literals, and fetched JSON results. """ - def __init__(self, array_type: AthenaArray, dialect: Any) -> None: - self.array_type = array_type + def __init__(self, type_: TypeEngine[Any], dialect: Any) -> None: + self.type_ = type_ self.dialect = dialect self._type_inspector = _ArrayTypeInspector(dialect) def bind(self, value: Any) -> Any: - return self._bind(value, self.array_type) + return self._bind(value, self.type_) def literal(self, value: Any) -> str: - return self._literal(value, self.array_type) + return self._literal(value, self.type_) def result(self, value: Any) -> Any: if value is None: @@ -250,7 +251,11 @@ def result(self, value: Any) -> Any: return value if isinstance(value, dict) and "_pyathena_array" in value: value = value["_pyathena_array"] - return self._decode(value, self.array_type, self.array_type.as_tuple) + return self._decode( + value, + self.type_, + self.type_.as_tuple if isinstance(self.type_, sqltypes.ARRAY) else False, + ) @staticmethod def _complex_values(value: Any, type_: TypeEngine[Any]): @@ -399,3 +404,242 @@ def _decode(self, value: Any, type_: TypeEngine[Any], as_tuple: bool = False) -> processor = type_.dialect_impl(self.dialect).result_processor(self.dialect, None) return processor(value) if processor else value return value + + +class _ArrayAssignmentType(types.TypeDecorator[Any]): + """Preserve declared element processors for ARRAY assignment values.""" + + impl = types.NullType + cache_ok = True + + def __init__(self, item_type): + super().__init__() + self.item_type = item_type + + def bind_processor(self, dialect): + processor = _ArrayValueProcessor(self.item_type, dialect) + + def process(value): + value = processor.bind(value) + if isinstance(value, (bytes, bytearray)): + return _ComplexParameter("FROM_HEX", (value.hex(),)) + return value + + return process + + def literal_processor(self, dialect): + return _ArrayValueProcessor(self.item_type, dialect).literal + + def bind_expression(self, bindvalue): + expression = self.item_type.bind_expression(bindvalue) + return bindvalue if expression is None else expression + + +class _ArrayWriteIndexType(types.TypeDecorator[int]): + """Reject non-integer and NULL bound ARRAY write indices.""" + + impl = types.Integer + cache_ok = True + + def process_bind_param(self, value, dialect): + if type(value) is not int: + raise ValueError("ARRAY write indices must be non-NULL integers") + return value + + def process_literal_param(self, value, dialect): + return self.process_bind_param(value, dialect) + + +class _ArrayUpdate(ColumnElement[Any]): + """Whole-column expression generated from one partial ARRAY assignment.""" + + __visit_name__ = "athena_array_update" + inherit_cache = True + _traverse_internals = [ # noqa: RUF012 + ("column", InternalTraversal.dp_clauseelement), + ("path", InternalTraversal.dp_clauseelement_list), + ("value", InternalTraversal.dp_clauseelement), + ("type", InternalTraversal.dp_type), + ] + + def __init__(self, column, path, value, value_type): + self.column = column + self.path = path + self.type = column.type + self.value_type = value_type + self.value = ( + value._with_binary_element_type( + _ArrayAssignmentType(value_type if value.type._isnull else value.type) + ) + if isinstance(value, BindParameter) + else value + ) + + @property + def _from_objects(self): + return self.column._from_objects + self.value._from_objects + + @classmethod + def rewrite(cls, statement, dialect): + inspector = _ArrayTypeInspector(dialect) + values = statement._ordered_values + if values is None: + values = list((statement._values or {}).items()) + rewritten = [] + seen = set() + partial = set() + for key, value in values: + base = key + path: list[Any] = [] + while isinstance(base, BinaryExpression) and base.operator is operators.getitem: + if inspector.array_type(base.left.type) is None: + break + path.insert(0, base.right) + base = base.left + name = base if isinstance(base, str) else getattr(base, "key", None) + if path: + if ( + not isinstance(base, Column) + or base.table is None + or base.table._deannotate() is not statement.table._deannotate() + ): + raise exc.CompileError("ARRAY updates require a column of the target table") + if name in seen: + raise exc.CompileError("Only one assignment per ARRAY column is supported") + if any(isinstance(index, Slice) for index in path[:-1]): + raise exc.CompileError("Only the final ARRAY update index can be a slice") + partial.add(name) + value_type = key.type + value = cls(base, path, value, value_type) + key = base + elif name in partial: + raise exc.CompileError("Only one assignment per ARRAY column is supported") + seen.add(name) + rewritten.append((key, value)) + if not partial: + return statement + result = statement._clone() + if statement._ordered_values is not None: + result._ordered_values = rewritten + else: + result._values = util.immutabledict(rewritten) + return result + + +class _ArrayUpdateCompiler: + """Render an ARRAY assignment by rebuilding its affected nested arrays.""" + + def __init__(self, compiler): + self.compiler = compiler + self._type_inspector = _ArrayTypeInspector(compiler.dialect) + + def process(self, expression, **kw): + compiler = self.compiler + value = expression.value + final_slice = isinstance(expression.path[-1], Slice) + if final_slice and ( + isinstance(value, Null) + or ( + isinstance(value, BindParameter) + and not value.required + and value.callable is None + and value.value is None + ) + ): + raise exc.CompileError("An ARRAY slice assignment requires a non-NULL array") + rhs = compiler.process(value, **kw) + rhs_type = compiler._complex_dml_type( + expression.value_type, implicit_bind=isinstance(value, BindParameter) + ) + rhs = f"CAST({rhs} AS {rhs_type})" + if final_slice: + # Reject SQL expressions that evaluate to NULL without issuing a second statement. + failure = ( + f"slice(CAST(ARRAY[] AS {rhs_type}), " + "CAST(concat('NULL ARRAY slice assignment', coalesce(CAST(cardinality(" + f"{rhs}) AS VARCHAR), '')) AS BIGINT), 0)" + ) + rhs = f"IF({rhs} IS NULL, {failure}, {rhs})" + return self._rebuild( + compiler.process(expression.column, **kw), expression.type, expression.path, rhs, **kw + ) + + def _index_sql(self, index: ColumnElement[Any], **kw): + compiler = self.compiler + if isinstance(index, Null): + raise exc.CompileError("ARRAY write indices must be non-NULL positive integers") + if ( + isinstance(index, BindParameter) + and not index.required + and index.callable is None + and (type(index.value) is not int or index.value <= 0) + ): + raise exc.CompileError( + "ARRAY write indices must be positive integers after normalization" + ) + if not isinstance(index.type, (types.Integer, types.NullType)) and not ( + isinstance(index, BindParameter) + and self._type_inspector.array_type(index.type) is not None + ): + raise exc.CompileError("ARRAY write indices must be integers") + + if isinstance(index, BindParameter): + index = index._with_binary_element_type(_ArrayWriteIndexType()) + sql = compiler.process(index, **kw) + failure = ( + "CAST(concat('Invalid ARRAY index: ', " + f"coalesce(CAST({sql} AS VARCHAR), 'NULL')) AS BIGINT)" + ) + return f"IF({sql} > 0, {sql}, {failure})" + + def _rebuild(self, array, array_type, path, rhs, **kw): + compiler = self.compiler + array_type = self._type_inspector.array_type(array_type) + if array_type is None: + raise exc.CompileError("Partial ARRAY updates require an ARRAY column type") + array_sql_type = compiler._complex_dml_type(array_type) + array = f"coalesce({array}, CAST(ARRAY[] AS {array_sql_type}))" + bound = path[0] + if isinstance(bound, Slice): + if not isinstance(bound.step, Null) and not ( + isinstance(bound.step, BindParameter) + and bound.step.unique + and type(bound.step.value) is int + and bound.step.value == 1 + ): + raise exc.CompileError("Athena ARRAY slices support only step=None or step=1") + start = "1" if isinstance(bound.start, Null) else self._index_sql(bound.start, **kw) + stop = ( + f"cardinality({array})" + if isinstance(bound.stop, Null) + else self._index_sql(bound.stop, **kw) + ) + prefix = f"slice({array}, 1, least({start} - 1, cardinality({array})))" + element_type = compiler._complex_dml_type(_ArrayTypeInspector.item_type(array_type)) + padding = ( + f"repeat(CAST(NULL AS {element_type}), " + f"CAST(greatest({start} - 1 - cardinality({array}), 0) AS INTEGER))" + ) + tail_start = f"greatest({start}, {stop} + 1)" + suffix = ( + f"slice({array}, {tail_start}, " + f"greatest(cardinality({array}) - {tail_start} + 1, 0))" + ) + return compiler._array_slice_step( + f"concat({prefix}, {padding}, {rhs}, {suffix})", bound.step, array_type, **kw + ) + index = self._index_sql(bound, **kw) + previous = f"element_at({array}, {index})" + replacement = ( + self._rebuild(previous, _ArrayTypeInspector.item_type(array_type), path[1:], rhs, **kw) + if len(path) > 1 + else rhs + ) + prefix = f"slice({array}, 1, least({index} - 1, cardinality({array})))" + element_type = compiler._complex_dml_type(_ArrayTypeInspector.item_type(array_type)) + padding = ( + f"repeat(CAST(NULL AS {element_type}), " + f"CAST(greatest({index} - 1 - cardinality({array}), 0) AS INTEGER))" + ) + suffix = f"slice({array}, {index} + 1, greatest(cardinality({array}) - {index}, 0))" + return f"concat({prefix}, {padding}, ARRAY[{replacement}], {suffix})" diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index 4de1543a..d2f0d573 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -33,8 +33,12 @@ AthenaPartitionTransform, AthenaRowFormatSerde, ) -from pyathena.sqlalchemy.array import _ArraySliceStepType, _ArrayTypeInspector -from pyathena.sqlalchemy._array_update import compile_array_update, rewrite_array_update +from pyathena.sqlalchemy.array import ( + _ArraySliceStepType, + _ArrayTypeInspector, + _ArrayUpdate, + _ArrayUpdateCompiler, +) from pyathena.sqlalchemy.preparer import AthenaDDLIdentifierPreparer from pyathena.sqlalchemy.types import ( AthenaMap, @@ -270,14 +274,15 @@ def _original_froms(elements): while element._is_clone_of is not None: element = element._is_clone_of yield element + def visit_update(self, update_stmt, visiting_cte=None, **kw): """Rewrite partial array assignments into one native Athena UPDATE.""" return super().visit_update( - rewrite_array_update(update_stmt), visiting_cte=visiting_cte, **kw + _ArrayUpdate.rewrite(update_stmt, self.dialect), visiting_cte=visiting_cte, **kw ) def visit_athena_array_update(self, expression, **kw): - return compile_array_update(self, expression, **kw) + return _ArrayUpdateCompiler(self).process(expression, **kw) def _array_lambda_name(self): names = { diff --git a/tests/pyathena/sqlalchemy/test_array.py b/tests/pyathena/sqlalchemy/test_array.py index e52498c1..5ed7d3c2 100644 --- a/tests/pyathena/sqlalchemy/test_array.py +++ b/tests/pyathena/sqlalchemy/test_array.py @@ -18,13 +18,16 @@ bindparam, cast, column, + func, literal, literal_column, select, text, types, + update, ) from sqlalchemy import exc as sa_exc +from sqlalchemy.orm import declarative_base from sqlalchemy.sql import sqltypes import pyathena @@ -83,6 +86,17 @@ def process_result_value(self, value, dialect): return tuple(value) if value is not None else None +class PrefixString(types.TypeDecorator): + impl = types.String + cache_ok = True + + def process_bind_param(self, value, dialect): + return f"prefix:{value}" + + def bind_expression(self, bindvalue): + return func.upper(bindvalue) + + class TestAthenaArray: def test_creation_with_default(self): array_type = AthenaArray() @@ -659,3 +673,186 @@ def test_array_rewrite_requires_labels_for_literal_expressions(self, expression) .compile(dialect=AthenaDialect()) ) assert sql.startswith("SELECT anon_1.value, json_format(") + + +class TestArrayUpdateCompiler: + @staticmethod + def _table(type_=None): + return Table( + "arrays", + MetaData(), + Column("id", Integer), + Column("items", type_ or AthenaArray(Integer)), + ) + + @pytest.mark.parametrize( + ("target", "value"), + [(1, 2), (4, None), (slice(2, 3), [4]), (slice(2, 2), []), (slice(None), [])], + ) + def test_partial_update_compiles_to_one_whole_column_assignment(self, target, value): + table = self._table() + statement = table.update().values({table.c["items"][target]: value}).where(table.c.id == 1) + original_key = statement._generate_cache_key().key + compiled = statement.compile(dialect=AthenaDialect()) + sql = str(compiled) + assert sql.startswith("UPDATE arrays SET items=") + assert "SET element_at" not in sql + assert "SELECT" not in sql + assert "WHERE arrays.id =" in sql + assert statement._generate_cache_key().key == original_key + parameters = { + name: compiled._bind_processors.get(name, lambda v: v)(value) + for name, value in compiled.params.items() + } + formatted = DefaultParameterFormatter().format(sql, parameters) + assert "ARRAY[" in formatted + + @pytest.mark.parametrize("index", [0, -1, None, 1.5, True]) + def test_invalid_partial_update_index(self, index): + table = self._table() + with pytest.raises(sa_exc.CompileError, match="indices"): + table.update().values({table.c["items"][index]: 1}).compile(dialect=AthenaDialect()) + + def test_multiple_updates_to_one_array_are_rejected(self): + table = self._table() + values = table.c["items"] + for assignments in ( + {values[1]: 2, values[2]: 3}, + {values: [], values[1]: 2}, + {values[1]: 2, "items": []}, + ): + with pytest.raises(sa_exc.CompileError, match="one assignment"): + table.update().values(assignments).compile(dialect=AthenaDialect()) + + def test_slice_null_and_nested_slice_rejected(self): + table = self._table(AthenaArray(Integer, dimensions=2)) + with pytest.raises(sa_exc.CompileError, match="non-NULL array"): + table.update().values({table.c["items"][1:2]: None}).compile(dialect=AthenaDialect()) + with pytest.raises(sa_exc.CompileError, match="final"): + table.update().values({table.c["items"][1:2][1]: [2]}).compile(dialect=AthenaDialect()) + + def test_bound_indices_and_values_are_reused_without_mutation(self): + table = self._table() + expression = table.c["items"][bindparam("index")] + statement = table.update().values({expression: bindparam("value"), table.c.id: 2}) + compiled = statement.compile(dialect=AthenaDialect()) + assert set(compiled.params) == {"index", "value", "id"} + assert compiled._bind_processors["index"](2) == 2 + with pytest.raises(ValueError, match="integers"): + compiled._bind_processors["index"](1.5) + assert str(statement.compile(dialect=AthenaDialect())) == str(compiled) + + def test_nested_and_zero_indexed_update(self): + table = self._table(AthenaArray(Integer, dimensions=2, zero_indexes=True)) + statement = table.update().values({table.c["items"][0][2]: 7}) + sql = str( + statement.compile(dialect=AthenaDialect(), compile_kwargs={"literal_binds": True}) + ) + assert "ARRAY[concat(" in sql + assert "sequence(" not in sql + assert "IF(1 > 0, 1," in sql + assert "IF(3 > 0, 3," in sql + + @pytest.mark.parametrize( + "array_type", + [ + types.ARRAY(Integer), + TupleArray(), + AthenaArray(Integer).with_variant(AthenaArray(Integer), "awsathena"), + ], + ) + @pytest.mark.parametrize("index", [1, bindparam("index")]) + def test_array_implementations_partial_update(self, array_type, index): + table = self._table(array_type) + sql = str( + table.update().values({table.c["items"][index]: 2}).compile(dialect=AthenaDialect()) + ) + assert "SET items=concat(" in sql + + @pytest.mark.parametrize("target", [1, slice(1, 2)]) + def test_decimal_assignment_requires_precision(self, target): + value = [Decimal("1.23")] if isinstance(target, slice) else Decimal("1.23") + table = self._table(AthenaArray(types.Numeric())) + with pytest.raises(sa_exc.CompileError, match="precision"): + table.update().values({table.c["items"][target]: value}).compile( + dialect=AthenaDialect() + ) + table = self._table(AthenaArray(types.Numeric(8, 2))) + sql = str( + table.update() + .values({table.c["items"][target]: value}) + .compile(dialect=AthenaDialect()) + ) + assert "DECIMAL(8, 2)" in sql + + def test_partial_update_requires_target_table_column(self): + table = self._table() + for column_ in (Column("items", AthenaArray(Integer)), self._table().c["items"]): + with pytest.raises(sa_exc.CompileError, match="target table"): + table.update().values({column_[1]: 2}).compile(dialect=AthenaDialect()) + + def test_write_index_expression_keeps_its_argument_types(self): + table = self._table() + index = func.length("abc") + statement = table.update().values({table.c["items"][index]: 9}) + compiled = statement.compile(dialect=AthenaDialect()) + params = { + name: compiled._bind_processors.get(name, lambda value: value)(value) + for name, value in compiled.params.items() + } + assert "length('abc')" in DefaultParameterFormatter().format(str(compiled), params) + + def test_ordered_partial_update_with_sql_expression(self): + table = self._table() + items = table.c["items"] + statement = table.update().ordered_values((items[2], items[1] + 1), (table.c.id, 2)) + compiled = str(statement.compile(dialect=AthenaDialect())) + assert compiled.index("SET items=") < compiled.index(", id=") + assert "element_at(arrays.items" in compiled + + def test_binary_element_assignment_uses_native_hex_parameter(self): + table = self._table(AthenaArray(types.BINARY)) + compiled = ( + table.update() + .values({table.c["items"][1]: b"\x00\xff"}) + .compile(dialect=AthenaDialect()) + ) + params = { + name: compiled._bind_processors.get(name, lambda value: value)(value) + for name, value in compiled.params.items() + } + assert "FROM_HEX('00ff')" in DefaultParameterFormatter().format(str(compiled), params) + + def test_explicit_assignment_type_and_callable_bindings(self): + table = self._table(AthenaArray(types.String)) + stmt = table.update().values( + { + table.c["items"][bindparam("index", callable_=lambda: 1)]: bindparam( + "value", type_=PrefixString(), callable_=lambda: "a" + ) + } + ) + compiled = stmt.compile(dialect=AthenaDialect()) + assert compiled._bind_processors["value"]("a") == "prefix:a" + assert "upper(%(value)s)" in str(compiled) + assert compiled.params["index"] == 1 + assert compiled.params["value"] == "a" + table.update().values( + {table.c["items"][1:2]: bindparam("values", callable_=lambda: ["a"])} + ).compile(dialect=AthenaDialect()) + + def test_orm_partial_update_and_renamed_attribute_conflicts(self): + base = declarative_base() + + class Model(base): + __tablename__ = "arrays" + id = Column(Integer, primary_key=True) + values = Column("stored", AthenaArray(Integer), key="db_key") + + sql = str(update(Model).values({Model.values[1]: 2}).compile(dialect=AthenaDialect())) + assert "UPDATE arrays SET stored=concat(" in sql + for whole in (Model.values, "values"): + with pytest.raises(sa_exc.CompileError, match="one assignment"): + update(Model).values({Model.values[1]: 2, whole: []}).compile( + dialect=AthenaDialect() + ) diff --git a/tests/pyathena/sqlalchemy/test_array_update.py b/tests/pyathena/sqlalchemy/test_array_update.py deleted file mode 100644 index ebb62d28..00000000 --- a/tests/pyathena/sqlalchemy/test_array_update.py +++ /dev/null @@ -1,170 +0,0 @@ -import pytest -from sqlalchemy import Column, Integer, MetaData, Table, bindparam, func, types, update -from sqlalchemy import exc as sa_exc -from sqlalchemy.orm import declarative_base - -from pyathena.formatter import DefaultParameterFormatter -from pyathena.sqlalchemy.base import AthenaDialect -from pyathena.sqlalchemy.types import AthenaArray - - -def array_table(type_=None): - return Table( - "arrays", MetaData(), Column("id", Integer), Column("items", type_ or AthenaArray(Integer)) - ) - - -@pytest.mark.parametrize( - ("target", "value"), - [(1, 2), (4, None), (slice(2, 3), [4]), (slice(2, 2), []), (slice(None), [])], -) -def test_partial_update_compiles_to_one_whole_column_assignment(target, value): - table = array_table() - statement = table.update().values({table.c["items"][target]: value}).where(table.c.id == 1) - original_key = statement._generate_cache_key().key - compiled = statement.compile(dialect=AthenaDialect()) - sql = str(compiled) - assert sql.startswith("UPDATE arrays SET items=") - assert "SET element_at" not in sql - assert "SELECT" not in sql - assert "WHERE arrays.id =" in sql - assert statement._generate_cache_key().key == original_key - parameters = { - name: compiled._bind_processors.get(name, lambda v: v)(value) - for name, value in compiled.params.items() - } - formatted = DefaultParameterFormatter().format(sql, parameters) - assert "ARRAY[" in formatted - - -@pytest.mark.parametrize("index", [0, -1, None, 1.5, True]) -def test_invalid_partial_update_index(index): - table = array_table() - with pytest.raises(sa_exc.CompileError, match="indices"): - table.update().values({table.c["items"][index]: 1}).compile(dialect=AthenaDialect()) - - -def test_multiple_updates_to_one_array_are_rejected(): - table = array_table() - values = table.c["items"] - for assignments in ( - {values[1]: 2, values[2]: 3}, - {values: [], values[1]: 2}, - {values[1]: 2, "items": []}, - ): - with pytest.raises(sa_exc.CompileError, match="one assignment"): - table.update().values(assignments).compile(dialect=AthenaDialect()) - - -def test_slice_null_and_nested_slice_rejected(): - table = array_table(AthenaArray(Integer, dimensions=2)) - with pytest.raises(sa_exc.CompileError, match="non-NULL array"): - table.update().values({table.c["items"][1:2]: None}).compile(dialect=AthenaDialect()) - with pytest.raises(sa_exc.CompileError, match="final"): - table.update().values({table.c["items"][1:2][1]: [2]}).compile(dialect=AthenaDialect()) - - -def test_bound_indices_and_values_are_reused_without_mutation(): - table = array_table() - expression = table.c["items"][bindparam("index")] - statement = table.update().values({expression: bindparam("value"), table.c.id: 2}) - compiled = statement.compile(dialect=AthenaDialect()) - assert set(compiled.params) == {"index", "value", "id"} - assert compiled._bind_processors["index"](2) == 2 - with pytest.raises(ValueError, match="integers"): - compiled._bind_processors["index"](1.5) - assert str(statement.compile(dialect=AthenaDialect())) == str(compiled) - - -def test_nested_and_zero_indexed_update(): - table = array_table(AthenaArray(Integer, dimensions=2, zero_indexes=True)) - statement = table.update().values({table.c["items"][0][2]: 7}) - sql = str(statement.compile(dialect=AthenaDialect(), compile_kwargs={"literal_binds": True})) - assert "ARRAY[concat(" in sql - assert "sequence(" not in sql - assert "IF(1 > 0, 1," in sql - assert "IF(3 > 0, 3," in sql - - -def test_generic_array_partial_update(): - table = array_table(types.ARRAY(Integer)) - sql = str(table.update().values({table.c["items"][1]: 2}).compile(dialect=AthenaDialect())) - assert "SET items=concat(" in sql - - -def test_write_index_expression_keeps_its_argument_types(): - table = array_table() - index = func.length("abc") - statement = table.update().values({table.c["items"][index]: 9}) - compiled = statement.compile(dialect=AthenaDialect()) - params = { - name: compiled._bind_processors.get(name, lambda value: value)(value) - for name, value in compiled.params.items() - } - assert "length('abc')" in DefaultParameterFormatter().format(str(compiled), params) - - -def test_ordered_partial_update_with_sql_expression(): - table = array_table() - items = table.c["items"] - statement = table.update().ordered_values((items[2], items[1] + 1), (table.c.id, 2)) - compiled = str(statement.compile(dialect=AthenaDialect())) - assert compiled.index("SET items=") < compiled.index(", id=") - assert "element_at(arrays.items" in compiled - - -def test_binary_element_assignment_uses_native_hex_parameter(): - table = array_table(AthenaArray(types.BINARY)) - compiled = ( - table.update().values({table.c["items"][1]: b"\x00\xff"}).compile(dialect=AthenaDialect()) - ) - params = { - name: compiled._bind_processors.get(name, lambda value: value)(value) - for name, value in compiled.params.items() - } - assert "FROM_HEX('00ff')" in DefaultParameterFormatter().format(str(compiled), params) - - -class PrefixString(types.TypeDecorator): - impl = types.String - cache_ok = True - - def process_bind_param(self, value, dialect): - return f"prefix:{value}" - - def bind_expression(self, bindvalue): - return func.upper(bindvalue) - - -def test_explicit_assignment_type_and_callable_bindings(): - table = array_table(AthenaArray(types.String)) - stmt = table.update().values( - { - table.c["items"][bindparam("index", callable_=lambda: 1)]: bindparam( - "value", type_=PrefixString(), callable_=lambda: "a" - ) - } - ) - compiled = stmt.compile(dialect=AthenaDialect()) - assert compiled._bind_processors["value"]("a") == "prefix:a" - assert "upper(%(value)s)" in str(compiled) - assert compiled.params["index"] == 1 - assert compiled.params["value"] == "a" - table.update().values( - {table.c["items"][1:2]: bindparam("values", callable_=lambda: ["a"])} - ).compile(dialect=AthenaDialect()) - - -def test_orm_partial_update_and_renamed_attribute_conflicts(): - base = declarative_base() - - class Model(base): - __tablename__ = "arrays" - id = Column(Integer, primary_key=True) - values = Column("stored", AthenaArray(Integer), key="db_key") - - sql = str(update(Model).values({Model.values[1]: 2}).compile(dialect=AthenaDialect())) - assert "UPDATE arrays SET stored=concat(" in sql - for whole in (Model.values, "values"): - with pytest.raises(sa_exc.CompileError, match="one assignment"): - update(Model).values({Model.values[1]: 2, whole: []}).compile(dialect=AthenaDialect()) diff --git a/tests/sqlalchemy/test_suite.py b/tests/sqlalchemy/test_suite.py index 2a12dd4d..33245bed 100644 --- a/tests/sqlalchemy/test_suite.py +++ b/tests/sqlalchemy/test_suite.py @@ -263,19 +263,40 @@ def test_expression_values_and_indices(self, connection, metadata): Column("id", Integer), Column("items", AthenaArray(Integer)), Column("binary_items", AthenaArray(types.BINARY)), + Column("tuple_items", _ArrayTuple()), + Column("decimal_items", AthenaArray(types.Numeric(8, 2))), ) table.create(connection) - connection.execute(table.insert().values(id=1, items=[1, 2, 3], binary_items=[b"abc"])) + connection.execute( + table.insert().values( + id=1, + items=[1, 2, 3], + binary_items=[b"abc"], + tuple_items=[1, 2], + decimal_items=[Decimal("1.23")], + ) + ) items = table.c["items"] connection.execute( table.update().ordered_values( (items[func.length("abc")], items[1] + 8), (table.c.binary_items[1], b"\x00\xff"), + (table.c.tuple_items[bindparam("tuple_index")], bindparam("tuple_value")), + (table.c.decimal_items[1], Decimal("4.56")), (table.c.id, 2), + ), + {"tuple_index": 1, "tuple_value": 5}, + ) + eq_(connection.execute(select(table.c.tuple_items)).scalar_one(), (5, 2)) + connection.execute( + table.update().values( + {items[1:2]: items[2:3].concat([4]), table.c.tuple_items[1:1]: [6, 7]} ) ) - connection.execute(table.update().values({items[1:2]: items[2:3].concat([4])})) - eq_(connection.execute(select(table)).one(), (2, [2, 9, 4, 9], [b"\x00\xff"])) + eq_( + connection.execute(select(table)).one(), + (2, [2, 9, 4, 9], [b"\x00\xff"], (6, 7, 2), [Decimal("4.56")]), + ) def test_orm_and_long_array_update(self, connection, metadata): table = Table( From 6280ee8386ea88e6006d2a7bf3eda3e6a8bc2a10 Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 21 Sep 2026 04:16:01 +0900 Subject: [PATCH 08/10] Align ARRAY update tests with classes and cover cached validation --- docs/sqlalchemy.md | 2 + pyathena/sqlalchemy/array.py | 4 +- tests/pyathena/sqlalchemy/test_array.py | 234 +++++++++++++----------- tests/sqlalchemy/test_suite.py | 34 +++- 4 files changed, 164 insertions(+), 110 deletions(-) diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index 7ae07a54..3bb6e09d 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -1152,6 +1152,8 @@ These resize rules are PyAthena-specific and do not promise full PostgreSQL arra Write indices and explicit slice boundaries must be non-NULL positive integers after normalization. A slice replacement must be a non-NULL array; use `[]` to delete elements. +Decimal partial updates require `Numeric(precision, scale)`, including when assigning SQL expressions. +Validation can occur during compilation, binding, or Athena execution, so the exception class can differ when a compiled statement is reused. Only `step=None` and `step=1` are supported, and only the final component of a nested update path may be a slice. PyAthena rejects multiple partial assignments to the same array column, or a partial assignment combined with a whole-column assignment to that column. Use one whole-array expression when an update needs several changes to the same array. diff --git a/pyathena/sqlalchemy/array.py b/pyathena/sqlalchemy/array.py index 8a78469c..26335cb7 100644 --- a/pyathena/sqlalchemy/array.py +++ b/pyathena/sqlalchemy/array.py @@ -548,9 +548,7 @@ def process(self, expression, **kw): ): raise exc.CompileError("An ARRAY slice assignment requires a non-NULL array") rhs = compiler.process(value, **kw) - rhs_type = compiler._complex_dml_type( - expression.value_type, implicit_bind=isinstance(value, BindParameter) - ) + rhs_type = compiler._complex_dml_type(expression.value_type, implicit_bind=True) rhs = f"CAST({rhs} AS {rhs_type})" if final_slice: # Reject SQL expressions that evaluate to NULL without issuing a second statement. diff --git a/tests/pyathena/sqlalchemy/test_array.py b/tests/pyathena/sqlalchemy/test_array.py index 5ed7d3c2..aef2d66e 100644 --- a/tests/pyathena/sqlalchemy/test_array.py +++ b/tests/pyathena/sqlalchemy/test_array.py @@ -32,6 +32,7 @@ import pyathena from pyathena.formatter import DefaultParameterFormatter +from pyathena.sqlalchemy.array import _ArrayWriteIndexType from pyathena.sqlalchemy.base import AthenaDialect from pyathena.sqlalchemy.types import ( ARRAY, @@ -97,6 +98,15 @@ def bind_expression(self, bindvalue): return func.upper(bindvalue) +def _array_update_table(type_=None): + return Table( + "arrays", + MetaData(), + Column("id", Integer), + Column("items", type_ or AthenaArray(Integer)), + ) + + class TestAthenaArray: def test_creation_with_default(self): array_type = AthenaArray() @@ -675,22 +685,116 @@ def test_array_rewrite_requires_labels_for_literal_expressions(self, expression) assert sql.startswith("SELECT anon_1.value, json_format(") -class TestArrayUpdateCompiler: - @staticmethod - def _table(type_=None): - return Table( - "arrays", - MetaData(), - Column("id", Integer), - Column("items", type_ or AthenaArray(Integer)), +class TestArrayUpdate: + def test_multiple_updates_to_one_array_are_rejected(self): + table = _array_update_table() + values = table.c["items"] + for assignments in ( + {values[1]: 2, values[2]: 3}, + {values: [], values[1]: 2}, + {values[1]: 2, "items": []}, + ): + with pytest.raises(sa_exc.CompileError, match="one assignment"): + table.update().values(assignments).compile(dialect=AthenaDialect()) + + def test_only_final_index_may_be_a_slice(self): + table = _array_update_table(AthenaArray(Integer, dimensions=2)) + with pytest.raises(sa_exc.CompileError, match="final"): + table.update().values({table.c["items"][1:2][1]: [2]}).compile(dialect=AthenaDialect()) + + def test_bound_indices_and_values_are_reused_without_mutation(self): + table = _array_update_table() + expression = table.c["items"][bindparam("index")] + statement = table.update().values({expression: bindparam("value"), table.c.id: 2}) + compiled = statement.compile(dialect=AthenaDialect()) + assert set(compiled.params) == {"index", "value", "id"} + assert str(statement.compile(dialect=AthenaDialect())) == str(compiled) + + def test_partial_update_requires_target_table_column(self): + table = _array_update_table() + for column_ in (Column("items", AthenaArray(Integer)), _array_update_table().c["items"]): + with pytest.raises(sa_exc.CompileError, match="target table"): + table.update().values({column_[1]: 2}).compile(dialect=AthenaDialect()) + + def test_ordered_partial_update_with_sql_expression(self): + table = _array_update_table() + items = table.c["items"] + statement = table.update().ordered_values((items[2], items[1] + 1), (table.c.id, 2)) + compiled = str(statement.compile(dialect=AthenaDialect())) + assert compiled.index("SET items=") < compiled.index(", id=") + assert "element_at(arrays.items" in compiled + + def test_orm_partial_update_and_renamed_attribute_conflicts(self): + base = declarative_base() + + class Model(base): + __tablename__ = "arrays" + id = Column(Integer, primary_key=True) + values = Column("stored", AthenaArray(Integer), key="db_key") + + sql = str(update(Model).values({Model.values[1]: 2}).compile(dialect=AthenaDialect())) + assert "UPDATE arrays SET stored=concat(" in sql + for whole in (Model.values, "values"): + with pytest.raises(sa_exc.CompileError, match="one assignment"): + update(Model).values({Model.values[1]: 2, whole: []}).compile( + dialect=AthenaDialect() + ) + + +class TestArrayAssignmentType: + def test_binary_element_assignment_uses_native_hex_parameter(self): + table = _array_update_table(AthenaArray(types.BINARY)) + compiled = ( + table.update() + .values({table.c["items"][1]: b"\x00\xff"}) + .compile(dialect=AthenaDialect()) ) + params = { + name: compiled._bind_processors.get(name, lambda value: value)(value) + for name, value in compiled.params.items() + } + assert "FROM_HEX('00ff')" in DefaultParameterFormatter().format(str(compiled), params) + + def test_explicit_assignment_type_and_callable_bindings(self): + table = _array_update_table(AthenaArray(types.String)) + stmt = table.update().values( + { + table.c["items"][bindparam("index", callable_=lambda: 1)]: bindparam( + "value", type_=PrefixString(), callable_=lambda: "a" + ) + } + ) + compiled = stmt.compile(dialect=AthenaDialect()) + assert compiled._bind_processors["value"]("a") == "prefix:a" + assert "upper(%(value)s)" in str(compiled) + assert compiled.params["index"] == 1 + assert compiled.params["value"] == "a" + table.update().values( + {table.c["items"][1:2]: bindparam("values", callable_=lambda: ["a"])} + ).compile(dialect=AthenaDialect()) + + +class TestArrayWriteIndexType: + @pytest.mark.parametrize("processor_name", ["bind_processor", "literal_processor"]) + @pytest.mark.parametrize("value", [None, True, 1.5, "1"]) + def test_rejects_non_integer_values(self, processor_name, value): + processor = getattr(_ArrayWriteIndexType(), processor_name)(AthenaDialect()) + with pytest.raises(ValueError, match="non-NULL integers"): + processor(value) + def test_bind_and_literal_processors(self): + type_ = _ArrayWriteIndexType() + assert type_.bind_processor(AthenaDialect())(2) == 2 + assert type_.literal_processor(AthenaDialect())(2) == "2" + + +class TestArrayUpdateCompiler: @pytest.mark.parametrize( ("target", "value"), [(1, 2), (4, None), (slice(2, 3), [4]), (slice(2, 2), []), (slice(None), [])], ) def test_partial_update_compiles_to_one_whole_column_assignment(self, target, value): - table = self._table() + table = _array_update_table() statement = table.update().values({table.c["items"][target]: value}).where(table.c.id == 1) original_key = statement._generate_cache_key().key compiled = statement.compile(dialect=AthenaDialect()) @@ -709,41 +813,12 @@ def test_partial_update_compiles_to_one_whole_column_assignment(self, target, va @pytest.mark.parametrize("index", [0, -1, None, 1.5, True]) def test_invalid_partial_update_index(self, index): - table = self._table() + table = _array_update_table() with pytest.raises(sa_exc.CompileError, match="indices"): table.update().values({table.c["items"][index]: 1}).compile(dialect=AthenaDialect()) - def test_multiple_updates_to_one_array_are_rejected(self): - table = self._table() - values = table.c["items"] - for assignments in ( - {values[1]: 2, values[2]: 3}, - {values: [], values[1]: 2}, - {values[1]: 2, "items": []}, - ): - with pytest.raises(sa_exc.CompileError, match="one assignment"): - table.update().values(assignments).compile(dialect=AthenaDialect()) - - def test_slice_null_and_nested_slice_rejected(self): - table = self._table(AthenaArray(Integer, dimensions=2)) - with pytest.raises(sa_exc.CompileError, match="non-NULL array"): - table.update().values({table.c["items"][1:2]: None}).compile(dialect=AthenaDialect()) - with pytest.raises(sa_exc.CompileError, match="final"): - table.update().values({table.c["items"][1:2][1]: [2]}).compile(dialect=AthenaDialect()) - - def test_bound_indices_and_values_are_reused_without_mutation(self): - table = self._table() - expression = table.c["items"][bindparam("index")] - statement = table.update().values({expression: bindparam("value"), table.c.id: 2}) - compiled = statement.compile(dialect=AthenaDialect()) - assert set(compiled.params) == {"index", "value", "id"} - assert compiled._bind_processors["index"](2) == 2 - with pytest.raises(ValueError, match="integers"): - compiled._bind_processors["index"](1.5) - assert str(statement.compile(dialect=AthenaDialect())) == str(compiled) - def test_nested_and_zero_indexed_update(self): - table = self._table(AthenaArray(Integer, dimensions=2, zero_indexes=True)) + table = _array_update_table(AthenaArray(Integer, dimensions=2, zero_indexes=True)) statement = table.update().values({table.c["items"][0][2]: 7}) sql = str( statement.compile(dialect=AthenaDialect(), compile_kwargs={"literal_binds": True}) @@ -763,21 +838,26 @@ def test_nested_and_zero_indexed_update(self): ) @pytest.mark.parametrize("index", [1, bindparam("index")]) def test_array_implementations_partial_update(self, array_type, index): - table = self._table(array_type) + table = _array_update_table(array_type) sql = str( table.update().values({table.c["items"][index]: 2}).compile(dialect=AthenaDialect()) ) assert "SET items=concat(" in sql @pytest.mark.parametrize("target", [1, slice(1, 2)]) - def test_decimal_assignment_requires_precision(self, target): + @pytest.mark.parametrize("expression", [False, True]) + def test_decimal_assignment_requires_precision(self, target, expression): value = [Decimal("1.23")] if isinstance(target, slice) else Decimal("1.23") - table = self._table(AthenaArray(types.Numeric())) + table = _array_update_table(AthenaArray(types.Numeric())) + if expression: + value = table.c["items"][target] with pytest.raises(sa_exc.CompileError, match="precision"): table.update().values({table.c["items"][target]: value}).compile( dialect=AthenaDialect() ) - table = self._table(AthenaArray(types.Numeric(8, 2))) + table = _array_update_table(AthenaArray(types.Numeric(8, 2))) + if expression: + value = table.c["items"][target] sql = str( table.update() .values({table.c["items"][target]: value}) @@ -785,14 +865,8 @@ def test_decimal_assignment_requires_precision(self, target): ) assert "DECIMAL(8, 2)" in sql - def test_partial_update_requires_target_table_column(self): - table = self._table() - for column_ in (Column("items", AthenaArray(Integer)), self._table().c["items"]): - with pytest.raises(sa_exc.CompileError, match="target table"): - table.update().values({column_[1]: 2}).compile(dialect=AthenaDialect()) - def test_write_index_expression_keeps_its_argument_types(self): - table = self._table() + table = _array_update_table() index = func.length("abc") statement = table.update().values({table.c["items"][index]: 9}) compiled = statement.compile(dialect=AthenaDialect()) @@ -802,57 +876,7 @@ def test_write_index_expression_keeps_its_argument_types(self): } assert "length('abc')" in DefaultParameterFormatter().format(str(compiled), params) - def test_ordered_partial_update_with_sql_expression(self): - table = self._table() - items = table.c["items"] - statement = table.update().ordered_values((items[2], items[1] + 1), (table.c.id, 2)) - compiled = str(statement.compile(dialect=AthenaDialect())) - assert compiled.index("SET items=") < compiled.index(", id=") - assert "element_at(arrays.items" in compiled - - def test_binary_element_assignment_uses_native_hex_parameter(self): - table = self._table(AthenaArray(types.BINARY)) - compiled = ( - table.update() - .values({table.c["items"][1]: b"\x00\xff"}) - .compile(dialect=AthenaDialect()) - ) - params = { - name: compiled._bind_processors.get(name, lambda value: value)(value) - for name, value in compiled.params.items() - } - assert "FROM_HEX('00ff')" in DefaultParameterFormatter().format(str(compiled), params) - - def test_explicit_assignment_type_and_callable_bindings(self): - table = self._table(AthenaArray(types.String)) - stmt = table.update().values( - { - table.c["items"][bindparam("index", callable_=lambda: 1)]: bindparam( - "value", type_=PrefixString(), callable_=lambda: "a" - ) - } - ) - compiled = stmt.compile(dialect=AthenaDialect()) - assert compiled._bind_processors["value"]("a") == "prefix:a" - assert "upper(%(value)s)" in str(compiled) - assert compiled.params["index"] == 1 - assert compiled.params["value"] == "a" - table.update().values( - {table.c["items"][1:2]: bindparam("values", callable_=lambda: ["a"])} - ).compile(dialect=AthenaDialect()) - - def test_orm_partial_update_and_renamed_attribute_conflicts(self): - base = declarative_base() - - class Model(base): - __tablename__ = "arrays" - id = Column(Integer, primary_key=True) - values = Column("stored", AthenaArray(Integer), key="db_key") - - sql = str(update(Model).values({Model.values[1]: 2}).compile(dialect=AthenaDialect())) - assert "UPDATE arrays SET stored=concat(" in sql - for whole in (Model.values, "values"): - with pytest.raises(sa_exc.CompileError, match="one assignment"): - update(Model).values({Model.values[1]: 2, whole: []}).compile( - dialect=AthenaDialect() - ) + def test_null_slice_assignment_rejected(self): + table = _array_update_table(AthenaArray(Integer, dimensions=2)) + with pytest.raises(sa_exc.CompileError, match="non-NULL array"): + table.update().values({table.c["items"][1:2]: None}).compile(dialect=AthenaDialect()) diff --git a/tests/sqlalchemy/test_suite.py b/tests/sqlalchemy/test_suite.py index 33245bed..99b1c7ea 100644 --- a/tests/sqlalchemy/test_suite.py +++ b/tests/sqlalchemy/test_suite.py @@ -265,6 +265,7 @@ def test_expression_values_and_indices(self, connection, metadata): Column("binary_items", AthenaArray(types.BINARY)), Column("tuple_items", _ArrayTuple()), Column("decimal_items", AthenaArray(types.Numeric(8, 2))), + Column("timestamp_items", AthenaArray(types.TIMESTAMP)), ) table.create(connection) connection.execute( @@ -274,6 +275,10 @@ def test_expression_values_and_indices(self, connection, metadata): binary_items=[b"abc"], tuple_items=[1, 2], decimal_items=[Decimal("1.23")], + timestamp_items=[ + _datetime(2024, 1, 1, microsecond=123000), + _datetime(2024, 1, 2, microsecond=456000), + ], ) ) items = table.c["items"] @@ -282,7 +287,8 @@ def test_expression_values_and_indices(self, connection, metadata): (items[func.length("abc")], items[1] + 8), (table.c.binary_items[1], b"\x00\xff"), (table.c.tuple_items[bindparam("tuple_index")], bindparam("tuple_value")), - (table.c.decimal_items[1], Decimal("4.56")), + (table.c.decimal_items[1], table.c.decimal_items[1] + Decimal("3.33")), + (table.c.timestamp_items[1], table.c.timestamp_items[2]), (table.c.id, 2), ), {"tuple_index": 1, "tuple_value": 5}, @@ -295,7 +301,14 @@ def test_expression_values_and_indices(self, connection, metadata): ) eq_( connection.execute(select(table)).one(), - (2, [2, 9, 4, 9], [b"\x00\xff"], (6, 7, 2), [Decimal("4.56")]), + ( + 2, + [2, 9, 4, 9], + [b"\x00\xff"], + (6, 7, 2), + [Decimal("4.56")], + [_datetime(2024, 1, 2, microsecond=456000)] * 2, + ), ) def test_orm_and_long_array_update(self, connection, metadata): @@ -331,6 +344,23 @@ class Record: finally: mapping.dispose() + def test_cached_literal_assignment_failures(self, connection, metadata): + table = Table( + "array_cached_invalid_updates", metadata, Column("items", AthenaArray(Integer)) + ) + table.create(connection) + connection.execute(table.insert().values(items=[1, 2])) + connection = connection.execution_options(compiled_cache={}) + items = table.c["items"] + connection.execute(table.update().values({items[1]: 3})) + with pytest.raises(sa_exc.DBAPIError, match="Invalid ARRAY index"): + connection.execute(table.update().values({items[0]: 4})) + eq_(connection.execute(select(items)).scalar_one(), [3, 2]) + connection.execute(table.update().values({items[1:2]: [7]})) + with pytest.raises(sa_exc.DBAPIError, match="NULL ARRAY slice assignment"): + connection.execute(table.update().values({items[1:2]: None})) + eq_(connection.execute(select(items)).scalar_one(), [7]) + def test_null_slice_binding_rejected(self, connection, metadata): table = Table("array_null_slice", metadata, Column("items", types.ARRAY(Integer))) table.create(connection) From e58f4e020d8a58a097e374063cf33233940c0b6f Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Mon, 21 Sep 2026 04:40:18 +0900 Subject: [PATCH 09/10] Preserve ARRAY assignment coverage across responsibility classes --- docs/sqlalchemy.md | 4 +- pyathena/sqlalchemy/array.py | 2 +- pyathena/sqlalchemy/compiler.py | 22 +++--- tests/pyathena/sqlalchemy/test_array.py | 90 +++++++++++++++---------- tests/sqlalchemy/test_suite.py | 4 +- 5 files changed, 71 insertions(+), 51 deletions(-) diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index 3bb6e09d..041906a0 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -1152,8 +1152,8 @@ These resize rules are PyAthena-specific and do not promise full PostgreSQL arra Write indices and explicit slice boundaries must be non-NULL positive integers after normalization. A slice replacement must be a non-NULL array; use `[]` to delete elements. -Decimal partial updates require `Numeric(precision, scale)`, including when assigning SQL expressions. -Validation can occur during compilation, binding, or Athena execution, so the exception class can differ when a compiled statement is reused. +Decimal partial updates require a declared `Numeric` precision, including SQL-expression assignments; specify a scale to retain fractional values. +Validation can raise `CompileError` during compilation, `StatementError` during binding, or `DBAPIError` from Athena; the exception class can differ when a compiled statement is reused. Only `step=None` and `step=1` are supported, and only the final component of a nested update path may be a slice. PyAthena rejects multiple partial assignments to the same array column, or a partial assignment combined with a whole-column assignment to that column. Use one whole-array expression when an update needs several changes to the same array. diff --git a/pyathena/sqlalchemy/array.py b/pyathena/sqlalchemy/array.py index 26335cb7..2a1f174e 100644 --- a/pyathena/sqlalchemy/array.py +++ b/pyathena/sqlalchemy/array.py @@ -548,7 +548,7 @@ def process(self, expression, **kw): ): raise exc.CompileError("An ARRAY slice assignment requires a non-NULL array") rhs = compiler.process(value, **kw) - rhs_type = compiler._complex_dml_type(expression.value_type, implicit_bind=True) + rhs_type = compiler._complex_dml_type(expression.value_type, require_precision=True) rhs = f"CAST({rhs} AS {rhs_type})" if final_slice: # Reject SQL expressions that evaluate to NULL without issuing a second statement. diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index d2f0d573..3ca46bfa 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -634,7 +634,7 @@ def visit_truediv_binary(self, binary, operator, **kw): def visit_cast(self, cast: Cast[Any], **kwargs): if isinstance(cast.type, (types.ARRAY, AthenaMap, AthenaStruct)): type_clause = self._complex_dml_type( - cast.type, implicit_bind=cast._annotations.get("_pyathena_array_bind", False) + cast.type, require_precision=cast._annotations.get("_pyathena_array_bind", False) ) return f"CAST({self.process(cast.clause, **kwargs)} AS {type_clause})" if (isinstance(cast.type, types.VARCHAR) and cast.type.length is None) or isinstance( @@ -656,27 +656,29 @@ def visit_cast(self, cast: Cast[Any], **kwargs): type_clause = cast.typeclause._compiler_dispatch(self, **kwargs) return f"CAST({cast.clause._compiler_dispatch(self, **kwargs)} AS {type_clause})" - def _complex_dml_type(self, type_, *, implicit_bind=False): + def _complex_dml_type(self, type_, *, require_precision=False): if isinstance(type_, types.TypeDecorator): return self._complex_dml_type( - self._array_type_inspector.decorator_impl(type_), implicit_bind=implicit_bind + self._array_type_inspector.decorator_impl(type_), + require_precision=require_precision, ) if isinstance(type_, types.NullType): raise exc.CompileError("Bound ARRAY values require an explicit element type") if isinstance(type_, types.ARRAY): item = self._complex_dml_type( - _ArrayTypeInspector.item_type(type_), implicit_bind=implicit_bind + _ArrayTypeInspector.item_type(type_), require_precision=require_precision ) return f"ARRAY({item})" if isinstance(type_, AthenaMap): - return ( - f"MAP({self._complex_dml_type(type_.key_type, implicit_bind=implicit_bind)}, " - f"{self._complex_dml_type(type_.value_type, implicit_bind=implicit_bind)})" + key_type = self._complex_dml_type(type_.key_type, require_precision=require_precision) + value_type = self._complex_dml_type( + type_.value_type, require_precision=require_precision ) + return f"MAP({key_type}, {value_type})" if isinstance(type_, AthenaStruct): fields = ", ".join( f"{self.preparer.quote(name)} " - f"{self._complex_dml_type(field_type, implicit_bind=implicit_bind)}" + f"{self._complex_dml_type(field_type, require_precision=require_precision)}" for name, field_type in type_.fields.items() ) return f"ROW({fields})" @@ -688,9 +690,9 @@ def _complex_dml_type(self, type_, *, implicit_bind=False): return "DOUBLE" if isinstance(type_, types.Float): return "REAL" - if implicit_bind and isinstance(type_, types.Numeric) and type_.precision is None: + if require_precision and isinstance(type_, types.Numeric) and type_.precision is None: raise exc.CompileError( - "ARRAY decimal binds require explicit Numeric precision; " + "ARRAY decimal values require explicit Numeric precision; " "specify precision and scale to avoid implicit rounding" ) return self.dialect.type_compiler_instance.process(type_) diff --git a/tests/pyathena/sqlalchemy/test_array.py b/tests/pyathena/sqlalchemy/test_array.py index aef2d66e..22a4624d 100644 --- a/tests/pyathena/sqlalchemy/test_array.py +++ b/tests/pyathena/sqlalchemy/test_array.py @@ -685,6 +685,45 @@ def test_array_rewrite_requires_labels_for_literal_expressions(self, expression) assert sql.startswith("SELECT anon_1.value, json_format(") +class TestArrayAssignmentType: + def test_binary_element_assignment_uses_native_hex_parameter(self): + table = _array_update_table(AthenaArray(types.BINARY)) + compiled = ( + table.update() + .values({table.c["items"][1]: b"\x00\xff"}) + .compile(dialect=AthenaDialect()) + ) + params = { + name: compiled._bind_processors.get(name, lambda value: value)(value) + for name, value in compiled.params.items() + } + assert "FROM_HEX('00ff')" in DefaultParameterFormatter().format(str(compiled), params) + + def test_explicit_assignment_type_and_callable_value(self): + table = _array_update_table(AthenaArray(types.String)) + stmt = table.update().values( + {table.c["items"][1]: bindparam("value", type_=PrefixString(), callable_=lambda: "a")} + ) + compiled = stmt.compile(dialect=AthenaDialect()) + assert compiled._bind_processors["value"]("a") == "prefix:a" + assert "upper(%(value)s)" in str(compiled) + assert compiled.params["value"] == "a" + + +class TestArrayWriteIndexType: + @pytest.mark.parametrize("processor_name", ["bind_processor", "literal_processor"]) + @pytest.mark.parametrize("value", [None, True, 1.5, "1"]) + def test_rejects_non_integer_values(self, processor_name, value): + processor = getattr(_ArrayWriteIndexType(), processor_name)(AthenaDialect()) + with pytest.raises(ValueError, match="non-NULL integers"): + processor(value) + + def test_bind_and_literal_processors(self): + type_ = _ArrayWriteIndexType() + assert type_.bind_processor(AthenaDialect())(2) == 2 + assert type_.literal_processor(AthenaDialect())(2) == "2" + + class TestArrayUpdate: def test_multiple_updates_to_one_array_are_rejected(self): table = _array_update_table() @@ -741,54 +780,31 @@ class Model(base): ) -class TestArrayAssignmentType: - def test_binary_element_assignment_uses_native_hex_parameter(self): - table = _array_update_table(AthenaArray(types.BINARY)) +class TestArrayUpdateCompiler: + def test_bound_index_uses_write_index_processor(self): + table = _array_update_table() compiled = ( table.update() - .values({table.c["items"][1]: b"\x00\xff"}) + .values({table.c["items"][bindparam("index")]: 9}) .compile(dialect=AthenaDialect()) ) - params = { - name: compiled._bind_processors.get(name, lambda value: value)(value) - for name, value in compiled.params.items() - } - assert "FROM_HEX('00ff')" in DefaultParameterFormatter().format(str(compiled), params) + assert compiled._bind_processors["index"](2) == 2 + for value in (1.5, True, None): + with pytest.raises(ValueError, match="non-NULL integers"): + compiled._bind_processors["index"](value) - def test_explicit_assignment_type_and_callable_bindings(self): + def test_callable_index_and_slice_value(self): table = _array_update_table(AthenaArray(types.String)) - stmt = table.update().values( - { - table.c["items"][bindparam("index", callable_=lambda: 1)]: bindparam( - "value", type_=PrefixString(), callable_=lambda: "a" - ) - } + compiled = ( + table.update() + .values({table.c["items"][bindparam("index", callable_=lambda: 1)]: "a"}) + .compile(dialect=AthenaDialect()) ) - compiled = stmt.compile(dialect=AthenaDialect()) - assert compiled._bind_processors["value"]("a") == "prefix:a" - assert "upper(%(value)s)" in str(compiled) assert compiled.params["index"] == 1 - assert compiled.params["value"] == "a" table.update().values( {table.c["items"][1:2]: bindparam("values", callable_=lambda: ["a"])} ).compile(dialect=AthenaDialect()) - -class TestArrayWriteIndexType: - @pytest.mark.parametrize("processor_name", ["bind_processor", "literal_processor"]) - @pytest.mark.parametrize("value", [None, True, 1.5, "1"]) - def test_rejects_non_integer_values(self, processor_name, value): - processor = getattr(_ArrayWriteIndexType(), processor_name)(AthenaDialect()) - with pytest.raises(ValueError, match="non-NULL integers"): - processor(value) - - def test_bind_and_literal_processors(self): - type_ = _ArrayWriteIndexType() - assert type_.bind_processor(AthenaDialect())(2) == 2 - assert type_.literal_processor(AthenaDialect())(2) == "2" - - -class TestArrayUpdateCompiler: @pytest.mark.parametrize( ("target", "value"), [(1, 2), (4, None), (slice(2, 3), [4]), (slice(2, 2), []), (slice(None), [])], @@ -877,6 +893,6 @@ def test_write_index_expression_keeps_its_argument_types(self): assert "length('abc')" in DefaultParameterFormatter().format(str(compiled), params) def test_null_slice_assignment_rejected(self): - table = _array_update_table(AthenaArray(Integer, dimensions=2)) + table = _array_update_table() with pytest.raises(sa_exc.CompileError, match="non-NULL array"): table.update().values({table.c["items"][1:2]: None}).compile(dialect=AthenaDialect()) diff --git a/tests/sqlalchemy/test_suite.py b/tests/sqlalchemy/test_suite.py index 99b1c7ea..75c1fab3 100644 --- a/tests/sqlalchemy/test_suite.py +++ b/tests/sqlalchemy/test_suite.py @@ -274,7 +274,7 @@ def test_expression_values_and_indices(self, connection, metadata): items=[1, 2, 3], binary_items=[b"abc"], tuple_items=[1, 2], - decimal_items=[Decimal("1.23")], + decimal_items=[Decimal("0.00")], timestamp_items=[ _datetime(2024, 1, 1, microsecond=123000), _datetime(2024, 1, 2, microsecond=456000), @@ -282,6 +282,8 @@ def test_expression_values_and_indices(self, connection, metadata): ) ) items = table.c["items"] + connection.execute(table.update().values({table.c.decimal_items[1]: Decimal("1.23")})) + eq_(connection.execute(select(table.c.decimal_items)).scalar_one(), [Decimal("1.23")]) connection.execute( table.update().ordered_values( (items[func.length("abc")], items[1] + 8), From 8f3f237ade70c8d70ed1a748a6562692722d664a Mon Sep 17 00:00:00 2001 From: laughingman7743 Date: Wed, 23 Sep 2026 11:01:17 +0900 Subject: [PATCH 10/10] Separate ARRAY element and slice update rendering --- pyathena/sqlalchemy/array.py | 76 +++++++++++++++++++----------------- 1 file changed, 41 insertions(+), 35 deletions(-) diff --git a/pyathena/sqlalchemy/array.py b/pyathena/sqlalchemy/array.py index 2a1f174e..09cefe69 100644 --- a/pyathena/sqlalchemy/array.py +++ b/pyathena/sqlalchemy/array.py @@ -599,45 +599,51 @@ def _rebuild(self, array, array_type, path, rhs, **kw): array = f"coalesce({array}, CAST(ARRAY[] AS {array_sql_type}))" bound = path[0] if isinstance(bound, Slice): - if not isinstance(bound.step, Null) and not ( - isinstance(bound.step, BindParameter) - and bound.step.unique - and type(bound.step.value) is int - and bound.step.value == 1 - ): - raise exc.CompileError("Athena ARRAY slices support only step=None or step=1") - start = "1" if isinstance(bound.start, Null) else self._index_sql(bound.start, **kw) - stop = ( - f"cardinality({array})" - if isinstance(bound.stop, Null) - else self._index_sql(bound.stop, **kw) - ) - prefix = f"slice({array}, 1, least({start} - 1, cardinality({array})))" - element_type = compiler._complex_dml_type(_ArrayTypeInspector.item_type(array_type)) - padding = ( - f"repeat(CAST(NULL AS {element_type}), " - f"CAST(greatest({start} - 1 - cardinality({array}), 0) AS INTEGER))" - ) - tail_start = f"greatest({start}, {stop} + 1)" - suffix = ( - f"slice({array}, {tail_start}, " - f"greatest(cardinality({array}) - {tail_start} + 1, 0))" - ) - return compiler._array_slice_step( - f"concat({prefix}, {padding}, {rhs}, {suffix})", bound.step, array_type, **kw - ) + return self._rebuild_slice(array, array_type, bound, rhs, **kw) + return self._rebuild_element(array, array_type, bound, path[1:], rhs, **kw) + + def _prefix_and_padding(self, array, start, array_type): + prefix = f"slice({array}, 1, least({start} - 1, cardinality({array})))" + element_type = self.compiler._complex_dml_type(_ArrayTypeInspector.item_type(array_type)) + padding = ( + f"repeat(CAST(NULL AS {element_type}), " + f"CAST(greatest({start} - 1 - cardinality({array}), 0) AS INTEGER))" + ) + return prefix, padding + + def _rebuild_slice(self, array, array_type, bound, rhs, **kw): + if not isinstance(bound.step, Null) and not ( + isinstance(bound.step, BindParameter) + and bound.step.unique + and type(bound.step.value) is int + and bound.step.value == 1 + ): + raise exc.CompileError("Athena ARRAY slices support only step=None or step=1") + start = "1" if isinstance(bound.start, Null) else self._index_sql(bound.start, **kw) + stop = ( + f"cardinality({array})" + if isinstance(bound.stop, Null) + else self._index_sql(bound.stop, **kw) + ) + prefix, padding = self._prefix_and_padding(array, start, array_type) + tail_start = f"greatest({start}, {stop} + 1)" + suffix = ( + f"slice({array}, {tail_start}, greatest(cardinality({array}) - {tail_start} + 1, 0))" + ) + return self.compiler._array_slice_step( + f"concat({prefix}, {padding}, {rhs}, {suffix})", bound.step, array_type, **kw + ) + + def _rebuild_element(self, array, array_type, bound, remaining_path, rhs, **kw): index = self._index_sql(bound, **kw) previous = f"element_at({array}, {index})" replacement = ( - self._rebuild(previous, _ArrayTypeInspector.item_type(array_type), path[1:], rhs, **kw) - if len(path) > 1 + self._rebuild( + previous, _ArrayTypeInspector.item_type(array_type), remaining_path, rhs, **kw + ) + if remaining_path else rhs ) - prefix = f"slice({array}, 1, least({index} - 1, cardinality({array})))" - element_type = compiler._complex_dml_type(_ArrayTypeInspector.item_type(array_type)) - padding = ( - f"repeat(CAST(NULL AS {element_type}), " - f"CAST(greatest({index} - 1 - cardinality({array}), 0) AS INTEGER))" - ) + prefix, padding = self._prefix_and_padding(array, index, array_type) suffix = f"slice({array}, {index} + 1, greatest(cardinality({array}) - {index}, 0))" return f"concat({prefix}, {padding}, ARRAY[{replacement}], {suffix})"