Skip to content

Commit fbcf8b8

Browse files
Preserve explicit ARRAY quantifier binds and clarify negation compatibility
1 parent eff2001 commit fbcf8b8

4 files changed

Lines changed: 26 additions & 6 deletions

File tree

‎docs/sqlalchemy.md‎

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1113,8 +1113,9 @@ Comparisons use SQL three-valued logic: NULL elements can produce NULL when no d
11131113
For an empty array, ANY is false and ALL is true.
11141114
A NULL array produces NULL for both.
11151115
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()`.
1116+
SQLAlchemy can rewrite negation into an element-wise comparison before dialect compilation, depending on the version and operand types.
1117+
For example, `~(any_(flags) == True)` can become an element-wise `!= True` comparison; older SQLAlchemy 2.0 releases also rewrite non-boolean comparisons this way.
1118+
To negate the whole match, always explicitly group the comparison first: `~(any_(flags) == True).self_group()`.
11181119
Quantifiers over subqueries retain their usual SQL compilation.
11191120

11201121
#### Querying ARRAY data

‎pyathena/sqlalchemy/compiler.py‎

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -318,7 +318,10 @@ def visit_binary(
318318
predicate.right = Column(variable, item_type)
319319
if isinstance(predicate.left, BindParameter) and (
320320
isinstance(item_type, types.ARRAY)
321-
or self._array_type_inspector.array_type(predicate.left.type) is not None
321+
or (
322+
predicate.left.type is aggregate.element.type
323+
and predicate.left.type._type_affinity is not types.ARRAY
324+
)
322325
):
323326
predicate.left = predicate.left._with_binary_element_type(item_type)
324327
if from_linter is not None and operators.is_comparison(binary.operator):

‎tests/pyathena/sqlalchemy/test_compiler.py‎

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -340,6 +340,15 @@ def test_multidimensional_array_quantifier_bind_type(self):
340340
items = column("items", AthenaArray(Integer, dimensions=2))
341341
assert "CAST(ARRAY[1, 2] AS ARRAY(INTEGER)) =" in self._compile_sql(items.any([1, 2]))
342342

343+
def test_quantifier_preserves_explicit_array_bind(self):
344+
items = column("items", AthenaArray(Integer))
345+
needle = literal([1, 2], items.type)
346+
sql = self._compile_sql(needle == any_(func.array_agg(items)))
347+
assert "CAST(ARRAY[1, 2] AS ARRAY(INTEGER)) = _pyathena_element_0" in sql
348+
unknown = column("unknown", AthenaArray(types.NullType()))
349+
sql = self._compile_sql(needle == any_(unknown))
350+
assert "CAST(ARRAY[1, 2] AS ARRAY(INTEGER)) = _pyathena_element_0" in sql
351+
343352
def test_subquery_any_remains_unchanged(self):
344353
sql = self._compile_sql(any_(select(column("item", Integer)).scalar_subquery()) == 2)
345354
assert "ANY (SELECT item)" in sql

‎tests/sqlalchemy/test_suite.py‎

Lines changed: 10 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -182,9 +182,10 @@ 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)],
185+
@sa_testing.combinations(
186+
(String(), String, ["a", "b"], "a"),
187+
(Integer(), Integer, [1, 2], 1),
188+
argnames="base_type,item_type,values,needle",
188189
)
189190
def test_variant_array_quantifier_scalar_bind(
190191
self, connection, base_type, item_type, values, needle
@@ -263,6 +264,12 @@ def test_dimensions_zero_indexes_and_bound_index(self, connection):
263264
eq_(connection.execute(select(array.any([1, 2]))).scalar_one(), True)
264265
plain = literal([1, 2], types.ARRAY(Integer))
265266
eq_(connection.execute(select(plain[bindparam("index")]), {"index": 2}).scalar_one(), 2)
267+
eq_(
268+
connection.execute(
269+
select(literal([1, 2], AthenaArray(Integer)) == any_(func.array_agg(plain)))
270+
).scalar_one(),
271+
True,
272+
)
266273

267274
def test_quantified_comparisons(self, connection):
268275
cases = [

0 commit comments

Comments
 (0)