Skip to content

Commit da595ac

Browse files
Fix ARRAY variant quantifier binds and document expression limits
1 parent 2676632 commit da595ac

5 files changed

Lines changed: 45 additions & 9 deletions

File tree

‎docs/sqlalchemy.md‎

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1104,14 +1104,17 @@ Slice stops are inclusive, following SQLAlchemy's SQL array convention.
11041104
Omitted boundaries mean the beginning or end of the array; explicit boundaries are clipped to the array.
11051105
A reversed slice returns an empty array, and slicing a NULL array returns NULL.
11061106
Only `step=None` and `step=1` are supported.
1107+
Slice SQL can evaluate the array and boundary expressions more than once; use deterministic expressions rather than volatile functions such as `random()` in slices.
11071108
Use `AthenaArray` for open-ended slices with `zero_indexes=True`; SQLAlchemy's generic `ARRAY` comparator requires explicit bounds in that mode.
11081109
Indexing a multidimensional array retains the remaining dimensions, while slicing retains its array type.
11091110

11101111
`any_(array)` and `all_(array)`, and the legacy `array.any(value)` and `array.all(value)` methods, compile to Athena's `any_match` and `all_match` functions.
11111112
Comparisons use SQL three-valued logic: NULL elements can produce NULL when no decisive true or false result exists.
11121113
For an empty array, ANY is false and ALL is true.
11131114
A NULL array produces NULL for both.
1114-
SQLAlchemy comparison flipping and negation are preserved.
1115+
SQLAlchemy comparison flipping is preserved.
1116+
For boolean comparisons, SQLAlchemy can turn `~(any_(flags) == True)` into an element-wise `!= True` comparison before dialect compilation.
1117+
To negate the whole match, explicitly group the comparison first: `~(any_(flags) == True).self_group()`.
11151118
Quantifiers over subqueries retain their usual SQL compilation.
11161119

11171120
#### Querying ARRAY data

‎pyathena/sqlalchemy/compiler.py‎

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -316,7 +316,10 @@ def visit_binary(
316316
predicate.left = Column(variable, item_type)
317317
else:
318318
predicate.right = Column(variable, item_type)
319-
if isinstance(predicate.left, BindParameter) and isinstance(item_type, types.ARRAY):
319+
if isinstance(predicate.left, BindParameter) and (
320+
isinstance(item_type, types.ARRAY)
321+
or self._array_type_inspector.array_type(predicate.left.type) is not None
322+
):
320323
predicate.left = predicate.left._with_binary_element_type(item_type)
321324
if from_linter is not None and operators.is_comparison(binary.operator):
322325
if lateral_from_linter is not None:

‎tests/pyathena/sqlalchemy/test_array.py‎

Lines changed: 11 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -277,16 +277,22 @@ def test_decorated_array_index_and_slice(self):
277277

278278
class TestArrayTypeInspector:
279279
@pytest.mark.parametrize(
280-
"type_", [TupleArray(), String().with_variant(TupleArray(), "awsathena")]
280+
("type_", "value", "expected"),
281+
[
282+
(TupleArray(), 2, "2"),
283+
(String().with_variant(TupleArray(), "awsathena"), 2, "2"),
284+
(Integer().with_variant(TupleArray(), "awsathena"), 2, "2"),
285+
(String().with_variant(AthenaArray(String), "awsathena"), "a", "'a'"),
286+
],
281287
)
282-
def test_decorated_and_variant_array_quantifiers(self, type_):
288+
def test_decorated_and_variant_array_quantifiers(self, type_, value, expected):
283289
items = column("items", type_)
284-
statement = select(any_(items) == 2, all_(items) > 0)
290+
statement = select(any_(items) == value, all_(items) > value)
285291
sql = str(
286292
statement.compile(dialect=AthenaDialect(), compile_kwargs={"literal_binds": True})
287293
)
288-
assert "any_match((items), _pyathena_element_0 -> 2 = _pyathena_element_0)" in sql
289-
assert "all_match((items), _pyathena_element_1 -> 0 < _pyathena_element_1)" in sql
294+
assert f"any_match((items), _pyathena_element_0 -> {expected} = _pyathena_element_0)" in sql
295+
assert f"all_match((items), _pyathena_element_1 -> {expected} < _pyathena_element_1)" in sql
290296

291297
@pytest.mark.parametrize(
292298
("type_", "ddl", "dml"),

‎tests/pyathena/sqlalchemy/test_compiler.py‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -362,6 +362,13 @@ def test_quantifier_boolean_left_operand(self):
362362
assert "any_match" in self._compile_sql(any_(flags) == True) # noqa: E712
363363
assert "all_match" in self._compile_sql(all_(flags) != False) # noqa: E712
364364
assert "IS DISTINCT FROM" in self._compile_sql(any_(flags).is_distinct_from(True))
365+
comparison = any_(flags) == True # noqa: E712
366+
assert self._compile_sql(~comparison) == (
367+
"any_match((flags), _pyathena_element_0 -> _pyathena_element_0 != true)"
368+
)
369+
assert self._compile_sql(~comparison.self_group()) == (
370+
"NOT (any_match((flags), _pyathena_element_0 -> _pyathena_element_0 = true))"
371+
)
365372

366373
def test_quantifier_join_linter_tracks_original_tables(self):
367374
left = table("left_table", column("value", Integer))

‎tests/sqlalchemy/test_suite.py‎

Lines changed: 19 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -182,6 +182,18 @@ def test_decorated_and_variant_arrays(self, connection):
182182
(1, (1, 2), True, True, (1, 2, 3, 4), True, (1, 2)),
183183
)
184184

185+
@pytest.mark.parametrize(
186+
("base_type", "item_type", "values", "needle"),
187+
[(String(), String, ["a", "b"], "a"), (Integer(), Integer, [1, 2], 1)],
188+
)
189+
def test_variant_array_quantifier_scalar_bind(
190+
self, connection, base_type, item_type, values, needle
191+
):
192+
array = literal(values, base_type.with_variant(AthenaArray(item_type), "awsathena"))
193+
statement = select(any_(array) == needle, all_(array) == needle)
194+
eq_(tuple(connection.execute(statement).one()), (True, False))
195+
eq_(tuple(connection.execute(statement).one()), (True, False))
196+
185197
def test_cached_steps_and_boolean_quantifiers(self, connection):
186198
inferred = func.array_agg(func.length(literal("abc")))
187199
eq_(connection.execute(select(inferred[1:2:1])).scalar_one(), [3])
@@ -191,9 +203,14 @@ def test_cached_steps_and_boolean_quantifiers(self, connection):
191203
connection.execute(select(array[1:2:2])).all()
192204
assert isinstance(error.value.orig, ValueError)
193205
flags = literal([True, False], AthenaArray(types.Boolean))
206+
comparison = any_(flags) == True # noqa: E712
194207
eq_(
195-
tuple(connection.execute(select(any_(flags) == True, all_(flags) == True)).one()), # noqa: E712
196-
(True, False),
208+
tuple(
209+
connection.execute(
210+
select(comparison, all_(flags) == True, ~comparison, ~comparison.self_group()) # noqa: E712
211+
).one()
212+
),
213+
(True, False, True, False),
197214
)
198215

199216
def test_index_slice_and_concat(self, connection):

0 commit comments

Comments
 (0)