mirror of
https://github.com/zulip/zulip.git
synced 2026-07-18 21:04:19 +08:00
Extract zerver/lib/fix_unreads.py.
This is a pure code move.
This commit is contained in:
parent
848c0803bd
commit
a2fe4178be
218
zerver/lib/fix_unreads.py
Normal file
218
zerver/lib/fix_unreads.py
Normal file
@ -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()
|
||||
@ -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."""
|
||||
|
||||
|
||||
@ -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)
|
||||
|
||||
Loading…
Reference in New Issue
Block a user