mirror of
https://github.com/apache/superset.git
synced 2026-09-01 04:51:23 +00:00
experiment: migrate Query.get() calls to session.get() for SQLAlchemy 2.0
session.query(Model).get(id) is deprecated under SQLAlchemy 2.0 in favor of session.get(Model, id). Migrate the 9 remaining call sites (flagged by review) and update the corresponding test mocks that asserted against the old session.query(...).get(...) call shape. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.8
parent
96cda701c3
commit
bc56ecbf15
@@ -155,7 +155,7 @@ def export_example( # noqa: C901
|
||||
|
||||
# Find the dashboard
|
||||
if dashboard_id:
|
||||
dashboard = db.session.query(Dashboard).get(dashboard_id)
|
||||
dashboard = db.session.get(Dashboard, dashboard_id)
|
||||
elif dashboard_slug:
|
||||
dashboard = db.session.query(Dashboard).filter_by(slug=dashboard_slug).first()
|
||||
else:
|
||||
|
||||
@@ -65,7 +65,7 @@ class DuplicateDatasetCommand(CreateMixin, BaseCommand):
|
||||
database_id = self._base_model.database_id
|
||||
table_name = self._properties["table_name"]
|
||||
editors = self._properties["editors"]
|
||||
database = db.session.query(Database).get(database_id)
|
||||
database = db.session.get(Database, database_id)
|
||||
if not database:
|
||||
raise SupersetErrorException(
|
||||
SupersetError(
|
||||
|
||||
@@ -69,7 +69,7 @@ def transpile_virtual_dataset_sql(config: dict[str, Any], database_id: int) -> N
|
||||
if not sql:
|
||||
return
|
||||
|
||||
database = db.session.query(Database).get(database_id)
|
||||
database = db.session.get(Database, database_id)
|
||||
if not database:
|
||||
logger.warning("Database %s not found, skipping SQL transpilation", database_id)
|
||||
return
|
||||
|
||||
@@ -66,7 +66,7 @@ class QueryEstimationCommand(BaseCommand):
|
||||
self._catalog = params.get("catalog")
|
||||
|
||||
def validate(self) -> None:
|
||||
self._database = db.session.query(Database).get(self._database_id)
|
||||
self._database = db.session.get(Database, self._database_id)
|
||||
if not self._database:
|
||||
raise SupersetErrorException(
|
||||
SupersetError(
|
||||
|
||||
@@ -630,7 +630,7 @@ class DatasetDAO(BaseDAO[SqlaTable]):
|
||||
dataset = DatasetDAO.find_by_id(dataset_id)
|
||||
if not dataset:
|
||||
return None
|
||||
return db.session.query(SqlMetric).get(metric_id)
|
||||
return db.session.get(SqlMetric, metric_id)
|
||||
|
||||
@staticmethod
|
||||
def get_table_by_name(database_id: int, table_name: str) -> SqlaTable | None:
|
||||
|
||||
@@ -81,7 +81,7 @@ def generate_preview_from_form_data(
|
||||
from superset.connectors.sqla.models import SqlaTable
|
||||
from superset.extensions import db
|
||||
|
||||
dataset = db.session.query(SqlaTable).get(dataset_id)
|
||||
dataset = db.session.get(SqlaTable, dataset_id)
|
||||
if not dataset:
|
||||
return ChartError(
|
||||
error=f"Dataset {dataset_id} not found", error_type="DatasetNotFound"
|
||||
|
||||
+1
-1
@@ -168,7 +168,7 @@ def upgrade(): # noqa: C901
|
||||
match_ds_id = re.match(r"\[None\]\.\[.*\]\(id:(\d+)\)", faulty_view_menu.name)
|
||||
if match_ds_id:
|
||||
dataset_id = int(match_ds_id.group(1))
|
||||
dataset = session.query(SqlaTable).get(dataset_id)
|
||||
dataset = session.get(SqlaTable, dataset_id)
|
||||
if dataset:
|
||||
try:
|
||||
new_view_menu = dataset.get_perm()
|
||||
|
||||
+1
-1
@@ -92,7 +92,7 @@ def upgrade():
|
||||
if "granularity" in params or "granularity_sqla" in params:
|
||||
continue
|
||||
|
||||
table = session.query(SqlaTable).get(slc.datasource_id)
|
||||
table = session.get(SqlaTable, slc.datasource_id)
|
||||
if not table:
|
||||
continue
|
||||
|
||||
|
||||
@@ -3008,7 +3008,7 @@ class SupersetSecurityManager( # pylint: disable=too-many-public-methods
|
||||
logger.warning(
|
||||
"Dataset has no database will retry with database_id to set permission"
|
||||
)
|
||||
database = self.session.query(Database).get(target.database_id)
|
||||
database = self.session.get(Database, target.database_id)
|
||||
dataset_perm = self.get_dataset_perm(
|
||||
target.id, target.table_name, database.database_name
|
||||
)
|
||||
|
||||
@@ -198,10 +198,8 @@ def test_duplicate_dataset_success() -> None:
|
||||
),
|
||||
patch("superset.commands.dataset.duplicate.security_manager.raise_for_access"),
|
||||
):
|
||||
with patch(
|
||||
"superset.commands.dataset.duplicate.db.session.query"
|
||||
) as mock_query:
|
||||
mock_query.return_value.get.return_value = mock_database
|
||||
with patch("superset.commands.dataset.duplicate.db.session.get") as mock_get:
|
||||
mock_get.return_value = mock_database
|
||||
with patch(
|
||||
"superset.commands.dataset.duplicate.DatasetDAO.validate_uniqueness",
|
||||
return_value=True,
|
||||
@@ -369,10 +367,8 @@ def test_duplicate_dataset_catalog_preserved() -> None:
|
||||
),
|
||||
patch("superset.commands.dataset.duplicate.security_manager.raise_for_access"),
|
||||
):
|
||||
with patch(
|
||||
"superset.commands.dataset.duplicate.db.session.query"
|
||||
) as mock_query:
|
||||
mock_query.return_value.get.return_value = mock_database
|
||||
with patch("superset.commands.dataset.duplicate.db.session.get") as mock_get:
|
||||
mock_get.return_value = mock_database
|
||||
with patch(
|
||||
"superset.commands.dataset.duplicate.DatasetDAO.validate_uniqueness",
|
||||
return_value=True,
|
||||
@@ -498,10 +494,8 @@ def test_duplicate_dataset_with_columns_and_metrics() -> None:
|
||||
),
|
||||
patch("superset.commands.dataset.duplicate.security_manager.raise_for_access"),
|
||||
):
|
||||
with patch(
|
||||
"superset.commands.dataset.duplicate.db.session.query"
|
||||
) as mock_query:
|
||||
mock_query.return_value.get.return_value = mock_database
|
||||
with patch("superset.commands.dataset.duplicate.db.session.get") as mock_get:
|
||||
mock_get.return_value = mock_database
|
||||
with patch(
|
||||
"superset.commands.dataset.duplicate.DatasetDAO.validate_uniqueness",
|
||||
return_value=True,
|
||||
|
||||
@@ -39,7 +39,7 @@ def test_transpile_virtual_dataset_sql_empty_sql():
|
||||
@patch("superset.commands.importers.v1.examples.db")
|
||||
def test_transpile_virtual_dataset_sql_database_not_found(mock_db):
|
||||
"""Test graceful handling when database is not found."""
|
||||
mock_db.session.query.return_value.get.return_value = None
|
||||
mock_db.session.get.return_value = None
|
||||
|
||||
config = {"table_name": "my_table", "sql": "SELECT * FROM foo"}
|
||||
original_sql = config["sql"]
|
||||
@@ -56,7 +56,7 @@ def test_transpile_virtual_dataset_sql_success(mock_transpile, mock_db):
|
||||
"""Test successful SQL transpilation with source engine."""
|
||||
mock_database = MagicMock()
|
||||
mock_database.db_engine_spec.engine = "mysql"
|
||||
mock_db.session.query.return_value.get.return_value = mock_database
|
||||
mock_db.session.get.return_value = mock_database
|
||||
|
||||
mock_transpile.return_value = "SELECT * FROM `foo`"
|
||||
|
||||
@@ -77,7 +77,7 @@ def test_transpile_virtual_dataset_sql_no_source_engine(mock_transpile, mock_db)
|
||||
"""Test transpilation when source_db_engine is not specified (legacy)."""
|
||||
mock_database = MagicMock()
|
||||
mock_database.db_engine_spec.engine = "mysql"
|
||||
mock_db.session.query.return_value.get.return_value = mock_database
|
||||
mock_db.session.get.return_value = mock_database
|
||||
|
||||
mock_transpile.return_value = "SELECT * FROM `foo`"
|
||||
|
||||
@@ -95,7 +95,7 @@ def test_transpile_virtual_dataset_sql_no_change(mock_transpile, mock_db):
|
||||
"""Test when transpilation returns same SQL (no dialect differences)."""
|
||||
mock_database = MagicMock()
|
||||
mock_database.db_engine_spec.engine = "postgresql"
|
||||
mock_db.session.query.return_value.get.return_value = mock_database
|
||||
mock_db.session.get.return_value = mock_database
|
||||
|
||||
original_sql = "SELECT * FROM foo"
|
||||
mock_transpile.return_value = original_sql
|
||||
@@ -118,7 +118,7 @@ def test_transpile_virtual_dataset_sql_error_fallback(mock_transpile, mock_db):
|
||||
|
||||
mock_database = MagicMock()
|
||||
mock_database.db_engine_spec.engine = "mysql"
|
||||
mock_db.session.query.return_value.get.return_value = mock_database
|
||||
mock_db.session.get.return_value = mock_database
|
||||
|
||||
mock_transpile.side_effect = QueryClauseValidationException("Parse error")
|
||||
|
||||
@@ -140,7 +140,7 @@ def test_transpile_virtual_dataset_sql_postgres_to_duckdb(mock_transpile, mock_d
|
||||
"""Test transpilation from PostgreSQL to DuckDB."""
|
||||
mock_database = MagicMock()
|
||||
mock_database.db_engine_spec.engine = "duckdb"
|
||||
mock_db.session.query.return_value.get.return_value = mock_database
|
||||
mock_db.session.get.return_value = mock_database
|
||||
|
||||
original_sql = """
|
||||
SELECT DATE_TRUNC('month', created_at) AS month, COUNT(*) AS cnt
|
||||
@@ -173,7 +173,7 @@ def test_transpile_virtual_dataset_sql_postgres_to_clickhouse(mock_transpile, mo
|
||||
"""
|
||||
mock_database = MagicMock()
|
||||
mock_database.db_engine_spec.engine = "clickhouse"
|
||||
mock_db.session.query.return_value.get.return_value = mock_database
|
||||
mock_db.session.get.return_value = mock_database
|
||||
|
||||
# PostgreSQL syntax
|
||||
original_sql = "SELECT DATE_TRUNC('month', created_at) AS month FROM orders"
|
||||
@@ -201,7 +201,7 @@ def test_transpile_virtual_dataset_sql_postgres_to_mysql(mock_transpile, mock_db
|
||||
"""
|
||||
mock_database = MagicMock()
|
||||
mock_database.db_engine_spec.engine = "mysql"
|
||||
mock_db.session.query.return_value.get.return_value = mock_database
|
||||
mock_db.session.get.return_value = mock_database
|
||||
|
||||
# PostgreSQL syntax with :: casting
|
||||
original_sql = "SELECT created_at::DATE AS date_only FROM orders"
|
||||
@@ -226,7 +226,7 @@ def test_transpile_virtual_dataset_sql_postgres_to_sqlite(mock_transpile, mock_d
|
||||
"""Test transpilation from PostgreSQL to SQLite."""
|
||||
mock_database = MagicMock()
|
||||
mock_database.db_engine_spec.engine = "sqlite"
|
||||
mock_db.session.query.return_value.get.return_value = mock_database
|
||||
mock_db.session.get.return_value = mock_database
|
||||
|
||||
original_sql = "SELECT * FROM orders WHERE created_at > NOW() - INTERVAL '7 days'"
|
||||
transpiled_sql = (
|
||||
|
||||
@@ -63,7 +63,7 @@ def test_validate_raises_when_database_not_found(
|
||||
mock_security_manager: MagicMock,
|
||||
) -> None:
|
||||
"""404 is raised before the access check when the database does not exist."""
|
||||
mock_db.session.query.return_value.get.return_value = None
|
||||
mock_db.session.get.return_value = None
|
||||
|
||||
command = QueryEstimationCommand(_make_params())
|
||||
with pytest.raises(SupersetErrorException) as exc_info:
|
||||
@@ -86,7 +86,7 @@ def test_validate_raises_when_database_access_denied(
|
||||
) -> None:
|
||||
"""SupersetSecurityException propagates when raise_for_access denies access."""
|
||||
mock_database = MagicMock()
|
||||
mock_db.session.query.return_value.get.return_value = mock_database
|
||||
mock_db.session.get.return_value = mock_database
|
||||
mock_security_manager.raise_for_access.side_effect = _security_exception()
|
||||
|
||||
command = QueryEstimationCommand(_make_params())
|
||||
@@ -111,7 +111,7 @@ def test_validate_succeeds_for_authorised_user(
|
||||
) -> None:
|
||||
"""validate() completes without error when access is granted."""
|
||||
mock_database = MagicMock()
|
||||
mock_db.session.query.return_value.get.return_value = mock_database
|
||||
mock_db.session.get.return_value = mock_database
|
||||
mock_security_manager.raise_for_access.return_value = None
|
||||
|
||||
command = QueryEstimationCommand(_make_params())
|
||||
@@ -136,7 +136,7 @@ def test_raise_for_access_called_with_correct_database(
|
||||
"""The database object fetched from the session is passed to raise_for_access."""
|
||||
mock_database = MagicMock()
|
||||
mock_database.id = 42
|
||||
mock_db.session.query.return_value.get.return_value = mock_database
|
||||
mock_db.session.get.return_value = mock_database
|
||||
mock_security_manager.raise_for_access.return_value = None
|
||||
|
||||
command = QueryEstimationCommand(_make_params(database_id=42))
|
||||
@@ -401,7 +401,7 @@ def test_run_wraps_raw_jinja_undefined_error(
|
||||
from jinja2.exceptions import UndefinedError
|
||||
|
||||
mock_database = MagicMock()
|
||||
mock_db.session.query.return_value.get.return_value = mock_database
|
||||
mock_db.session.get.return_value = mock_database
|
||||
mock_security_manager.raise_for_access.return_value = None
|
||||
mock_get_template_processor.return_value.process_template.side_effect = (
|
||||
UndefinedError("'foo' is undefined")
|
||||
|
||||
@@ -1024,7 +1024,7 @@ class TestChartDataCommandValidation:
|
||||
"superset.common.query_context_factory.QueryContextFactory"
|
||||
) as mock_factory,
|
||||
):
|
||||
mock_db.session.query.return_value.get.return_value = mock_dataset
|
||||
mock_db.session.get.return_value = mock_dataset
|
||||
mock_factory.return_value.create.return_value = MagicMock()
|
||||
|
||||
from superset.mcp_service.chart.preview_utils import (
|
||||
@@ -1073,7 +1073,7 @@ class TestChartDataCommandValidation:
|
||||
"superset.common.query_context_factory.QueryContextFactory"
|
||||
) as mock_factory,
|
||||
):
|
||||
mock_db.session.query.return_value.get.return_value = mock_dataset
|
||||
mock_db.session.get.return_value = mock_dataset
|
||||
mock_factory.return_value.create.return_value = MagicMock()
|
||||
|
||||
from superset.mcp_service.chart.preview_utils import (
|
||||
|
||||
Reference in New Issue
Block a user