151 lines
6.9 KiB
Python
151 lines
6.9 KiB
Python
"""Additional synthetic protocol audit regressions; no native client access."""
|
|
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
from unittest import mock
|
|
|
|
import ai_config
|
|
from protocol_engine import ProtocolBot, model_context_text, session_key
|
|
from test_native_protocol_pipeline import Database, Sender, message, ACCOUNT, CONV
|
|
from wecom_native_sender import DeliveryUnknown
|
|
|
|
|
|
class ProtocolAuditBoundaries(unittest.TestCase):
|
|
def setUp(self):
|
|
self.temp = tempfile.TemporaryDirectory()
|
|
self.addCleanup(self.temp.cleanup)
|
|
self.root = Path(self.temp.name)
|
|
self.enterContext(mock.patch('protocol_engine.application_data_dir', return_value=self.root))
|
|
self.enterContext(mock.patch('queue_log.application_data_dir', return_value=self.root))
|
|
self.db = Database()
|
|
self.sender = Sender(self.db)
|
|
self.generate = mock.Mock(return_value='已收到您的文字问题')
|
|
self.bot = self.make_bot()
|
|
|
|
def make_bot(self, account=ACCOUNT):
|
|
bot = ProtocolBot(db=self.db, sender=self.sender, identity={'accountId': account, 'pid': 1},
|
|
path=self.root/'queue.json', generate=self.generate, log=lambda _:None)
|
|
bot.message_batch_window_seconds = 0
|
|
bot.send_delay_seconds = 0
|
|
bot._needs_review = lambda state: (False, '')
|
|
return bot
|
|
|
|
def test_image_then_related_text_stays_manual(self):
|
|
self.db.messages = [message(1, self=True, text='历史回复')]
|
|
pic = message(2, text='[图片]')
|
|
pic['content_type'] = 3
|
|
self.db.add(pic)
|
|
self.db.add(message(3, text='这张图片应该怎么处理?'))
|
|
self.bot.poll_once()
|
|
self.generate.assert_not_called()
|
|
self.assertEqual(self.sender.calls, [])
|
|
self.assertEqual(self.bot._pending[session_key(ACCOUNT, CONV)]['stage'], 'error')
|
|
|
|
def test_old_media_before_own_reply_does_not_block_new_text(self):
|
|
pic = message(1, text='[图片]')
|
|
pic['content_type'] = 3
|
|
self.db.messages = [pic, message(2, self=True, text='人工已处理图片')]
|
|
self.db.add(message(3, text='现在营业时间是几点?'))
|
|
self.bot.poll_once()
|
|
self.generate.assert_called_once()
|
|
self.assertEqual(len(self.sender.calls), 1)
|
|
|
|
def test_restart_restores_uncertain_receipt_pause(self):
|
|
self.sender.error = DeliveryUnknown('synthetic receipt timeout')
|
|
self.db.add(message(2))
|
|
self.bot.poll_once()
|
|
self.assertTrue(self.bot.needs_attention)
|
|
restarted = self.make_bot()
|
|
self.assertTrue(restarted.needs_attention)
|
|
self.sender.error = None
|
|
self.db.add(message(3, conv='S:100_201'))
|
|
restarted.poll_once()
|
|
self.assertEqual(len(self.sender.calls), 1)
|
|
|
|
def test_uncertain_receipt_for_other_account_does_not_pause_current(self):
|
|
self.sender.error = DeliveryUnknown('synthetic receipt timeout')
|
|
self.db.add(message(2))
|
|
self.bot.poll_once()
|
|
self.assertFalse(self.make_bot(account='101').needs_attention)
|
|
|
|
def test_context_enabled_retains_both_roles_and_unanswered_batch(self):
|
|
messages = [message(i, self=i%2 == 0, text=f'entry-{i}') for i in range(1, 9)]
|
|
messages += [message(9, text='目前的问题'), message(10, text='补充条件')]
|
|
with mock.patch.object(ai_config, 'AI_CONTEXT_ENABLED', True), \
|
|
mock.patch.object(ai_config, 'AI_CONTEXT_MAX_ROUNDS', 2):
|
|
text = model_context_text({'display_name': '客户甲', 'messages': messages})
|
|
self.assertNotIn('entry-6', text)
|
|
self.assertIn('entry-7', text)
|
|
self.assertIn('entry-8', text)
|
|
self.assertIn('目前的问题', text)
|
|
self.assertIn('补充条件', text)
|
|
self.assertIn('我 ', text)
|
|
self.assertIn('客户甲 ', text)
|
|
|
|
def test_unanswered_batch_not_lost_with_short_context_limit(self):
|
|
messages = [message(1, self=True, text='旧回答')]
|
|
messages += [message(i, text=f'问题-{i}') for i in range(2, 20)]
|
|
with mock.patch.object(ai_config, 'AI_CONTEXT_ENABLED', True), \
|
|
mock.patch.object(ai_config, 'AI_CONTEXT_MAX_ROUNDS', 1):
|
|
text = model_context_text({'display_name': '客户甲', 'messages': messages})
|
|
self.assertNotIn('旧回答', text)
|
|
for i in range(2, 20):
|
|
self.assertIn(f'问题-{i}', text)
|
|
|
|
def test_cancelled_context_history_only_keeps_unanswered(self):
|
|
messages = [message(1, text='旧问题'), message(2, self=True, text='旧回复'),
|
|
message(3, text='新问题'), message(4, text='新补充')]
|
|
with mock.patch.object(ai_config, 'AI_CONTEXT_ENABLED', False):
|
|
text = model_context_text({'display_name': '客户甲', 'messages': messages})
|
|
self.assertNotIn('旧问题', text)
|
|
self.assertNotIn('旧回复', text)
|
|
self.assertIn('新问题', text)
|
|
self.assertIn('新补充', text)
|
|
|
|
|
|
def test_overflowing_unanswered_batch_cannot_answer_from_truncated_tail(self):
|
|
self.db.messages = [message(i, text=f'连续问题-{i}') for i in range(1,102)]
|
|
self.db.events = [self.db.messages[-1]]
|
|
self.bot.poll_once()
|
|
self.generate.assert_not_called()
|
|
self.assertEqual(self.sender.calls, [])
|
|
self.assertIn('100', self.bot._pending[session_key(ACCOUNT, CONV)]['last_error'])
|
|
|
|
def test_exactly_100_unanswered_messages_preserve_every_part(self):
|
|
self.db.messages = [message(i, text=f'连续问题-{i}') for i in range(1,101)]
|
|
self.db.events = [self.db.messages[-1]]
|
|
self.bot.poll_once()
|
|
context = self.generate.call_args.args[0]
|
|
self.assertEqual(len(context['messages']), 100)
|
|
self.assertEqual(context['messages'][0]['content'], '连续问题-1')
|
|
self.assertEqual(len(self.sender.calls), 1)
|
|
|
|
def test_own_reply_within_long_history_makes_current_batch_complete(self):
|
|
self.db.messages = [message(i, text=f'历史-{i}') for i in range(1,202)]
|
|
self.db.messages[-3] = message(199, self=True, text='人工已回复')
|
|
self.db.events = [self.db.messages[-1]]
|
|
self.bot.poll_once()
|
|
self.generate.assert_called_once()
|
|
self.assertEqual(len(self.sender.calls), 1)
|
|
|
|
def test_existing_approved_draft_with_media_is_not_sent(self):
|
|
pic = message(2, text='[图片]')
|
|
pic['content_type'] = 3
|
|
latest = message(3, text='请看图')
|
|
self.db.messages = [pic, latest]
|
|
self.bot.message_batch_window_seconds = 60
|
|
self.db.events = [latest]
|
|
self.bot.poll_once()
|
|
state = self.bot._pending[session_key(ACCOUNT, CONV)]
|
|
state.update(reply_text='过时批准草稿', approved=True, awaiting_review=True, ready_at=0)
|
|
self.bot.poll_once()
|
|
self.assertFalse(state['approved'])
|
|
self.assertFalse(state['awaiting_review'])
|
|
self.assertEqual(state['reply_text'], '')
|
|
self.assertEqual(self.sender.calls, [])
|
|
|
|
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main()
|