mirror of
https://github.com/galaxyproject/galaxy.git
synced 2026-09-24 16:30:27 +08:00
Move test utils out of model test into common.py
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user