diff --git a/superset/cli/export_example.py b/superset/cli/export_example.py index 5c14b62dab3..cbf4bf6c589 100644 --- a/superset/cli/export_example.py +++ b/superset/cli/export_example.py @@ -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: diff --git a/superset/commands/dataset/duplicate.py b/superset/commands/dataset/duplicate.py index 85cfe290412..fb1cbb7048e 100644 --- a/superset/commands/dataset/duplicate.py +++ b/superset/commands/dataset/duplicate.py @@ -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( diff --git a/superset/commands/importers/v1/examples.py b/superset/commands/importers/v1/examples.py index 773fc3c1bcf..9cc929f6f14 100644 --- a/superset/commands/importers/v1/examples.py +++ b/superset/commands/importers/v1/examples.py @@ -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 diff --git a/superset/commands/sql_lab/estimate.py b/superset/commands/sql_lab/estimate.py index fd68e201c63..b711efeaaea 100644 --- a/superset/commands/sql_lab/estimate.py +++ b/superset/commands/sql_lab/estimate.py @@ -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( diff --git a/superset/daos/dataset.py b/superset/daos/dataset.py index a1df794984d..7b4b6db2ce8 100644 --- a/superset/daos/dataset.py +++ b/superset/daos/dataset.py @@ -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: diff --git a/superset/mcp_service/chart/preview_utils.py b/superset/mcp_service/chart/preview_utils.py index 1adddad0daf..c074d4357ff 100644 --- a/superset/mcp_service/chart/preview_utils.py +++ b/superset/mcp_service/chart/preview_utils.py @@ -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" diff --git a/superset/migrations/versions/2020-09-24_12-04_3fbbc6e8d654_fix_data_access_permissions_for_virtual_.py b/superset/migrations/versions/2020-09-24_12-04_3fbbc6e8d654_fix_data_access_permissions_for_virtual_.py index 5eabe3f1618..8a72975efc5 100644 --- a/superset/migrations/versions/2020-09-24_12-04_3fbbc6e8d654_fix_data_access_permissions_for_virtual_.py +++ b/superset/migrations/versions/2020-09-24_12-04_3fbbc6e8d654_fix_data_access_permissions_for_virtual_.py @@ -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() diff --git a/superset/migrations/versions/2021-02-04_09-34_070c043f2fdb_add_granularity_to_charts_where_missing.py b/superset/migrations/versions/2021-02-04_09-34_070c043f2fdb_add_granularity_to_charts_where_missing.py index 7e3d7895ef7..6c8fc02e10c 100644 --- a/superset/migrations/versions/2021-02-04_09-34_070c043f2fdb_add_granularity_to_charts_where_missing.py +++ b/superset/migrations/versions/2021-02-04_09-34_070c043f2fdb_add_granularity_to_charts_where_missing.py @@ -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 diff --git a/superset/security/manager.py b/superset/security/manager.py index f1f3f807819..fea0e47961f 100644 --- a/superset/security/manager.py +++ b/superset/security/manager.py @@ -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 ) diff --git a/tests/unit_tests/commands/dataset/test_duplicate.py b/tests/unit_tests/commands/dataset/test_duplicate.py index f03b921c949..aeb6efc176b 100644 --- a/tests/unit_tests/commands/dataset/test_duplicate.py +++ b/tests/unit_tests/commands/dataset/test_duplicate.py @@ -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, diff --git a/tests/unit_tests/commands/importers/v1/examples_test.py b/tests/unit_tests/commands/importers/v1/examples_test.py index be1134a7ad8..71bd1dcac1c 100644 --- a/tests/unit_tests/commands/importers/v1/examples_test.py +++ b/tests/unit_tests/commands/importers/v1/examples_test.py @@ -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 = ( diff --git a/tests/unit_tests/commands/sql_lab/test_estimate.py b/tests/unit_tests/commands/sql_lab/test_estimate.py index c59f2f7bc68..0c77bed89ce 100644 --- a/tests/unit_tests/commands/sql_lab/test_estimate.py +++ b/tests/unit_tests/commands/sql_lab/test_estimate.py @@ -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") diff --git a/tests/unit_tests/mcp_service/chart/tool/test_get_chart_data.py b/tests/unit_tests/mcp_service/chart/tool/test_get_chart_data.py index 8cb3b6ef714..d764bd153c3 100644 --- a/tests/unit_tests/mcp_service/chart/tool/test_get_chart_data.py +++ b/tests/unit_tests/mcp_service/chart/tool/test_get_chart_data.py @@ -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 (