From 3b0c4589b2189a26e593282c7bcfaaef42f9e7e0 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 09:55:54 +0000 Subject: [PATCH 1/2] Bind typed ARRAY/MAP values and LargeBinary correctly The warehouse ignores the elements of every Thrift ARRAY/MAP parameter encoding, so DatabricksArray/DatabricksMap values were persisted as empty collections without an error, and a NULL collection could not be bound. Send a typed value as one JSON string rebuilt with from_json(..., FAILFAST) (maps parsed with string keys and CAST to the declared type). LargeBinary binds raised AttributeError (the DB-API module has no Binary()) and raw bytes were sent as an ARRAY parameter; bind them as a hex string decoded with unhex(). --- src/databricks/sqlalchemy/_types.py | 156 +++++++++++++++++++--- src/databricks/sqlalchemy/base.py | 1 + tests/test_local/test_collection_binds.py | 58 ++++++++ 3 files changed, 194 insertions(+), 21 deletions(-) create mode 100644 tests/test_local/test_collection_binds.py diff --git a/src/databricks/sqlalchemy/_types.py b/src/databricks/sqlalchemy/_types.py index c180404..589cadd 100644 --- a/src/databricks/sqlalchemy/_types.py +++ b/src/databricks/sqlalchemy/_types.py @@ -426,14 +426,10 @@ def __init__(self, item_type): self.item_type = item_type() if isinstance(item_type, type) else item_type def bind_processor(self, dialect): - item_processor = self.item_type.bind_processor(dialect) - if item_processor is None: - item_processor = identity_processor + return _collection_bind_processor(self) - def process(value): - return [item_processor(val) for val in value] - - return process + def bind_expression(self, bindvalue): + return _collection_bind_expression(self, bindvalue) @compiles(DatabricksArray, "databricks") @@ -458,21 +454,10 @@ def __init__(self, key_type, value_type): self.value_type = value_type() if isinstance(value_type, type) else value_type def bind_processor(self, dialect): - key_processor = self.key_type.bind_processor(dialect) - value_processor = self.value_type.bind_processor(dialect) - - if key_processor is None: - key_processor = identity_processor - if value_processor is None: - value_processor = identity_processor + return _collection_bind_processor(self) - def process(value): - return { - key_processor(key): value_processor(value) - for key, value in value.items() - } - - return process + def bind_expression(self, bindvalue): + return _collection_bind_expression(self, bindvalue) @compiles(DatabricksMap, "databricks") @@ -537,3 +522,132 @@ def process(value): @compiles(DatabricksVariant, "databricks") def compile_variant(type_, compiler, **kw): return "VARIANT" + + +# --- Binding ARRAY/MAP values ------------------------------------------------- +# +# The warehouse ignores the elements of every Thrift ARRAY/MAP parameter +# encoding (native ArrayParameter/MapParameter binds persist as empty +# collections, without an error). A typed value is therefore sent as ONE JSON +# STRING parameter and rebuilt in SQL with from_json. JSON object keys are +# strings, so maps parse as MAP and are CAST to the declared type. +# FAILFAST makes malformed elements fail the statement instead of becoming NULL. +# The SQL does not depend on the element count and uses no lambdas, so it works +# in multi-row INSERT ... VALUES and executemany. + +_FROM_JSON_OPTIONS = "map('mode', 'FAILFAST')" + + +def _json_scalar(item): + import decimal + import math + from datetime import date + + if item is None: + return "null" + if isinstance(item, bool): + return "true" if item else "false" + if isinstance(item, int): + return str(item) + if isinstance(item, decimal.Decimal): + if not item.is_finite(): + raise ValueError("non-finite Decimal in a Databricks collection bind") + return format(item, "f") + if isinstance(item, float): + if not math.isfinite(item): + raise ValueError("non-finite float in a Databricks collection bind") + return repr(item) + if isinstance(item, str): + return json.dumps(item, ensure_ascii=False) + if isinstance(item, (datetime, date)): + return json.dumps(item.isoformat()) + raise TypeError(f"cannot bind {type(item).__name__} inside a Databricks collection") + + +def _json_key(key): + if key is None: + raise ValueError("NULL map key in a Databricks collection bind") + if isinstance(key, str): + return key + if isinstance(key, bool): + return "true" if key else "false" + if isinstance(key, (datetime,)) or hasattr(key, "isoformat"): + return key.isoformat() + return _json_scalar(key) + + +def _collection_json(value, type_): + if value is None: + return "null" + if isinstance(type_, DatabricksArray): + return "[" + ",".join(_collection_json(v, type_.item_type) for v in value) + "]" + if isinstance(type_, DatabricksMap): + return ( + "{" + + ",".join( + json.dumps(_json_key(k), ensure_ascii=False) + + ":" + + _collection_json(v, type_.value_type) + for k, v in dict(value).items() + ) + + "}" + ) + return _json_scalar(value) + + +def _json_schema(type_): + if isinstance(type_, DatabricksArray): + return f"ARRAY<{_json_schema(type_.item_type)}>" + if isinstance(type_, DatabricksMap): + return f"MAP" + from databricks.sqlalchemy import DatabricksDialect + + return DatabricksDialect().type_compiler_instance.process(type_) + + +def _has_map(type_): + if isinstance(type_, DatabricksMap): + return True + return isinstance(type_, DatabricksArray) and _has_map(type_.item_type) + + +def _collection_bind_processor(type_): + def process(value): + return None if value is None else _collection_json(value, type_) + + return process + + +def _collection_bind_expression(type_, bindvalue): + parsed = expression.func.from_json( + bindvalue, + expression.literal_column(f"'{_json_schema(type_)}'"), + expression.literal_column(_FROM_JSON_OPTIONS), + ) + if _has_map(type_): + return expression.cast(parsed, type_) + return expression.type_coerce(parsed, type_) + + +class DatabricksBinary(TypeDecorator): + """LargeBinary values bind as a hex STRING decoded with unhex(). + + The connector has no DB-API Binary() constructor (LargeBinary's default bind + processor raises AttributeError), treats bytes as a Sequence (sending an + ARRAY parameter) and the warehouse rejects BINARY as a parameter type. + """ + + impl = sqlalchemy.types.LargeBinary + cache_ok = True + + def bind_processor(self, dialect): + def process(value): + return None if value is None else bytes(value).hex() + + return process + + def bind_expression(self, bindvalue): + return expression.func.unhex(bindvalue, type_=self) + + def process_result_value(self, value, dialect): + return None if value is None else bytes(value) diff --git a/src/databricks/sqlalchemy/base.py b/src/databricks/sqlalchemy/base.py index bcdd6a8..905cc1c 100644 --- a/src/databricks/sqlalchemy/base.py +++ b/src/databricks/sqlalchemy/base.py @@ -82,6 +82,7 @@ class DatabricksDialect(default.DefaultDialect): sqlalchemy.types.Time: dialect_type_impl.DatabricksTimeType, sqlalchemy.types.String: dialect_type_impl.DatabricksStringType, sqlalchemy.types.Uuid: dialect_type_impl.DatabricksUUID, + sqlalchemy.types._Binary: dialect_type_impl.DatabricksBinary, } # SQLAlchemy requires that a table with no primary key diff --git a/tests/test_local/test_collection_binds.py b/tests/test_local/test_collection_binds.py new file mode 100644 index 0000000..76bdbaa --- /dev/null +++ b/tests/test_local/test_collection_binds.py @@ -0,0 +1,58 @@ +"""Typed ARRAY/MAP and binary values bind as JSON / hex strings rebuilt in SQL.""" + +from datetime import date, datetime +from decimal import Decimal + +import sqlalchemy as sa + +from databricks.sqlalchemy import DatabricksArray, DatabricksMap + +engine = sa.create_engine("databricks://token:x@host?http_path=p") +dialect = engine.dialect + + +def compile_insert(*columns): + table = sa.Table("t", sa.MetaData(), *columns) + return str(table.insert().compile(dialect=dialect)) + + +def processed(type_, value): + return type_.dialect_impl(dialect).bind_processor(dialect)(value) + + +def test_array_renders_from_json(): + sql = compile_insert(sa.Column("a", DatabricksArray(sa.Numeric(38, 18)))) + assert "from_json(:`a`, 'ARRAY', map('mode', 'FAILFAST'))" in sql + assert "CAST" not in sql + + +def test_map_is_parsed_with_string_keys_and_cast(): + sql = compile_insert(sa.Column("m", DatabricksMap(sa.Integer, sa.String))) + assert ( + "CAST(from_json(:`m`, 'MAP', map('mode', 'FAILFAST')) " + "AS MAP)" + ) in sql + + +def test_values_serialize_exactly(): + assert ( + processed(DatabricksArray(sa.Numeric(38, 18)), [Decimal("1E-18"), None]) + == "[0.000000000000000001,null]" + ) + assert processed(DatabricksArray(sa.BigInteger), [9007199254740993]) == ( + "[9007199254740993]" + ) + assert processed(DatabricksMap(sa.Integer, sa.String), {7: "雪"}) == '{"7":"雪"}' + assert processed(DatabricksArray(sa.Date), [date(2024, 2, 29)]) == ( + '["2024-02-29"]' + ) + assert processed(DatabricksArray(sa.DateTime), [datetime(2024, 1, 1, 1, 2, 3, 4)]) == ( + '["2024-01-01T01:02:03.000004"]' + ) + assert processed(DatabricksArray(sa.String), None) is None + + +def test_binary_binds_as_hex_with_unhex(): + sql = compile_insert(sa.Column("b", sa.LargeBinary)) + assert "unhex(:`b`)" in sql + assert processed(sa.LargeBinary(), b"\x00\xff") == "00ff" From dcefdf1eb47f53aff6c51cc58cc114892c48ec6f Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Wed, 30 Sep 2026 04:11:20 +0000 Subject: [PATCH 2/2] Apply element bind processors inside ARRAY/MAP values Serializing collection values to JSON skipped each element type's own bind processor, so DatabricksArray(Uuid), DatabricksArray(Time) and DatabricksMap(Uuid, ...) raised TypeError. Run the element's and key's dialect bind processor before serializing. Numeric has no bind processor here, so decimals still serialize exactly. --- src/databricks/sqlalchemy/_types.py | 37 +++++++++++++++++------ tests/test_local/test_collection_binds.py | 14 +++++++++ 2 files changed, 42 insertions(+), 9 deletions(-) diff --git a/src/databricks/sqlalchemy/_types.py b/src/databricks/sqlalchemy/_types.py index 589cadd..c000c6f 100644 --- a/src/databricks/sqlalchemy/_types.py +++ b/src/databricks/sqlalchemy/_types.py @@ -426,7 +426,7 @@ def __init__(self, item_type): self.item_type = item_type() if isinstance(item_type, type) else item_type def bind_processor(self, dialect): - return _collection_bind_processor(self) + return _collection_bind_processor(self, dialect) def bind_expression(self, bindvalue): return _collection_bind_expression(self, bindvalue) @@ -454,7 +454,7 @@ def __init__(self, key_type, value_type): self.value_type = value_type() if isinstance(value_type, type) else value_type def bind_processor(self, dialect): - return _collection_bind_processor(self) + return _collection_bind_processor(self, dialect) def bind_expression(self, bindvalue): return _collection_bind_expression(self, bindvalue) @@ -576,23 +576,42 @@ def _json_key(key): return _json_scalar(key) -def _collection_json(value, type_): +def _leaf_bind_processor(type_, dialect): + """The element type's own (dialect-adapted) bind processor, if any. + + Applied before JSON serialization so element types that convert their + values (Uuid, Time, custom TypeDecorators, ...) keep working inside a + collection, as they did before values were sent as JSON. + """ + return type_.dialect_impl(dialect).bind_processor(dialect) + + +def _collection_json(value, type_, dialect): if value is None: return "null" if isinstance(type_, DatabricksArray): - return "[" + ",".join(_collection_json(v, type_.item_type) for v in value) + "]" + return ( + "[" + + ",".join(_collection_json(v, type_.item_type, dialect) for v in value) + + "]" + ) if isinstance(type_, DatabricksMap): + key_processor = _leaf_bind_processor(type_.key_type, dialect) return ( "{" + ",".join( - json.dumps(_json_key(k), ensure_ascii=False) + json.dumps( + _json_key(key_processor(k) if key_processor else k), + ensure_ascii=False, + ) + ":" - + _collection_json(v, type_.value_type) + + _collection_json(v, type_.value_type, dialect) for k, v in dict(value).items() ) + "}" ) - return _json_scalar(value) + processor = _leaf_bind_processor(type_, dialect) + return _json_scalar(processor(value) if processor else value) def _json_schema(type_): @@ -611,9 +630,9 @@ def _has_map(type_): return isinstance(type_, DatabricksArray) and _has_map(type_.item_type) -def _collection_bind_processor(type_): +def _collection_bind_processor(type_, dialect): def process(value): - return None if value is None else _collection_json(value, type_) + return None if value is None else _collection_json(value, type_, dialect) return process diff --git a/tests/test_local/test_collection_binds.py b/tests/test_local/test_collection_binds.py index 76bdbaa..ff4c6ed 100644 --- a/tests/test_local/test_collection_binds.py +++ b/tests/test_local/test_collection_binds.py @@ -56,3 +56,17 @@ def test_binary_binds_as_hex_with_unhex(): sql = compile_insert(sa.Column("b", sa.LargeBinary)) assert "unhex(:`b`)" in sql assert processed(sa.LargeBinary(), b"\x00\xff") == "00ff" + + +def test_element_bind_processors_are_applied(): + import uuid + from datetime import time + + u = uuid.UUID(int=1) + assert processed(DatabricksArray(sa.Uuid), [u, None]) == f'["{u}",null]' + assert processed(DatabricksArray(sa.Time), [time(1, 2, 3)]) == '["01:02:03"]' + assert processed(DatabricksMap(sa.Uuid, sa.Integer), {u: 1}) == f'{{"{u}":1}}' + # exact numeric serialization is unchanged (Numeric has no bind processor here) + assert processed(DatabricksArray(sa.Numeric(38, 18)), [Decimal("1E-18")]) == ( + "[0.000000000000000001]" + )