From afbab2f377d34feda212de755829e85d5169552a Mon Sep 17 00:00:00 2001 From: Sergey Golitsynskiy Date: Thu, 21 Oct 2021 19:38:53 -0400 Subject: [PATCH] Move test utils out of model test into common.py --- test/unit/data/model/mapping/common.py | 103 ++++++++++++++++ .../data/model/mapping/test_model_mapping.py | 115 ++---------------- 2 files changed, 115 insertions(+), 103 deletions(-) diff --git a/test/unit/data/model/mapping/common.py b/test/unit/data/model/mapping/common.py index 3c72213ef0f..8441cfb5274 100644 --- a/test/unit/data/model/mapping/common.py +++ b/test/unit/data/model/mapping/common.py @@ -1,3 +1,96 @@ +from contextlib import contextmanager +from uuid import uuid4 + +from sqlalchemy import ( + delete, + select, + UniqueConstraint, +) + + +def dbcleanup_wrapper(session, obj, where_clause=None): + with dbcleanup(session, obj, where_clause): + yield obj + + +@contextmanager +def dbcleanup(session, obj, where_clause=None): + """ + Use the session to store obj in database; delete from database on exit, bypassing the session. + + If obj does not have an id field, a SQLAlchemy WHERE clause should be provided to construct + a custom select statement. + """ + return_id = where_clause is None + + try: + obj_id = persist(session, obj, return_id) + yield obj_id + finally: + table = obj.__table__ + if where_clause is None: + where_clause = _get_default_where_clause(type(obj), obj_id) + stmt = delete(table).where(where_clause) + session.execute(stmt) + + +def persist(session, obj, return_id=True): + """ + Use the session to store obj in database, then remove obj from session, + so that on a subsequent load from the database we get a clean instance. + """ + session.add(obj) + session.flush() + obj_id = obj.id if return_id else None # save this before obj is expunged + session.expunge(obj) + return obj_id + + +def delete_from_database(session, objects): + """ + Delete each object in objects from database. + May be called at the end of a test if use of a context manager is impractical. + (Assume all objects have the id field as their primary key.) + """ + # Ensure we have a list of objects (check for list explicitly: a model can be iterable) + if not isinstance(objects, list): + objects = [objects] + + for obj in objects: + table = obj.__table__ + stmt = delete(table).where(table.c.id == obj.id) + session.execute(stmt) + + +def get_stored_obj(session, cls, obj_id=None, where_clause=None, unique=False): + # Either obj_id or where_clause must be provided, but not both + assert bool(obj_id) ^ (where_clause is not None) + if where_clause is None: + where_clause = _get_default_where_clause(cls, obj_id) + stmt = select(cls).where(where_clause) + result = session.execute(stmt) + # unique() is required if result contains joint eager loads against collections + # https://gerrit.sqlalchemy.org/c/sqlalchemy/sqlalchemy/+/2253 + if unique: + result = result.unique() + return result.scalar_one() + + +def has_unique_constraint(table, fields): + for constraint in table.constraints: + if isinstance(constraint, UniqueConstraint): + col_names = {c.name for c in constraint.columns} + if set(fields) == col_names: + return True + + +def has_index(table, fields): + for index in table.indexes: + col_names = {c.name for c in index.columns} + if set(fields) == col_names: + return True + + def collection_consists_of_objects(collection, *objects): """ Returns True iff list(collection) == list(objects), where object equality is determined @@ -17,3 +110,13 @@ def collection_consists_of_objects(collection, *objects): if item1.id is None or item2.id is None or item1.id != item2.id: return False return True + + +def get_unique_value(): + """Generate unique values to accommodate unique constraints.""" + return uuid4().hex + + +def _get_default_where_clause(cls, obj_id): + where_clause = cls.__table__.c.id == obj_id + return where_clause diff --git a/test/unit/data/model/mapping/test_model_mapping.py b/test/unit/data/model/mapping/test_model_mapping.py index ca2cd079071..b720a1ba536 100644 --- a/test/unit/data/model/mapping/test_model_mapping.py +++ b/test/unit/data/model/mapping/test_model_mapping.py @@ -60,19 +60,23 @@ class TestPlanet(BaseTest): # BaseTest is a base class; we need it to get the t See other model tests in this module for examples of more complex setups. """ -from contextlib import contextmanager from datetime import datetime, timedelta from uuid import UUID, uuid4 import pytest -from sqlalchemy import ( - delete, - select, - UniqueConstraint, -) import galaxy.model.mapping as mapping -from .common import collection_consists_of_objects +from .common import ( + collection_consists_of_objects, + dbcleanup, + dbcleanup_wrapper, + delete_from_database, + get_stored_obj, + get_unique_value, + has_index, + has_unique_constraint, + persist, +) class BaseTest: @@ -7866,102 +7870,7 @@ def workflow_step_factory(model, workflow): return make_instance -# Test utilities - - -def dbcleanup_wrapper(session, obj, where_clause=None): - with dbcleanup(session, obj, where_clause): - yield obj - - -@contextmanager -def dbcleanup(session, obj, where_clause=None): - """ - Use the session to store obj in database; delete from database on exit, bypassing the session. - - If obj does not have an id field, a SQLAlchemy WHERE clause should be provided to construct - a custom select statement. - """ - return_id = where_clause is None - - try: - obj_id = persist(session, obj, return_id) - yield obj_id - finally: - table = obj.__table__ - if where_clause is None: - where_clause = _get_default_where_clause(type(obj), obj_id) - stmt = delete(table).where(where_clause) - session.execute(stmt) - - -def persist(session, obj, return_id=True): - """ - Use the session to store obj in database, then remove obj from session, - so that on a subsequent load from the database we get a clean instance. - """ - session.add(obj) - session.flush() - obj_id = obj.id if return_id else None # save this before obj is expunged - session.expunge(obj) - return obj_id - - -def delete_from_database(session, objects): - """ - Delete each object in objects from database. - May be called at the end of a test if use of a context manager is impractical. - (Assume all objects have the id field as their primary key.) - """ - # Ensure we have a list of objects (check for list explicitly: a model can be iterable) - if not isinstance(objects, list): - objects = [objects] - - for obj in objects: - table = obj.__table__ - stmt = delete(table).where(table.c.id == obj.id) - session.execute(stmt) - - -def get_stored_obj(session, cls, obj_id=None, where_clause=None, unique=False): - # Either obj_id or where_clause must be provided, but not both - assert bool(obj_id) ^ (where_clause is not None) - if where_clause is None: - where_clause = _get_default_where_clause(cls, obj_id) - stmt = select(cls).where(where_clause) - result = session.execute(stmt) - # unique() is required if result contains joint eager loads against collections - # https://gerrit.sqlalchemy.org/c/sqlalchemy/sqlalchemy/+/2253 - if unique: - result = result.unique() - return result.scalar_one() - - -def _get_default_where_clause(cls, obj_id): - where_clause = cls.__table__.c.id == obj_id - return where_clause - - -def has_unique_constraint(table, fields): - for constraint in table.constraints: - if isinstance(constraint, UniqueConstraint): - col_names = {c.name for c in constraint.columns} - if set(fields) == col_names: - return True - - -def has_index(table, fields): - for index in table.indexes: - col_names = {c.name for c in index.columns} - if set(fields) == col_names: - return True - - -def get_unique_value(): - """Generate unique values to accommodate unique constraints.""" - return uuid4().hex - - +# Test helpers def _run_average_rating_test(session, obj, user, obj_rating_association_factory): # obj has been expunged; to access its deferred properties, # it needs to be added back to the session.