mirror of
https://github.com/apache/superset.git
synced 2026-08-12 11:11:01 +00:00
refactor: Ensure Flask framework leverages the Flask-SQLAlchemy session (Phase II) (#26909)
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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"}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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(
|
||||
{
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
|
||||
@@ -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()
|
||||
)
|
||||
|
||||
@@ -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)]),
|
||||
(
|
||||
|
||||
Reference in New Issue
Block a user