diff --git a/zerver/lib/soft_deactivation.py b/zerver/lib/soft_deactivation.py index b9ff4ac8d4..84112aa6cf 100644 --- a/zerver/lib/soft_deactivation.py +++ b/zerver/lib/soft_deactivation.py @@ -1,12 +1,127 @@ from __future__ import absolute_import +from collections import defaultdict from django.db import transaction - -from zerver.models import UserProfile, UserMessage, RealmAuditLog - from django.utils.timezone import now as timezone_now +from typing import DefaultDict, List -from typing import List +from zerver.models import UserProfile, UserMessage, RealmAuditLog, \ + Subscription, Message, Recipient + +def find_and_store_to_insert_stream_msgs(user_profile, + all_stream_messages, + all_stream_subscription_logs, + all_messages_to_insert): + # type: (UserProfile, DefaultDict[int, List[Message]], DefaultDict[int, List[RealmAuditLog]], List[UserMessage]) -> None + def store_user_message_to_insert(message): + # type: (Message) -> None + message = UserMessage(user_profile=user_profile, + message_id=message['id'], flags=0) + all_messages_to_insert.append(message) + + for (stream_id, stream_messages) in all_stream_messages.items(): + stream_subscription_logs = all_stream_subscription_logs[stream_id] + + for log_entry in stream_subscription_logs: + if len(stream_messages) == 0: + continue + if log_entry.event_type == 'subscription_deactivated': + for stream_message in stream_messages: + if stream_message['id'] <= log_entry.event_last_message_id: + store_user_message_to_insert(stream_message) + else: + break + elif log_entry.event_type in ('subscription_activated', + 'subscription_created'): + initial_msg_count = len(stream_messages) + for i, stream_message in enumerate(stream_messages): + if stream_message['id'] > log_entry.event_last_message_id: + stream_messages = stream_messages[i:] + break + final_msg_count = len(stream_messages) + if initial_msg_count == final_msg_count: + if stream_messages[-1]['id'] <= log_entry.event_last_message_id: + stream_messages = [] + else: + raise AssertionError('%s is not a Subscription Event.' % (log_entry.event_type)) + + if len(stream_messages) > 0: + # We do this check for last event since if the last subscription + # event was a subscription_deactivated then we don't want to create + # UserMessage rows for any of the remaining messages. + if stream_subscription_logs[-1].event_type in ( + 'subscription_activated', + 'subscription_created'): + for stream_message in stream_messages: + store_user_message_to_insert(stream_message) + +def add_missing_messages(user_profile): + # type: (UserProfile) -> None + # This list will store all the messages for which we eventually will create + # UserMessage table rows by doing a bulk insert. + all_messages_to_insert = [] # type: List[UserMessage] + + all_stream_subs = list(Subscription.objects.select_related('recipient').filter( + user_profile=user_profile, + recipient__type=Recipient.STREAM).values('recipient', 'recipient__type_id')) + + # For Stream messages we need to check messages against data from + # RealmAuditLog for visibility to user. So we fetch the subscription logs. + stream_ids = [sub['recipient__type_id'] for sub in all_stream_subs] + events = ['subscription_created', 'subscription_deactivated', 'subscription_activated'] + subscription_logs = list(RealmAuditLog.objects.select_related( + 'modified_stream').filter( + modified_user=user_profile, + modified_stream__id__in=stream_ids, + event_type__in=events).order_by('event_last_message_id')) + + all_stream_subscription_logs = defaultdict(list) # type: DefaultDict[int, List] + for log in subscription_logs: + all_stream_subscription_logs[log.modified_stream.id].append(log) + + recipient_ids = [] + for sub in all_stream_subs: + stream_subscription_logs = all_stream_subscription_logs[sub['recipient__type_id']] + if (stream_subscription_logs[-1].event_type == 'subscription_deactivated' and + stream_subscription_logs[-1].event_last_message_id < user_profile.last_active_message_id): + # We are going to short circuit this iteration as its no use + # iterating since user unsubscribed before soft-deactivation + continue + recipient_ids.append(sub['recipient']) + + all_stream_msgs = list(Message.objects.select_related( + 'recipient').filter( + recipient__id__in=recipient_ids, + id__gt=user_profile.last_active_message_id).order_by('id').values( + 'id', 'recipient__type_id')) + already_created_um_objs = list(UserMessage.objects.select_related( + 'message').filter( + user_profile=user_profile, + message__recipient__type=Recipient.STREAM, + message__id__gt=user_profile.last_active_message_id).values( + 'message__id')) + already_created_ums = set([obj['message__id'] for obj in already_created_um_objs]) + + # Filter those messages for which UserMessage rows have been already created + all_stream_msgs = [msg for msg in all_stream_msgs + if msg['id'] not in already_created_ums] + + stream_messages = defaultdict(list) # type: DefaultDict[int, List] + for msg in all_stream_msgs: + stream_messages[msg['recipient__type_id']].append(msg) + + # Calling this function to filter out stream messages based upon + # subscription logs and then store all UserMessage objects for bulk insert + # This function does not perform any SQL related task and gets all the data + # required for its operation in its params. + find_and_store_to_insert_stream_msgs(user_profile, + stream_messages, + all_stream_subscription_logs, + all_messages_to_insert) + + # Doing a bulk create for all the UserMessage objects stored for creation. + if all_messages_to_insert: + UserMessage.objects.bulk_create(all_messages_to_insert) def do_soft_deactivate_user(user_profile): # type: (UserProfile) -> None diff --git a/zerver/tests/test_messages.py b/zerver/tests/test_messages.py index c1b3572f8d..7b825c17f1 100644 --- a/zerver/tests/test_messages.py +++ b/zerver/tests/test_messages.py @@ -1,6 +1,6 @@ # -*- coding: utf-8 -*- from __future__ import absolute_import -from django.db.models import Q +from django.db.models import Q, Max from django.conf import settings from django.http import HttpResponse from django.test import TestCase, override_settings @@ -29,11 +29,14 @@ from zerver.lib.test_classes import ( ZulipTestCase, ) +from zerver.lib.soft_deactivation import add_missing_messages, do_soft_deactivate_users + from zerver.models import ( MAX_MESSAGE_LENGTH, MAX_SUBJECT_LENGTH, - Message, Realm, Recipient, Stream, UserMessage, UserProfile, Attachment, RealmDomain, - get_realm, get_realm_by_email_domain, get_stream, get_system_bot, get_user, - Reaction, sew_messages_and_reactions, flush_per_request_caches + Message, Realm, Recipient, Stream, UserMessage, UserProfile, Attachment, + RealmAuditLog, RealmDomain, get_realm, get_realm_by_email_domain, + get_stream, get_recipient, get_system_bot, get_user, Reaction, + sew_messages_and_reactions, flush_per_request_caches ) from zerver.lib.actions import ( @@ -44,7 +47,6 @@ from zerver.lib.actions import ( extract_recipients, do_create_user, get_client, - get_recipient, ) from zerver.lib.upload import create_attachment @@ -2137,3 +2139,140 @@ class DeleteMessageTest(ZulipTestCase): self.client_delete('/json/messages/{msg_id}'.format(msg_id=msg_id)) result = self.client_delete('/json/messages/{msg_id}'.format(msg_id=msg_id)) self.assert_json_error(result, "Invalid message(s)") + +class SoftDeactivationMessageTest(ZulipTestCase): + + def test_add_missing_messages(self): + # type: () -> None + recipient_list = [self.example_email("hamlet"), self.example_email("iago")] + for email in recipient_list: + self.subscribe_to_stream(email, "Denmark") + + sender = self.example_user('iago') + realm = sender.realm + sending_client = make_client(name="test suite") + stream_name = 'Denmark' + stream = get_stream(stream_name, realm) + subject = 'foo' + + def send_fake_message(message_content, stream): + # type: (str, Stream) -> Message + recipient = get_recipient(Recipient.STREAM, stream.id) + return Message.objects.create(sender = sender, + recipient = recipient, + subject = subject, + content = message_content, + pub_date = timezone_now(), + sending_client = sending_client) + + long_term_idle_user = self.example_user('hamlet') + do_soft_deactivate_users([long_term_idle_user]) + + # Test that add_missing_messages() in simplest case of adding a + # message for which UserMessage row doesn't exist for this user. + sent_message = send_fake_message('Test Message 1', stream) + idle_user_msg_list = get_user_messages(long_term_idle_user) + idle_user_msg_count = len(idle_user_msg_list) + self.assertNotEqual(idle_user_msg_list[-1], sent_message) + with queries_captured() as queries: + add_missing_messages(long_term_idle_user) + self.assert_length(queries, 5) + idle_user_msg_list = get_user_messages(long_term_idle_user) + self.assertEqual(len(idle_user_msg_list), idle_user_msg_count + 1) + self.assertEqual(idle_user_msg_list[-1], sent_message) + + # Test that add_missing_messages() only adds messages that aren't + # already present in the UserMessage table. This test works on the + # fact that previous test just above this added a message but didn't + # updated the last_active_message_id field for the user. + sent_message = send_fake_message('Test Message 2', stream) + idle_user_msg_list = get_user_messages(long_term_idle_user) + idle_user_msg_count = len(idle_user_msg_list) + self.assertNotEqual(idle_user_msg_list[-1], sent_message) + with queries_captured() as queries: + add_missing_messages(long_term_idle_user) + self.assert_length(queries, 5) + idle_user_msg_list = get_user_messages(long_term_idle_user) + self.assertEqual(len(idle_user_msg_list), idle_user_msg_count + 1) + self.assertEqual(idle_user_msg_list[-1], sent_message) + + # Test UserMessage rows are created correctly in case of stream + # Subscription was altered by admin while user was away. + + # Test for a public stream. + sent_message_list = [] + sent_message_list.append(send_fake_message('Test Message 3', stream)) + # Alter subscription to stream. + self.unsubscribe_from_stream(long_term_idle_user.email, stream_name, realm) + send_fake_message('Test Message 4', stream) + self.subscribe_to_stream(long_term_idle_user.email, stream_name, realm) + sent_message_list.append(send_fake_message('Test Message 5', stream)) + sent_message_list.reverse() + idle_user_msg_list = get_user_messages(long_term_idle_user) + idle_user_msg_count = len(idle_user_msg_list) + for sent_message in sent_message_list: + self.assertNotEqual(idle_user_msg_list.pop(), sent_message) + with queries_captured() as queries: + add_missing_messages(long_term_idle_user) + self.assert_length(queries, 5) + idle_user_msg_list = get_user_messages(long_term_idle_user) + self.assertEqual(len(idle_user_msg_list), idle_user_msg_count + 2) + for sent_message in sent_message_list: + self.assertEqual(idle_user_msg_list.pop(), sent_message) + + # Test consecutive subscribe/unsubscribe in a public stream + sent_message_list = [] + + sent_message_list.append(send_fake_message('Test Message 6', stream)) + # Unsubscribe from stream and then immediately subscribe back again. + self.unsubscribe_from_stream(long_term_idle_user.email, stream_name, realm) + self.subscribe_to_stream(long_term_idle_user.email, stream_name, realm) + sent_message_list.append(send_fake_message('Test Message 7', stream)) + # Again unsubscribe from stream and send a message. + # This will make sure that if initially in a unsubscribed state + # a consecutive subscribe/unsubscribe doesn't misbehave. + self.unsubscribe_from_stream(long_term_idle_user.email, stream_name, realm) + send_fake_message('Test Message 8', stream) + # Do a subscribe and unsubscribe immediately. + self.subscribe_to_stream(long_term_idle_user.email, stream_name, realm) + self.unsubscribe_from_stream(long_term_idle_user.email, stream_name, realm) + + sent_message_list.reverse() + idle_user_msg_list = get_user_messages(long_term_idle_user) + idle_user_msg_count = len(idle_user_msg_list) + for sent_message in sent_message_list: + self.assertNotEqual(idle_user_msg_list.pop(), sent_message) + with queries_captured() as queries: + add_missing_messages(long_term_idle_user) + self.assert_length(queries, 5) + idle_user_msg_list = get_user_messages(long_term_idle_user) + self.assertEqual(len(idle_user_msg_list), idle_user_msg_count + 2) + for sent_message in sent_message_list: + self.assertEqual(idle_user_msg_list.pop(), sent_message) + # Note: At this point in this test we have long_term_idle_user + # unsubscribed from the 'Denmark' stream. + + # Test for a Private Stream. + stream_name = "Core" + private_stream = self.make_stream('Core', invite_only=True) + self.subscribe_to_stream(self.example_email("iago"), stream_name) + sent_message_list = [] + send_fake_message('Test Message 9', private_stream) + self.subscribe_to_stream(self.example_email("hamlet"), stream_name) + sent_message_list.append(send_fake_message('Test Message 10', private_stream)) + self.unsubscribe_from_stream(long_term_idle_user.email, stream_name, realm) + send_fake_message('Test Message 11', private_stream) + self.subscribe_to_stream(long_term_idle_user.email, stream_name, realm) + sent_message_list.append(send_fake_message('Test Message 12', private_stream)) + sent_message_list.reverse() + idle_user_msg_list = get_user_messages(long_term_idle_user) + idle_user_msg_count = len(idle_user_msg_list) + for sent_message in sent_message_list: + self.assertNotEqual(idle_user_msg_list.pop(), sent_message) + with queries_captured() as queries: + add_missing_messages(long_term_idle_user) + self.assert_length(queries, 5) + idle_user_msg_list = get_user_messages(long_term_idle_user) + self.assertEqual(len(idle_user_msg_list), idle_user_msg_count + 2) + for sent_message in sent_message_list: + self.assertEqual(idle_user_msg_list.pop(), sent_message)