diff --git a/zerver/lib/fix_unreads.py b/zerver/lib/fix_unreads.py new file mode 100644 index 0000000000..f8ba6891fb --- /dev/null +++ b/zerver/lib/fix_unreads.py @@ -0,0 +1,218 @@ +from __future__ import absolute_import +from __future__ import print_function + +import time + +from typing import Callable, List, TypeVar +from psycopg2.extensions import cursor +CursorObj = TypeVar('CursorObj', bound=cursor) + +from django.db import connection + +from zerver.lib.topic_mutes import build_topic_mute_checker +from zerver.models import UserProfile + +def update_unread_flags(cursor, user_message_ids): + # type: (CursorObj, List[int]) -> None + um_id_list = ', '.join(str(id) for id in user_message_ids) + query = ''' + UPDATE zerver_usermessage + SET flags = flags | 1 + WHERE id IN (%s) + ''' % (um_id_list,) + + cursor.execute(query) + + +def get_timing(message, f): + # type: (str, Callable) -> None + start = time.time() + print(message) + f() + elapsed = time.time() - start + print('elapsed time: %.03f\n' % (elapsed,)) + + +def fix_unsubscribed(cursor, user_profile): + # type: (CursorObj, UserProfile) -> None + + recipient_ids = [] + + def find_recipients(): + # type: () -> None + query = ''' + SELECT + zerver_subscription.recipient_id + FROM + zerver_subscription + INNER JOIN zerver_recipient ON ( + zerver_recipient.id = zerver_subscription.recipient_id + ) + WHERE ( + zerver_subscription.user_profile_id = '%s' AND + zerver_recipient.type = 2 AND + (NOT zerver_subscription.active) + ) + ''' + cursor.execute(query, [user_profile.id]) + rows = cursor.fetchall() + for row in rows: + recipient_ids.append(row[0]) + print(recipient_ids) + + get_timing( + 'get recipients', + find_recipients + ) + + if not recipient_ids: + return + + user_message_ids = [] + + def find(): + # type: () -> None + recips = ', '.join(str(id) for id in recipient_ids) + + query = ''' + SELECT + zerver_usermessage.id + FROM + zerver_usermessage + INNER JOIN zerver_message ON ( + zerver_message.id = zerver_usermessage.message_id + ) + WHERE ( + zerver_usermessage.user_profile_id = %s AND + (zerver_usermessage.flags & 1) = 0 AND + zerver_message.recipient_id in (%s) + ) + ''' % (user_profile.id, recips) + + print(''' + EXPLAIN analyze''' + query.rstrip() + ';') + + cursor.execute(query) + rows = cursor.fetchall() + for row in rows: + user_message_ids.append(row[0]) + print('rows found: %d' % (len(user_message_ids),)) + + get_timing( + 'finding unread messages for non-active streams', + find + ) + + if not user_message_ids: + return + + def fix(): + # type: () -> None + update_unread_flags(cursor, user_message_ids) + + get_timing( + 'fixing unread messages for non-active streams', + fix + ) + +def fix_pre_pointer(cursor, user_profile): + # type: (CursorObj, UserProfile) -> None + + pointer = user_profile.pointer + + if not pointer: + return + + recipient_ids = [] + + def find_non_muted_recipients(): + # type: () -> None + query = ''' + SELECT + zerver_subscription.recipient_id + FROM + zerver_subscription + INNER JOIN zerver_recipient ON ( + zerver_recipient.id = zerver_subscription.recipient_id + ) + WHERE ( + zerver_subscription.user_profile_id = '%s' AND + zerver_recipient.type = 2 AND + zerver_subscription.in_home_view AND + zerver_subscription.active + ) + ''' + cursor.execute(query, [user_profile.id]) + rows = cursor.fetchall() + for row in rows: + recipient_ids.append(row[0]) + print(recipient_ids) + + get_timing( + 'find_non_muted_recipients', + find_non_muted_recipients + ) + + if not recipient_ids: + return + + user_message_ids = [] + + def find_old_ids(): + # type: () -> None + recips = ', '.join(str(id) for id in recipient_ids) + + is_topic_muted = build_topic_mute_checker(user_profile) + + query = ''' + SELECT + zerver_usermessage.id, + zerver_message.recipient_id, + zerver_message.subject + FROM + zerver_usermessage + INNER JOIN zerver_message ON ( + zerver_message.id = zerver_usermessage.message_id + ) + WHERE ( + zerver_usermessage.user_profile_id = %s AND + zerver_usermessage.message_id <= %s AND + (zerver_usermessage.flags & 1) = 0 AND + zerver_message.recipient_id in (%s) + ) + ''' % (user_profile.id, pointer, recips) + + print(''' + EXPLAIN analyze''' + query.rstrip() + ';') + + cursor.execute(query) + rows = cursor.fetchall() + for (um_id, recipient_id, topic) in rows: + if not is_topic_muted(recipient_id, topic): + user_message_ids.append(um_id) + print('rows found: %d' % (len(user_message_ids),)) + + get_timing( + 'finding pre-pointer messages that are not muted', + find_old_ids + ) + + if not user_message_ids: + return + + def fix(): + # type: () -> None + update_unread_flags(cursor, user_message_ids) + + get_timing( + 'fixing unread messages for pre-pointer non-muted messages', + fix + ) + +def fix(user_profile): + # type: (UserProfile) -> None + print('\n---\nFixing %s:' % (user_profile.email,)) + with connection.cursor() as cursor: + fix_unsubscribed(cursor, user_profile) + fix_pre_pointer(cursor, user_profile) + connection.commit() diff --git a/zerver/management/commands/fix_unreads.py b/zerver/management/commands/fix_unreads.py index c4f92aeb2e..445b547b39 100644 --- a/zerver/management/commands/fix_unreads.py +++ b/zerver/management/commands/fix_unreads.py @@ -2,230 +2,20 @@ from __future__ import absolute_import from __future__ import print_function import sys -import time -import ujson -from typing import Any, Callable, Dict, List, Set, Text, TypeVar -from psycopg2.extensions import cursor -CursorObj = TypeVar('CursorObj', bound=cursor) +from typing import Any, List, Text from argparse import ArgumentParser from django.core.management.base import CommandError -from django.db import connection from zerver.lib.management import ZulipBaseCommand -from zerver.lib.topic_mutes import build_topic_mute_checker +from zerver.lib.fix_unreads import fix from zerver.models import ( Realm, UserProfile ) -def update_unread_flags(cursor, user_message_ids): - # type: (CursorObj, List[int]) -> None - um_id_list = ', '.join(str(id) for id in user_message_ids) - query = ''' - UPDATE zerver_usermessage - SET flags = flags | 1 - WHERE id IN (%s) - ''' % (um_id_list,) - - cursor.execute(query) - - -def get_timing(message, f): - # type: (str, Callable) -> None - start = time.time() - print(message) - f() - elapsed = time.time() - start - print('elapsed time: %.03f\n' % (elapsed,)) - - -def fix_unsubscribed(cursor, user_profile): - # type: (CursorObj, UserProfile) -> None - - recipient_ids = [] - - def find_recipients(): - # type: () -> None - query = ''' - SELECT - zerver_subscription.recipient_id - FROM - zerver_subscription - INNER JOIN zerver_recipient ON ( - zerver_recipient.id = zerver_subscription.recipient_id - ) - WHERE ( - zerver_subscription.user_profile_id = '%s' AND - zerver_recipient.type = 2 AND - (NOT zerver_subscription.active) - ) - ''' - cursor.execute(query, [user_profile.id]) - rows = cursor.fetchall() - for row in rows: - recipient_ids.append(row[0]) - print(recipient_ids) - - get_timing( - 'get recipients', - find_recipients - ) - - if not recipient_ids: - return - - user_message_ids = [] - - def find(): - # type: () -> None - recips = ', '.join(str(id) for id in recipient_ids) - - query = ''' - SELECT - zerver_usermessage.id - FROM - zerver_usermessage - INNER JOIN zerver_message ON ( - zerver_message.id = zerver_usermessage.message_id - ) - WHERE ( - zerver_usermessage.user_profile_id = %s AND - (zerver_usermessage.flags & 1) = 0 AND - zerver_message.recipient_id in (%s) - ) - ''' % (user_profile.id, recips) - - print(''' - EXPLAIN analyze''' + query.rstrip() + ';') - - cursor.execute(query) - rows = cursor.fetchall() - for row in rows: - user_message_ids.append(row[0]) - print('rows found: %d' % (len(user_message_ids),)) - - get_timing( - 'finding unread messages for non-active streams', - find - ) - - if not user_message_ids: - return - - def fix(): - # type: () -> None - update_unread_flags(cursor, user_message_ids) - - get_timing( - 'fixing unread messages for non-active streams', - fix - ) - -def fix_pre_pointer(cursor, user_profile): - # type: (CursorObj, UserProfile) -> None - - pointer = user_profile.pointer - - if not pointer: - return - - recipient_ids = [] - - def find_non_muted_recipients(): - # type: () -> None - query = ''' - SELECT - zerver_subscription.recipient_id - FROM - zerver_subscription - INNER JOIN zerver_recipient ON ( - zerver_recipient.id = zerver_subscription.recipient_id - ) - WHERE ( - zerver_subscription.user_profile_id = '%s' AND - zerver_recipient.type = 2 AND - zerver_subscription.in_home_view AND - zerver_subscription.active - ) - ''' - cursor.execute(query, [user_profile.id]) - rows = cursor.fetchall() - for row in rows: - recipient_ids.append(row[0]) - print(recipient_ids) - - get_timing( - 'find_non_muted_recipients', - find_non_muted_recipients - ) - - if not recipient_ids: - return - - user_message_ids = [] - - def find_old_ids(): - # type: () -> None - recips = ', '.join(str(id) for id in recipient_ids) - - is_topic_muted = build_topic_mute_checker(user_profile) - - query = ''' - SELECT - zerver_usermessage.id, - zerver_message.recipient_id, - zerver_message.subject - FROM - zerver_usermessage - INNER JOIN zerver_message ON ( - zerver_message.id = zerver_usermessage.message_id - ) - WHERE ( - zerver_usermessage.user_profile_id = %s AND - zerver_usermessage.message_id <= %s AND - (zerver_usermessage.flags & 1) = 0 AND - zerver_message.recipient_id in (%s) - ) - ''' % (user_profile.id, pointer, recips) - - print(''' - EXPLAIN analyze''' + query.rstrip() + ';') - - cursor.execute(query) - rows = cursor.fetchall() - for (um_id, recipient_id, topic) in rows: - if not is_topic_muted(recipient_id, topic): - user_message_ids.append(um_id) - print('rows found: %d' % (len(user_message_ids),)) - - get_timing( - 'finding pre-pointer messages that are not muted', - find_old_ids - ) - - if not user_message_ids: - return - - def fix(): - # type: () -> None - update_unread_flags(cursor, user_message_ids) - - get_timing( - 'fixing unread messages for pre-pointer non-muted messages', - fix - ) - -def fix(user_profile): - # type: (UserProfile) -> None - print('\n---\nFixing %s:' % (user_profile.email,)) - with connection.cursor() as cursor: - fix_unsubscribed(cursor, user_profile) - fix_pre_pointer(cursor, user_profile) - connection.commit() - class Command(ZulipBaseCommand): help = """Fix problems related to unread counts.""" diff --git a/zerver/tests/test_unread.py b/zerver/tests/test_unread.py index a67929611b..6b4a263f00 100644 --- a/zerver/tests/test_unread.py +++ b/zerver/tests/test_unread.py @@ -1,16 +1,36 @@ # -*- coding: utf-8 -*-AA from __future__ import absolute_import -from typing import Any, Dict, List, Mapping +from typing import Any, Dict, List, Mapping, Text + +from django.db import connection from zerver.models import ( - get_user, Recipient, UserMessage, get_stream, get_realm + get_realm, + get_recipient, + get_stream, + get_user, + Recipient, + Stream, + Subscription, + UserMessage, ) -from zerver.lib.test_helpers import tornado_redirected_to_list +from zerver.lib.fix_unreads import ( + fix, + fix_pre_pointer, + fix_unsubscribed, +) +from zerver.lib.test_helpers import ( + get_subscription, + tornado_redirected_to_list, +) from zerver.lib.test_classes import ( ZulipTestCase, ) +from zerver.lib.topic_mutes import add_topic_mute + +import mock import ujson class PointerTest(ZulipTestCase): @@ -316,3 +336,128 @@ class UnreadCountTests(ZulipTestCase): "topic_name": invalid_topic_name, }) self.assert_json_error(result, 'No such topic \'abc\'') + +class FixUnreadTests(ZulipTestCase): + def test_fix_unreads(self): + # type: () -> None + user = self.example_user('hamlet') + realm = get_realm('zulip') + + def send_message(stream_name, topic_name): + # type: (Text, Text) -> int + msg_id = self.send_message( + self.example_email("othello"), + stream_name, + Recipient.STREAM, + subject=topic_name) + um = UserMessage.objects.get( + user_profile=user, + message_id=msg_id) + return um.id + + def assert_read(user_message_id): + # type: (int) -> None + um = UserMessage.objects.get(id=user_message_id) + self.assertTrue(um.flags.read) + + def assert_unread(user_message_id): + # type: (int) -> None + um = UserMessage.objects.get(id=user_message_id) + self.assertFalse(um.flags.read) + + def mute_stream(stream_name): + # type: (Text) -> None + stream = get_stream(stream_name, realm) + recipient = Recipient.objects.get(type_id=stream.id, type=Recipient.STREAM) + subscription = Subscription.objects.get( + user_profile=user, + recipient=recipient + ) + subscription.in_home_view = False + subscription.save() + + def mute_topic(stream_name, topic_name): + # type: (Text, Text) -> None + stream = get_stream(stream_name, realm) + recipient = get_recipient(Recipient.STREAM, stream.id) + + add_topic_mute( + user_profile=user, + stream_id=stream.id, + recipient_id=recipient.id, + topic_name=topic_name, + ) + + def force_unsubscribe(stream_name): + # type: (Text) -> None + ''' + We don't want side effects here, since the eventual + unsubscribe path may mark messages as read, defeating + the test setup here. + ''' + sub = get_subscription(stream_name, user) + sub.active = False + sub.save() + + # The data setup here is kind of funny, because some of these + # conditions should not actually happen in practice going forward, + # but we may have had bad data from the past. + + mute_stream('Denmark') + mute_topic('Verona', 'muted_topic') + + um_normal_id = send_message('Verona', 'normal') + um_muted_topic_id = send_message('Verona', 'muted_topic') + um_muted_stream_id = send_message('Denmark', 'whatever') + + user.pointer = self.get_last_message().id + user.save() + + um_post_pointer_id = send_message('Verona', 'muted_topic') + + self.subscribe(user, 'temporary') + um_unsubscribed_id = send_message('temporary', 'whatever') + force_unsubscribe('temporary') + + # verify data setup + assert_unread(um_normal_id) + assert_unread(um_muted_topic_id) + assert_unread(um_muted_stream_id) + assert_unread(um_post_pointer_id) + assert_unread(um_unsubscribed_id) + + with connection.cursor() as cursor: + fix_pre_pointer(cursor, user) + + # The only message that should have been fixed is the "normal" + # unumuted message before the pointer. + assert_read(um_normal_id) + + # We don't "fix" any messages that are either muted or after the + # pointer, because they can be legitimately unread. + assert_unread(um_muted_topic_id) + assert_unread(um_muted_stream_id) + assert_unread(um_post_pointer_id) + assert_unread(um_unsubscribed_id) + + # fix unsubscribed + with connection.cursor() as cursor: + fix_unsubscribed(cursor, user) + + # Most messages don't change. + assert_unread(um_muted_topic_id) + assert_unread(um_muted_stream_id) + assert_unread(um_post_pointer_id) + + # The unsubscribed entry should change. + assert_read(um_unsubscribed_id) + + # test idempotency + with mock.patch('zerver.lib.fix_unreads.connection.commit'): + fix(user) + + assert_read(um_normal_id) + assert_unread(um_muted_topic_id) + assert_unread(um_muted_stream_id) + assert_unread(um_post_pointer_id) + assert_read(um_unsubscribed_id)