diff --git a/docs/sqlalchemy.md b/docs/sqlalchemy.md index edff37a7..041906a0 100644 --- a/docs/sqlalchemy.md +++ b/docs/sqlalchemy.md @@ -1118,6 +1118,47 @@ 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. + +For an existing Iceberg table named `orders` with an ARRAY column `item_ids`: + +```python +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(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. +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. +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. +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. +Partial updates can evaluate indices, boundaries, and replacement SQL expressions more than once; use deterministic expressions. + #### 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/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.py b/pyathena/sqlalchemy/array.py index cae77efb..09cefe69 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,246 @@ 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, 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. + 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): + 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), remaining_path, rhs, **kw + ) + if remaining_path + else rhs + ) + 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})" diff --git a/pyathena/sqlalchemy/compiler.py b/pyathena/sqlalchemy/compiler.py index feb54129..3ca46bfa 100644 --- a/pyathena/sqlalchemy/compiler.py +++ b/pyathena/sqlalchemy/compiler.py @@ -33,7 +33,12 @@ AthenaPartitionTransform, AthenaRowFormatSerde, ) -from pyathena.sqlalchemy.array import _ArraySliceStepType, _ArrayTypeInspector +from pyathena.sqlalchemy.array import ( + _ArraySliceStepType, + _ArrayTypeInspector, + _ArrayUpdate, + _ArrayUpdateCompiler, +) from pyathena.sqlalchemy.preparer import AthenaDDLIdentifierPreparer from pyathena.sqlalchemy.types import ( AthenaMap, @@ -270,6 +275,15 @@ def _original_froms(elements): 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( + _ArrayUpdate.rewrite(update_stmt, self.dialect), visiting_cte=visiting_cte, **kw + ) + + def visit_athena_array_update(self, expression, **kw): + return _ArrayUpdateCompiler(self).process(expression, **kw) + def _array_lambda_name(self): names = { str( @@ -620,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( @@ -642,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})" @@ -674,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 e52498c1..22a4624d 100644 --- a/tests/pyathena/sqlalchemy/test_array.py +++ b/tests/pyathena/sqlalchemy/test_array.py @@ -18,17 +18,21 @@ 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 from pyathena.formatter import DefaultParameterFormatter +from pyathena.sqlalchemy.array import _ArrayWriteIndexType from pyathena.sqlalchemy.base import AthenaDialect from pyathena.sqlalchemy.types import ( ARRAY, @@ -83,6 +87,26 @@ 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) + + +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() @@ -659,3 +683,216 @@ def test_array_rewrite_requires_labels_for_literal_expressions(self, expression) .compile(dialect=AthenaDialect()) ) 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() + 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 TestArrayUpdateCompiler: + def test_bound_index_uses_write_index_processor(self): + table = _array_update_table() + compiled = ( + table.update() + .values({table.c["items"][bindparam("index")]: 9}) + .compile(dialect=AthenaDialect()) + ) + 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_callable_index_and_slice_value(self): + table = _array_update_table(AthenaArray(types.String)) + compiled = ( + table.update() + .values({table.c["items"][bindparam("index", callable_=lambda: 1)]: "a"}) + .compile(dialect=AthenaDialect()) + ) + assert compiled.params["index"] == 1 + table.update().values( + {table.c["items"][1:2]: bindparam("values", callable_=lambda: ["a"])} + ).compile(dialect=AthenaDialect()) + + @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 = _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()) + 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 = _array_update_table() + with pytest.raises(sa_exc.CompileError, match="indices"): + table.update().values({table.c["items"][index]: 1}).compile(dialect=AthenaDialect()) + + def test_nested_and_zero_indexed_update(self): + 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}) + ) + 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 = _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)]) + @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 = _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 = _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}) + .compile(dialect=AthenaDialect()) + ) + assert "DECIMAL(8, 2)" in sql + + def test_write_index_expression_keeps_its_argument_types(self): + table = _array_update_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_null_slice_assignment_rejected(self): + 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 64949284..75c1fab3 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 @@ -161,6 +163,217 @@ 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_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)), + Column("tuple_items", _ArrayTuple()), + Column("decimal_items", AthenaArray(types.Numeric(8, 2))), + Column("timestamp_items", AthenaArray(types.TIMESTAMP)), + ) + table.create(connection) + connection.execute( + table.insert().values( + id=1, + items=[1, 2, 3], + binary_items=[b"abc"], + tuple_items=[1, 2], + decimal_items=[Decimal("0.00")], + timestamp_items=[ + _datetime(2024, 1, 1, microsecond=123000), + _datetime(2024, 1, 2, microsecond=456000), + ], + ) + ) + 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), + (table.c.binary_items[1], b"\x00\xff"), + (table.c.tuple_items[bindparam("tuple_index")], bindparam("tuple_value")), + (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}, + ) + 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]} + ) + ) + eq_( + connection.execute(select(table)).one(), + ( + 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): + 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_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) + 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",)