Do not start a new db transaction

This commit is contained in:
John Davis
2023-06-14 15:18:27 -04:00
parent fafcfbb0dd
commit da9aa2bb39
2 changed files with 116 additions and 115 deletions
+103 -102
View File
@@ -34,6 +34,7 @@ from galaxy.model import (
UserNotificationAssociation,
UserRoleAssociation,
)
from galaxy.model.base import transaction
from galaxy.model.scoped_session import galaxy_scoped_session
from galaxy.schema.notifications import (
BroadcastNotificationCreateRequest,
@@ -115,11 +116,11 @@ class NotificationManager:
self.ensure_notifications_enabled()
recipient_users = self.recipient_resolver.resolve(request.recipients)
notifications_sent = len(recipient_users)
with self.sa_session.begin():
notification = self._create_notification_model(request.notification)
self.sa_session.add(notification)
self._send_to_users(notification, recipient_users)
notification = self._create_notification_model(request.notification)
self.sa_session.add(notification)
self._send_to_users(notification, recipient_users)
with transaction(self.sa_session):
self.sa_session.commit()
return notification, notifications_sent
def _send_to_users(self, notification: Notification, users: List[User]):
@@ -140,9 +141,10 @@ class NotificationManager:
This kind of notification is not explicitly associated with any specific user but it is accessible by all users.
"""
self.ensure_notifications_enabled()
with self.sa_session.begin():
notification = self._create_notification_model(request)
self.sa_session.add(notification)
notification = self._create_notification_model(request)
self.sa_session.add(notification)
with transaction(self.sa_session):
self.sa_session.commit()
return notification
def get_user_notification(self, user: User, notification_id: int, active_only: Optional[bool] = True):
@@ -228,44 +230,46 @@ class NotificationManager:
) -> int:
"""Updates a batch of notifications associated with the user using the requested values."""
updated_row_count = 0
with self.sa_session.begin():
stmt = update(UserNotificationAssociation).where(
and_(
UserNotificationAssociation.user_id == user.id,
UserNotificationAssociation.notification_id.in_(notification_ids),
)
stmt = update(UserNotificationAssociation).where(
and_(
UserNotificationAssociation.user_id == user.id,
UserNotificationAssociation.notification_id.in_(notification_ids),
)
if request.seen is not None:
seen_time = self._now if request.seen else None
stmt = stmt.values(seen_time=seen_time)
if request.deleted is not None:
stmt = stmt.values(deleted=request.deleted)
result = self.sa_session.execute(stmt)
updated_row_count = result.rowcount
)
if request.seen is not None:
seen_time = self._now if request.seen else None
stmt = stmt.values(seen_time=seen_time)
if request.deleted is not None:
stmt = stmt.values(deleted=request.deleted)
result = self.sa_session.execute(stmt)
updated_row_count = result.rowcount
with transaction(self.sa_session):
self.sa_session.commit()
return updated_row_count
def update_broadcasted_notification(self, notification_id: int, request: NotificationBroadcastUpdateRequest) -> int:
"""Updates a single broadcasted notification with the requested values."""
updated_row_count = 0
with self.sa_session.begin():
stmt = update(Notification).where(
and_(
Notification.id == notification_id,
Notification.category == MandatoryNotificationCategory.broadcast,
)
stmt = update(Notification).where(
and_(
Notification.id == notification_id,
Notification.category == MandatoryNotificationCategory.broadcast,
)
if request.source is not None:
stmt = stmt.values(source=request.source)
if request.variant is not None:
stmt = stmt.values(variant=request.variant)
if request.publication_time is not None:
stmt = stmt.values(publication_time=request.publication_time)
if request.expiration_time is not None:
stmt = stmt.values(expiration_time=request.expiration_time)
if request.content is not None:
stmt = stmt.values(content=request.content.json())
result = self.sa_session.execute(stmt)
updated_row_count = result.rowcount
)
if request.source is not None:
stmt = stmt.values(source=request.source)
if request.variant is not None:
stmt = stmt.values(variant=request.variant)
if request.publication_time is not None:
stmt = stmt.values(publication_time=request.publication_time)
if request.expiration_time is not None:
stmt = stmt.values(expiration_time=request.expiration_time)
if request.content is not None:
stmt = stmt.values(content=request.content.json())
result = self.sa_session.execute(stmt)
updated_row_count = result.rowcount
with transaction(self.sa_session):
self.sa_session.commit()
return updated_row_count
def get_user_notification_preferences(self, user: User) -> UserNotificationPreferences:
@@ -287,8 +291,9 @@ class NotificationManager:
"""Updates the user's notification preferences with the requested changes."""
notification_preferences = self.get_user_notification_preferences(user)
notification_preferences.update(request.preferences)
with self.sa_session.begin():
user.preferences[NOTIFICATION_PREFERENCES_SECTION_NAME] = notification_preferences.json()
user.preferences[NOTIFICATION_PREFERENCES_SECTION_NAME] = notification_preferences.json()
with transaction(self.sa_session):
self.sa_session.commit()
return notification_preferences
def cleanup_expired_notifications(self) -> CleanupResultSummary:
@@ -298,74 +303,70 @@ class NotificationManager:
deleted_notifications_count = 0
deleted_associations_count = 0
execution_options_for_delete = {"synchronize_session": "fetch"}
with self.sa_session.begin():
is_not_deleted = and_(
UserNotificationAssociation.deleted.is_(False),
)
has_expired = Notification.expiration_time <= self._now
is_not_deleted = and_(
UserNotificationAssociation.deleted.is_(False),
)
has_expired = Notification.expiration_time <= self._now
# Find those notification ids that have expired
non_expired_notifications_query = (
select(Notification.id)
.where(has_expired)
.where(~Notification.user_notification_associations.any(is_not_deleted))
)
non_expired_notification_ids = (
self.sa_session.execute(non_expired_notifications_query).scalars().fetchall()
)
# Find those notification ids that have expired
non_expired_notifications_query = (
select(Notification.id)
.where(has_expired)
.where(~Notification.user_notification_associations.any(is_not_deleted))
)
non_expired_notification_ids = self.sa_session.execute(non_expired_notifications_query).scalars().fetchall()
# Delete all notifications and associations that have expired
delete_expired_associations_query = delete(UserNotificationAssociation).where(
UserNotificationAssociation.notification_id.in_(non_expired_notification_ids)
)
result = self.sa_session.execute(
delete_expired_associations_query, execution_options=execution_options_for_delete
)
deleted_associations_count += result.rowcount
# Delete all notifications and associations that have expired
delete_expired_associations_query = delete(UserNotificationAssociation).where(
UserNotificationAssociation.notification_id.in_(non_expired_notification_ids)
)
result = self.sa_session.execute(
delete_expired_associations_query, execution_options=execution_options_for_delete
)
deleted_associations_count += result.rowcount
delete_expired_notifications_query = delete(Notification).where(
Notification.id.in_(non_expired_notification_ids)
)
result = self.sa_session.execute(
delete_expired_notifications_query, execution_options=execution_options_for_delete
)
deleted_notifications_count += result.rowcount
delete_expired_notifications_query = delete(Notification).where(
Notification.id.in_(non_expired_notification_ids)
)
result = self.sa_session.execute(
delete_expired_notifications_query, execution_options=execution_options_for_delete
)
deleted_notifications_count += result.rowcount
# Find those notification ids that have expired
expired_notifications_query = (
select(Notification.id)
.where(has_expired)
.where(Notification.user_notification_associations.any(is_not_deleted))
)
expired_notification_ids = self.sa_session.execute(expired_notifications_query).scalars()
# Find those notification ids that have expired
expired_notifications_query = (
select(Notification.id)
.where(has_expired)
.where(Notification.user_notification_associations.any(is_not_deleted))
)
expired_notification_ids = self.sa_session.execute(expired_notifications_query).scalars()
# Delete those associations that did expire
expired_associations_query = (
select(UserNotificationAssociation)
.where(UserNotificationAssociation.notification_id.in_(expired_notification_ids))
# Delete those associations that did expire
expired_associations_query = select(UserNotificationAssociation).where(
UserNotificationAssociation.notification_id.in_(expired_notification_ids)
)
delete_expired_associations_query = delete(UserNotificationAssociation).where(
UserNotificationAssociation.id.in_(
select(UserNotificationAssociation.id).select_from(expired_associations_query.subquery())
)
delete_expired_associations_query = delete(UserNotificationAssociation).where(
UserNotificationAssociation.id.in_(
select(UserNotificationAssociation.id).select_from(
expired_associations_query.subquery()
)
)
)
result = self.sa_session.execute(
delete_expired_associations_query, execution_options=execution_options_for_delete
)
deleted_associations_count += result.rowcount
)
result = self.sa_session.execute(
delete_expired_associations_query, execution_options=execution_options_for_delete
)
deleted_associations_count += result.rowcount
# Delete broadcasted
delete_expired_broadcasted_notifications_query = (
delete(Notification)
.where(has_expired)
.where(Notification.category == MandatoryNotificationCategory.broadcast)
)
result = self.sa_session.execute(
delete_expired_broadcasted_notifications_query, execution_options=execution_options_for_delete
)
deleted_notifications_count += result.rowcount
# Delete broadcasted
delete_expired_broadcasted_notifications_query = (
delete(Notification)
.where(has_expired)
.where(Notification.category == MandatoryNotificationCategory.broadcast)
)
result = self.sa_session.execute(
delete_expired_broadcasted_notifications_query, execution_options=execution_options_for_delete
)
deleted_notifications_count += result.rowcount
with transaction(self.sa_session):
self.sa_session.commit()
return CleanupResultSummary(deleted_notifications_count, deleted_associations_count)
def _create_notification_model(self, payload: NotificationCreateData):
@@ -443,19 +443,19 @@ class TestNotificationRecipientResolver(NotificationsBaseTestCase):
def _create_test_group(self, name: str, users: List[User], roles: List[Role]):
sa_session = self.trans.sa_session
with sa_session.begin():
group = Group(name=name)
sa_session.add(group)
self.trans.app.security_agent.set_entity_group_associations(groups=[group], roles=roles, users=users)
return group
group = Group(name=name)
sa_session.add(group)
self.trans.app.security_agent.set_entity_group_associations(groups=[group], roles=roles, users=users)
sa_session.flush()
return group
def _create_test_role(self, name: str, users: List[User], groups: List[Group]):
sa_session = self.trans.sa_session
with sa_session.begin():
role = Role(name=name)
sa_session.add(role)
for user in users:
self.trans.app.security_agent.associate_user_role(user, role)
for group in groups:
self.trans.app.security_agent.associate_group_role(group, role)
return role
role = Role(name=name)
sa_session.add(role)
for user in users:
self.trans.app.security_agent.associate_user_role(user, role)
for group in groups:
self.trans.app.security_agent.associate_group_role(group, role)
sa_session.flush()
return role