Compare commits

...
8 changed files with 181 additions and 8 deletions
+18 -7
View File
@@ -1228,7 +1228,11 @@ class TableColumn(AuditMixinNullable, ImportExportMixin, CertificationMixin, Mod
expression = self._validate_stored_expression(expression)
col = literal_column(expression, type_=type_)
else:
col = column(self.column_name, type_=type_)
identifier = db_engine_spec.prepare_identifier(
cast(str, self.column_name),
normalize_columns=bool(getattr(self.table, "normalize_columns", False)),
)
col = column(identifier, type_=type_)
col = self.database.make_sqla_column_compatible(col, label)
return col
@@ -1254,12 +1258,15 @@ class TableColumn(AuditMixinNullable, ImportExportMixin, CertificationMixin, Mod
pdf = self.python_date_format
is_epoch = pdf in ("epoch_s", "epoch_ms")
column_spec = self.db_engine_spec.get_column_spec(
self.type, db_extra=self.db_extra
)
db_engine_spec = self.db_engine_spec
column_spec = db_engine_spec.get_column_spec(self.type, db_extra=self.db_extra)
type_ = column_spec.sqla_type if column_spec else DateTime
if not self.expression and not time_grain and not is_epoch:
sqla_col = column(self.column_name, type_=type_)
identifier = db_engine_spec.prepare_identifier(
cast(str, self.column_name),
normalize_columns=bool(getattr(self.table, "normalize_columns", False)),
)
sqla_col = column(identifier, type_=type_)
return self.database.make_sqla_column_compatible(sqla_col, label)
if expression := self.expression:
if template_processor:
@@ -1284,8 +1291,12 @@ class TableColumn(AuditMixinNullable, ImportExportMixin, CertificationMixin, Mod
expression = self._validate_stored_expression(expression)
col = literal_column(expression, type_=type_)
else:
col = column(self.column_name, type_=type_)
time_expr = self.db_engine_spec.get_timestamp_expr(col, pdf, time_grain)
identifier = db_engine_spec.prepare_identifier(
cast(str, self.column_name),
normalize_columns=bool(getattr(self.table, "normalize_columns", False)),
)
col = column(identifier, type_=type_)
time_expr = db_engine_spec.get_timestamp_expr(col, pdf, time_grain)
return self.database.make_sqla_column_compatible(time_expr, label)
@property
+13
View File
@@ -2808,6 +2808,19 @@ class BaseEngineSpec: # pylint: disable=too-many-public-methods
return name
@classmethod
def prepare_identifier(
cls,
name: str,
normalize_columns: bool = False,
) -> str:
"""
Prepare a physical identifier for SQLAlchemy column construction.
The default preserves SQLAlchemy's automatic identifier-quoting behavior.
"""
return name
@classmethod
def quote_table(cls, table: Table, dialect: Dialect) -> str:
"""
+12
View File
@@ -33,6 +33,7 @@ from marshmallow import fields, Schema
from sqlalchemy import text, types
from sqlalchemy.engine.reflection import Inspector
from sqlalchemy.engine.url import URL
from sqlalchemy.sql import quoted_name
from superset.constants import TimeGrain
from superset.databases.utils import make_url_safe
@@ -98,6 +99,17 @@ class SnowflakeEngineSpec(PostgresBaseEngineSpec):
supports_catalog = supports_dynamic_catalog = supports_cross_catalog_queries = True
supports_grouping_sets = True
@classmethod
def prepare_identifier(
cls,
name: str,
normalize_columns: bool = False,
) -> str:
"""Preserve exact-case physical identifiers when columns are not normalized."""
if normalize_columns:
return name
return quoted_name(name, quote=True)
metadata = {
"description": "Snowflake is a cloud-native data warehouse.",
"logo": "snowflake.svg",
+5 -1
View File
@@ -3899,7 +3899,11 @@ class ExploreMixin: # pylint: disable=too-many-public-methods
expression = self._validate_stored_expression(expression)
col = literal_column(expression, type_=type_)
else:
col = sa.column(tbl_column.column_name, type_=type_)
identifier = db_engine_spec.prepare_identifier(
cast(str, tbl_column.column_name),
normalize_columns=bool(self.normalize_columns),
)
col = sa.column(identifier, type_=type_)
col = self.make_sqla_column_compatible(col, label)
return col
@@ -21,6 +21,7 @@ import pandas as pd
import pytest
from pytest_mock import MockerFixture
from sqlalchemy import create_engine
from sqlalchemy.dialects import sqlite
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm.session import Session
@@ -49,6 +50,70 @@ from superset.superset_typing import QueryObjectDict
from superset.utils import json
def test_get_sqla_col_quotes_snowflake_case_sensitive_identifier(
mocker: MockerFixture,
) -> None:
"""Snowflake physical columns retain their exact reflected case in generated SQL."""
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
database = Database(database_name="db", sqlalchemy_uri="sqlite://")
mocker.patch.object(
Database,
"get_db_engine_spec",
return_value=SnowflakeEngineSpec,
)
table = SqlaTable(
table_name="bug_test",
database=database,
normalize_columns=False,
)
tbl_column = TableColumn(column_name="id", type="INTEGER", table=table)
rendered = str(
tbl_column.get_sqla_col().compile(
dialect=sqlite.dialect(),
compile_kwargs={"literal_binds": True},
)
)
assert rendered == '"id"'
@pytest.mark.parametrize("time_grain", [None, "P1D"])
def test_get_timestamp_expression_quotes_snowflake_case_sensitive_identifier(
mocker: MockerFixture,
time_grain: str | None,
) -> None:
"""Snowflake timestamp paths quote exact-case physical columns."""
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
database = Database(database_name="db", sqlalchemy_uri="sqlite://")
mocker.patch.object(
Database,
"get_db_engine_spec",
return_value=SnowflakeEngineSpec,
)
table = SqlaTable(
table_name="bug_test",
database=database,
normalize_columns=False,
)
tbl_column = TableColumn(
column_name="created_at",
type="TIMESTAMP",
table=table,
)
rendered = str(
tbl_column.get_timestamp_expression(time_grain=time_grain).compile(
dialect=sqlite.dialect(),
compile_kwargs={"literal_binds": True},
)
)
assert '"created_at"' in rendered
def test_query_bubbles_errors(mocker: MockerFixture) -> None:
"""
Test that the `query` method bubbles exceptions correctly.
@@ -291,6 +291,13 @@ def test_get_default_catalog(mocker: MockerFixture) -> None:
assert BaseEngineSpec.get_default_catalog(database) is None
def test_prepare_identifier_returns_name_unchanged() -> None:
name = "physical_column"
assert BaseEngineSpec.prepare_identifier(name, normalize_columns=False) is name
assert BaseEngineSpec.prepare_identifier(name, normalize_columns=True) is name
def test_quote_table() -> None:
"""
Test the `quote_table` function.
@@ -24,6 +24,7 @@ from unittest import mock
import pytest
from pytest_mock import MockerFixture
from sqlalchemy.engine.url import make_url
from sqlalchemy.sql import quoted_name
from superset.errors import ErrorLevel, SupersetError, SupersetErrorType
from superset.utils import json
@@ -31,6 +32,32 @@ from tests.unit_tests.db_engine_specs.utils import assert_convert_dttm
from tests.unit_tests.fixtures.common import dttm # noqa: F401
@pytest.mark.parametrize("name", ["lowercase", "UPPERCASE"])
def test_prepare_identifier_quotes_exact_case_names(name: str) -> None:
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
identifier = SnowflakeEngineSpec.prepare_identifier(
name,
normalize_columns=False,
)
assert isinstance(identifier, quoted_name)
assert str(identifier) == name
assert identifier.quote is True
def test_prepare_identifier_preserves_normalized_name() -> None:
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
name = "lowercase"
identifier = SnowflakeEngineSpec.prepare_identifier(
name,
normalize_columns=True,
)
assert identifier is name
@pytest.mark.parametrize(
"target_type,expected_result",
[
+34
View File
@@ -4552,6 +4552,40 @@ def test_simple_metric_quotes_column_requiring_quoting(database: Database) -> No
)
def test_convert_tbl_column_quotes_snowflake_case_sensitive_identifier(
database: Database,
mocker: MockerFixture,
) -> None:
"""The chart query-object path quotes exact-case Snowflake physical columns."""
from superset.connectors.sqla.models import SqlaTable, TableColumn
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
from superset.models.core import Database
mocker.patch.object(
Database,
"get_db_engine_spec",
return_value=SnowflakeEngineSpec,
)
table = SqlaTable(
database=database,
table_name="bug_test",
normalize_columns=False,
)
tbl_column = TableColumn(column_name="name", type="VARCHAR", table=table)
with database.get_sqla_engine() as engine:
dialect = engine.dialect
rendered = str(
table.convert_tbl_column_to_sqla_col(tbl_column).compile(
dialect=dialect,
compile_kwargs={"literal_binds": True},
)
)
assert rendered == '"name"'
@pytest.mark.parametrize(
"native_type",
[