mirror of
https://github.com/galaxyproject/galaxy.git
synced 2026-09-24 16:30:27 +08:00
Merge pull request #12064 from ic4f/dev_declarative3
Declarative mappings + tests for all models
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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:
|
||||
|
||||
+4
-2674
File diff suppressed because it is too large
Load Diff
@@ -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__)
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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__)
|
||||
|
||||
|
||||
@@ -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 (
|
||||
|
||||
@@ -3,6 +3,7 @@ boltons
|
||||
docutils
|
||||
markupsafe
|
||||
packaging
|
||||
pycryptodome
|
||||
pyyaml
|
||||
requests
|
||||
routes
|
||||
|
||||
@@ -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
@@ -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
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user