refactor: Ensure Flask framework leverages the Flask-SQLAlchemy session (Phase II) (#26909)

This commit is contained in:
John Bodley
2024-02-14 06:20:15 +13:00
committed by GitHub
parent 827864b939
commit 847ed3f5b0
96 changed files with 656 additions and 730 deletions
+12 -18
View File
@@ -187,24 +187,22 @@ class SupersetTestCase(TestCase):
except ImportError:
return False
def get_or_create(self, cls, criteria, session, **kwargs):
obj = session.query(cls).filter_by(**criteria).first()
def get_or_create(self, cls, criteria, **kwargs):
obj = db.session.query(cls).filter_by(**criteria).first()
if not obj:
obj = cls(**criteria)
obj.__dict__.update(**kwargs)
session.add(obj)
session.commit()
db.session.add(obj)
db.session.commit()
return obj
def login(self, username="admin", password="general"):
return login(self.client, username, password)
def get_slice(
self, slice_name: str, session: Session, expunge_from_session: bool = True
) -> Slice:
slc = session.query(Slice).filter_by(slice_name=slice_name).one()
def get_slice(self, slice_name: str, expunge_from_session: bool = True) -> Slice:
slc = db.session.query(Slice).filter_by(slice_name=slice_name).one()
if expunge_from_session:
session.expunge_all()
db.session.expunge_all()
return slc
@staticmethod
@@ -353,7 +351,6 @@ class SupersetTestCase(TestCase):
return self.get_or_create(
cls=models.Database,
criteria={"database_name": database_name},
session=db.session,
sqlalchemy_uri="sqlite:///:memory:",
id=db_id,
extra=extra,
@@ -375,7 +372,6 @@ class SupersetTestCase(TestCase):
database = self.get_or_create(
cls=models.Database,
criteria={"database_name": database_name},
session=db.session,
sqlalchemy_uri="db_for_macros_testing://user@host:8080/hive",
id=db_id,
)
@@ -398,8 +394,7 @@ class SupersetTestCase(TestCase):
db.session.commit()
def get_dash_by_slug(self, dash_slug):
sesh = db.session()
return sesh.query(Dashboard).filter_by(slug=dash_slug).first()
return db.session.query(Dashboard).filter_by(slug=dash_slug).first()
def get_assert_metric(self, uri: str, func_name: str) -> Response:
"""
@@ -522,11 +517,10 @@ class SupersetTestCase(TestCase):
@contextmanager
def db_insert_temp_object(obj: DeclarativeMeta):
"""Insert a temporary object in database; delete when done."""
session = db.session
try:
session.add(obj)
session.commit()
db.session.add(obj)
db.session.commit()
yield obj
finally:
session.delete(obj)
session.commit()
db.session.delete(obj)
db.session.commit()
+2 -2
View File
@@ -46,7 +46,7 @@ class TestCache(SupersetTestCase):
app.config["DATA_CACHE_CONFIG"] = {"CACHE_TYPE": "NullCache"}
cache_manager.init_app(app)
slc = self.get_slice("Top 10 Girl Name Share", db.session)
slc = self.get_slice("Top 10 Girl Name Share")
json_endpoint = "/superset/explore_json/{}/{}/".format(
slc.datasource_type, slc.datasource_id
)
@@ -73,7 +73,7 @@ class TestCache(SupersetTestCase):
}
cache_manager.init_app(app)
slc = self.get_slice("Top 10 Girl Name Share", db.session)
slc = self.get_slice("Top 10 Girl Name Share")
json_endpoint = "/superset/explore_json/{}/{}/".format(
slc.datasource_type, slc.datasource_id
)
+5 -5
View File
@@ -453,7 +453,7 @@ class TestChartApi(ApiOwnersTestCaseMixin, InsertChartMixin, SupersetTestCase):
"""
Chart API: Test create chart
"""
dashboards_ids = get_dashboards_ids(db, ["world_health", "births"])
dashboards_ids = get_dashboards_ids(["world_health", "births"])
admin_id = self.get_user("admin").id
chart_data = {
"slice_name": "name1",
@@ -1736,7 +1736,7 @@ class TestChartApi(ApiOwnersTestCaseMixin, InsertChartMixin, SupersetTestCase):
@pytest.mark.usefixtures("load_birth_names_dashboard_with_slices")
def test_warm_up_cache(self, slice_name):
self.login()
slc = self.get_slice(slice_name, db.session)
slc = self.get_slice(slice_name)
rv = self.client.put("/api/v1/chart/warm_up_cache", json={"chart_id": slc.id})
self.assertEqual(rv.status_code, 200)
data = json.loads(rv.data.decode("utf-8"))
@@ -1815,7 +1815,7 @@ class TestChartApi(ApiOwnersTestCaseMixin, InsertChartMixin, SupersetTestCase):
@pytest.mark.usefixtures("load_birth_names_dashboard_with_slices")
def test_warm_up_cache_error(self) -> None:
self.login()
slc = self.get_slice("Pivot Table v2", db.session)
slc = self.get_slice("Pivot Table v2")
with mock.patch.object(ChartDataCommand, "run") as mock_run:
mock_run.side_effect = ChartDataQueryFailedError(
@@ -1843,7 +1843,7 @@ class TestChartApi(ApiOwnersTestCaseMixin, InsertChartMixin, SupersetTestCase):
@pytest.mark.usefixtures("load_birth_names_dashboard_with_slices")
def test_warm_up_cache_no_query_context(self) -> None:
self.login()
slc = self.get_slice("Pivot Table v2", db.session)
slc = self.get_slice("Pivot Table v2")
with mock.patch.object(Slice, "get_query_context") as mock_get_query_context:
mock_get_query_context.return_value = None
@@ -1866,7 +1866,7 @@ class TestChartApi(ApiOwnersTestCaseMixin, InsertChartMixin, SupersetTestCase):
@pytest.mark.usefixtures("load_birth_names_dashboard_with_slices")
def test_warm_up_cache_no_datasource(self) -> None:
self.login()
slc = self.get_slice("Top 10 Girl Name Share", db.session)
slc = self.get_slice("Top 10 Girl Name Share")
with mock.patch.object(
Slice,
@@ -413,7 +413,7 @@ class TestChartWarmUpCacheCommand(SupersetTestCase):
@pytest.mark.usefixtures("load_birth_names_dashboard_with_slices")
def test_warm_up_cache(self):
slc = self.get_slice("Top 10 Girl Name Share", db.session)
slc = self.get_slice("Top 10 Girl Name Share")
result = ChartWarmUpCacheCommand(slc.id, None, None).run()
self.assertEqual(
result, {"chart_id": slc.id, "viz_error": None, "viz_status": "success"}
+6 -7
View File
@@ -135,7 +135,7 @@ class TestCore(SupersetTestCase):
@pytest.mark.usefixtures("load_birth_names_dashboard_with_slices")
def test_viz_cache_key(self):
self.login(username="admin")
slc = self.get_slice("Top 10 Girl Name Share", db.session)
slc = self.get_slice("Top 10 Girl Name Share")
viz = slc.viz
qobj = viz.query_obj()
@@ -175,7 +175,7 @@ class TestCore(SupersetTestCase):
def test_save_slice(self):
self.login(username="admin")
slice_name = f"Energy Sankey"
slice_id = self.get_slice(slice_name, db.session).id
slice_id = self.get_slice(slice_name).id
copy_name_prefix = "Test Sankey"
copy_name = f"{copy_name_prefix}[save]{random.random()}"
tbl_id = self.table_ids.get("energy_usage")
@@ -242,7 +242,6 @@ class TestCore(SupersetTestCase):
self.login(username="admin")
slc = self.get_slice(
slice_name="Top 10 Girl Name Share",
session=db.session,
expunge_from_session=False,
)
slc_data_attributes = slc.data.keys()
@@ -356,7 +355,7 @@ class TestCore(SupersetTestCase):
)
def test_warm_up_cache(self):
self.login()
slc = self.get_slice("Top 10 Girl Name Share", db.session)
slc = self.get_slice("Top 10 Girl Name Share")
data = self.get_json_resp(f"/superset/warm_up_cache?slice_id={slc.id}")
self.assertEqual(
data, [{"slice_id": slc.id, "viz_error": None, "viz_status": "success"}]
@@ -381,7 +380,7 @@ class TestCore(SupersetTestCase):
@pytest.mark.usefixtures("load_birth_names_dashboard_with_slices")
def test_warm_up_cache_error(self) -> None:
self.login()
slc = self.get_slice("Pivot Table v2", db.session)
slc = self.get_slice("Pivot Table v2")
with mock.patch.object(
ChartDataCommand,
@@ -406,7 +405,7 @@ class TestCore(SupersetTestCase):
self.login("admin")
store_cache_keys = app.config["STORE_CACHE_KEYS_IN_METADATA_DB"]
app.config["STORE_CACHE_KEYS_IN_METADATA_DB"] = True
slc = self.get_slice("Top 10 Girl Name Share", db.session)
slc = self.get_slice("Top 10 Girl Name Share")
self.get_json_resp(f"/superset/warm_up_cache?slice_id={slc.id}")
ck = db.session.query(CacheKey).order_by(CacheKey.id.desc()).first()
assert ck.datasource_uid == f"{slc.table.id}__table"
@@ -1172,7 +1171,7 @@ class TestCore(SupersetTestCase):
random_key = "random_key"
mock_command.return_value = random_key
slice_name = f"Energy Sankey"
slice_id = self.get_slice(slice_name, db.session).id
slice_id = self.get_slice(slice_name).id
form_data = {"slice_id": slice_id, "viz_type": "line", "datasource": "1__table"}
rv = self.client.get(
f"/superset/explore/?form_data={quote(json.dumps(form_data))}"
@@ -1661,7 +1661,7 @@ class TestDashboardApi(ApiOwnersTestCaseMixin, InsertChartMixin, SupersetTestCas
Dashboard API: Test dashboard export
"""
self.login(username="admin")
dashboards_ids = get_dashboards_ids(db, ["world_health", "births"])
dashboards_ids = get_dashboards_ids(["world_health", "births"])
uri = f"api/v1/dashboard/export/?q={prison.dumps(dashboards_ids)}"
rv = self.get_assert_metric(uri, "export")
@@ -1699,7 +1699,7 @@ class TestDashboardApi(ApiOwnersTestCaseMixin, InsertChartMixin, SupersetTestCas
"""
Dashboard API: Test dashboard export
"""
dashboards_ids = get_dashboards_ids(db, ["world_health", "births"])
dashboards_ids = get_dashboards_ids(["world_health", "births"])
uri = f"api/v1/dashboard/export/?q={prison.dumps(dashboards_ids)}"
self.login(username="admin")
@@ -22,6 +22,7 @@ from flask.ctx import AppContext
from flask_appbuilder.security.sqla.models import User
from sqlalchemy.orm import Session
from superset import db
from superset.commands.dashboard.exceptions import DashboardAccessDeniedError
from superset.commands.temporary_cache.entry import Entry
from superset.extensions import cache_manager
@@ -40,15 +41,13 @@ UPDATED_VALUE = json.dumps({"test": "updated value"})
@pytest.fixture
def dashboard_id(app_context: AppContext, load_world_bank_dashboard_with_slices) -> int:
session: Session = app_context.app.appbuilder.get_session
dashboard = session.query(Dashboard).filter_by(slug="world_health").one()
dashboard = db.session.query(Dashboard).filter_by(slug="world_health").one()
return dashboard.id
@pytest.fixture
def admin_id(app_context: AppContext) -> int:
session: Session = app_context.app.appbuilder.get_session
admin = session.query(User).filter_by(username="admin").one_or_none()
admin = db.session.query(User).filter_by(username="admin").one_or_none()
return admin.id
@@ -42,10 +42,8 @@ STATE = {
@pytest.fixture
def dashboard_id(load_world_bank_dashboard_with_slices) -> int:
with app.app_context() as ctx:
session: Session = ctx.app.appbuilder.get_session
dashboard = session.query(Dashboard).filter_by(slug="world_health").one()
return dashboard.id
dashboard = db.session.query(Dashboard).filter_by(slug="world_health").one()
return dashboard.id
@pytest.fixture
@@ -38,8 +38,6 @@ from tests.integration_tests.dashboards.dashboard_test_utils import (
logger = logging.getLogger(__name__)
session = db.session
inserted_dashboards_ids = []
inserted_databases_ids = []
inserted_sqltables_ids = []
@@ -99,9 +97,9 @@ def create_dashboard(
def insert_model(dashboard: Model) -> None:
session.add(dashboard)
session.commit()
session.refresh(dashboard)
db.session.add(dashboard)
db.session.commit()
db.session.refresh(dashboard)
def create_slice_to_db(
@@ -193,7 +191,7 @@ def delete_all_inserted_objects() -> None:
def delete_all_inserted_dashboards():
try:
dashboards_to_delete: list[Dashboard] = (
session.query(Dashboard)
db.session.query(Dashboard)
.filter(Dashboard.id.in_(inserted_dashboards_ids))
.all()
)
@@ -204,7 +202,7 @@ def delete_all_inserted_dashboards():
logger.error(f"failed to delete {dashboard.id}", exc_info=True)
raise ex
if len(inserted_dashboards_ids) > 0:
session.commit()
db.session.commit()
inserted_dashboards_ids.clear()
except Exception as ex2:
logger.error("delete_all_inserted_dashboards failed", exc_info=True)
@@ -216,25 +214,25 @@ def delete_dashboard(dashboard: Dashboard, do_commit: bool = False) -> None:
delete_dashboard_roles_associations(dashboard)
delete_dashboard_users_associations(dashboard)
delete_dashboard_slices_associations(dashboard)
session.delete(dashboard)
db.session.delete(dashboard)
if do_commit:
session.commit()
db.session.commit()
def delete_dashboard_users_associations(dashboard: Dashboard) -> None:
session.execute(
db.session.execute(
dashboard_user.delete().where(dashboard_user.c.dashboard_id == dashboard.id)
)
def delete_dashboard_roles_associations(dashboard: Dashboard) -> None:
session.execute(
db.session.execute(
DashboardRoles.delete().where(DashboardRoles.c.dashboard_id == dashboard.id)
)
def delete_dashboard_slices_associations(dashboard: Dashboard) -> None:
session.execute(
db.session.execute(
dashboard_slices.delete().where(dashboard_slices.c.dashboard_id == dashboard.id)
)
@@ -242,7 +240,7 @@ def delete_dashboard_slices_associations(dashboard: Dashboard) -> None:
def delete_all_inserted_slices():
try:
slices_to_delete: list[Slice] = (
session.query(Slice).filter(Slice.id.in_(inserted_slices_ids)).all()
db.session.query(Slice).filter(Slice.id.in_(inserted_slices_ids)).all()
)
for slice in slices_to_delete:
try:
@@ -251,7 +249,7 @@ def delete_all_inserted_slices():
logger.error(f"failed to delete {slice.id}", exc_info=True)
raise ex
if len(inserted_slices_ids) > 0:
session.commit()
db.session.commit()
inserted_slices_ids.clear()
except Exception as ex2:
logger.error("delete_all_inserted_slices failed", exc_info=True)
@@ -261,19 +259,19 @@ def delete_all_inserted_slices():
def delete_slice(slice_: Slice, do_commit: bool = False) -> None:
logger.info(f"deleting slice{slice_.id}")
delete_slice_users_associations(slice_)
session.delete(slice_)
db.session.delete(slice_)
if do_commit:
session.commit()
db.session.commit()
def delete_slice_users_associations(slice_: Slice) -> None:
session.execute(slice_user.delete().where(slice_user.c.slice_id == slice_.id))
db.session.execute(slice_user.delete().where(slice_user.c.slice_id == slice_.id))
def delete_all_inserted_tables():
try:
tables_to_delete: list[SqlaTable] = (
session.query(SqlaTable)
db.session.query(SqlaTable)
.filter(SqlaTable.id.in_(inserted_sqltables_ids))
.all()
)
@@ -284,7 +282,7 @@ def delete_all_inserted_tables():
logger.error(f"failed to delete {table.id}", exc_info=True)
raise ex
if len(inserted_sqltables_ids) > 0:
session.commit()
db.session.commit()
inserted_sqltables_ids.clear()
except Exception as ex2:
logger.error("delete_all_inserted_tables failed", exc_info=True)
@@ -294,32 +292,32 @@ def delete_all_inserted_tables():
def delete_sqltable(table: SqlaTable, do_commit: bool = False) -> None:
logger.info(f"deleting table{table.id}")
delete_table_users_associations(table)
session.delete(table)
db.session.delete(table)
if do_commit:
session.commit()
db.session.commit()
def delete_table_users_associations(table: SqlaTable) -> None:
session.execute(
db.session.execute(
sqlatable_user.delete().where(sqlatable_user.c.table_id == table.id)
)
def delete_all_inserted_dbs():
try:
dbs_to_delete: list[Database] = (
session.query(Database)
databases_to_delete: list[Database] = (
db.session.query(Database)
.filter(Database.id.in_(inserted_databases_ids))
.all()
)
for db in dbs_to_delete:
for database in databases_to_delete:
try:
delete_database(db, False)
delete_database(database, False)
except Exception as ex:
logger.error(f"failed to delete {db.id}", exc_info=True)
logger.error(f"failed to delete {database.id}", exc_info=True)
raise ex
if len(inserted_databases_ids) > 0:
session.commit()
db.session.commit()
inserted_databases_ids.clear()
except Exception as ex2:
logger.error("delete_all_inserted_databases failed", exc_info=True)
@@ -328,6 +326,6 @@ def delete_all_inserted_dbs():
def delete_database(database: Database, do_commit: bool = False) -> None:
logger.info(f"deleting database{database.id}")
session.delete(database)
db.session.delete(database)
if do_commit:
session.commit()
db.session.commit()
@@ -1365,12 +1365,11 @@ class TestDatabaseApi(SupersetTestCase):
"""
Database API: Test get select star with datasource access
"""
session = db.session
table = SqlaTable(
schema="main", table_name="ab_permission", database=get_main_database()
)
session.add(table)
session.commit()
db.session.add(table)
db.session.commit()
tmp_table_perm = security_manager.find_permission_view_menu(
"datasource_access", table.get_perm()
@@ -1732,15 +1731,14 @@ class TestDatabaseApi(SupersetTestCase):
with self.create_app().app_context():
main_db = get_main_database()
main_db.allow_file_upload = True
session = db.session
table = SqlaTable(
schema="public",
table_name="ab_permission",
database=get_main_database(),
)
session.add(table)
session.commit()
db.session.add(table)
db.session.commit()
tmp_table_perm = security_manager.find_permission_view_menu(
"datasource_access", table.get_perm()
)
@@ -1748,7 +1748,6 @@ class TestDatasetApi(SupersetTestCase):
assert rv.status_code == 200
cli_export = export_to_dict(
session=db.session,
recursive=True,
back_references=False,
include_defaults=False,
+16 -20
View File
@@ -79,7 +79,6 @@ class TestDatasource(SupersetTestCase):
def test_always_filter_main_dttm(self):
self.login(username="admin")
session = db.session
database = get_example_database()
sql = f"SELECT DATE() as default_dttm, DATE() as additional_dttm, 1 as metric;"
@@ -115,8 +114,8 @@ class TestDatasource(SupersetTestCase):
sql=sql,
)
session.add(table)
session.commit()
db.session.add(table)
db.session.commit()
table.always_filter_main_dttm = False
result = str(table.get_sqla_query(**query_obj).sqla_query.whereclause)
@@ -126,27 +125,26 @@ class TestDatasource(SupersetTestCase):
result = str(table.get_sqla_query(**query_obj).sqla_query.whereclause)
assert "default_dttm" in result and "additional_dttm" in result
session.delete(table)
session.commit()
db.session.delete(table)
db.session.commit()
def test_external_metadata_for_virtual_table(self):
self.login(username="admin")
session = db.session
table = SqlaTable(
table_name="dummy_sql_table",
database=get_example_database(),
schema=get_example_default_schema(),
sql="select 123 as intcol, 'abc' as strcol",
)
session.add(table)
session.commit()
db.session.add(table)
db.session.commit()
table = self.get_table(name="dummy_sql_table")
url = f"/datasource/external_metadata/table/{table.id}/"
resp = self.get_json_resp(url)
assert {o.get("column_name") for o in resp} == {"intcol", "strcol"}
session.delete(table)
session.commit()
db.session.delete(table)
db.session.commit()
@pytest.mark.usefixtures("load_birth_names_dashboard_with_slices")
def test_external_metadata_by_name_for_physical_table(self):
@@ -171,15 +169,14 @@ class TestDatasource(SupersetTestCase):
def test_external_metadata_by_name_for_virtual_table(self):
self.login(username="admin")
session = db.session
table = SqlaTable(
table_name="dummy_sql_table",
database=get_example_database(),
schema=get_example_default_schema(),
sql="select 123 as intcol, 'abc' as strcol",
)
session.add(table)
session.commit()
db.session.add(table)
db.session.commit()
tbl = self.get_table(name="dummy_sql_table")
params = prison.dumps(
@@ -195,8 +192,8 @@ class TestDatasource(SupersetTestCase):
url = f"/datasource/external_metadata_by_name/?q={params}"
resp = self.get_json_resp(url)
assert {o.get("column_name") for o in resp} == {"intcol", "strcol"}
session.delete(tbl)
session.commit()
db.session.delete(tbl)
db.session.commit()
def test_external_metadata_by_name_from_sqla_inspector(self):
self.login(username="admin")
@@ -265,7 +262,6 @@ class TestDatasource(SupersetTestCase):
def test_external_metadata_for_virtual_table_template_params(self):
self.login(username="admin")
session = db.session
table = SqlaTable(
table_name="dummy_sql_table_with_template_params",
database=get_example_database(),
@@ -273,15 +269,15 @@ class TestDatasource(SupersetTestCase):
sql="select {{ foo }} as intcol",
template_params=json.dumps({"foo": "123"}),
)
session.add(table)
session.commit()
db.session.add(table)
db.session.commit()
table = self.get_table(name="dummy_sql_table_with_template_params")
url = f"/datasource/external_metadata/table/{table.id}/"
resp = self.get_json_resp(url)
assert {o.get("column_name") for o in resp} == {"intcol"}
session.delete(table)
session.commit()
db.session.delete(table)
db.session.commit()
def test_external_metadata_for_malicious_virtual_table(self):
self.login(username="admin")
@@ -33,10 +33,10 @@ class TestDatabricksDbEngineSpec(TestDbEngineSpec):
assert get_engine_spec("databricks", "pyhive").engine == "databricks"
def test_extras_without_ssl(self):
db = mock.Mock()
db.extra = default_db_extra
db.server_cert = None
extras = DatabricksNativeEngineSpec.get_extra_params(db)
database = mock.Mock()
database.extra = default_db_extra
database.server_cert = None
extras = DatabricksNativeEngineSpec.get_extra_params(database)
assert extras == {
"engine_params": {
"connect_args": {
@@ -50,12 +50,12 @@ class TestDatabricksDbEngineSpec(TestDbEngineSpec):
}
def test_extras_with_ssl_custom(self):
db = mock.Mock()
db.extra = default_db_extra.replace(
database = mock.Mock()
database.extra = default_db_extra.replace(
'"engine_params": {}',
'"engine_params": {"connect_args": {"ssl": "1"}}',
)
db.server_cert = ssl_certificate
extras = DatabricksNativeEngineSpec.get_extra_params(db)
database.server_cert = ssl_certificate
extras = DatabricksNativeEngineSpec.get_extra_params(database)
connect_args = extras["engine_params"]["connect_args"]
assert connect_args["ssl"] == "1"
@@ -337,14 +337,14 @@ def test_fetch_data_success(fetch_data_mock):
@mock.patch("superset.db_engine_specs.hive.HiveEngineSpec._latest_partition_from_df")
def test_where_latest_partition(mock_method):
mock_method.return_value = ("01-01-19", 1)
db = mock.Mock()
db.get_indexes = mock.Mock(return_value=[{"column_names": ["ds", "hour"]}])
db.get_extra = mock.Mock(return_value={})
db.get_df = mock.Mock()
database = mock.Mock()
database.get_indexes = mock.Mock(return_value=[{"column_names": ["ds", "hour"]}])
database.get_extra = mock.Mock(return_value={})
database.get_df = mock.Mock()
columns = [{"name": "ds"}, {"name": "hour"}]
with app.app_context():
result = HiveEngineSpec.where_latest_partition(
"test_table", "test_schema", db, select(), columns
"test_table", "test_schema", database, select(), columns
)
query_result = str(result.compile(compile_kwargs={"literal_binds": True}))
assert "SELECT \nWHERE ds = '01-01-19' AND hour = 1" == query_result
@@ -353,11 +353,11 @@ def test_where_latest_partition(mock_method):
@mock.patch("superset.db_engine_specs.presto.PrestoEngineSpec.latest_partition")
def test_where_latest_partition_super_method_exception(mock_method):
mock_method.side_effect = Exception()
db = mock.Mock()
database = mock.Mock()
columns = [{"name": "ds"}, {"name": "hour"}]
with app.app_context():
result = HiveEngineSpec.where_latest_partition(
"test_table", "test_schema", db, select(), columns
"test_table", "test_schema", database, select(), columns
)
assert result is None
mock_method.assert_called()
@@ -119,29 +119,29 @@ class TestPostgresDbEngineSpec(TestDbEngineSpec):
assert "postgres" in backends
def test_extras_without_ssl(self):
db = mock.Mock()
db.extra = default_db_extra
db.server_cert = None
extras = PostgresEngineSpec.get_extra_params(db)
database = mock.Mock()
database.extra = default_db_extra
database.server_cert = None
extras = PostgresEngineSpec.get_extra_params(database)
assert "connect_args" not in extras["engine_params"]
def test_extras_with_ssl_default(self):
db = mock.Mock()
db.extra = default_db_extra
db.server_cert = ssl_certificate
extras = PostgresEngineSpec.get_extra_params(db)
database = mock.Mock()
database.extra = default_db_extra
database.server_cert = ssl_certificate
extras = PostgresEngineSpec.get_extra_params(database)
connect_args = extras["engine_params"]["connect_args"]
assert connect_args["sslmode"] == "verify-full"
assert "sslrootcert" in connect_args
def test_extras_with_ssl_custom(self):
db = mock.Mock()
db.extra = default_db_extra.replace(
database = mock.Mock()
database.extra = default_db_extra.replace(
'"engine_params": {}',
'"engine_params": {"connect_args": {"sslmode": "verify-ca"}}',
)
db.server_cert = ssl_certificate
extras = PostgresEngineSpec.get_extra_params(db)
database.server_cert = ssl_certificate
extras = PostgresEngineSpec.get_extra_params(database)
connect_args = extras["engine_params"]["connect_args"]
assert connect_args["sslmode"] == "verify-ca"
assert "sslrootcert" in connect_args
@@ -550,13 +550,17 @@ class TestPrestoDbEngineSpec(TestDbEngineSpec):
self.assertEqual(actual_expanded_cols, expected_expanded_cols)
def test_presto_extra_table_metadata(self):
db = mock.Mock()
db.get_indexes = mock.Mock(return_value=[{"column_names": ["ds", "hour"]}])
db.get_extra = mock.Mock(return_value={})
database = mock.Mock()
database.get_indexes = mock.Mock(
return_value=[{"column_names": ["ds", "hour"]}]
)
database.get_extra = mock.Mock(return_value={})
df = pd.DataFrame({"ds": ["01-01-19"], "hour": [1]})
db.get_df = mock.Mock(return_value=df)
database.get_df = mock.Mock(return_value=df)
PrestoEngineSpec.get_create_view = mock.Mock(return_value=None)
result = PrestoEngineSpec.extra_table_metadata(db, "test_table", "test_schema")
result = PrestoEngineSpec.extra_table_metadata(
database, "test_table", "test_schema"
)
assert result["partitions"]["cols"] == ["ds", "hour"]
assert result["partitions"]["latest"] == {"ds": "01-01-19", "hour": 1}
@@ -43,11 +43,10 @@ class TestDictImportExport(SupersetTestCase):
def delete_imports(cls):
with app.app_context():
# Imported data clean up
session = db.session
for table in session.query(SqlaTable):
for table in db.session.query(SqlaTable):
if DBREF in table.params_dict:
session.delete(table)
session.commit()
db.session.delete(table)
db.session.commit()
@classmethod
def setUpClass(cls):
@@ -124,7 +123,7 @@ class TestDictImportExport(SupersetTestCase):
def test_import_table_no_metadata(self):
table, dict_table = self.create_table("pure_table", id=ID_PREFIX + 1)
new_table = SqlaTable.import_from_dict(db.session, dict_table)
new_table = SqlaTable.import_from_dict(dict_table)
db.session.commit()
imported_id = new_table.id
imported = self.get_table_by_id(imported_id)
@@ -139,7 +138,7 @@ class TestDictImportExport(SupersetTestCase):
cols_uuids=[uuid4()],
metric_names=["metric1"],
)
imported_table = SqlaTable.import_from_dict(db.session, dict_table)
imported_table = SqlaTable.import_from_dict(dict_table)
db.session.commit()
imported = self.get_table_by_id(imported_table.id)
self.assert_table_equals(table, imported)
@@ -156,7 +155,7 @@ class TestDictImportExport(SupersetTestCase):
cols_uuids=[uuid4(), uuid4()],
metric_names=["m1", "m2"],
)
imported_table = SqlaTable.import_from_dict(db.session, dict_table)
imported_table = SqlaTable.import_from_dict(dict_table)
db.session.commit()
imported = self.get_table_by_id(imported_table.id)
self.assert_table_equals(table, imported)
@@ -166,7 +165,7 @@ class TestDictImportExport(SupersetTestCase):
table, dict_table = self.create_table(
"table_override", id=ID_PREFIX + 3, cols_names=["col1"], metric_names=["m1"]
)
imported_table = SqlaTable.import_from_dict(db.session, dict_table)
imported_table = SqlaTable.import_from_dict(dict_table)
db.session.commit()
table_over, dict_table_over = self.create_table(
"table_override",
@@ -174,7 +173,7 @@ class TestDictImportExport(SupersetTestCase):
cols_names=["new_col1", "col2", "col3"],
metric_names=["new_metric1"],
)
imported_over_table = SqlaTable.import_from_dict(db.session, dict_table_over)
imported_over_table = SqlaTable.import_from_dict(dict_table_over)
db.session.commit()
imported_over = self.get_table_by_id(imported_over_table.id)
@@ -195,7 +194,7 @@ class TestDictImportExport(SupersetTestCase):
table, dict_table = self.create_table(
"table_override", id=ID_PREFIX + 3, cols_names=["col1"], metric_names=["m1"]
)
imported_table = SqlaTable.import_from_dict(db.session, dict_table)
imported_table = SqlaTable.import_from_dict(dict_table)
db.session.commit()
table_over, dict_table_over = self.create_table(
"table_override",
@@ -204,7 +203,7 @@ class TestDictImportExport(SupersetTestCase):
metric_names=["new_metric1"],
)
imported_over_table = SqlaTable.import_from_dict(
session=db.session, dict_rep=dict_table_over, sync=["metrics", "columns"]
dict_rep=dict_table_over, sync=["metrics", "columns"]
)
db.session.commit()
@@ -229,7 +228,7 @@ class TestDictImportExport(SupersetTestCase):
cols_names=["new_col1", "col2", "col3"],
metric_names=["new_metric1"],
)
imported_table = SqlaTable.import_from_dict(db.session, dict_table)
imported_table = SqlaTable.import_from_dict(dict_table)
db.session.commit()
copy_table, dict_copy_table = self.create_table(
"copy_cat",
@@ -237,7 +236,7 @@ class TestDictImportExport(SupersetTestCase):
cols_names=["new_col1", "col2", "col3"],
metric_names=["new_metric1"],
)
imported_copy_table = SqlaTable.import_from_dict(db.session, dict_copy_table)
imported_copy_table = SqlaTable.import_from_dict(dict_copy_table)
db.session.commit()
self.assertEqual(imported_table.id, imported_copy_table.id)
self.assert_table_equals(copy_table, self.get_table_by_id(imported_table.id))
@@ -250,7 +249,6 @@ class TestDictImportExport(SupersetTestCase):
self.delete_fake_db()
cli_export = export_to_dict(
session=db.session,
recursive=True,
back_references=False,
include_defaults=False,
+4 -6
View File
@@ -21,6 +21,7 @@ import pytest
from flask_appbuilder.security.sqla.models import User
from sqlalchemy.orm import Session
from superset import db
from superset.commands.explore.form_data.state import TemporaryExploreState
from superset.connectors.sqla.models import SqlaTable
from superset.explore.exceptions import DatasetAccessDeniedError
@@ -39,25 +40,22 @@ FORM_DATA = {"test": "test value"}
@pytest.fixture
def chart_id(load_world_bank_dashboard_with_slices) -> int:
with app.app_context() as ctx:
session: Session = ctx.app.appbuilder.get_session
chart = session.query(Slice).filter_by(slice_name="World's Population").one()
chart = db.session.query(Slice).filter_by(slice_name="World's Population").one()
return chart.id
@pytest.fixture
def admin_id() -> int:
with app.app_context() as ctx:
session: Session = ctx.app.appbuilder.get_session
admin = session.query(User).filter_by(username="admin").one()
admin = db.session.query(User).filter_by(username="admin").one()
return admin.id
@pytest.fixture
def dataset() -> int:
with app.app_context() as ctx:
session: Session = ctx.app.appbuilder.get_session
dataset = (
session.query(SqlaTable)
db.session.query(SqlaTable)
.filter_by(table_name="wb_health_population")
.first()
)
@@ -21,6 +21,7 @@ import pytest
from flask_appbuilder.security.sqla.models import User
from sqlalchemy.orm import Session
from superset import db
from superset.commands.dataset.exceptions import DatasetAccessDeniedError
from superset.commands.explore.form_data.state import TemporaryExploreState
from superset.connectors.sqla.models import SqlaTable
@@ -41,25 +42,22 @@ UPDATED_FORM_DATA = json.dumps({"test": "updated value"})
@pytest.fixture
def chart_id(load_world_bank_dashboard_with_slices) -> int:
with app.app_context() as ctx:
session: Session = ctx.app.appbuilder.get_session
chart = session.query(Slice).filter_by(slice_name="World's Population").one()
chart = db.session.query(Slice).filter_by(slice_name="World's Population").one()
return chart.id
@pytest.fixture
def admin_id() -> int:
with app.app_context() as ctx:
session: Session = ctx.app.appbuilder.get_session
admin = session.query(User).filter_by(username="admin").one()
admin = db.session.query(User).filter_by(username="admin").one()
return admin.id
@pytest.fixture
def datasource() -> int:
with app.app_context() as ctx:
session: Session = ctx.app.appbuilder.get_session
dataset = (
session.query(SqlaTable)
db.session.query(SqlaTable)
.filter_by(table_name="wb_health_population")
.first()
)
@@ -45,22 +45,22 @@ class TestCreateFormDataCommand(SupersetTestCase):
schema=get_example_default_schema(),
sql="select 123 as intcol, 'abc' as strcol",
)
session = db.session
session.add(dataset)
session.commit()
db.session.add(dataset)
db.session.commit()
yield dataset
# rollback
session.delete(dataset)
session.commit()
db.session.delete(dataset)
db.session.commit()
@pytest.fixture()
def create_slice(self):
with self.create_app().app_context():
session = db.session
dataset = (
session.query(SqlaTable).filter_by(table_name="dummy_sql_table").first()
db.session.query(SqlaTable)
.filter_by(table_name="dummy_sql_table")
.first()
)
slice = Slice(
datasource_id=dataset.id,
@@ -69,34 +69,32 @@ class TestCreateFormDataCommand(SupersetTestCase):
slice_name="slice_name",
)
session.add(slice)
session.commit()
db.session.add(slice)
db.session.commit()
yield slice
# rollback
session.delete(slice)
session.commit()
db.session.delete(slice)
db.session.commit()
@pytest.fixture()
def create_query(self):
with self.create_app().app_context():
session = db.session
query = Query(
sql="select 1 as foo;",
client_id="sldkfjlk",
database=get_example_database(),
)
session.add(query)
session.commit()
db.session.add(query)
db.session.commit()
yield query
# rollback
session.delete(query)
session.commit()
db.session.delete(query)
db.session.commit()
@patch("superset.security.manager.g")
@pytest.mark.usefixtures("create_dataset", "create_slice")
@@ -38,8 +38,7 @@ from tests.integration_tests.test_app import app
@pytest.fixture
def chart(app_context, load_world_bank_dashboard_with_slices) -> Slice:
session: Session = app_context.app.appbuilder.get_session
chart = session.query(Slice).filter_by(slice_name="World's Population").one()
chart = db.session.query(Slice).filter_by(slice_name="World's Population").one()
return chart
@@ -43,22 +43,22 @@ class TestCreatePermalinkDataCommand(SupersetTestCase):
schema=get_example_default_schema(),
sql="select 123 as intcol, 'abc' as strcol",
)
session = db.session
session.add(dataset)
session.commit()
db.session.add(dataset)
db.session.commit()
yield dataset
# rollback
session.delete(dataset)
session.commit()
db.session.delete(dataset)
db.session.commit()
@pytest.fixture()
def create_slice(self):
with self.create_app().app_context():
session = db.session
dataset = (
session.query(SqlaTable).filter_by(table_name="dummy_sql_table").first()
db.session.query(SqlaTable)
.filter_by(table_name="dummy_sql_table")
.first()
)
slice = Slice(
datasource_id=dataset.id,
@@ -67,34 +67,32 @@ class TestCreatePermalinkDataCommand(SupersetTestCase):
slice_name="slice_name",
)
session.add(slice)
session.commit()
db.session.add(slice)
db.session.commit()
yield slice
# rollback
session.delete(slice)
session.commit()
db.session.delete(slice)
db.session.commit()
@pytest.fixture()
def create_query(self):
with self.create_app().app_context():
session = db.session
query = Query(
sql="select 1 as foo;",
client_id="sldkfjlk",
database=get_example_database(),
)
session.add(query)
session.commit()
db.session.add(query)
db.session.commit()
yield query
# rollback
session.delete(query)
session.commit()
db.session.delete(query)
db.session.commit()
@patch("superset.security.manager.g")
@pytest.mark.usefixtures("create_dataset", "create_slice")
@@ -177,7 +177,6 @@ def load_dataset_with_columns() -> Generator[SqlaTable, None, None]:
with app.app_context():
engine = create_engine(app.config["SQLALCHEMY_DATABASE_URI"], echo=True)
meta = MetaData()
session = db.session
students = Table(
"students",
@@ -196,8 +195,8 @@ def load_dataset_with_columns() -> Generator[SqlaTable, None, None]:
)
column = TableColumn(table_id=dataset.id, column_name="name")
dataset.columns = [column]
session.add(dataset)
session.commit()
db.session.add(dataset)
db.session.commit()
yield dataset
# cleanup
@@ -205,8 +204,8 @@ def load_dataset_with_columns() -> Generator[SqlaTable, None, None]:
if students_table is not None:
base = declarative_base()
# needed for sqlite
session.commit()
db.session.commit()
base.metadata.drop_all(engine, [students_table], checkfirst=True)
session.delete(dataset)
session.delete(column)
session.commit()
db.session.delete(dataset)
db.session.delete(column)
db.session.commit()
@@ -53,17 +53,16 @@ from .base_tests import SupersetTestCase
def delete_imports():
with app.app_context():
# Imported data clean up
session = db.session
for slc in session.query(Slice):
for slc in db.session.query(Slice):
if "remote_id" in slc.params_dict:
session.delete(slc)
for dash in session.query(Dashboard):
db.session.delete(slc)
for dash in db.session.query(Dashboard):
if "remote_id" in dash.params_dict:
session.delete(dash)
for table in session.query(SqlaTable):
db.session.delete(dash)
for table in db.session.query(SqlaTable):
if "remote_id" in table.params_dict:
session.delete(table)
session.commit()
db.session.delete(table)
db.session.commit()
@pytest.fixture(autouse=True, scope="module")
@@ -66,6 +66,5 @@ def key_value_entry() -> Generator[KeyValueEntry, None, None]:
@pytest.fixture
def admin() -> User:
with app.app_context() as ctx:
session: Session = ctx.app.appbuilder.get_session
admin = session.query(User).filter_by(username="admin").one()
admin = db.session.query(User).filter_by(username="admin").one()
return admin
@@ -230,15 +230,14 @@ class TestGuestUserDatasourceAccess(SupersetTestCase):
schema=get_example_default_schema(),
sql="select 123 as intcol, 'abc' as strcol",
)
session = db.session
session.add(dataset)
session.commit()
db.session.add(dataset)
db.session.commit()
yield dataset
# rollback
session.delete(dataset)
session.commit()
db.session.delete(dataset)
db.session.commit()
def setUp(self) -> None:
self.dash = self.get_dash_by_slug("births")
@@ -258,11 +257,9 @@ class TestGuestUserDatasourceAccess(SupersetTestCase):
],
}
)
self.chart = self.get_slice("Girls", db.session, expunge_from_session=False)
self.chart = self.get_slice("Girls", expunge_from_session=False)
self.datasource = self.chart.datasource
self.other_chart = self.get_slice(
"Treemap", db.session, expunge_from_session=False
)
self.other_chart = self.get_slice("Treemap", expunge_from_session=False)
self.other_datasource = self.other_chart.datasource
self.native_filter_datasource = (
db.session.query(SqlaTable).filter_by(table_name="dummy_sql_table").first()
@@ -245,11 +245,10 @@ def test_migrate_role(
logger.info(description)
with create_old_role(pvm_map, external_pvms) as old_role:
role_name = old_role.name
session = db.session
# Run migrations
add_pvms(session, new_pvms)
migrate_roles(session, pvm_map)
add_pvms(db.session, new_pvms)
migrate_roles(db.session, pvm_map)
role = db.session.query(Role).filter(Role.name == role_name).one_or_none()
for old_pvm, new_pvms in pvm_map.items():
@@ -74,8 +74,6 @@ class TestRowLevelSecurity(SupersetTestCase):
BASE_FILTER_REGEX = re.compile(r"gender = 'boy'")
def setUp(self):
session = db.session
# Create roles
self.role_ab = security_manager.add_role(self.NAME_AB_ROLE)
self.role_q = security_manager.add_role(self.NAME_Q_ROLE)
@@ -83,13 +81,13 @@ class TestRowLevelSecurity(SupersetTestCase):
gamma_user.roles.append(self.role_ab)
gamma_user.roles.append(self.role_q)
self.create_user_with_roles("NoRlsRoleUser", ["Gamma"])
session.commit()
db.session.commit()
# Create regular RowLevelSecurityFilter (energy_usage, unicode_test)
self.rls_entry1 = RowLevelSecurityFilter()
self.rls_entry1.name = "rls_entry1"
self.rls_entry1.tables.extend(
session.query(SqlaTable)
db.session.query(SqlaTable)
.filter(SqlaTable.table_name.in_(["energy_usage", "unicode_test"]))
.all()
)
@@ -104,7 +102,7 @@ class TestRowLevelSecurity(SupersetTestCase):
self.rls_entry2 = RowLevelSecurityFilter()
self.rls_entry2.name = "rls_entry2"
self.rls_entry2.tables.extend(
session.query(SqlaTable)
db.session.query(SqlaTable)
.filter(SqlaTable.table_name.in_(["birth_names"]))
.all()
)
@@ -118,7 +116,7 @@ class TestRowLevelSecurity(SupersetTestCase):
self.rls_entry3 = RowLevelSecurityFilter()
self.rls_entry3.name = "rls_entry3"
self.rls_entry3.tables.extend(
session.query(SqlaTable)
db.session.query(SqlaTable)
.filter(SqlaTable.table_name.in_(["birth_names"]))
.all()
)
@@ -132,7 +130,7 @@ class TestRowLevelSecurity(SupersetTestCase):
self.rls_entry4 = RowLevelSecurityFilter()
self.rls_entry4.name = "rls_entry4"
self.rls_entry4.tables.extend(
session.query(SqlaTable)
db.session.query(SqlaTable)
.filter(SqlaTable.table_name.in_(["birth_names"]))
.all()
)
@@ -145,15 +143,14 @@ class TestRowLevelSecurity(SupersetTestCase):
db.session.commit()
def tearDown(self):
session = db.session
session.delete(self.rls_entry1)
session.delete(self.rls_entry2)
session.delete(self.rls_entry3)
session.delete(self.rls_entry4)
session.delete(security_manager.find_role("NameAB"))
session.delete(security_manager.find_role("NameQ"))
session.delete(self.get_user("NoRlsRoleUser"))
session.commit()
db.session.delete(self.rls_entry1)
db.session.delete(self.rls_entry2)
db.session.delete(self.rls_entry3)
db.session.delete(self.rls_entry4)
db.session.delete(security_manager.find_role("NameAB"))
db.session.delete(security_manager.find_role("NameQ"))
db.session.delete(self.get_user("NoRlsRoleUser"))
db.session.commit()
@pytest.fixture()
def create_dataset(self):
+2 -2
View File
@@ -1704,11 +1704,11 @@ class TestSecurityManager(SupersetTestCase):
mock_is_owner,
):
births = self.get_dash_by_slug("births")
girls = self.get_slice("Girls", db.session, expunge_from_session=False)
girls = self.get_slice("Girls", expunge_from_session=False)
birth_names = girls.datasource
world_health = self.get_dash_by_slug("world_health")
treemap = self.get_slice("Treemap", db.session, expunge_from_session=False)
treemap = self.get_slice("Treemap", expunge_from_session=False)
births.json_metadata = json.dumps(
{
+2 -4
View File
@@ -434,8 +434,6 @@ class TestSqlLab(SupersetTestCase):
Test query api with can_access_all_queries perm added to
gamma and make sure all queries show up.
"""
session = db.session
# Add all_query_access perm to Gamma user
all_queries_view = security_manager.find_permission_view_menu(
"all_query_access", "all_query_access"
@@ -444,7 +442,7 @@ class TestSqlLab(SupersetTestCase):
security_manager.add_permission_role(
security_manager.find_role("gamma_sqllab"), all_queries_view
)
session.commit()
db.session.commit()
# Test search_queries for Admin user
self.run_some_queries()
@@ -461,7 +459,7 @@ class TestSqlLab(SupersetTestCase):
security_manager.find_role("gamma_sqllab"), all_queries_view
)
session.commit()
db.session.commit()
def test_query_admin_can_access_all_queries(self) -> None:
"""
+19 -19
View File
@@ -114,10 +114,10 @@ def test_template_hive(app_context: AppContext, mocker: MockFixture) -> None:
"superset.jinja_context.HiveTemplateProcessor.latest_partition"
)
lp_mock.return_value = "the_latest"
db = mock.Mock()
db.backend = "hive"
database = mock.Mock()
database.backend = "hive"
template = "{{ hive.latest_partition('my_table') }}"
tp = get_template_processor(database=db)
tp = get_template_processor(database=database)
assert tp.process_template(template) == "the_latest"
@@ -126,15 +126,15 @@ def test_template_trino(app_context: AppContext, mocker: MockFixture) -> None:
"superset.jinja_context.TrinoTemplateProcessor.latest_partition"
)
lp_mock.return_value = "the_latest"
db = mock.Mock()
db.backend = "trino"
database = mock.Mock()
database.backend = "trino"
template = "{{ trino.latest_partition('my_table') }}"
tp = get_template_processor(database=db)
tp = get_template_processor(database=database)
assert tp.process_template(template) == "the_latest"
# Backwards compatibility if migrating from Presto.
template = "{{ presto.latest_partition('my_table') }}"
tp = get_template_processor(database=db)
tp = get_template_processor(database=database)
assert tp.process_template(template) == "the_latest"
@@ -154,9 +154,9 @@ def test_custom_process_template(app_context: AppContext, mocker: MockFixture) -
"tests.integration_tests.superset_test_custom_template_processors.datetime"
)
mock_dt.utcnow = mock.Mock(return_value=datetime(1970, 1, 1))
db = mock.Mock()
db.backend = "db_for_macros_testing"
tp = get_template_processor(database=db)
database = mock.Mock()
database.backend = "db_for_macros_testing"
tp = get_template_processor(database=database)
template = "SELECT '$DATE()'"
assert tp.process_template(template) == f"SELECT '1970-01-01'"
@@ -168,28 +168,28 @@ def test_custom_process_template(app_context: AppContext, mocker: MockFixture) -
def test_custom_get_template_kwarg(app_context: AppContext) -> None:
"""Test macro passed as kwargs when getting template processor
works in custom template processor."""
db = mock.Mock()
db.backend = "db_for_macros_testing"
database = mock.Mock()
database.backend = "db_for_macros_testing"
template = "$foo()"
tp = get_template_processor(database=db, foo=lambda: "bar")
tp = get_template_processor(database=database, foo=lambda: "bar")
assert tp.process_template(template) == "bar"
def test_custom_template_kwarg(app_context: AppContext) -> None:
"""Test macro passed as kwargs when processing template
works in custom template processor."""
db = mock.Mock()
db.backend = "db_for_macros_testing"
database = mock.Mock()
database.backend = "db_for_macros_testing"
template = "$foo()"
tp = get_template_processor(database=db)
tp = get_template_processor(database=database)
assert tp.process_template(template, foo=lambda: "bar") == "bar"
def test_custom_template_processors_overwrite(app_context: AppContext) -> None:
"""Test template processor for presto gets overwritten by custom one."""
db = mock.Mock()
db.backend = "db_for_macros_testing"
tp = get_template_processor(database=db)
database = mock.Mock()
database.backend = "db_for_macros_testing"
tp = get_template_processor(database=database)
template = "SELECT '{{ datetime(2017, 1, 1).isoformat() }}'"
assert tp.process_template(template) == template
@@ -15,12 +15,11 @@
# specific language governing permissions and limitations
# under the License.
from flask_appbuilder import SQLA
from superset import db
from superset.models.dashboard import Dashboard
def get_dashboards_ids(db: SQLA, dashboard_slugs: list[str]) -> list[int]:
def get_dashboards_ids(dashboard_slugs: list[str]) -> list[int]:
result = (
db.session.query(Dashboard.id).filter(Dashboard.slug.in_(dashboard_slugs)).all()
)
+2 -2
View File
@@ -898,7 +898,7 @@ class TestUtils(SupersetTestCase):
def test_log_this(self) -> None:
# TODO: Add additional scenarios.
self.login(username="admin")
slc = self.get_slice("Top 10 Girl Name Share", db.session)
slc = self.get_slice("Top 10 Girl Name Share")
dashboard_id = 1
assert slc.viz is not None
@@ -956,7 +956,7 @@ class TestUtils(SupersetTestCase):
@pytest.mark.usefixtures("load_birth_names_dashboard_with_slices")
def test_extract_dataframe_dtypes(self):
slc = self.get_slice("Girls", db.session)
slc = self.get_slice("Girls")
cols: tuple[tuple[str, GenericDataType, list[Any]], ...] = (
("dt", GenericDataType.TEMPORAL, [date(2021, 2, 4), date(2021, 2, 4)]),
(