mirror of
https://github.com/apache/superset.git
synced 2026-08-19 14:41:16 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d1fd9cc5f5 | ||
|
|
979857675f | ||
|
|
d5f5c40817 |
@@ -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
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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",
|
||||
[
|
||||
|
||||
@@ -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",
|
||||
[
|
||||
|
||||
Reference in New Issue
Block a user