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
74 changes: 57 additions & 17 deletions src/databricks/sqlalchemy/_parse.py
Original file line number Diff line number Diff line change
Expand Up @@ -336,33 +336,73 @@ def parse_numeric_type_precision_and_scale(type_name_str):
If type_name_str is "DECIMAL(18,5) returns sqlalchemy.types.Numeric(18,5)
"""

pattern = re.compile(r"DECIMAL\((\d+,\d+)\)")
pattern = re.compile(r"DECIMAL\(\s*(\d+)\s*,\s*(\d+)\s*\)", re.IGNORECASE)
match = re.search(pattern, type_name_str)
precision_and_scale = match.group(1)
precision, scale = tuple(precision_and_scale.split(","))
if match is None:
return sqlalchemy.types.Numeric()
precision, scale = match.groups()

return sqlalchemy.types.Numeric(int(precision), int(scale))


def _split_top_level(text: str) -> List[str]:
"""Split ``text`` on commas that are not nested inside <>, () or backticks."""
parts, depth, quoted, start = [], 0, False, 0
for index, char in enumerate(text):
if char == "`":
quoted = not quoted
elif quoted:
continue
elif char in "<(":
depth += 1
elif char in ">)":
depth -= 1
elif char == "," and depth == 0:
parts.append(text[start:index])
start = index + 1
parts.append(text[start:])
return [part.strip() for part in parts]


def parse_type_name(type_name: str) -> sqlalchemy.types.TypeEngine:
"""Return the SQLAlchemy type for a Databricks TYPE_NAME such as
``DECIMAL(10,2)``, ``ARRAY<STRING>`` or ``MAP<INT, ARRAY<BIGINT>>``.

ARRAY and MAP are parsed recursively into DatabricksArray / DatabricksMap.
A type that is not in GET_COLUMNS_TYPE_MAP (e.g. TIME, INTERVAL, VOID) is
reflected as NullType with a warning, like other SQLAlchemy dialects do,
instead of raising and failing reflection of the whole table.
"""
type_name = type_name.strip()
match = re.match(r"^(\w+)\s*<(.*)>$", type_name, re.DOTALL)
if match:
outer, inner = match.group(1).lower(), match.group(2)
args = _split_top_level(inner)
if outer == "array" and len(args) == 1:
return type_overrides.DatabricksArray(parse_type_name(args[0]))
if outer == "map" and len(args) == 2:
return type_overrides.DatabricksMap(
parse_type_name(args[0]), parse_type_name(args[1])
)

base_match = re.match(r"^\w+", type_name)
raw_type = base_match.group(0).lower() if base_match else ""
if raw_type == "decimal":
return parse_numeric_type_precision_and_scale(type_name)
if raw_type not in GET_COLUMNS_TYPE_MAP:
sqlalchemy.util.warn(
f"Did not recognize type '{type_name}' of a Databricks column"
)
return sqlalchemy.types.NullType()
return GET_COLUMNS_TYPE_MAP[raw_type]()


def parse_column_info_from_tgetcolumnsresponse(thrift_resp_row) -> ReflectedColumn:
"""Returns a dictionary of the ReflectedColumn schema parsed from
a single of the result of a TGetColumnsRequest thrift RPC
"""

pat = re.compile(r"^\w+")

# This method assumes a valid TYPE_NAME field in the response.
# TODO: add error handling in case TGetColumnsResponse format changes

_raw_col_type = re.search(pat, thrift_resp_row.TYPE_NAME).group(0).lower() # type: ignore
_col_type = GET_COLUMNS_TYPE_MAP[_raw_col_type]

if _raw_col_type == "decimal":
final_col_type = parse_numeric_type_precision_and_scale(
thrift_resp_row.TYPE_NAME
)
else:
final_col_type = _col_type
final_col_type = parse_type_name(thrift_resp_row.TYPE_NAME)

# See comments about autoincrement in test_suite.py
# Since Databricks SQL doesn't currently support inline AUTOINCREMENT declarations
Expand Down
60 changes: 60 additions & 0 deletions tests/test_local/test_parsing.py
Original file line number Diff line number Diff line change
Expand Up @@ -252,3 +252,63 @@ def test_multilevel_map_type_parsing(internal_type):
internal_type.compile(dialect=dialect)
)
assert actual_parsed == expected_parsed


@pytest.mark.parametrize(
"type_name, expected",
[
("ARRAY<STRING>", "ARRAY<STRING>"),
("ARRAY<ARRAY<BIGINT>>", "ARRAY<ARRAY<BIGINT>>"),
("MAP<STRING, INT>", "MAP<STRING,INT>"),
("MAP<INT, ARRAY<DECIMAL(10,2)>>", "MAP<INT,ARRAY<DECIMAL(10, 2)>>"),
("ARRAY<MAP<STRING, TIMESTAMP>>", "ARRAY<MAP<STRING,TIMESTAMP>>"),
("DECIMAL(38,18)", "DECIMAL(38, 18)"),
("BIGINT", "BIGINT"),
],
)
def test_parse_type_name_nested(type_name, expected):
from databricks.sqlalchemy._parse import parse_type_name

assert parse_type_name(type_name).compile(dialect=dialect) == expected


@pytest.mark.parametrize(
"type_name", ["TIME(6)", "INTERVAL DAY TO SECOND", "VOID", "GEOMETRY(4326)"]
)
def test_parse_type_name_unknown_type_does_not_raise(type_name):
from sqlalchemy.exc import SAWarning
from sqlalchemy.types import NullType

from databricks.sqlalchemy._parse import (
parse_column_info_from_tgetcolumnsresponse,
parse_type_name,
)

with pytest.warns(SAWarning, match="Did not recognize type"):
assert isinstance(parse_type_name(type_name), NullType)

class Row:
TYPE_NAME = type_name
COLUMN_NAME = "c"
NULLABLE = 1
COLUMN_DEF = None
REMARKS = None

with pytest.warns(SAWarning):
column = parse_column_info_from_tgetcolumnsresponse(Row())
assert column["name"] == "c" and isinstance(column["type"], NullType)


@pytest.mark.parametrize(
"type_name, expected",
[
("DECIMAL(10, 2)", "DECIMAL(10, 2)"),
("decimal(10,2)", "DECIMAL(10, 2)"),
("array<decimal(38, 18)>", "ARRAY<DECIMAL(38, 18)>"),
("DECIMAL", "DECIMAL"),
],
)
def test_parse_type_name_decimal_spelling(type_name, expected):
from databricks.sqlalchemy._parse import parse_type_name

assert parse_type_name(type_name).compile(dialect=dialect) == expected