diff --git a/admin_tests/nodes/test_views.py b/admin_tests/nodes/test_views.py index 9607a1a0c99..92ad6c01858 100644 --- a/admin_tests/nodes/test_views.py +++ b/admin_tests/nodes/test_views.py @@ -602,14 +602,14 @@ def setUp(self): self.contr1 = UserFactory() self.contr2 = UserFactory() self.contr3 = UserFactory() - - pre_moderation_draft = DraftRegistrationFactory( - title='pre-moderation-registration', - description='some description', - registration_schema=get_default_metaschema(), - provider=RegistrationProviderFactory(reviews_workflow='pre-moderation'), - creator=self.user - ) + with capture_notifications(): + pre_moderation_draft = DraftRegistrationFactory( + title='pre-moderation-registration', + description='some description', + registration_schema=get_default_metaschema(), + provider=RegistrationProviderFactory(reviews_workflow='pre-moderation'), + creator=self.user + ) self._add_contributor(pre_moderation_draft, permissions.ADMIN, self.contr1) self._add_contributor(pre_moderation_draft, permissions.ADMIN, self.contr2) self._add_contributor(pre_moderation_draft, permissions.ADMIN, self.contr3) diff --git a/api/subscriptions/serializers.py b/api/subscriptions/serializers.py index 1b7e6449833..4cfdf132885 100644 --- a/api/subscriptions/serializers.py +++ b/api/subscriptions/serializers.py @@ -25,14 +25,13 @@ class SubscriptionSerializer(JSONAPISerializer): source='message_frequency', required=True, ) - - class Meta: - type_ = 'subscription' - links = LinksField({ 'self': 'get_absolute_url', }) + class Meta: + type_ = 'subscription' + def get_absolute_url(self, obj): return obj.absolute_api_v2_url diff --git a/api_tests/draft_nodes/views/test_draft_node_detail.py b/api_tests/draft_nodes/views/test_draft_node_detail.py index 777d77a948c..4044973294e 100644 --- a/api_tests/draft_nodes/views/test_draft_node_detail.py +++ b/api_tests/draft_nodes/views/test_draft_node_detail.py @@ -7,6 +7,7 @@ AuthUserFactory, ProjectFactory ) +from tests.utils import capture_notifications @pytest.mark.django_db @@ -21,7 +22,8 @@ def user_two(self): return AuthUserFactory() def test_detail_response(self, app, user, user_two): - draft_reg = DraftRegistrationFactory(initiator=user) + with capture_notifications(): + draft_reg = DraftRegistrationFactory(initiator=user) draft_reg.add_contributor(user_two) draft_reg.save() diff --git a/api_tests/draft_nodes/views/test_draft_node_draft_registrations_list.py b/api_tests/draft_nodes/views/test_draft_node_draft_registrations_list.py index 63d3934f258..be115b0ad4c 100644 --- a/api_tests/draft_nodes/views/test_draft_node_draft_registrations_list.py +++ b/api_tests/draft_nodes/views/test_draft_node_draft_registrations_list.py @@ -6,6 +6,7 @@ AuthUserFactory, ) from osf.utils.permissions import WRITE +from tests.utils import capture_notifications @pytest.mark.django_db @@ -21,9 +22,10 @@ def user_write_contrib(self): @pytest.fixture() def draft_registration(self, user, user_write_contrib): - draft_reg = DraftRegistrationFactory( - initiator=user - ) + with capture_notifications(): + draft_reg = DraftRegistrationFactory( + initiator=user + ) draft_reg.add_contributor( user_write_contrib, permissions=WRITE) diff --git a/api_tests/draft_nodes/views/test_draft_node_files_lists.py b/api_tests/draft_nodes/views/test_draft_node_files_lists.py index bad74618834..c9a61fac75b 100644 --- a/api_tests/draft_nodes/views/test_draft_node_files_lists.py +++ b/api_tests/draft_nodes/views/test_draft_node_files_lists.py @@ -18,6 +18,7 @@ from addons.github.tests.factories import GitHubAccountFactory from api.base.utils import waterbutler_api_url_for from api_tests import utils as api_utils +from tests.utils import capture_notifications from website import settings @@ -25,7 +26,8 @@ class TestDraftNodeProvidersList(ApiTestCase): def setUp(self): super().setUp() self.user = AuthUserFactory() - self.draft_reg = DraftRegistrationFactory(creator=self.user) + with capture_notifications(): + self.draft_reg = DraftRegistrationFactory(creator=self.user) self.draft_node = self.draft_reg.branched_from self.url = f'/{API_BASE}draft_nodes/{self.draft_node._id}/files/' @@ -149,7 +151,8 @@ class TestNodeFilesList(ApiTestCase): def setUp(self): super().setUp() self.user = AuthUserFactory() - self.draft_reg = DraftRegistrationFactory(creator=self.user) + with capture_notifications(): + self.draft_reg = DraftRegistrationFactory(creator=self.user) self.draft_node = self.draft_reg.branched_from self.private_url = '/{}draft_nodes/{}/files/'.format( API_BASE, self.draft_node._id) @@ -489,7 +492,8 @@ class TestNodeFilesListFiltering(ApiTestCase): def setUp(self): super().setUp() self.user = AuthUserFactory() - self.draft_reg = DraftRegistrationFactory(creator=self.user) + with capture_notifications(): + self.draft_reg = DraftRegistrationFactory(creator=self.user) self.draft_node = self.draft_reg.branched_from # Prep HTTP mocks prepare_mock_wb_response( @@ -631,7 +635,8 @@ class TestNodeFilesListPagination(ApiTestCase): def setUp(self): super().setUp() self.user = AuthUserFactory() - self.draft_reg = DraftRegistrationFactory(creator=self.user) + with capture_notifications(): + self.draft_reg = DraftRegistrationFactory(creator=self.user) self.draft_node = self.draft_reg.branched_from def add_github(self): @@ -703,7 +708,8 @@ class TestDraftNodeStorageProviderDetail(ApiTestCase): def setUp(self): super().setUp() self.user = AuthUserFactory() - self.draft_reg = DraftRegistrationFactory(initiator=self.user) + with capture_notifications(): + self.draft_reg = DraftRegistrationFactory(initiator=self.user) self.draft_node = self.draft_reg.branched_from self.private_url = '/{}draft_nodes/{}/files/providers/osfstorage/'.format( API_BASE, self.draft_node._id) diff --git a/api_tests/draft_registrations/views/test_draft_registration_contributor_detail.py b/api_tests/draft_registrations/views/test_draft_registration_contributor_detail.py index 0c2dce3501b..88e97db2c99 100644 --- a/api_tests/draft_registrations/views/test_draft_registration_contributor_detail.py +++ b/api_tests/draft_registrations/views/test_draft_registration_contributor_detail.py @@ -14,6 +14,7 @@ AuthUserFactory ) from osf.utils import permissions +from tests.utils import capture_notifications @pytest.fixture() @@ -41,9 +42,10 @@ def project_public(self, user, title, description, category): @pytest.fixture() def project_private(self, user, title, description, category): # Defining "private project" as a draft reg, overriding TestContributorDetail - draft = DraftRegistrationFactory( - initiator=user, - ) + with capture_notifications(): + draft = DraftRegistrationFactory( + initiator=user, + ) return draft @pytest.fixture() @@ -108,7 +110,8 @@ class TestDraftContributorOrdering(TestNodeContributorOrdering): @pytest.fixture() def project(self, user, contribs): # Overrides TestNodeContributorOrdering - project = DraftRegistrationFactory(initiator=user, title='hey') + with capture_notifications(): + project = DraftRegistrationFactory(initiator=user, title='hey') for contrib in contribs: if contrib._id != user._id: project.add_contributor( @@ -145,7 +148,8 @@ class TestDraftRegistrationContributorUpdate(TestNodeContributorUpdate): @pytest.fixture() def project(self, user, contrib): # Overrides TestNodeContributorUpdate - draft = DraftRegistrationFactory(creator=user) + with capture_notifications(): + draft = DraftRegistrationFactory(creator=user) draft.add_contributor( contrib, permissions=permissions.WRITE, @@ -175,12 +179,14 @@ def contrib(self): @pytest.fixture() def project(self, user, contrib): # Overrides TestNodeContributorPartialUpdate - project = DraftRegistrationFactory(creator=user) - project.add_contributor( - contrib, - permissions=permissions.WRITE, - visible=True, - save=True) + with capture_notifications(): + project = DraftRegistrationFactory(creator=user) + project.add_contributor( + contrib, + permissions=permissions.WRITE, + visible=True, + save=True + ) return project @pytest.fixture() @@ -226,7 +232,8 @@ class TestDraftContributorDelete(TestNodeContributorDelete): @pytest.fixture() def project(self, user, user_write_contrib): # Overrides TestNodeContributorDelete - project = DraftRegistrationFactory(creator=user) + with capture_notifications(): + project = DraftRegistrationFactory(creator=user) project.add_contributor( user_write_contrib, permissions=permissions.WRITE, @@ -265,7 +272,8 @@ def user_non_biblio_contrib(self): @pytest.fixture() def draft_registration(self, user, user_non_biblio_contrib): # Overrides TestNodeContributorDelete - project = DraftRegistrationFactory(creator=user) + with capture_notifications(): + project = DraftRegistrationFactory(creator=user) project.add_contributor( user, permissions=permissions.ADMIN, diff --git a/api_tests/draft_registrations/views/test_draft_registration_contributor_list.py b/api_tests/draft_registrations/views/test_draft_registration_contributor_list.py index f896022f149..7d7d3082cf6 100644 --- a/api_tests/draft_registrations/views/test_draft_registration_contributor_list.py +++ b/api_tests/draft_registrations/views/test_draft_registration_contributor_list.py @@ -49,12 +49,13 @@ def project_public(self, user, title, description, category): @pytest.fixture() def project_private(self, user, title, description, category): - return DraftRegistrationFactory( - title=title, - description=description, - category=category, - initiator=user - ) + with capture_notifications(): + return DraftRegistrationFactory( + title=title, + description=description, + category=category, + initiator=user + ) class TestDraftRegistrationContributorList(DraftRegistrationCRUDTestCase, TestNodeContributorList): @@ -338,9 +339,10 @@ class TestDraftContributorBulkUpdated(DraftRegistrationCRUDTestCase, TestNodeCon def project_public( self, user, user_two, user_three, title, description, category): - project_public = DraftRegistrationFactory( - initiator=user - ) + with capture_notifications(): + project_public = DraftRegistrationFactory( + initiator=user + ) project_public.add_contributor( user_two, permissions=permissions.READ, @@ -355,9 +357,13 @@ def project_public( def project_private( self, user, user_two, user_three, title, description, category): - project_private = DraftRegistrationFactory( - initiator=user - ) + + try: + with capture_notifications(): + project_private = DraftRegistrationFactory(initiator=user) + except AssertionError: # No message sent + project_private = DraftRegistrationFactory(initiator=user) + project_private.add_contributor( user_two, permissions=permissions.READ, @@ -382,9 +388,10 @@ class TestDraftRegistrationContributorBulkPartialUpdate(DraftRegistrationCRUDTes def project_public( self, user, user_two, user_three, title, description, category): - project_public = DraftRegistrationFactory( - initiator=user - ) + with capture_notifications(): + project_public = DraftRegistrationFactory( + initiator=user + ) project_public.add_contributor( user_two, permissions=permissions.READ, @@ -399,9 +406,11 @@ def project_public( def project_private( self, user, user_two, user_three, title, description, category): - project_private = DraftRegistrationFactory( - initiator=user - ) + try: + with capture_notifications(): + project_private = DraftRegistrationFactory(initiator=user) + except AssertionError: # No message sent + project_private = DraftRegistrationFactory(initiator=user) project_private.add_contributor( user_two, permissions=permissions.READ, @@ -436,9 +445,10 @@ def url_private(self, project_private): def project_public( self, user, user_two, user_three, title, description, category): - project_public = DraftRegistrationFactory( - initiator=user - ) + with capture_notifications(): + project_public = DraftRegistrationFactory( + initiator=user + ) project_public.add_contributor( user_two, permissions=permissions.READ, @@ -453,9 +463,11 @@ def project_public( def project_private( self, user, user_two, user_three, title, description, category): - project_private = DraftRegistrationFactory( - initiator=user - ) + try: + with capture_notifications(): + project_private = DraftRegistrationFactory(initiator=user) + except AssertionError: # No message sent + project_private = DraftRegistrationFactory(initiator=user) project_private.add_contributor( user_two, permissions=permissions.READ, @@ -472,7 +484,8 @@ def project_private( class TestDraftRegistrationContributorFiltering(DraftRegistrationCRUDTestCase, TestNodeContributorFiltering): @pytest.fixture() def project(self, user): - return DraftRegistrationFactory(initiator=user) + with capture_notifications(): + return DraftRegistrationFactory(initiator=user) @pytest.fixture() def url(self, project): diff --git a/api_tests/draft_registrations/views/test_draft_registration_detail.py b/api_tests/draft_registrations/views/test_draft_registration_detail.py index 2106f87fb5a..0e5c7c56481 100644 --- a/api_tests/draft_registrations/views/test_draft_registration_detail.py +++ b/api_tests/draft_registrations/views/test_draft_registration_detail.py @@ -16,6 +16,7 @@ SubjectFactory, ProjectFactory, ) +from tests.utils import capture_notifications from website.settings import API_DOMAIN @@ -63,8 +64,8 @@ def test_detail_view_returns_editable_fields( assert 'contributors' in relationships def test_detail_view_returns_editable_fields_no_specified_node(self, app, user): - - draft_registration = DraftRegistrationFactory(initiator=user, branched_from=None) + with capture_notifications(): + draft_registration = DraftRegistrationFactory(initiator=user, branched_from=None) url = f'{API_DOMAIN}{API_BASE}draft_registrations/{draft_registration._id}/' res = app.get(url, auth=user.auth, expect_errors=True) @@ -271,11 +272,12 @@ def url_draft_registrations(self, project_public, draft_registration): @pytest.fixture() def draft_registration(self, user, user_read_contrib, user_write_contrib, project_public, schema): - draft_registration = DraftRegistrationFactory( - initiator=user, - registration_schema=schema, - branched_from=None - ) + with capture_notifications(): + draft_registration = DraftRegistrationFactory( + initiator=user, + registration_schema=schema, + branched_from=None + ) draft_registration.add_contributor( user_write_contrib, permissions=WRITE) @@ -293,11 +295,12 @@ def schema_open_ended(self): @pytest.fixture def draft_registration_open_ended(self, user, schema_open_ended): - return DraftRegistrationFactory( - initiator=user, - registration_schema=schema_open_ended, - branched_from=None - ) + with capture_notifications(): + return DraftRegistrationFactory( + initiator=user, + registration_schema=schema_open_ended, + branched_from=None + ) @pytest.fixture() def url_draft_registration_open_ended(self, draft_registration_open_ended): @@ -498,11 +501,12 @@ def url_draft_registrations(self, project_public, draft_registration): @pytest.fixture() def draft_registration(self, user, user_read_contrib, user_write_contrib, project_public, schema): - draft_registration = DraftRegistrationFactory( - initiator=user, - registration_schema=schema, - branched_from=None - ) + with capture_notifications(): + draft_registration = DraftRegistrationFactory( + initiator=user, + registration_schema=schema, + branched_from=None + ) draft_registration.add_contributor( user_write_contrib, permissions=WRITE) @@ -540,11 +544,12 @@ def url_draft_registrations(self, project_public, draft_registration): @pytest.fixture() def draft_registration(self, user, user_read_contrib, user_write_contrib, project_public, schema): - draft_registration = DraftRegistrationFactory( - initiator=user, - registration_schema=schema, - branched_from=None - ) + with capture_notifications(): + draft_registration = DraftRegistrationFactory( + initiator=user, + registration_schema=schema, + branched_from=None + ) draft_registration.add_contributor( user_write_contrib, permissions=WRITE) diff --git a/api_tests/draft_registrations/views/test_draft_registration_institutions_list.py b/api_tests/draft_registrations/views/test_draft_registration_institutions_list.py index 19c308dc937..6b0def92d02 100644 --- a/api_tests/draft_registrations/views/test_draft_registration_institutions_list.py +++ b/api_tests/draft_registrations/views/test_draft_registration_institutions_list.py @@ -3,6 +3,8 @@ from api.base.settings.defaults import API_BASE from api_tests.nodes.views.test_node_institutions_list import TestNodeInstitutionList from osf_tests.factories import DraftRegistrationFactory, AuthUserFactory +from tests.utils import capture_notifications +from tests.utils import capture_notifications @pytest.fixture() @@ -20,7 +22,10 @@ class TestDraftRegistrationInstitutionList(TestNodeInstitutionList): @pytest.fixture() def node_one(self, institution, user): # Overrides TestNodeInstitutionList - draft = DraftRegistrationFactory(initiator=user) + with capture_notifications(): + draft = DraftRegistrationFactory(initiator=user) + with capture_notifications(): + draft = DraftRegistrationFactory(initiator=user) draft.affiliated_institutions.add(institution) draft.save() return draft @@ -28,7 +33,8 @@ def node_one(self, institution, user): @pytest.fixture() def node_two(self, user): # Overrides TestNodeInstitutionList - return DraftRegistrationFactory(initiator=user) + with capture_notifications(): + return DraftRegistrationFactory(initiator=user) @pytest.fixture() def node_one_url(self, node_one): diff --git a/api_tests/draft_registrations/views/test_draft_registration_relationship_institutions.py b/api_tests/draft_registrations/views/test_draft_registration_relationship_institutions.py index babd6daf398..5dee3506165 100644 --- a/api_tests/draft_registrations/views/test_draft_registration_relationship_institutions.py +++ b/api_tests/draft_registrations/views/test_draft_registration_relationship_institutions.py @@ -4,6 +4,7 @@ from osf_tests.factories import DraftRegistrationFactory, AuthUserFactory, InstitutionFactory from osf.utils import permissions +from tests.utils import capture_notifications @pytest.mark.django_db @@ -16,7 +17,8 @@ def resource_factory(self): @pytest.fixture() def node(self, user, write_contrib, read_contrib): # Overrides TestNodeRelationshipInstitutions - draft = DraftRegistrationFactory(initiator=user) + with capture_notifications(): + draft = DraftRegistrationFactory(initiator=user) draft.add_contributor( write_contrib, permissions=permissions.WRITE) @@ -443,7 +445,8 @@ def test_delete_user_is_read_only( def test_delete_user_is_admin_but_not_affiliated_with_inst( self, app, institution_one, resource_factory, create_payload, make_resource_url): user = AuthUserFactory() - node = resource_factory(creator=user) + with capture_notifications(): + node = resource_factory(creator=user) node.affiliated_institutions.add(institution_one) node.save() assert institution_one in node.affiliated_institutions.all() diff --git a/api_tests/draft_registrations/views/test_draft_registration_relationship_subjects.py b/api_tests/draft_registrations/views/test_draft_registration_relationship_subjects.py index 4e982ad0275..709a1b8297d 100644 --- a/api_tests/draft_registrations/views/test_draft_registration_relationship_subjects.py +++ b/api_tests/draft_registrations/views/test_draft_registration_relationship_subjects.py @@ -6,13 +6,15 @@ from osf_tests.factories import ( DraftRegistrationFactory ) +from tests.utils import capture_notifications @pytest.mark.django_db class TestDraftRegistrationRelationshipSubjects(SubjectsRelationshipMixin): @pytest.fixture() def resource(self, user_admin_contrib, user_write_contrib, user_read_contrib): - draft = DraftRegistrationFactory(creator=user_admin_contrib) + with capture_notifications(): + draft = DraftRegistrationFactory(creator=user_admin_contrib) draft.add_contributor(user_write_contrib, permissions=WRITE) draft.add_contributor(user_read_contrib, permissions=READ) draft.save() diff --git a/api_tests/draft_registrations/views/test_draft_registration_subjects_list.py b/api_tests/draft_registrations/views/test_draft_registration_subjects_list.py index 97b2ffb001d..c26d8bfd6d7 100644 --- a/api_tests/draft_registrations/views/test_draft_registration_subjects_list.py +++ b/api_tests/draft_registrations/views/test_draft_registration_subjects_list.py @@ -6,12 +6,15 @@ from osf_tests.factories import ( DraftRegistrationFactory, ) +from tests.utils import capture_notifications + class TestDraftRegistrationSubjectsList(SubjectsListMixin): @pytest.fixture() def resource(self, user_admin_contrib, user_write_contrib, user_read_contrib): # Overrides SubjectsListMixin - draft = DraftRegistrationFactory(initiator=user_admin_contrib) + with capture_notifications(): + draft = DraftRegistrationFactory(initiator=user_admin_contrib) draft.add_contributor(user_write_contrib, permissions=WRITE) draft.add_contributor(user_read_contrib, permissions=READ) draft.save() diff --git a/api_tests/nodes/views/test_node_contributors_detail.py b/api_tests/nodes/views/test_node_contributors_detail.py index 5dc600279ba..63268adf369 100644 --- a/api_tests/nodes/views/test_node_contributors_detail.py +++ b/api_tests/nodes/views/test_node_contributors_detail.py @@ -182,7 +182,10 @@ def test_detail_includes_is_curator( other_contributor = AuthUserFactory() project_public.add_contributor( - other_contributor, auth=Auth(user), save=True) + other_contributor, + auth=Auth(user), + save=True + ) other_contributor_detail = self.make_resource_url(project_public._id, other_contributor._id) @@ -191,7 +194,12 @@ def test_detail_includes_is_curator( curator_contributor = AuthUserFactory() project_public.add_contributor( - curator_contributor, auth=Auth(user), save=True, make_curator=True, visible=False) + curator_contributor, + auth=Auth(user), + save=True, + make_curator=True, + visible=False + ) curator_contributor_detail = self.make_resource_url(project_public._id, curator_contributor._id) diff --git a/api_tests/nodes/views/test_node_contributors_list.py b/api_tests/nodes/views/test_node_contributors_list.py index 952ffd09878..1e13454312c 100644 --- a/api_tests/nodes/views/test_node_contributors_list.py +++ b/api_tests/nodes/views/test_node_contributors_list.py @@ -136,8 +136,7 @@ def test_permissions_work_with_many_users( project_private.add_contributor( user, - permissions=perm, - notification_type=False + permissions=perm ) users[perm].append(user._id) @@ -244,8 +243,7 @@ def test_disabled_contributors_contain_names_under_meta( ): project_public.add_contributor( user_two, - save=True, - notification_type=False + save=True ) user_two.is_disabled = True @@ -297,7 +295,9 @@ def test_unregistered_contributor_field_is_null_if_account_claimed(self, app, us def test_unregistered_contributors_show_up_as_name_associated_with_project(self, app, user): project = ProjectFactory(creator=user, is_public=True) project.add_unregistered_contributor( - 'Robert Jackson', 'robert@gmail.com', auth=Auth(user) + 'Robert Jackson', + 'robert@gmail.com', + auth=Auth(user) ) url = f'/{API_BASE}nodes/{project._id}/contributors/' res = app.get(url, auth=user.auth, expect_errors=True) @@ -314,7 +314,9 @@ def test_unregistered_contributors_show_up_as_name_associated_with_project(self, project_two = ProjectFactory(creator=user, is_public=True) project_two.add_unregistered_contributor( - 'Bob Jackson', 'robert@gmail.com', auth=Auth(user) + 'Bob Jackson', + 'robert@gmail.com', + auth=Auth(user) ) url = f'/{API_BASE}nodes/{project_two._id}/contributors/' res = app.get(url, auth=user.auth, expect_errors=True) @@ -341,7 +343,10 @@ def test_contributors_order_is_the_same_over_multiple_requests( else: visible = False project_public.add_contributor( - new_user, visible=visible, auth=Auth(project_public.creator), save=True + new_user, + visible=visible, + auth=Auth(project_public.creator), + save=True ) req_one = app.get(f'{url_public}?page=2', auth=Auth(project_public.creator)) req_two = app.get(f'{url_public}?page=2', auth=Auth(project_public.creator)) @@ -634,7 +639,10 @@ def test_adds_contributor_public_project_non_admin( url_public, ): project_public.add_contributor( - user_two, permissions=permissions.WRITE, auth=Auth(user), save=True + user_two, + permissions=permissions.WRITE, + auth=Auth(user), + save=True ) res = app.post_json_api( url_public, data_user_three, auth=user_two.auth, expect_errors=True @@ -805,7 +813,11 @@ def test_adds_none_permission_contributor_private_project_admin_uses_default_per def test_adds_already_existing_contributor_private_project_admin( self, app, user, user_two, project_private, data_user_two, url_private ): - project_private.add_contributor(user_two, auth=Auth(user), save=True) + project_private.add_contributor( + user_two, + auth=Auth(user), + save=True + ) project_private.reload() res = app.post_json_api( @@ -840,7 +852,9 @@ def test_adds_contributor_private_project_non_admin( url_private, ): project_private.add_contributor( - user_two, permissions=permissions.WRITE, auth=Auth(user) + user_two, + permissions=permissions.WRITE, + auth=Auth(user) ) res = app.post_json_api( url_private, data_user_three, auth=user_two.auth, expect_errors=True @@ -966,7 +980,9 @@ def test_add_unregistered_contributor_already_contributor( ): name, email = fake.name(), fake_email() project_public.add_unregistered_contributor( - auth=Auth(user), fullname=name, email=email + auth=Auth(user), + fullname=name, + email=email ) payload = { 'data': { @@ -1040,7 +1056,10 @@ def test_add_contributor_set_index_out_of_range( user_contrib_one = UserFactory() project_public.add_contributor(user_contrib_one, save=True) user_contrib_two = UserFactory() - project_public.add_contributor(user_contrib_two, save=True) + project_public.add_contributor( + user_contrib_two, + save=True + ) payload = { 'data': { 'type': 'contributors', @@ -1426,7 +1445,10 @@ def test_node_contributor_bulk_create_contributor_exists( self, app, user, user_two, project_public, payload_one, payload_two, url_public ): project_public.add_contributor( - user_two, permissions=permissions.READ, visible=True, save=True + user_two, + permissions=permissions.READ, + visible=True, + save=True ) res = app.post_json_api( url_public, @@ -1498,7 +1520,9 @@ def test_node_contributor_bulk_create_errors( # test_node_contributor_bulk_create_logged_in_read_only_contrib_private_project project_private.add_contributor( - user_two, permissions=permissions.READ, save=True + user_two, + permissions=permissions.READ, + save=True ) res = app.post_json_api( url_private, @@ -1655,10 +1679,16 @@ def project_private(self, user, user_two, user_three, title, description, catego creator=user, ) project_private.add_contributor( - user_two, permissions=permissions.READ, visible=True, save=True + user_two, + permissions=permissions.READ, + visible=True, + save=True ) project_private.add_contributor( - user_three, permissions=permissions.READ, visible=True, save=True + user_three, + permissions=permissions.READ, + visible=True, + save=True ) return project_private @@ -2108,10 +2138,16 @@ def project_private(self, user, user_two, user_three, title, description, catego creator=user, ) project_private.add_contributor( - user_two, permissions=permissions.READ, visible=True, save=True + user_two, + permissions=permissions.READ, + visible=True, + save=True ) project_private.add_contributor( - user_three, permissions=permissions.READ, visible=True, save=True + user_three, + permissions=permissions.READ, + visible=True, + save=True ) return project_private @@ -2483,10 +2519,16 @@ def project_private(self, user, user_two, user_three, title, description, catego creator=user, ) project_private.add_contributor( - user_two, permissions=permissions.READ, visible=True, save=True + user_two, + permissions=permissions.READ, + visible=True, + save=True ) project_private.add_contributor( - user_three, permissions=permissions.READ, visible=True, save=True + user_three, + permissions=permissions.READ, + visible=True, + save=True ) return project_private @@ -2842,8 +2884,15 @@ def test_filtering(self, app, user, url, project): user_two = AuthUserFactory() user_three = AuthUserFactory() - project.add_contributor(user_two, permissions.WRITE) - project.add_contributor(user_three, permissions.READ, visible=False) + project.add_contributor( + user_two, + permissions.WRITE + ) + project.add_contributor( + user_three, + permissions.READ, + visible=False + ) # test_filtering_node_with_only_bibliographic_contributors # no filter diff --git a/api_tests/nodes/views/test_node_detail.py b/api_tests/nodes/views/test_node_detail.py index 901e83b26ef..aa1ce491afe 100644 --- a/api_tests/nodes/views/test_node_detail.py +++ b/api_tests/nodes/views/test_node_detail.py @@ -148,8 +148,7 @@ def test_return_private_project_details_logged_in_write_contributor( project_private.add_contributor( contributor=user_two, auth=Auth(user), - save=True, - notification_type=False + save=True ) res = app.get(url_private, auth=user_two.auth) assert res.status_code == 200 @@ -519,8 +518,7 @@ def test_current_user_permissions(self, app, user, url_public, project_public, u project_public.add_contributor( new_user, permissions=permissions.WRITE, - auth=Auth(project_public.creator), - notification_type=False + auth=Auth(project_public.creator) ) res = app.get(url, auth=new_user.auth) assert res.json['data']['attributes']['current_user_permissions'] == [permissions.WRITE, permissions.READ] diff --git a/api_tests/nodes/views/test_node_detail_delete.py b/api_tests/nodes/views/test_node_detail_delete.py index fa205fb3531..9c2fe078968 100644 --- a/api_tests/nodes/views/test_node_detail_delete.py +++ b/api_tests/nodes/views/test_node_detail_delete.py @@ -71,8 +71,7 @@ def test_deletes_invalid_node( def test_deletes_private_node_logged_in_read_only_contributor(self, app, user_two, project_private, url_private): project_private.add_contributor( user_two, - permissions=permissions.READ, - notification_type=False + permissions=permissions.READ ) project_private.save() res = app.delete( @@ -88,8 +87,7 @@ def test_deletes_private_node_logged_in_read_only_contributor(self, app, user_tw def test_deletes_private_node_logged_in_write_contributor(self, app, user_two, project_private, url_private): project_private.add_contributor( user_two, - permissions=permissions.WRITE, - notification_type=False + permissions=permissions.WRITE ) project_private.save() res = app.delete( diff --git a/api_tests/nodes/views/test_node_detail_license.py b/api_tests/nodes/views/test_node_detail_license.py index c34db2f64c5..e4cd4ff3d8b 100644 --- a/api_tests/nodes/views/test_node_detail_license.py +++ b/api_tests/nodes/views/test_node_detail_license.py @@ -77,14 +77,12 @@ def project_private( project_private.add_contributor( user_admin, permissions=permissions.CREATOR_PERMISSIONS, - save=True, - notification_type=False + save=True ) project_private.add_contributor( user, permissions=permissions.DEFAULT_CONTRIBUTOR_PERMISSIONS, - save=True, - notification_type=False + save=True ) project_private.node_license = NodeLicenseRecordFactory( node_license=node_license, diff --git a/api_tests/nodes/views/test_node_detail_tags.py b/api_tests/nodes/views/test_node_detail_tags.py index 5c0a4a16cd5..e793b6a0f67 100644 --- a/api_tests/nodes/views/test_node_detail_tags.py +++ b/api_tests/nodes/views/test_node_detail_tags.py @@ -51,14 +51,12 @@ def project_private(self, user, user_admin): project_private.add_contributor( user_admin, permissions=permissions.CREATOR_PERMISSIONS, - save=True, - notification_type=False + save=True ) project_private.add_contributor( user, permissions=permissions.DEFAULT_CONTRIBUTOR_PERMISSIONS, - save=True, - notification_type=False + save=True ) # Sets private project storage cache to avoid need for retries in tests updating public status key = cache_settings.STORAGE_USAGE_KEY.format(target_id=project_private._id) diff --git a/api_tests/nodes/views/test_node_detail_update.py b/api_tests/nodes/views/test_node_detail_update.py index 0c6dc09c79f..ca6ee537722 100644 --- a/api_tests/nodes/views/test_node_detail_update.py +++ b/api_tests/nodes/views/test_node_detail_update.py @@ -92,8 +92,7 @@ def test_cannot_make_project_public_if_non_admin_contributor( project_private.add_contributor( non_admin, permissions=permissions.WRITE, - auth=Auth(project_private.creator), - notification_type=False + auth=Auth(project_private.creator) ) project_private.save() res = app.patch_json( @@ -113,8 +112,7 @@ def test_can_make_project_public_if_admin_contributor( project_private.add_contributor( admin_user, permissions=permissions.ADMIN, - auth=Auth(project_private.creator), - notification_type=False + auth=Auth(project_private.creator) ) project_private.save() with capture_notifications(): diff --git a/api_tests/nodes/views/test_node_list.py b/api_tests/nodes/views/test_node_list.py index 35a8ebed143..72cbb883ef1 100644 --- a/api_tests/nodes/views/test_node_list.py +++ b/api_tests/nodes/views/test_node_list.py @@ -1606,7 +1606,12 @@ def test_create_component_with_tags(self, app, user_one, title, category): def test_create_component_inherit_contributors_with_blocked_email( self, app, user_one, title, category): parent_project = ProjectFactory(creator=user_one) - parent_project.add_unregistered_contributor(fullname='far', email='foo@bar.baz', permissions=permissions.READ, auth=Auth(user=user_one)) + parent_project.add_unregistered_contributor( + fullname='far', + email='foo@bar.baz', + permissions=permissions.READ, + auth=Auth(user=user_one) + ) contributor = parent_project.contributors.filter(fullname='far').first() contributor.username = 'foo@example.com' contributor.save() diff --git a/api_tests/nodes/views/test_node_relationship_institutions.py b/api_tests/nodes/views/test_node_relationship_institutions.py index fa0eeca1edb..52a1cf6ee36 100644 --- a/api_tests/nodes/views/test_node_relationship_institutions.py +++ b/api_tests/nodes/views/test_node_relationship_institutions.py @@ -70,8 +70,7 @@ def affiliated_admin(self, node, institution_one): admin.save() node.add_contributor( admin, - permissions=permissions.ADMIN, - notification_type=False + permissions=permissions.ADMIN ) return admin @@ -85,13 +84,11 @@ def node(self, user, write_contrib, read_contrib): project = NodeFactory(creator=user) project.add_contributor( write_contrib, - permissions=permissions.WRITE, - notification_type=False + permissions=permissions.WRITE ) project.add_contributor( read_contrib, - permissions=permissions.READ, - notification_type=False + permissions=permissions.READ ) project.save() return project diff --git a/api_tests/preprints/views/test_preprint_contributors_list.py b/api_tests/preprints/views/test_preprint_contributors_list.py index 39b26063091..0b6d29f2cb6 100644 --- a/api_tests/preprints/views/test_preprint_contributors_list.py +++ b/api_tests/preprints/views/test_preprint_contributors_list.py @@ -3001,8 +3001,7 @@ def test_filtering_node_with_non_bibliographic_contributor( non_bibliographic_contrib = UserFactory() preprint.add_contributor( non_bibliographic_contrib, - visible=False, - notification_type=False + visible=False ) preprint.save() diff --git a/api_tests/preprints/views/test_preprint_subjects_list.py b/api_tests/preprints/views/test_preprint_subjects_list.py index 2cf98dbea34..4603ecfaa90 100644 --- a/api_tests/preprints/views/test_preprint_subjects_list.py +++ b/api_tests/preprints/views/test_preprint_subjects_list.py @@ -16,13 +16,11 @@ def resource(self, user_admin_contrib, user_write_contrib, user_read_contrib): preprint.subjects.clear() preprint.add_contributor( user_write_contrib, - permissions=WRITE, - notification_type=False + permissions=WRITE ) preprint.add_contributor( user_read_contrib, - permissions=READ, - notification_type=False + permissions=READ ) preprint.save() return preprint diff --git a/api_tests/registrations/views/test_registration_list.py b/api_tests/registrations/views/test_registration_list.py index faa78c3a72d..e4677b993cf 100644 --- a/api_tests/registrations/views/test_registration_list.py +++ b/api_tests/registrations/views/test_registration_list.py @@ -1597,7 +1597,8 @@ def test_need_admin_perms_on_draft( user_two = AuthUserFactory() # User is an admin contributor on draft registration but not on node - draft_registration = DraftRegistrationFactory(creator=user_two, registration_schema=schema) + with capture_notifications(): + draft_registration = DraftRegistrationFactory(creator=user_two, registration_schema=schema) draft_registration.add_contributor(user, permissions.ADMIN) draft_registration.branched_from.add_contributor(user, permissions.WRITE) payload_ver['data']['attributes']['draft_registration_id'] = draft_registration._id @@ -1610,7 +1611,8 @@ def test_need_admin_perms_on_draft( assert res.status_code == 201 # User is admin on draft and node - draft_registration = DraftRegistrationFactory(creator=user, registration_schema=schema) + with capture_notifications(): + draft_registration = DraftRegistrationFactory(creator=user, registration_schema=schema) assert draft_registration.branched_from.is_admin_contributor(user) is True assert draft_registration.has_permission(user, permissions.ADMIN) is True payload_ver['data']['attributes']['draft_registration_id'] = draft_registration._id diff --git a/api_tests/registrations/views/test_registration_relationship_institutions.py b/api_tests/registrations/views/test_registration_relationship_institutions.py index 033d59af697..28dfb5b3f61 100644 --- a/api_tests/registrations/views/test_registration_relationship_institutions.py +++ b/api_tests/registrations/views/test_registration_relationship_institutions.py @@ -18,13 +18,11 @@ def node(self, user, write_contrib, read_contrib): registration = RegistrationFactory(creator=user) registration.add_contributor( write_contrib, - permissions=permissions.WRITE, - notification_type=False + permissions=permissions.WRITE ) registration.add_contributor( read_contrib, - permissions=permissions.READ, - notification_type=False + permissions=permissions.READ ) registration.save() return registration diff --git a/api_tests/requests/mixins.py b/api_tests/requests/mixins.py index 3258ee2d51d..3242648a005 100644 --- a/api_tests/requests/mixins.py +++ b/api_tests/requests/mixins.py @@ -135,8 +135,7 @@ def auto_withdrawable_pre_mod_preprint(self, admin, write_contrib, pre_mod_provi pre.add_contributor( contributor=write_contrib, permissions=permissions.WRITE, - save=True, - notification_type=False + save=True ) return pre @@ -151,7 +150,6 @@ def post_mod_preprint(self, admin, write_contrib, post_mod_provider): contributor=write_contrib, permissions=permissions.WRITE, save=True, - notification_type=False ) return post @@ -165,8 +163,7 @@ def none_mod_preprint(self, admin, write_contrib, none_mod_provider): preprint.add_contributor( contributor=write_contrib, permissions=permissions.WRITE, - save=True, - notification_type=False + save=True ) return preprint diff --git a/api_tests/users/views/test_user_claim.py b/api_tests/users/views/test_user_claim.py index 5cbbd5fc6e9..a52574f48f7 100644 --- a/api_tests/users/views/test_user_claim.py +++ b/api_tests/users/views/test_user_claim.py @@ -37,8 +37,7 @@ def unreg_user(self, referrer, project): return project.add_unregistered_contributor( 'David Davidson', 'david@david.son', - auth=Auth(referrer), - notification_type=False + auth=Auth(referrer) ) @pytest.fixture() diff --git a/api_tests/wikis/views/test_wiki_detail.py b/api_tests/wikis/views/test_wiki_detail.py index bde47f758f2..1d2d9ab1c9e 100644 --- a/api_tests/wikis/views/test_wiki_detail.py +++ b/api_tests/wikis/views/test_wiki_detail.py @@ -76,8 +76,7 @@ def user_write_contributor(self, project_public, project_private): user = AuthUserFactory() project_public.add_contributor( user, - permissions=permissions.WRITE, - notification_type=False + permissions=permissions.WRITE ) project_private.add_contributor( user, diff --git a/osf/management/commands/migrate_notifications.py b/osf/management/commands/migrate_notifications.py index c80565c8036..6e1786cd880 100644 --- a/osf/management/commands/migrate_notifications.py +++ b/osf/management/commands/migrate_notifications.py @@ -2,20 +2,27 @@ import time import signal from contextlib import contextmanager + from django.contrib.contenttypes.models import ContentType from django.core.management.base import BaseCommand, CommandError from django.db import transaction + from osf.models import NotificationType, NotificationSubscription from osf.models.notifications import NotificationSubscriptionLegacy from osf.management.commands.populate_notification_types import populate_notification_types +from tqdm import tqdm logger = logging.getLogger(__name__) +TIMEOUT_SECONDS = 3600 # 60 minutes timeout +BATCH_SIZE = 1000 # Default batch size + FREQ_MAP = { 'none': 'none', 'email_digest': 'weekly', 'email_transactional': 'instantly', } + EVENT_NAME_TO_NOTIFICATION_TYPE = { # Provider notifications 'new_pending_withdraw_requests': NotificationType.Type.PROVIDER_NEW_PENDING_WITHDRAW_REQUESTS, @@ -40,10 +47,6 @@ } -TIMEOUT_SECONDS = 3600 # 60 minutes timeout -BATCH_SIZE = 1000 # batch size, can be changed - - @contextmanager def time_limit(seconds): def signal_handler(signum, frame): @@ -56,38 +59,16 @@ def signal_handler(signum, frame): finally: signal.alarm(0) -def migrate_legacy_notification_subscriptions( - dry_run=False, - batch_size=1000, - timeout_per_batch=300, - default_frequency='none' -): - logger.info('Starting legacy notification subscription migration...') +def iter_batches(first_id: int, last_id: int, batch_size: int): + """Yield [start_id, end_id] ranges for batching.""" + for start in range(first_id, last_id + 1, batch_size): + yield start, min(start + batch_size - 1, last_id) - PROVIDER_BASED_LEGACY_NOTIFICATION_TYPES = [ - 'new_pending_submissions', - 'new_pending_withdraw_requests', - 'reviews_submission_confirmation', - 'reviews_moderator_submission_confirmation', - 'reviews_reject_confirmation', - 'reviews_accept_confirmation', - 'reviews_resubmission_confirmation', - 'reviews_comment_edited', - 'contributor_added_preprint', - 'confirm_email_moderation', - 'moderator_added', - 'confirm_email_preprints', - 'user_invite_preprint', - ] - def timeout_handler(signum, frame): - raise TimeoutError('Batch processing timed out') - - # Notification type IDs - notiftype_map = dict(NotificationType.objects.values_list('name', 'id')) - # Cache existing keys - existing_keys = set( +def build_existing_keys(): + """Fetch already migrated subscription keys to prevent duplicates.""" + return set( ( user_id, content_type_id, @@ -95,113 +76,102 @@ def timeout_handler(signum, frame): notification_type_id, ) for user_id, content_type_id, object_id, notification_type_id in - NotificationSubscription.objects.all().values_list( + NotificationSubscription.objects.values_list( 'user_id', 'content_type_id', 'object_id', 'notification_type_id' ) ) - total = NotificationSubscriptionLegacy.objects.count() - created = 0 - skipped = 0 - last_id = 0 + +def migrate_legacy_notification_subscriptions( + dry_run=False, + batch_size=BATCH_SIZE, + default_frequency='none', + start_id=0, +): + logger.info('Starting legacy notification subscription migration...') + + legacy_qs = NotificationSubscriptionLegacy.objects.filter(id__gte=start_id).order_by('id') + total = legacy_qs.count() + if total == 0: + logger.info('No legacy subscriptions to migrate.') + return + + notiftype_map = dict(NotificationType.objects.values_list('name', 'id')) + existing_keys = build_existing_keys() + + created, skipped = 0, 0 + content_type_cache = {} + + first_id, last_id = legacy_qs.first().id, legacy_qs.last().id start_time_total = time.time() - while True: - batch_start_time = time.time() - # Fetch the next chunk directly from DB + for batch_range in tqdm(list(iter_batches(first_id, last_id, batch_size)), desc='Processing', unit='batch'): batch = list( NotificationSubscriptionLegacy.objects - .filter(id__gt=last_id) - .order_by('id')[:batch_size] + .filter(id__range=batch_range) + .order_by('id') + .select_related('provider', 'node', 'user') ) if not batch: - break - subscriptions_to_create = [] + continue - signal.signal(signal.SIGALRM, timeout_handler) - signal.alarm(timeout_per_batch) + subscriptions_to_create = [] - try: - for legacy in batch: - event_name = legacy.event_name - if event_name in PROVIDER_BASED_LEGACY_NOTIFICATION_TYPES: - subscribed_object = legacy.provider - if not subscribed_object: - skipped += 1 - continue - elif legacy.node: - subscribed_object = legacy.node - elif legacy.user: - subscribed_object = legacy.user - else: - skipped += 1 - continue - - content_type = ContentType.objects.get_for_model(subscribed_object.__class__) - notif_enum = EVENT_NAME_TO_NOTIFICATION_TYPE.get(event_name) - if not notif_enum: - skipped += 1 - continue - - notification_type_id = notiftype_map.get(notif_enum) - if not notification_type_id: - skipped += 1 - continue - - key = ( - legacy.user_id, - content_type.id, - int(subscribed_object.id), - notification_type_id, - ) - if key in existing_keys: - skipped += 1 - continue - - frequency = 'weekly' if getattr(legacy, 'email_digest', False) else default_frequency - - if dry_run: - created += 1 - else: - subscriptions_to_create.append(NotificationSubscription( - notification_type_id=notification_type_id, - user_id=legacy.user_id, - content_type=content_type, - object_id=subscribed_object.id, - message_frequency=frequency, - )) - existing_keys.add(key) - - if not dry_run and subscriptions_to_create: - with transaction.atomic(): - NotificationSubscription.objects.bulk_create( - subscriptions_to_create, - ignore_conflicts=True, - ) - created += len(subscriptions_to_create) - - # Logging ETA - batch_time = time.time() - batch_start_time - processed = last_id + len(batch) - rate = processed / (time.time() - start_time_total) - eta = (total - processed) / rate if rate else 0 - - logger.info( - f"Processed batch {last_id}-{last_id + len(batch)} " - f"in {batch_time:.2f}s | " - f"Progress {processed}/{total} ({processed / total:.1%}) | " - f"ETA ~ {eta / 60:.1f} min" + for legacy in batch: + event_name = legacy.event_name + subscribed_object = legacy.provider or legacy.node or legacy.user + if not subscribed_object: + skipped += 1 + continue + + model_class = subscribed_object.__class__ + if model_class not in content_type_cache: + content_type_cache[model_class] = ContentType.objects.get_for_model(model_class) + content_type = content_type_cache[model_class] + + notif_enum = EVENT_NAME_TO_NOTIFICATION_TYPE.get(event_name) + if not notif_enum: + skipped += 1 + continue + + notification_type_id = notiftype_map.get(notif_enum) + if not notification_type_id: + skipped += 1 + continue + + key = ( + legacy.user_id, + content_type.id, + int(subscribed_object.id), + notification_type_id, ) + if key in existing_keys: + skipped += 1 + continue + + frequency = 'weekly' if getattr(legacy, 'email_digest', False) else default_frequency + + if dry_run: + created += 1 + else: + subscriptions_to_create.append(NotificationSubscription( + notification_type_id=notification_type_id, + user_id=legacy.user_id, + content_type=content_type, + object_id=subscribed_object.id, + message_frequency=frequency, + )) + existing_keys.add(key) + + if not dry_run and subscriptions_to_create: + with transaction.atomic(): + NotificationSubscription.objects.bulk_create( + subscriptions_to_create, + ignore_conflicts=True, + ) + created += len(subscriptions_to_create) - except TimeoutError: - logger.error(f"Batch {last_id}-{last_id + len(batch)} timed out, skipping.") - skipped += len(batch) - except Exception as e: - logger.exception(f"Batch {last_id}-{last_id + len(batch)} failed: {e}") - skipped += len(batch) - finally: - signal.alarm(0) - last_id = batch[-1].id + logger.info(f"Processed batch {batch_range[0]}-{batch_range[1]} (Created: {created}, Skipped: {skipped})") elapsed_total = time.time() - start_time_total logger.info( @@ -210,6 +180,16 @@ def timeout_handler(signum, frame): ) +def run_migration(dry_run: bool, batch_size: int, start_id: int): + """Main entry point for command and tests.""" + with time_limit(TIMEOUT_SECONDS): + if not dry_run: + with transaction.atomic(): + logger.info('Populating notification types...') + populate_notification_types(None, {}) + migrate_legacy_notification_subscriptions(dry_run=dry_run, batch_size=batch_size, start_id=start_id) + + class Command(BaseCommand): help = 'Migrate legacy NotificationSubscriptionLegacy objects to new Notification app models.' @@ -220,19 +200,27 @@ def add_arguments(self, parser): 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): - if not dry_run: - with transaction.atomic(): - logger.info('Populating notification types...') - populate_notification_types(args, options) + parser.add_argument( + '--batch-size', + type=int, + default=BATCH_SIZE, + help=f'Batch size (default: {BATCH_SIZE})' + ) - with transaction.atomic(): - migrate_legacy_notification_subscriptions(dry_run=dry_run) + parser.add_argument( + '--start-id', + type=int, + default=0, + help='Start migrating from this ID' + ) + def handle(self, *args, **options): + try: + run_migration( + dry_run=options['dry_run'], + batch_size=options['batch_size'], + start_id=options['start_id'], + ) except TimeoutError: logger.error('Migration timed out. Rolling back changes.') raise CommandError('Migration failed due to timeout') diff --git a/osf_tests/factories.py b/osf_tests/factories.py index b472856e8a4..bc0e9dad8c0 100644 --- a/osf_tests/factories.py +++ b/osf_tests/factories.py @@ -571,8 +571,7 @@ def _create(cls, *args, **kwargs): user=initiator, schema=registration_schema, data=registration_metadata, - provider=provider, - notification_type=False + provider=provider ) if title: draft.title = title diff --git a/osf_tests/management_commands/test_migration_registration_responses.py b/osf_tests/management_commands/test_migration_registration_responses.py index 880590717d1..e1ebab1e658 100644 --- a/osf_tests/management_commands/test_migration_registration_responses.py +++ b/osf_tests/management_commands/test_migration_registration_responses.py @@ -3,6 +3,7 @@ from osf.models import RegistrationSchema from osf_tests.factories import DraftRegistrationFactory, RegistrationFactory +from tests.utils import capture_notifications from osf.management.commands.migrate_registration_responses import ( migrate_draft_registrations, @@ -1603,36 +1604,38 @@ class TestMigrateDraftRegistrationRegistrationResponses: @pytest.fixture() def draft_osf_standard(self, osf_standard_schema): - draft = DraftRegistrationFactory( - registration_schema=osf_standard_schema, - registration_metadata={ - 'looked': { - 'comments': [], - 'value': 'Yes', - 'extra': [] - }, - 'datacompletion': { - 'comments': [], - 'value': 'No, data collection has not begun', - 'extra': [] - }, - 'comments': { - 'comments': [], - 'value': 'more comments', - 'extra': [] + with capture_notifications(): + draft = DraftRegistrationFactory( + registration_schema=osf_standard_schema, + registration_metadata={ + 'looked': { + 'comments': [], + 'value': 'Yes', + 'extra': [] + }, + 'datacompletion': { + 'comments': [], + 'value': 'No, data collection has not begun', + 'extra': [] + }, + 'comments': { + 'comments': [], + 'value': 'more comments', + 'extra': [] + } } - } - ) + ) draft.registration_responses = {} draft.save() return draft @pytest.fixture() def empty_draft_osf_standard(self, osf_standard_schema): - draft = DraftRegistrationFactory( - registration_schema=osf_standard_schema, - registration_metadata={} - ) + with capture_notifications(): + draft = DraftRegistrationFactory( + registration_schema=osf_standard_schema, + registration_metadata={} + ) draft.registration_responses = {} draft.registration_responses_migrated = False draft.save() @@ -1640,10 +1643,11 @@ def empty_draft_osf_standard(self, osf_standard_schema): @pytest.fixture() def draft_prereg(self, prereg_schema): - draft = DraftRegistrationFactory( - registration_schema=prereg_schema, - registration_metadata=prereg_registration_metadata - ) + with capture_notifications(): + draft = DraftRegistrationFactory( + registration_schema=prereg_schema, + registration_metadata=prereg_registration_metadata + ) draft.registration_responses = {} draft.registration_responses_migrated = False draft.save() @@ -1651,10 +1655,11 @@ def draft_prereg(self, prereg_schema): @pytest.fixture() def draft_veer(self, veer_schema): - draft = DraftRegistrationFactory( - registration_schema=veer_schema, - registration_metadata=veer_registration_metadata - ) + with capture_notifications(): + draft = DraftRegistrationFactory( + registration_schema=veer_schema, + registration_metadata=veer_registration_metadata + ) draft.registration_responses = {} draft.registration_responses_migrated = False draft.save() @@ -1856,26 +1861,27 @@ class TestMigrateRegistrationRegistrationResponses: @pytest.fixture() def reg_osf_standard(self, osf_standard_schema): - draft = DraftRegistrationFactory( - registration_schema=osf_standard_schema, - registration_metadata={ - 'looked': { - 'comments': [], - 'value': 'Yes', - 'extra': [] - }, - 'datacompletion': { - 'comments': [], - 'value': 'No, data collection has not begun', - 'extra': [] - }, - 'comments': { - 'comments': [], - 'value': 'more comments', - 'extra': [] + with capture_notifications(): + draft = DraftRegistrationFactory( + registration_schema=osf_standard_schema, + registration_metadata={ + 'looked': { + 'comments': [], + 'value': 'Yes', + 'extra': [] + }, + 'datacompletion': { + 'comments': [], + 'value': 'No, data collection has not begun', + 'extra': [] + }, + 'comments': { + 'comments': [], + 'value': 'more comments', + 'extra': [] + } } - } - ) + ) return RegistrationFactory( schema=osf_standard_schema, draft_registration=draft, @@ -1884,10 +1890,11 @@ def reg_osf_standard(self, osf_standard_schema): @pytest.fixture() def reg_prereg(self, prereg_schema): - draft = DraftRegistrationFactory( - registration_schema=prereg_schema, - registration_metadata=prereg_registration_metadata - ) + with capture_notifications(): + draft = DraftRegistrationFactory( + registration_schema=prereg_schema, + registration_metadata=prereg_registration_metadata + ) return RegistrationFactory( schema=prereg_schema, draft_registration=draft, @@ -1896,10 +1903,11 @@ def reg_prereg(self, prereg_schema): @pytest.fixture() def reg_veer(self, veer_schema): - draft = DraftRegistrationFactory( - registration_metadata=veer_registration_metadata, - registration_schema=veer_schema, - ) + with capture_notifications(): + draft = DraftRegistrationFactory( + registration_metadata=veer_registration_metadata, + registration_schema=veer_schema, + ) return RegistrationFactory( schema=veer_schema, draft_registration=draft, diff --git a/osf_tests/test_draft_registration.py b/osf_tests/test_draft_registration.py index 03691166fa7..bcf98f08f0d 100644 --- a/osf_tests/test_draft_registration.py +++ b/osf_tests/test_draft_registration.py @@ -38,13 +38,15 @@ def draft_registration(project): class TestDraftRegistrations: # copied from tests/test_registrations/test_models.py def test_factory(self): - draft = factories.DraftRegistrationFactory() + with capture_notifications(): + draft = factories.DraftRegistrationFactory() assert draft.branched_from is not None assert draft.initiator is not None assert draft.registration_schema is not None user = factories.UserFactory() - draft = factories.DraftRegistrationFactory(initiator=user) + with capture_notifications(): + draft = factories.DraftRegistrationFactory(initiator=user) assert draft.initiator == user node = factories.ProjectFactory() @@ -55,7 +57,8 @@ def test_factory(self): # Pick an arbitrary v2 schema schema = RegistrationSchema.objects.filter(schema_version=2).first() data = {'some': 'data'} - draft = factories.DraftRegistrationFactory(registration_schema=schema, registration_metadata=data) + with capture_notifications(): + draft = factories.DraftRegistrationFactory(registration_schema=schema, registration_metadata=data) assert draft.registration_schema == schema assert draft.registration_metadata == data @@ -384,8 +387,7 @@ def test_remove_unregistered_conributor_removes_unclaimed_record(self, draft_reg new_user = draft_registration.add_unregistered_contributor( fullname='David Davidson', email='david@davidson.com', - auth=auth, - notification_type=False + auth=auth ) draft_registration.save() assert draft_registration.is_contributor(new_user) # sanity check diff --git a/osf_tests/test_management_commands.py b/osf_tests/test_management_commands.py index 26e34601648..caabe2de267 100644 --- a/osf_tests/test_management_commands.py +++ b/osf_tests/test_management_commands.py @@ -28,6 +28,7 @@ from osf.management.commands.data_storage_usage import ( process_usages, ) +from tests.utils import capture_notifications # Using powers of two so that any combination of file sizes will give a unique total @@ -409,7 +410,8 @@ def active_draft_registration_multiple_contributor(self, project, initiator, dra @pytest.fixture() def no_project_draft_registration(self, initiator): - return DraftRegistrationFactory() + with capture_notifications(): + return DraftRegistrationFactory() def test_draft_reg_to_sync_retrieval( self, app, active_draft_registration, inactive_draft_registration, active_draft_registration_multiple_contributor, no_project_draft_registration): diff --git a/osf_tests/test_migration_sql.py b/osf_tests/test_migration_sql.py index b4b7c770903..2c05f981443 100644 --- a/osf_tests/test_migration_sql.py +++ b/osf_tests/test_migration_sql.py @@ -2,6 +2,7 @@ from django.db import connection import pytest +from tests.utils import capture_notifications from . import factories from osf.models import DraftRegistration from osf.models.registrations import DraftRegistrationGroupObjectPermission @@ -15,7 +16,8 @@ class TestMigrationSQL197: @pytest.mark.django_db def test_remove_draft_auth_groups(self): - draft_reg = factories.DraftRegistrationFactory() + with capture_notifications(): + draft_reg = factories.DraftRegistrationFactory() draft_reg.save() assert (len(draft_reg.group_objects)) with connection.cursor() as cursor: @@ -30,13 +32,15 @@ def test_add_draft_read_write_admin_auth_groups(self): cursor.execute(drop_draft_reg_group_object_permission_table) cursor.execute(remove_draft_auth_groups) cursor.execute(add_draft_read_write_admin_auth_groups) - draft_reg = factories.DraftRegistrationFactory() + with capture_notifications(): + draft_reg = factories.DraftRegistrationFactory() draft_reg.save() assert (len(draft_reg.group_objects)) @pytest.mark.django_db def test_drop_draft_reg_group_object_permission_table(self): - draft_registration = factories.DraftRegistrationFactory() + with capture_notifications(): + draft_registration = factories.DraftRegistrationFactory() draft_registration.save() draft_reg_group_obj_perm = DraftRegistrationGroupObjectPermission.objects.filter(content_object=draft_registration)[0] draft_reg_group_obj_perm_id = draft_reg_group_obj_perm.id @@ -50,6 +54,7 @@ def test_add_permissions_to_draft_registration_groups(self): with connection.cursor() as cursor: cursor.execute(drop_draft_reg_group_object_permission_table) cursor.execute(add_permissions_to_draft_registration_groups) - draft_reg = factories.DraftRegistrationFactory() + with capture_notifications(): + draft_reg = factories.DraftRegistrationFactory() draft_reg.save() assert (DraftRegistrationGroupObjectPermission.objects.filter(content_object=draft_reg).exists()) diff --git a/osf_tests/test_node.py b/osf_tests/test_node.py index 3dc919f81ff..4645f4594e0 100644 --- a/osf_tests/test_node.py +++ b/osf_tests/test_node.py @@ -1734,8 +1734,7 @@ def test_is_contributor_unregistered(self, project, auth): project.add_unregistered_contributor( fullname='David Davidson', email=unreg.username, - auth=auth, - notification_type=False + auth=auth ) project.save() assert project.is_contributor(unreg) is True diff --git a/osf_tests/test_registrations.py b/osf_tests/test_registrations.py index 86e2208a301..c69e511698d 100644 --- a/osf_tests/test_registrations.py +++ b/osf_tests/test_registrations.py @@ -595,17 +595,19 @@ def test_validate_good_doi(self): class TestRegistrationMixin: @pytest.fixture() def draft_prereg(self, prereg_schema): - return factories.DraftRegistrationFactory( - registration_schema=prereg_schema, - registration_metadata={}, - ) + with capture_notifications(): + return factories.DraftRegistrationFactory( + registration_schema=prereg_schema, + registration_metadata={}, + ) @pytest.fixture() def draft_veer(self, veer_schema): - return factories.DraftRegistrationFactory( - registration_schema=veer_schema, - registration_metadata={}, - ) + with capture_notifications(): + return factories.DraftRegistrationFactory( + registration_schema=veer_schema, + registration_metadata={}, + ) @pytest.fixture() def prereg_schema(self): diff --git a/osf_tests/test_user.py b/osf_tests/test_user.py index 6cddff997c0..f7be0b0e3df 100644 --- a/osf_tests/test_user.py +++ b/osf_tests/test_user.py @@ -376,18 +376,21 @@ def test_merge_preprints(self, user): def test_merge_drafts(self, user): user2 = AuthUserFactory() - draft_one = DraftRegistrationFactory(creator=user, title='draft_one') - - draft_two = DraftRegistrationFactory(title='draft_two') + with capture_notifications(): + draft_one = DraftRegistrationFactory(creator=user, title='draft_one') + draft_two = DraftRegistrationFactory(title='draft_two') draft_two.add_contributor(user2) - draft_three = DraftRegistrationFactory(title='draft_three', creator=user2) + with capture_notifications(): + draft_three = DraftRegistrationFactory(title='draft_three', creator=user2) draft_three.add_contributor(user, visible=False) - draft_four = DraftRegistrationFactory(title='draft_four') + with capture_notifications(): + draft_four = DraftRegistrationFactory(title='draft_four') draft_four.add_contributor(user2, permissions=permissions.READ, visible=False) - draft_five = DraftRegistrationFactory(title='draft_five') + with capture_notifications(): + draft_five = DraftRegistrationFactory(title='draft_five') draft_five.add_contributor(user2, permissions=permissions.READ, visible=False) draft_five.add_contributor(user, permissions=permissions.WRITE, visible=True) @@ -742,8 +745,7 @@ def test_display_full_name_unregistered(self): project.add_unregistered_contributor( fullname=name, email=u.username, - auth=Auth(project.creator), - notification_type=False + auth=Auth(project.creator) ) project.save() u.reload() @@ -756,8 +758,7 @@ def test_repeat_add_same_unreg_user_with_diff_name(self): project.add_unregistered_contributor( fullname=old_name, email=unreg_user.username, - auth=Auth(project.creator), - notification_type=False + auth=Auth(project.creator) ) project.save() unreg_user.reload() @@ -771,8 +772,7 @@ def test_repeat_add_same_unreg_user_with_diff_name(self): project.add_unregistered_contributor( fullname=new_name, email=unreg_user.username, - auth=Auth(project.creator), - notification_type=False + auth=Auth(project.creator) ) project.save() unreg_user.reload() @@ -2103,8 +2103,7 @@ def project_user_is_only_admin(self, user): 'lisa', 'lisafrank@cos.io', permissions=permissions.ADMIN, - auth=Auth(user), - notification_type=False + auth=Auth(user) ) project.save() return project diff --git a/scripts/tests/test_fix_registration_unclaimed_records.py b/scripts/tests/test_fix_registration_unclaimed_records.py index 10d6bc065b0..90be96241f3 100644 --- a/scripts/tests/test_fix_registration_unclaimed_records.py +++ b/scripts/tests/test_fix_registration_unclaimed_records.py @@ -30,8 +30,7 @@ def contributor_unregistered(self, user, auth, project): ret = project.add_unregistered_contributor( fullname='Jason Kelece', email='burds@eagles.com', - auth=auth, - notification_type=False + auth=auth ) project.save() return ret @@ -41,8 +40,7 @@ def contributor_unregistered_no_email(self, user, auth, project): ret = project.add_unregistered_contributor( fullname='Big Play Slay', email='', - auth=auth, - notification_type=False + auth=auth ) project.save() return ret diff --git a/tests/test_claim_views.py b/tests/test_claim_views.py index c5a0f3cc5e2..5acbb48cd77 100644 --- a/tests/test_claim_views.py +++ b/tests/test_claim_views.py @@ -54,20 +54,17 @@ def setUp(self): self.project_with_source_tag.add_unregistered_contributor( fullname=self.given_name, email=self.given_email, - auth=Auth(user=self.referrer), - notification_type=False + auth=Auth(user=self.referrer) ) self.preprint_with_source_tag.add_unregistered_contributor( fullname=self.given_name, email=self.given_email, - auth=Auth(user=self.referrer), - notification_type=False + auth=Auth(user=self.referrer) ) self.user = self.project.add_unregistered_contributor( fullname=self.given_name, email=self.given_email, - auth=Auth(user=self.referrer), - notification_type=False + auth=Auth(user=self.referrer) ) self.project.save() @@ -79,8 +76,7 @@ def test_claim_user_already_registered_redirects_to_claim_user_registered(self): unregistered_user = self.project.add_unregistered_contributor( fullname=name, email=None, - auth=Auth(user=self.referrer), - notification_type=False + auth=Auth(user=self.referrer) ) assert unregistered_user in self.project.contributors @@ -128,8 +124,7 @@ def test_claim_user_already_registered_secondary_email_redirects_to_claim_user_r unregistered_user = self.project.add_unregistered_contributor( fullname=name, email=None, - auth=Auth(user=self.referrer), - notification_type=False + auth=Auth(user=self.referrer) ) assert unregistered_user in self.project.contributors @@ -175,8 +170,7 @@ def test_claim_user_invited_with_no_email_posts_to_claim_form(self): invited_user = self.project.add_unregistered_contributor( fullname=given_name, email=None, - auth=Auth(user=self.referrer), - notification_type=False + auth=Auth(user=self.referrer) ) self.project.save() diff --git a/website/maintenance.py b/website/maintenance.py index 2424651d758..98359540cfb 100644 --- a/website/maintenance.py +++ b/website/maintenance.py @@ -42,19 +42,11 @@ def set_maintenance(message, level=1, start=None, end=None): return {'start': state.start, 'end': state.end} - -class InFailedSqlTransaction: - pass - - def get_maintenance(): """Get the current start and end times for the maintenance state. Return None if there is no current maintenance state. """ - try: - maintenance = MaintenanceState.objects.all().first() - except InFailedSqlTransaction: - return None + maintenance = MaintenanceState.objects.all().first() return MaintenanceStateSerializer(maintenance).data if maintenance else None def unset_maintenance(): diff --git a/website/notifications/views.py b/website/notifications/views.py index 45aed44f50e..fa01e283674 100644 --- a/website/notifications/views.py +++ b/website/notifications/views.py @@ -300,7 +300,7 @@ def serialize_event(user, subscription=None, node=None, event_description=None): 'event': { 'title': event_description, 'description': {}[event_type], - 'notificationType': notification_type, + 'notificationType': notification_type.name.title(), 'parent_notification_type': get_parent_notification_type(node, event_type, user) }, 'kind': 'event', @@ -318,14 +318,14 @@ def get_parent_notification_type(node, event, user): :return: str notification type (e.g. 'email_transactional') """ AbstractNode = apps.get_model('osf.AbstractNode') - NotificationSubscriptionLegacy = apps.get_model('osf.NotificationSubscriptionLegacy') + NotificationSubscription = apps.get_model('osf.NotificationSubscription') if node and isinstance(node, AbstractNode) and node.parent_node and node.parent_node.has_permission(user, READ): parent = node.parent_node key = to_subscription_key(parent._id, event) try: - subscription = NotificationSubscriptionLegacy.objects.get(_id=key) - except NotificationSubscriptionLegacy.DoesNotExist: + subscription = NotificationSubscription.objects.get(_id=key) + except NotificationSubscription.DoesNotExist: return get_parent_notification_type(parent, event, user) for notification_type in NOTIFICATION_TYPES: diff --git a/website/project/views/contributor.py b/website/project/views/contributor.py index 562dc97afb0..1bcd6c1b545 100644 --- a/website/project/views/contributor.py +++ b/website/project/views/contributor.py @@ -592,14 +592,15 @@ def check_email_throttle( from osf.models import NotificationSubscription from datetime import timedelta # Check for an active subscription for this contributor and this node - subscription, created = NotificationSubscription.objects.get_or_create( + subscription = NotificationSubscription.objects.filter( user=user, notification_type=notification_type.instance, ) - if created: - return False # No subscription means no previous notifications, so no throttling - # Check the most recent Notification for this subscription - return subscription.created > timezone.now() - timedelta(seconds=throttle) + if not subscription: + return False + else: # No subscription means no previous notifications, so no throttling + # Check the most recent Notification for this subscription + return subscription.created > timezone.now() - timedelta(seconds=throttle) @contributor_added.connect