"""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()