Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
130 changes: 130 additions & 0 deletions osf/management/commands/rollback_notifications.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,130 @@
import logging
from django.core.management.base import BaseCommand, CommandError
from django.db import transaction
import signal
from contextlib import contextmanager
from osf.models import NotificationSubscription, NotificationSubscriptionLegacy, NotificationType

logger = logging.getLogger(__name__)

# Reverse maps
FREQ_MAP_ROLLBACK = {
'none': 'none',
'weekly': 'email_digest',
'instantly': 'email_transactional',
}

NOTIFICATION_TYPE_TO_EVENT_NAME = {
# Provider notifications
NotificationType.Type.PROVIDER_NEW_PENDING_WITHDRAW_REQUESTS: 'new_pending_withdraw_requests',
NotificationType.Type.PROVIDER_CONTRIBUTOR_ADDED_PREPRINT: 'contributor_added_preprint',
NotificationType.Type.PROVIDER_NEW_PENDING_SUBMISSIONS: 'new_pending_submissions',
NotificationType.Type.PROVIDER_MODERATOR_ADDED: 'moderator_added',
NotificationType.Type.PROVIDER_REVIEWS_SUBMISSION_CONFIRMATION: 'reviews_submission_confirmation',
NotificationType.Type.PROVIDER_REVIEWS_RESUBMISSION_CONFIRMATION: 'reviews_resubmission_confirmation',
NotificationType.Type.PROVIDER_CONFIRM_EMAIL_MODERATION: 'confirm_email_moderation',

# Node notifications
NotificationType.Type.NODE_FILE_UPDATED: 'file_updated',

# Collection submissions
NotificationType.Type.COLLECTION_SUBMISSION_SUBMITTED: 'collection_submission_submitted',
NotificationType.Type.COLLECTION_SUBMISSION_ACCEPTED: 'collection_submission_accepted',
NotificationType.Type.COLLECTION_SUBMISSION_REJECTED: 'collection_submission_rejected',
NotificationType.Type.COLLECTION_SUBMISSION_REMOVED_ADMIN: 'collection_submission_removed_admin',
NotificationType.Type.COLLECTION_SUBMISSION_REMOVED_MODERATOR: 'collection_submission_removed_moderator',
NotificationType.Type.COLLECTION_SUBMISSION_REMOVED_PRIVATE: 'collection_submission_removed_private',
NotificationType.Type.COLLECTION_SUBMISSION_CANCEL: 'collection_submission_cancel',
}

TIMEOUT_SECONDS = 60 * 60 # 60 minutes timeout


@contextmanager
def time_limit(seconds):
def signal_handler(signum, frame):
raise TimeoutError('Migration timed out')

signal.signal(signal.SIGALRM, signal_handler)
signal.alarm(seconds)
try:
yield
finally:
signal.alarm(0)

def rollback_notification_subscriptions(dry_run=False):
migrated = list(NotificationSubscription.objects.select_related('notification_type', 'content_type'))
if not migrated:
logger.info('No NotificationSubscription objects found to rollback.')
return

recreated_count = 0
for sub in migrated:
notif_type_enum = sub.notification_type.name if sub.notification_type else None
event_name = NOTIFICATION_TYPE_TO_EVENT_NAME.get(notif_type_enum)

if not event_name:
logger.warning(f"Skipping rollback for subscription {sub.id}, unmapped notification type {notif_type_enum}")
continue

legacy_freq = FREQ_MAP_ROLLBACK.get(sub.message_frequency, 'none')

content_type = sub.content_type
model_class = content_type.model_class()
subscribed_object = model_class.objects.filter(id=sub.object_id).first()

if not subscribed_object:
logger.warning(f"Skipping rollback for subscription {sub.id}, missing subscribed object.")
continue

legacy_id = f"{sub.subscribed_object._id}_{event_name}"

if not dry_run:
obj, created = NotificationSubscriptionLegacy.objects.get_or_create(
_id=legacy_id,
event_name=event_name,
user_id=subscribed_object if content_type.model == 'osfuser' else None,
node=subscribed_object if content_type.model == 'abstractnode' else None,
provider=subscribed_object if content_type.model == 'abstractprovider' else None,
)

if sub.user:
if legacy_freq == 'email_digest':
obj.email_digest.add(sub.user)
elif legacy_freq == 'email_transactional':
obj.email_transactional.add(sub.user)
else:
obj.none.add(sub.user)

recreated_count += 1

if not dry_run:
logger.info(f"Rollback complete: recreated {recreated_count} legacy entries.")
else:
logger.info(f"[Dry Run] Would recreate {recreated_count} NotificationSubscriptionLegacy entries.")


class Command(BaseCommand):
help = 'Rollback NotificationSubscription objects back into NotificationSubscriptionLegacy.'

def add_arguments(self, parser):
parser.add_argument(
'--dry-run',
action='store_true',
help='Run migration in dry-run mode (no DB changes will be committed).'
)

def handle(self, *args, **options):
dry_run = options['dry_run']

try:
with time_limit(TIMEOUT_SECONDS):
with transaction.atomic():
rollback_notification_subscriptions(dry_run=dry_run)

except TimeoutError:
logger.error('Migration timed out. Rolling back changes.')
raise CommandError('Migration failed due to timeout')
except Exception as e:
logger.exception('Migration failed. Rolling back changes.')
raise CommandError(str(e))
110 changes: 110 additions & 0 deletions osf_tests/management_commands/test_rollback_notifications.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,110 @@
import pytest

from osf.models import RegistrationProvider
from osf_tests.factories import (
AuthUserFactory,
PreprintProviderFactory,
ProjectFactory,
)
from osf.models import (
NotificationType,
NotificationSubscription,
NotificationSubscriptionLegacy
)
from osf.management.commands.populate_notification_types import populate_notification_types
from osf.management.commands.rollback_notifications import rollback_notification_subscriptions


@pytest.mark.django_db
class TestNotificationSubscriptionRollback:

@pytest.fixture(autouse=True)
def notification_types(self):
return populate_notification_types()

@pytest.fixture()
def user(self):
return AuthUserFactory()

@pytest.fixture()
def users(self):
return {
'none': AuthUserFactory(),
'weekly': AuthUserFactory(),
'instantly': AuthUserFactory(),
}

@pytest.fixture()
def provider(self):
return PreprintProviderFactory()

@pytest.fixture()
def provider2(self):
return PreprintProviderFactory()

@pytest.fixture()
def node(self):
return ProjectFactory()

def create_sub(self, event_name, users, subscribed_object):
for message_frequency, user in users.items():
NotificationSubscription.objects.create(
notification_type=NotificationType.objects.get(name=event_name),
user=user,
message_frequency=message_frequency,
subscribed_object=subscribed_object,
)

def test_rollback_provider_subscription(self, users, provider, provider2):
self.create_sub(event_name=NotificationType.Type.PROVIDER_NEW_PENDING_SUBMISSIONS, users=users, subscribed_object=provider)
self.create_sub(event_name=NotificationType.Type.PROVIDER_NEW_PENDING_SUBMISSIONS, users=users, subscribed_object=provider2)
self.create_sub(event_name=NotificationType.Type.PROVIDER_NEW_PENDING_SUBMISSIONS, users=users, subscribed_object=RegistrationProvider.get_default())
rollback_notification_subscriptions()

provider_sub = NotificationSubscriptionLegacy.objects.filter(_id=f'{provider._id}_new_pending_submissions')
assert provider_sub.count() == 1
assert provider_sub.first().provider == provider
assert provider_sub.first().email_transactional.count() == 1
assert provider_sub.first().email_digest.count() == 1
assert provider_sub.first().none.count() == 1

provider2_sub = NotificationSubscriptionLegacy.objects.filter(_id=f'{provider2._id}_new_pending_submissions')
assert provider2_sub.count() == 1
assert provider2_sub.first().provider == provider2
assert provider2_sub.first().email_transactional.count() == 1
assert provider2_sub.first().email_digest.count() == 1
assert provider2_sub.first().none.count() == 1

default_provider_sub = NotificationSubscriptionLegacy.objects.filter(_id=f'{RegistrationProvider.get_default()._id}_new_pending_submissions')
assert default_provider_sub.count() == 1
assert default_provider_sub.first().provider == RegistrationProvider.get_default()
assert default_provider_sub.first().email_transactional.count() == 1
assert default_provider_sub.first().email_digest.count() == 1
assert default_provider_sub.first().none.count() == 1

def test_rollback_node_subscription(self, users, node):
self.create_sub(NotificationType.Type.NODE_FILE_UPDATED, users, subscribed_object=node)
rollback_notification_subscriptions()
node_sub = NotificationSubscriptionLegacy.objects.filter(_id=f'{node._id}_file_updated')
assert node_sub.count() == 1
assert node_sub.first().node == node
assert node_sub.first().email_transactional.count() == 1
assert node_sub.first().email_digest.count() == 1
assert node_sub.first().none.count() == 1

def test_multiple_subscriptions_no_old_types(self, users, user, provider, node):
assert not NotificationSubscription.objects.filter(user=user)
self.create_sub(NotificationType.Type.NODE_FORK_COMPLETED, users, subscribed_object=node)
rollback_notification_subscriptions()
assert not NotificationSubscriptionLegacy.objects.filter(node=node)

def test_idempotent_migration(self, users, user, node, provider):
self.create_sub(NotificationType.Type.NODE_FILE_UPDATED, users, subscribed_object=node)
rollback_notification_subscriptions()
rollback_notification_subscriptions()
node_sub = NotificationSubscriptionLegacy.objects.filter(_id=f'{node._id}_file_updated')
assert node_sub.count() == 1
assert node_sub.first().node == node
assert node_sub.first().email_transactional.count() == 1
assert node_sub.first().email_digest.count() == 1
assert node_sub.first().none.count() == 1
Loading