diff --git a/src/databricks/sqlalchemy/_types.py b/src/databricks/sqlalchemy/_types.py index c180404..c000c6f 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, dialect) - 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, dialect) - 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,151 @@ 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 _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, dialect) for v in value) + + "]" + ) + if isinstance(type_, DatabricksMap): + key_processor = _leaf_bind_processor(type_.key_type, dialect) + return ( + "{" + + ",".join( + json.dumps( + _json_key(key_processor(k) if key_processor else k), + ensure_ascii=False, + ) + + ":" + + _collection_json(v, type_.value_type, dialect) + for k, v in dict(value).items() + ) + + "}" + ) + processor = _leaf_bind_processor(type_, dialect) + return _json_scalar(processor(value) if processor else 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_, dialect): + def process(value): + return None if value is None else _collection_json(value, type_, dialect) + + 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..ff4c6ed --- /dev/null +++ b/tests/test_local/test_collection_binds.py @@ -0,0 +1,72 @@ +"""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" + + +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]" + )