From 85bfa9478c05de4d5e43ff03d1bda6a42321ce7c Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 09:53:25 +0000 Subject: [PATCH 1/2] Reflect nested ARRAY/MAP types and tolerate unknown column types get_columns mapped only the first word of TYPE_NAME: ARRAY and MAP columns reflected as String, and any type missing from GET_COLUMNS_TYPE_MAP (e.g. TIME, INTERVAL, VOID) raised KeyError and failed reflection of the whole table or view. Parse TYPE_NAME recursively into DatabricksArray/DatabricksMap and reflect unrecognised types as NullType with a warning, as other SQLAlchemy dialects do. --- src/databricks/sqlalchemy/_parse.py | 67 +++++++++++++++++++++++------ tests/test_local/test_parsing.py | 45 +++++++++++++++++++ 2 files changed, 98 insertions(+), 14 deletions(-) diff --git a/src/databricks/sqlalchemy/_parse.py b/src/databricks/sqlalchemy/_parse.py index 37a6cc4..3b1372c 100644 --- a/src/databricks/sqlalchemy/_parse.py +++ b/src/databricks/sqlalchemy/_parse.py @@ -344,25 +344,64 @@ def parse_numeric_type_precision_and_scale(type_name_str): 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`` or ``MAP>``. + + 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 diff --git a/tests/test_local/test_parsing.py b/tests/test_local/test_parsing.py index 026b6a4..1c1758d 100644 --- a/tests/test_local/test_parsing.py +++ b/tests/test_local/test_parsing.py @@ -252,3 +252,48 @@ 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", "ARRAY"), + ("ARRAY>", "ARRAY>"), + ("MAP", "MAP"), + ("MAP>", "MAP>"), + ("ARRAY>", "ARRAY>"), + ("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) From 76343882e709807c4c7b6f72cd6eff51351d6fe3 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Wed, 30 Sep 2026 04:11:18 +0000 Subject: [PATCH 2/2] Parse DECIMAL precision and scale in any case and spacing A nested DECIMAL element spelled with a space or in lower case, e.g. ARRAY, raised AttributeError and failed reflection of the whole table. Match case-insensitively, tolerate whitespace, and fall back to Numeric() for a bare DECIMAL. --- src/databricks/sqlalchemy/_parse.py | 7 ++++--- tests/test_local/test_parsing.py | 15 +++++++++++++++ 2 files changed, 19 insertions(+), 3 deletions(-) diff --git a/src/databricks/sqlalchemy/_parse.py b/src/databricks/sqlalchemy/_parse.py index 3b1372c..bc33d96 100644 --- a/src/databricks/sqlalchemy/_parse.py +++ b/src/databricks/sqlalchemy/_parse.py @@ -336,10 +336,11 @@ 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)) diff --git a/tests/test_local/test_parsing.py b/tests/test_local/test_parsing.py index 1c1758d..2c196a6 100644 --- a/tests/test_local/test_parsing.py +++ b/tests/test_local/test_parsing.py @@ -297,3 +297,18 @@ class Row: 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", "ARRAY"), + ("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