Merge pull request #12064 from ic4f/dev_declarative3

Declarative mappings + tests for all models
This commit is contained in:
Marius van den Beek
2021-09-08 12:11:03 +02:00
committed by GitHub
18 changed files with 10903 additions and 3212 deletions
+1 -1
View File
@@ -49,7 +49,7 @@ class RatableManagerMixin:
# TODO?: update and create to RatingsManager (if not overkill)
rating = self.rating(item, user, as_int=False)
if not rating:
rating = self.rating_assoc(user=user)
rating = self.rating_assoc(user, item)
self.associate(rating, item)
rating.rating = value
+1 -1
View File
@@ -237,7 +237,7 @@ class UserManager(base.ModelManager, deletable.PurgableManagerMixin):
return schema.BootstrapAdminUser()
sa_session = sa_session or self.app.model.session
try:
provided_key = sa_session.query(self.app.model.APIKeys).filter(self.app.model.APIKeys.table.c.key == api_key).one()
provided_key = sa_session.query(self.app.model.APIKeys).filter(self.app.model.APIKeys.key == api_key).one()
except NoResultFound:
raise exceptions.AuthenticationFailed('Provided API key is not valid.')
if provided_key.user.deleted:
+2933 -278
View File
File diff suppressed because it is too large Load Diff
+1 -4
View File
@@ -43,10 +43,7 @@ class UsesItemRatings:
if not item_rating:
# User has not yet rated item; create rating.
item_rating_assoc_class = self._get_item_rating_assoc_class(item, webapp_model=webapp_model)
item_rating = item_rating_assoc_class()
item_rating.user = user
item_rating.set_item(item)
item_rating.rating = rating
item_rating = item_rating_assoc_class(user, item, rating)
db_session.add(item_rating)
db_session.flush()
elif item_rating.rating != rating:
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -18,6 +18,7 @@ from sqlalchemy.orm import object_session
from sqlalchemy.orm.attributes import flag_modified
import galaxy.model
from galaxy.security.object_wrapper import sanitize_lists_to_string
from galaxy.util import (
form_builder,
listify,
@@ -26,7 +27,6 @@ from galaxy.util import (
unicodify,
)
from galaxy.util.json import safe_dumps
from galaxy.util.object_wrapper import sanitize_lists_to_string
log = logging.getLogger(__name__)
+7 -7
View File
@@ -358,28 +358,28 @@ class GalaxyTagHandler(TagHandler):
TagHandler.__init__(self, sa_session)
self.item_tag_assoc_info["History"] = ItemTagAssocInfo(model.History,
model.HistoryTagAssociation,
model.HistoryTagAssociation.table.c.history_id)
model.HistoryTagAssociation.history_id)
self.item_tag_assoc_info["HistoryDatasetAssociation"] = \
ItemTagAssocInfo(model.HistoryDatasetAssociation,
model.HistoryDatasetAssociationTagAssociation,
model.HistoryDatasetAssociationTagAssociation.table.c.history_dataset_association_id)
model.HistoryDatasetAssociationTagAssociation.history_dataset_association_id)
self.item_tag_assoc_info["HistoryDatasetCollectionAssociation"] = \
ItemTagAssocInfo(model.HistoryDatasetCollectionAssociation,
model.HistoryDatasetCollectionTagAssociation,
model.HistoryDatasetCollectionTagAssociation.table.c.history_dataset_collection_id)
model.HistoryDatasetCollectionTagAssociation.history_dataset_collection_id)
self.item_tag_assoc_info["LibraryDatasetDatasetAssociation"] = \
ItemTagAssocInfo(model.LibraryDatasetDatasetAssociation,
model.LibraryDatasetDatasetAssociationTagAssociation,
model.LibraryDatasetDatasetAssociationTagAssociation.table.c.library_dataset_dataset_association_id)
model.LibraryDatasetDatasetAssociationTagAssociation.library_dataset_dataset_association_id)
self.item_tag_assoc_info["Page"] = ItemTagAssocInfo(model.Page,
model.PageTagAssociation,
model.PageTagAssociation.table.c.page_id)
model.PageTagAssociation.page_id)
self.item_tag_assoc_info["StoredWorkflow"] = ItemTagAssocInfo(model.StoredWorkflow,
model.StoredWorkflowTagAssociation,
model.StoredWorkflowTagAssociation.table.c.stored_workflow_id)
model.StoredWorkflowTagAssociation.stored_workflow_id)
self.item_tag_assoc_info["Visualization"] = ItemTagAssocInfo(model.Visualization,
model.VisualizationTagAssociation,
model.VisualizationTagAssociation.table.c.visualization_id)
model.VisualizationTagAssociation.visualization_id)
class GalaxyTagHandlerSession(GalaxyTagHandler):
+2 -2
View File
@@ -144,7 +144,7 @@ class DatabaseQuotaAgent(QuotaAgent):
return self._default_quota(self.model.DefaultQuotaAssociation.types.REGISTERED)
def _default_quota(self, default_type):
dqa = self.sa_session.query(self.model.DefaultQuotaAssociation).filter(self.model.DefaultQuotaAssociation.table.c.type == default_type).first()
dqa = self.sa_session.query(self.model.DefaultQuotaAssociation).filter(self.model.DefaultQuotaAssociation.type == default_type).first()
if not dqa:
return None
if dqa.quota.bytes < 0:
@@ -161,7 +161,7 @@ class DatabaseQuotaAgent(QuotaAgent):
for gqa in quota.groups:
self.sa_session.delete(gqa)
# Find the old default, assign the new quota if it exists
dqa = self.sa_session.query(self.model.DefaultQuotaAssociation).filter(self.model.DefaultQuotaAssociation.table.c.type == default_type).first()
dqa = self.sa_session.query(self.model.DefaultQuotaAssociation).filter(self.model.DefaultQuotaAssociation.type == default_type).first()
if dqa:
dqa.quota = quota
# Or create if necessary
@@ -23,6 +23,8 @@ from types import (
TracebackType,
)
import sqlalchemy
NoneType = type(None)
NotImplementedType = type(NotImplemented)
EllipsisType = type(Ellipsis)
@@ -261,6 +263,14 @@ class SafeStringWrapper:
return self.__safe_string_wrapper_function__(getattr(self.unsanitized, name))
def __setattr__(self, name, value):
# A class mapped declaratively is a subclass of DeclarativeMeta. It will check at creation time
# if self has _sa_instance_state set, and if not, it'll try to set it. This happens BEFORE self.__init__
# has been called, so self.unsanitized does not exists, which raises an AttributeError.
# To avoid this, as well as to avoid SQLAlchemy state to be set on SafeStringWrapper,
# we simply ignore this call.
if isinstance(value, sqlalchemy.orm.state.InstanceState):
return
if name in SafeStringWrapper.__NO_WRAP_NAMES__:
return object.__setattr__(self, name, value)
return setattr(self.unsanitized, name, value)
+1 -2
View File
@@ -4,11 +4,11 @@ import os
import shlex
import tempfile
from galaxy import model
from galaxy.files import ProvidesUserFileSourcesUserContext
from galaxy.job_execution.setup import ensure_configs_directory
from galaxy.model.none_like import NoneDataset
from galaxy.security.object_wrapper import wrap_with_safe_string
from galaxy.tools import global_tool_errors
from galaxy.tools.parameters import (
visit_input_values,
@@ -42,7 +42,6 @@ from galaxy.util import (
unicodify,
)
from galaxy.util.bunch import Bunch
from galaxy.util.object_wrapper import wrap_with_safe_string
from galaxy.util.template import fill_template
from galaxy.work.context import WorkRequestContext
+1 -1
View File
@@ -6,12 +6,12 @@ from functools import total_ordering
from galaxy import exceptions
from galaxy.model.none_like import NoneDataset
from galaxy.security.object_wrapper import wrap_with_safe_string
from galaxy.tools.parameters.wrapped_json import (
data_collection_input_to_staging_path_and_source_path,
data_input_to_staging_path_and_source_path,
)
from galaxy.util import filesystem_safe_string
from galaxy.util.object_wrapper import wrap_with_safe_string
log = logging.getLogger(__name__)
+1 -1
View File
@@ -8,9 +8,9 @@ from json import loads
from galaxy.util import (
galaxy_directory,
sanitize_lists_to_string,
unicodify,
)
from galaxy.util.object_wrapper import sanitize_lists_to_string
log = logging.getLogger(__name__)
@@ -5,7 +5,7 @@ from html.parser import HTMLParser
from http.client import HTTPConnection
from markupsafe import escape
from sqlalchemy import and_
from sqlalchemy import and_, desc
from sqlalchemy.orm import eagerload, joinedload, lazyload, undefer
from sqlalchemy.sql import expression
@@ -20,7 +20,6 @@ from galaxy.managers.workflows import (
WorkflowUpdateOptions,
)
from galaxy.model.item_attrs import UsesItemRatings
from galaxy.model.mapping import desc
from galaxy.security.validate_user_input import validate_publicname
from galaxy.tools.parameters.basic import workflow_building_modes
from galaxy.util import (
+1
View File
@@ -3,6 +3,7 @@ boltons
docutils
markupsafe
packaging
pycryptodome
pyyaml
requests
routes
+16 -18
View File
@@ -130,52 +130,50 @@ class MappingTests(BaseModelTestCase):
def test_ratings(self):
model = self.model
u = model.User(email="rater@example.com", password="password")
user_email = "rater@example.com"
u = model.User(email=user_email, password="password")
self.persist(u)
def persist_and_check_rating(rating_class, **kwds):
rating_association = rating_class()
rating_association.rating = 5
rating_association.user = u
for key, value in kwds.items():
setattr(rating_association, key, value)
def persist_and_check_rating(rating_class, item):
rating = 5
rating_association = rating_class(u, item, rating)
self.persist(rating_association)
self.expunge()
stored_annotation = self.query(rating_class).all()[0]
assert stored_annotation.rating == 5
assert stored_annotation.user.email == "rater@example.com"
stored_rating = self.query(rating_class).all()[0]
assert stored_rating.rating == rating
assert stored_rating.user.email == user_email
sw = model.StoredWorkflow()
sw.user = u
self.persist(sw)
persist_and_check_rating(model.StoredWorkflowRatingAssociation, stored_workflow=sw)
persist_and_check_rating(model.StoredWorkflowRatingAssociation, sw)
h = model.History(name="History for Rating", user=u)
self.persist(h)
persist_and_check_rating(model.HistoryRatingAssociation, history=h)
persist_and_check_rating(model.HistoryRatingAssociation, h)
d1 = model.HistoryDatasetAssociation(extension="txt", history=h, create_dataset=True, sa_session=model.session)
self.persist(d1)
persist_and_check_rating(model.HistoryDatasetAssociationRatingAssociation, hda=d1)
persist_and_check_rating(model.HistoryDatasetAssociationRatingAssociation, d1)
page = model.Page()
page.user = u
self.persist(page)
persist_and_check_rating(model.PageRatingAssociation, page=page)
persist_and_check_rating(model.PageRatingAssociation, page)
visualization = model.Visualization()
visualization.user = u
self.persist(visualization)
persist_and_check_rating(model.VisualizationRatingAssociation, visualization=visualization)
persist_and_check_rating(model.VisualizationRatingAssociation, visualization)
dataset_collection = model.DatasetCollection(collection_type="paired")
history_dataset_collection = model.HistoryDatasetCollectionAssociation(collection=dataset_collection)
self.persist(history_dataset_collection)
persist_and_check_rating(model.HistoryDatasetCollectionRatingAssociation, history_dataset_collection=history_dataset_collection)
persist_and_check_rating(model.HistoryDatasetCollectionRatingAssociation, history_dataset_collection)
library_dataset_collection = model.LibraryDatasetCollectionAssociation(collection=dataset_collection)
self.persist(library_dataset_collection)
persist_and_check_rating(model.LibraryDatasetCollectionRatingAssociation, library_dataset_collection=library_dataset_collection)
persist_and_check_rating(model.LibraryDatasetCollectionRatingAssociation, library_dataset_collection)
def test_display_name(self):
@@ -255,7 +253,7 @@ class MappingTests(BaseModelTestCase):
sw = model.StoredWorkflow()
sw.user = u
tag_and_test(sw, model.StoredWorkflowTagAssociation, "tagged_workflows")
tag_and_test(sw, model.StoredWorkflowTagAssociation, "tagged_stored_workflows")
h = model.History(name="History for Tagging", user=u)
tag_and_test(h, model.HistoryTagAssociation, "tagged_histories")
File diff suppressed because it is too large Load Diff
-220
View File
@@ -1,220 +0,0 @@
import pytest
from sqlalchemy import (
delete,
select,
UniqueConstraint,
)
import galaxy.model.mapping as mapping
@pytest.fixture(scope='module')
def model():
db_uri = 'sqlite:///:memory:'
return mapping.init('/tmp', db_uri, create_tables=True)
@pytest.fixture
def session(model):
Session = model.session
yield Session()
Session.remove() # Ensures we get a new session for each test
def test_Group_table(model):
tbl = model.Group.__table__
assert tbl.name == 'galaxy_group'
def test_Group(model, session):
cls = model.Group
name = 'a'
obj = cls(name)
persist(session, obj)
stmt = select(cls)
stored_obj = session.execute(stmt).scalar_one()
assert stored_obj.id
assert stored_obj.create_time
assert stored_obj.update_time
assert stored_obj.name == name
assert stored_obj.deleted is False
cleanup(session, cls)
def test_Quota_table(model):
tbl = model.Quota.__table__
assert tbl.name == 'quota'
def test_Quota(model, session):
cls = model.Quota
name, description = 'a', 'b'
obj = cls(name, description)
persist(session, obj)
stmt = select(cls)
stored_obj = session.execute(stmt).scalar_one()
assert stored_obj.id
assert stored_obj.create_time
assert stored_obj.update_time
assert stored_obj.name == name
assert stored_obj.description == description
assert stored_obj.bytes == 0
assert stored_obj.operation == '='
assert stored_obj.deleted is False
cleanup(session, cls)
def test_Role_table(model):
tbl = model.Role.__table__
assert tbl.name == 'role'
def test_Role(model, session):
cls = model.Role
name, description = 'a', 'b'
obj = cls(name, description)
persist(session, obj)
stmt = select(cls)
stored_obj = session.execute(stmt).scalar_one()
assert stored_obj.id
assert stored_obj.create_time
assert stored_obj.update_time
assert stored_obj.name == name
assert stored_obj.description == description
assert stored_obj.type == model.Role.types.SYSTEM
assert stored_obj.deleted is False
cleanup(session, cls)
def test_WorkerProcess_table(model):
tbl = model.WorkerProcess.__table__
assert tbl.name == 'worker_process'
assert has_unique_constraint(tbl, ('server_name', 'hostname'))
def test_WorkerProcess(model, session):
cls = model.WorkerProcess
server_name, hostname = 'a', 'b'
obj = cls(server_name, hostname)
persist(session, obj)
stmt = select(cls)
stored_obj = session.execute(stmt).scalar_one()
assert stored_obj.id
assert stored_obj.server_name == server_name
assert stored_obj.hostname == hostname
assert stored_obj.pid is None
assert stored_obj.update_time
cleanup(session, cls)
def test_PSAAssociation_table(model):
tbl = model.PSAAssociation.__table__
assert tbl.name == 'psa_association'
def test_PSAAssociation(model, session):
cls = model.PSAAssociation
server_url, handle, secret, issued, lifetime, assoc_type = 'a', 'b', 'c', 1, 2, 'd'
obj = cls(server_url, handle, secret, issued, lifetime, assoc_type)
persist(session, obj)
stmt = select(cls)
stored_obj = session.execute(stmt).scalar_one()
assert stored_obj.id
assert stored_obj.server_url == server_url
assert stored_obj.handle == handle
assert stored_obj.secret == secret
assert stored_obj.issued == issued
assert stored_obj.lifetime == lifetime
assert stored_obj.assoc_type == assoc_type
cleanup(session, cls)
def test_PSACode_table(model):
tbl = model.PSACode.__table__
assert tbl.name == 'psa_code'
assert has_unique_constraint(tbl, ('code', 'email'))
def test_PSACode(model, session):
cls = model.PSACode
email, code = 'a', 'b'
obj = cls(email, code)
persist(session, obj)
stmt = select(cls)
stored_obj = session.execute(stmt).scalar_one()
assert stored_obj.id
assert stored_obj.email == email
assert stored_obj.code == code
cleanup(session, cls)
def test_PSANonce_table(model):
tbl = model.PSANonce.__table__
assert tbl.name == 'psa_nonce'
def test_PSANonce(model, session):
cls = model.PSANonce
server_url, timestamp, salt = 'a', 1, 'b'
obj = cls(server_url, timestamp, salt)
persist(session, obj)
stmt = select(cls)
stored_obj = session.execute(stmt).scalar_one()
assert stored_obj.id
assert stored_obj.server_url
assert stored_obj.timestamp == timestamp
assert stored_obj.salt == salt
cleanup(session, cls)
def test_PSAPartial_table(model):
tbl = model.PSAPartial.__table__
assert tbl.name == 'psa_partial'
def test_PSAPartial(model, session):
cls = model.PSAPartial
token, data, next_step, backend = 'a', 'b', 1, 'c'
obj = cls(token, data, next_step, backend)
persist(session, obj)
stmt = select(cls)
stored_obj = session.execute(stmt).scalar_one()
assert stored_obj.id
assert stored_obj.token == token
assert stored_obj.data == data
assert stored_obj.next_step == next_step
assert stored_obj.backend == backend
cleanup(session, cls)
def persist(session, obj):
session.add(obj)
session.flush()
def cleanup(session, cls):
session.execute(delete(cls))
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
+230
View File
@@ -0,0 +1,230 @@
"""
This module contains tests for the utility functions in the test_mapping module.
"""
from contextlib import contextmanager
import pytest
from sqlalchemy import (
Column,
create_engine,
Index,
Integer,
select,
UniqueConstraint
)
from sqlalchemy.exc import NoResultFound
from sqlalchemy.orm import registry, Session
from . test_mapping import (
are_same_entity_collections,
dbcleanup,
dbcleanup_wrapper,
delete_from_database,
get_stored_obj,
has_index,
has_unique_constraint,
persist,
)
def test_persist(session):
"""
Verify item is stored in database.
We call persist() with default arg value: `return_id=True`. Passing `False` results
in the same functionality except the return value of the funtion is `None`.
"""
instance = Foo()
instance_id = persist(session, instance) # store instance in database
assert instance_id == 1 # id was assigned
assert instance not in session # instance was expunged from session
stored_instance = _get_stored_instance_by_id(session, Foo, instance_id)
assert stored_instance.id == instance.id # instance can be retrieved by id
assert stored_instance is not instance # retrieved instance is not the same object
def test_get_stored_obj_must_have_obj_id_xor_where_clause():
"""
Verify that function must be called with either obj_id or where_clause, but not both.
"""
with pytest.raises(AssertionError):
get_stored_obj(None, None, obj_id=1, where_clause='a')
with pytest.raises(AssertionError):
get_stored_obj(None, None, obj_id=None, where_clause=None)
def test_get_stored_obj_by_id(session):
"""Verify item is retrieved from database by id."""
instance = Foo()
id = persist(session, instance)
stored_instance = get_stored_obj(session, Foo, id)
assert stored_instance.id == instance.id
def test_get_stored_obj_by_where_clause(session):
"""Verify item is retrieved from database by a custom WHERE clause."""
instance1 = Foo()
instance2 = Foo()
id1 = persist(session, instance1)
id2 = persist(session, instance2)
where_clause = Foo.__table__.c.id == id1
stored_instance = get_stored_obj(session, Foo, where_clause=where_clause)
assert stored_instance.id == instance1.id
where_clause = Foo.__table__.c.id == id2
stored_instance = get_stored_obj(session, Foo, where_clause=where_clause)
assert stored_instance.id == instance2.id
def test_delete_from_database_one_item(session):
"""Verify item is deleted from database."""
instance = Foo()
id = persist(session, instance) # store instance in database
stored_instance = _get_stored_instance_by_id(session, Foo, id)
assert stored_instance is not None # instance is present in database
delete_from_database(session, instance) # delete instance from database
with pytest.raises(NoResultFound): # instance no longer present in database
stored_instance = _get_stored_instance_by_id(session, Foo, id)
def test_delete_from_database_multiple_items(session):
"""Verify multiple items are deleted from database."""
instance1 = Foo()
instance2 = Foo()
id1 = persist(session, instance1)
id2 = persist(session, instance2)
delete_from_database(session, [instance1, instance2])
with pytest.raises(NoResultFound):
_get_stored_instance_by_id(session, Foo, id1)
with pytest.raises(NoResultFound):
_get_stored_instance_by_id(session, Foo, id2)
def test_dbcleanup_by_id(session):
instance = Foo()
with dbcleanup(session, instance) as instance_id:
stored_instance = _get_stored_instance_by_id(session, Foo, instance_id)
assert stored_instance # has been stored in the database
with pytest.raises(NoResultFound): # has been deleted from the database
_get_stored_instance_by_id(session, Foo, instance_id)
def test_dbcleanup_by_where_clause(session):
instance = Foo()
where_clause = Foo.__table__.c.id == 1
with dbcleanup(session, instance, where_clause):
stored_instance = _get_stored_instance_by_id(session, Foo, where_clause)
assert stored_instance # has been stored in the database
with pytest.raises(NoResultFound): # has been deleted from the database
_get_stored_instance_by_id(session, Foo, stored_instance.id)
def test_dbcleanup_wrapper(session):
"""
Verify dbcleanup_wrapper has same effect as dbcleanup context manager,
and yields the object instance.
"""
@contextmanager # we need a context manager to similate scope
def managed(a, b):
yield from dbcleanup_wrapper(a, b)
instance = Foo()
# this will call dbcleanup_wrapper that stores instance in the database and returns a reference to it
with managed(session, instance) as instance2:
assert instance2 is instance # instance2 should be a reference to instance
stored_instance = _get_stored_instance_by_id(session, Foo, instance.id)
assert stored_instance # has been stored in the database
# on exit we expect the entity to be deleted from the database
with pytest.raises(NoResultFound): # has been deleted from the database
_get_stored_instance_by_id(session, Foo, instance.id)
def test_has_index(session):
assert has_index(Bar.__table__, ('field1',))
assert not has_index(Foo.__table__, ('field1',))
def test_has_unique_constraint(session):
assert has_unique_constraint(Bar.__table__, ('field2',))
assert not has_unique_constraint(Foo.__table__, ('field1',))
def test_are_same_entity_collections(session):
foo1 = Foo()
foo2 = Foo()
foo3 = Foo()
persist(session, foo1)
persist(session, foo2)
persist(session, foo3)
stored_foo1 = _get_stored_instance_by_id(session, Foo, foo1.id)
stored_foo2 = _get_stored_instance_by_id(session, Foo, foo2.id)
stored_foo3 = _get_stored_instance_by_id(session, Foo, foo3.id)
expected = [foo1, foo2]
assert are_same_entity_collections([stored_foo1, stored_foo2], expected)
assert are_same_entity_collections([stored_foo2, stored_foo1], expected)
assert not are_same_entity_collections([stored_foo1, stored_foo3], expected)
assert not are_same_entity_collections([stored_foo1, stored_foo1, stored_foo2], expected)
# Test utilities
mapper_registry = registry()
@mapper_registry.mapped
class Foo:
__tablename__ = 'foo'
id = Column(Integer, primary_key=True)
field1 = Column(Integer)
@mapper_registry.mapped
class Bar:
__tablename__ = 'bar'
id = Column(Integer, primary_key=True)
field1 = Column(Integer)
field2 = Column(Integer)
__table_args__ = (
Index('ix', 'field1'),
UniqueConstraint('field2'),
)
@pytest.fixture(scope='module')
def engine():
"""Create connection engine. (Fixture is module-scoped)."""
db_uri = 'sqlite:///:memory:' # We only need sqlite for these tests
return create_engine(db_uri)
@pytest.fixture(scope='module')
def init(engine):
"""Create database objects. (Fixture is module-scoped)."""
mapper_registry.metadata.create_all(engine)
@pytest.fixture
def session(init, engine):
"""Provides a basic session object, closed on exit."""
with Session(engine) as s:
yield s
def _get_stored_instance_by_id(session, cls_, id):
statement = select(Foo).where(cls_.__table__.c.id == id)
return session.execute(statement).scalar_one()