Skip to content
11 changes: 11 additions & 0 deletions CHANGES.md
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,17 @@
columns need migrating to `NUMERIC`
- Types: Added `numeric` and `numeric_array` to the reflected type map, where
they previously resolved to an abstract type that could not be compiled
- Compiler: Fixed array indexes rendering as object keys, as in `arr['1']`.
Added support for array slices, and for columns and expressions as indexes
- Compiler: Fixed cached statements on `sa.JSON` and `ObjectArray` columns
reusing the subscript of an earlier statement, which returned values of
another key. Fixed object keys containing a quote
- BREAKING: Compiler: A statement with a literal subscript on an `sa.JSON`,
`ObjectArray` or `ARRAY` column can't run with `executemany()` on a cached
engine. For that, use `execution_options(compiled_cache=None)`
- BREAKING: Compiler: An integer subscript on a nested object renders as an
array position. Write the key as a string, as in `obj['n']['1']`
- Compiler: `bool`, `float` and `None` subscripts raise `CompileError`

## 2026/06/22 0.43.1
- Compiler: Fixed `AttributeError: 'CrateCompilerSA20' object has no attribute
Expand Down
15 changes: 15 additions & 0 deletions docs/working-with-types.rst
Original file line number Diff line number Diff line change
Expand Up @@ -203,6 +203,21 @@ all values of that field of all objects in that object array:
>>> query.all()
[([1, 2, 3],), (None,), (None,)]

Select an element of an array by its position, or a part of it by a slice.
Positions start at 1, and slices include both bounds. A position of 0 or less,
or past the end of the array, returns ``NULL``:

>>> query = session.query(Character.more_details[1]['foo']).order_by(Character.name)
>>> query.all()
[(1,), (None,), (None,)]

>>> query = session.query(Character.more_details['foo'][2:3])
>>> query.filter_by(name='Arthur Dent').all()
[([2, 3],)]

Positions and slices work the same way on ``sa.ARRAY`` columns. With
``sa.ARRAY(..., zero_indexes=True)``, positions start at 0, as in Python.


Geospatial types
================
Expand Down
114 changes: 111 additions & 3 deletions src/sqlalchemy_cratedb/compiler.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,15 @@
import sqlalchemy as sa
from sqlalchemy.dialects.postgresql.base import RESERVED_WORDS as POSTGRESQL_RESERVED_WORDS
from sqlalchemy.dialects.postgresql.base import PGCompiler
from sqlalchemy.sql import compiler
from sqlalchemy.sql import compiler, operators
from sqlalchemy.sql.elements import (
BinaryExpression,
BindParameter,
Grouping,
Null,
Slice,
TypeCoerce,
)
from sqlalchemy.types import String

from .sa_version import SA_1_4, SA_VERSION
Expand Down Expand Up @@ -335,6 +343,72 @@ def visit_JSONB(self, type_, **kw):
return "OBJECT"


def _slice_bound(value):
if value is None:
return ""
if isinstance(value, bool) or not isinstance(value, int):
raise sa.exc.CompileError(f"CrateDB array slices take integer bounds, not {value!r}")
return str(value)


class _SubscriptType(sa.types.TypeEngine):
"""
Renders an array index, an array slice, or an object key as a literal.

With `int_as_key`, an integer renders as an object key.
"""

cache_ok = True
# Reject `None`, instead of rendering it as `NULL`.
should_evaluate_none = True

def __init__(self, int_as_key=False):
self.int_as_key = int_as_key

def literal_processor(self, dialect):
def process(value):
if isinstance(value, slice):
if value.step is not None:
raise sa.exc.CompileError("CrateDB array slices do not support a step")
return "%s:%s" % (_slice_bound(value.start), _slice_bound(value.stop))
if isinstance(value, bool) or not isinstance(value, (int, str)):
raise sa.exc.CompileError(
f"CrateDB subscripts take an integer index or a string key, not {value!r}"
)
if isinstance(value, int) and not self.int_as_key:
return str(value)
return "'%s'" % str(value).replace("'", "''")

return process


def _is_top_level_object(element):
"""
Whether `element` is an object, and not a value taken out of one, which
may be an array. CrateDB does not accept an array index on an object.
"""
while isinstance(element, Grouping):
element = element.element
if isinstance(element, BinaryExpression) and element.operator in (
operators.getitem,
operators.json_getitem_op,
):
return False
return isinstance(element.type, sa.types.JSON)


def _unwrap_bindparam(element):
"""
Return the bind parameter inside `element`, if it renders as a bare one.
"""
inner = element
while isinstance(inner, Grouping):
inner = inner.element
if isinstance(inner, TypeCoerce):
inner = inner.typed_expression
return inner if isinstance(inner, BindParameter) else element


class CrateCompiler(compiler.SQLCompiler):
def visit_typeclause(self, typeclause, **kw):
"""
Expand All @@ -345,11 +419,45 @@ def visit_typeclause(self, typeclause, **kw):
visit_on_conflict_do_update = PGCompiler.visit_on_conflict_do_update
_on_conflict_target = PGCompiler._on_conflict_target

def _render_subscript(self, binary, **kw):
left = self.process(binary.left, **kw)
index = _unwrap_bindparam(binary.right)
if isinstance(index, Slice):
return "%s[%s]" % (left, self._render_slice(index, **kw))
if not isinstance(index, BindParameter):
# CrateDB accepts bind parameters inside an expression.
return "%s[%s]" % (left, self.process(index, **kw))

# CrateDB does not accept a bind parameter as an array index or object
# key, so render it as a literal.
type_ = _SubscriptType(int_as_key=_is_top_level_object(binary.left))
if not index.required:
type_.literal_processor(self.dialect)(index.effective_value)
elif SA_VERSION < SA_1_4:
raise sa.exc.CompileError(
"A subscript taking its value on execution needs SQLAlchemy 1.4 or later"
)
if index.required or getattr(self, "cache_key", None) is not None:
kw["literal_execute"] = True
else:
kw["literal_binds"] = True
index = sa.type_coerce(index, type_).typed_expression
return "%s[%s]" % (left, self.process(index, **kw))

def _render_slice(self, slice_, **kw):
# CrateDB accepts bind parameters as slice bounds.
if not isinstance(slice_.step, Null):
raise sa.exc.CompileError("CrateDB array slices do not support a step")
return ":".join(
"" if isinstance(bound, Null) else self.process(bound, **kw)
for bound in (slice_.start, slice_.stop)
)

def visit_getitem_binary(self, binary, operator, **kw):
return "{0}['{1}']".format(self.process(binary.left, **kw), binary.right.value)
return self._render_subscript(binary, **kw)

def visit_json_getitem_op_binary(self, binary, operator, _cast_applied=False, **kw):
return "{0}['{1}']".format(self.process(binary.left, **kw), binary.right.value)
return self._render_subscript(binary, **kw)

def visit_any(self, element, **kw):
return "%s%sANY (%s)" % (
Expand Down
2 changes: 1 addition & 1 deletion tests/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@
from .function_test import SqlAlchemyFunctionTest
from .insert_from_select_test import SqlAlchemyInsertFromSelectTest
from .match_test import SqlAlchemyMatchTest
from .query_caching import SqlAlchemyQueryCompilationCaching
from .query_caching_test import SqlAlchemyQueryCompilationCaching
from .update_test import SqlAlchemyUpdateTest
from .warnings_test import SqlAlchemyWarningsTest

Expand Down
139 changes: 138 additions & 1 deletion tests/array_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@
from sqlalchemy.orm import Session
from sqlalchemy.sql import operators

from sqlalchemy_cratedb import SA_1_4, SA_VERSION
from sqlalchemy_cratedb import SA_1_4, SA_VERSION, ObjectArray

try:
from sqlalchemy.orm import declarative_base
Expand Down Expand Up @@ -57,10 +57,24 @@ class User(Base):

self.User = User
self.session = Session(bind=self.engine)
self.subscripts = sa.Table(
"subscripts",
self.metadata,
sa.Column("idx", sa.Integer),
sa.Column("arr0", sa.ARRAY(sa.Integer, zero_indexes=True)),
sa.Column("arr2d", sa.ARRAY(sa.Integer, dimensions=2)),
sa.Column("data_list", ObjectArray),
)

def assertSQL(self, expected_str, actual_expr):
self.assertEqual(expected_str, str(actual_expr).replace("\n", ""))

def assertWhere(self, expected_where, clause):
s = self.session.query(self.subscripts.c.idx).filter(clause)
self.assertSQL(
"SELECT subscripts.idx AS subscripts_idx FROM subscripts WHERE " + expected_where, s
)

@skipIf(SA_VERSION < SA_1_4, "`as_generic` not available with SQLAlchemy 1.3")
def test_as_generic(self):
t1 = sa.Table(
Expand Down Expand Up @@ -116,6 +130,129 @@ def test_any_with_operator(self):
s,
)

def test_index(self):
s = self.session.query(self.User.name).filter(self.User.scores[1] == 5)
self.assertSQL(
"SELECT users.name AS users_name FROM users WHERE users.scores[1] = %(param_1)s", s
)

def test_index_zero_indexes(self):
self.assertWhere("subscripts.arr0[1] = %(param_1)s", self.subscripts.c.arr0[0] == 5)

def test_index_nested(self):
self.assertWhere("subscripts.arr2d[1][2] = %(param_1)s", self.subscripts.c.arr2d[1][2] == 5)

def test_index_column(self):
t = self.subscripts
self.assertWhere(
"subscripts.arr0[(subscripts.idx + %(idx_1)s)] = %(param_1)s", t.c.arr0[t.c.idx] == 5
)

def test_index_expression(self):
s = self.session.query(self.User.name).filter(
self.User.scores[sa.func.char_length(self.User.name) + 1] == 5
)
self.assertSQL(
"SELECT users.name AS users_name FROM users "
"WHERE users.scores[(char_length(users.name) + %(char_length_1)s)] = %(param_1)s",
s,
)

def test_index_expression_bindparam(self):
s = self.session.query(self.User.name).filter(
self.User.scores[sa.func.char_length(self.User.name) + sa.bindparam("off")] == 5
)
self.assertSQL(
"SELECT users.name AS users_name FROM users "
"WHERE users.scores[(char_length(users.name) + %(off)s)] = %(param_1)s",
s,
)

def test_index_zero_indexes_bindparam(self):
self.assertWhere(
"subscripts.arr0[(%(pos)s + %(param_1)s)] = %(param_2)s",
self.subscripts.c.arr0[sa.bindparam("pos")] == 5,
)

def test_index_type_coerce(self):
s = self.session.query(self.User.name).filter(
self.User.scores[sa.type_coerce(sa.literal(2), sa.Integer)] == 5
)
self.assertSQL(
"SELECT users.name AS users_name FROM users WHERE users.scores[2] = %(param_1)s", s
)

@skipIf(SA_VERSION < SA_1_4, "SQLAlchemy 1.3 has no `literal_execute`")
def test_index_bindparam(self):
s = self.session.query(self.User.name).filter(self.User.scores[sa.bindparam("pos")] == 5)
self.assertSQL(
"SELECT users.name AS users_name FROM users "
"WHERE users.scores[__[POSTCOMPILE_pos]] = %(param_1)s",
s,
)

@skipIf(SA_VERSION >= SA_1_4, "SQLAlchemy 1.4+ renders the parameter on execution")
def test_index_bindparam_sa13(self):
s = self.session.query(self.User.name).filter(self.User.scores[sa.bindparam("pos")] == 5)
with self.assertRaises(sa.exc.CompileError) as cm:
str(s)
self.assertEqual(
"A subscript taking its value on execution needs SQLAlchemy 1.4 or later",
str(cm.exception),
)

def test_index_invalid(self):
s = self.session.query(self.User.name).filter(self.User.scores[1.5] == 5)
with self.assertRaises(sa.exc.CompileError) as cm:
str(s)
self.assertEqual(
"CrateDB subscripts take an integer index or a string key, not 1.5", str(cm.exception)
)

def test_slice(self):
s = self.session.query(self.User.name).filter(self.User.scores[1:2] == [5])
self.assertSQL(
"SELECT users.name AS users_name FROM users "
"WHERE users.scores[%(scores_1)s:%(scores_2)s] = %(param_1)s",
s,
)

def test_slice_open(self):
s = self.session.query(self.User.name).filter(self.User.scores[:2] == [5])
self.assertSQL(
"SELECT users.name AS users_name FROM users "
"WHERE users.scores[:%(scores_1)s] = %(param_1)s",
s,
)
s = self.session.query(self.User.name).filter(self.User.scores[2:] == [5])
self.assertSQL(
"SELECT users.name AS users_name FROM users "
"WHERE users.scores[%(scores_1)s:] = %(param_1)s",
s,
)

def test_slice_step(self):
s = self.session.query(self.User.name).filter(self.User.scores[1:3:2] == [5])
with self.assertRaises(sa.exc.CompileError) as cm:
str(s)
self.assertEqual("CrateDB array slices do not support a step", str(cm.exception))

def test_update_element(self):
stmt = sa.update(self.User.__table__).values({self.User.scores[1]: 99})
self.assertSQL("UPDATE users SET scores[1] = %(param_1)s", stmt.compile(bind=self.engine))

def test_object_array_key(self):
self.assertWhere(
"%(param_1)s = ANY (subscripts.data_list['foo'])",
self.subscripts.c.data_list["foo"].any(1),
)

def test_object_array_index_key(self):
self.assertWhere(
"subscripts.data_list[1]['foo'] = %(param_1)s",
self.subscripts.c.data_list[1]["foo"] == 1,
)

def test_multidimensional_arrays(self):
t1 = sa.Table(
"t",
Expand Down
Loading
Loading