Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
175 changes: 154 additions & 21 deletions src/databricks/sqlalchemy/_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand All @@ -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")
Expand Down Expand Up @@ -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<STRING, V> 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<STRING, {_json_schema(type_.value_type)}>"
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)
1 change: 1 addition & 0 deletions src/databricks/sqlalchemy/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
72 changes: 72 additions & 0 deletions tests/test_local/test_collection_binds.py
Original file line number Diff line number Diff line change
@@ -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<DECIMAL(38, 18)>', 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<STRING, STRING>', map('mode', 'FAILFAST')) "
"AS MAP<INT,STRING>)"
) 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]"
)