61 lines
2.9 KiB
Python
61 lines
2.9 KiB
Python
"""Malformed DB timestamps must not crash receive or context construction."""
|
|
from pathlib import Path
|
|
from contextlib import closing
|
|
import sqlite3
|
|
import tempfile
|
|
import unittest
|
|
|
|
from wxwork_db import WXWorkDB
|
|
|
|
|
|
class DatabaseTimestampAudit(unittest.TestCase):
|
|
def setUp(self):
|
|
self.temp = tempfile.TemporaryDirectory()
|
|
self.addCleanup(self.temp.cleanup)
|
|
self.root = Path(self.temp.name)
|
|
data = self.root / 'WXWork' / '100' / 'Data'
|
|
data.mkdir(parents=True)
|
|
with closing(sqlite3.connect(data/'user.db')) as conn, conn:
|
|
conn.execute('CREATE TABLE user_table(id TEXT,name TEXT,real_name TEXT,account TEXT)')
|
|
conn.execute("INSERT INTO user_table VALUES('200','客户甲','','')")
|
|
self.message_path = data / 'message.db'
|
|
with closing(sqlite3.connect(self.message_path)) as conn, conn:
|
|
conn.execute('CREATE TABLE message_table(sender_id TEXT,conversation_id TEXT,content_type INT,send_time,content TEXT,server_id TEXT)')
|
|
self.db = WXWorkDB(str(self.root/'WXWork'), {}, str(self.root/'cache'))
|
|
self.addCleanup(self.db.close)
|
|
|
|
def row(self, stamp):
|
|
return {'sender_id':'200', 'conversation_id':'M:200', 'content_type':2,
|
|
'send_time':stamp, 'content':'合成消息', 'server_id':'1', '__rowid':1}
|
|
|
|
def test_nonfinite_timestamps_are_not_messages(self):
|
|
for value in (float('nan'), float('inf'), float('-inf'), 'NaN', 'Infinity'):
|
|
with self.subTest(value=str(value)):
|
|
self.assertIsNone(self.db._parse_message('100', self.row(value)))
|
|
|
|
def test_out_of_range_timestamps_are_not_messages(self):
|
|
for value in (1e300, -1e300, '1e300', 10**200):
|
|
with self.subTest(value=str(value)):
|
|
self.assertIsNone(self.db._parse_message('100', self.row(value)))
|
|
|
|
def test_seconds_and_milliseconds_preserve_chronology(self):
|
|
seconds = 1789520400
|
|
for value in (seconds, seconds*1000, str(seconds)):
|
|
with self.subTest(value=value):
|
|
self.assertEqual(self.db._parse_message('100', self.row(value))['send_time'], seconds)
|
|
|
|
def test_bad_rows_do_not_crash_context_or_skip_later_valid_messages(self):
|
|
with closing(sqlite3.connect(self.message_path)) as conn, conn:
|
|
for i, value in enumerate((1789520400, float('inf'), 'NaN', 1e300, 1789520401),1):
|
|
conn.execute("INSERT INTO message_table VALUES('200','M:200',2,?,?,?)",
|
|
(value, f'合成消息-{i}', str(i)))
|
|
messages = self.db.get_new_messages(0)
|
|
self.assertEqual([m['content'] for m in messages], ['合成消息-1', '合成消息-5'])
|
|
context = self.db.get_conversation_context('客户甲')
|
|
self.assertEqual([m['content'] for m in context['messages']], ['合成消息-1', '合成消息-5'])
|
|
self.assertEqual(context['last_message']['server_id'], '5')
|
|
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main()
|