181 lines
12 KiB
Python
181 lines
12 KiB
Python
"""Offline end-to-end request checks for AI context and provider adaptation."""
|
|
import asyncio
|
|
import copy
|
|
import json
|
|
import sys
|
|
import types
|
|
import unittest
|
|
from unittest import mock
|
|
|
|
import ai_chat
|
|
import model_protocol as mp
|
|
from ai_chat import Provider
|
|
|
|
|
|
class AIContextAuditTest(unittest.TestCase):
|
|
def setUp(self):
|
|
self.stack = __import__('contextlib').ExitStack()
|
|
self.addCleanup(self.stack.close)
|
|
self.stack.enter_context(mock.patch.multiple(
|
|
ai_chat.ai_config, AI_CONTEXT_ENABLED=True, AI_CONTEXT_MAX_ROUNDS=5,
|
|
AI_MCP_ENABLED=False, AI_USE_VISION=False, AI_DEVELOPMENT_MODE=False))
|
|
self.stack.enter_context(mock.patch('requests.sessions.Session.request', side_effect=AssertionError('live HTTP forbidden')))
|
|
self.stack.enter_context(mock.patch('urllib.request.urlopen', side_effect=AssertionError('live HTTP forbidden')))
|
|
self.stack.enter_context(mock.patch('ai_chat._system_prompt', return_value='规则:已发布知识KB-742支持周六服务。'))
|
|
self.history = [
|
|
{'role': 'user', 'content': '我问的是蓝色方案。', 'ts': 1},
|
|
{'role': 'assistant', 'content': '蓝色方案每月三百元。', 'ts': 2},
|
|
]
|
|
|
|
def provider(self, kind='openai', model='selected-model'):
|
|
return Provider(kind=kind, base_url='https://audit.invalid/v1', api_key='fake', model=model)
|
|
|
|
def test_text_request_preserves_roles_order_and_selected_provider(self):
|
|
for kind in ('openai', 'claude', 'gateway'):
|
|
with self.subTest(kind=kind), mock.patch('ai_chat._chat_completion', return_value={'content': '可以办理。'}) as call:
|
|
provider = self.provider(kind)
|
|
ai_chat.call_ai_text('那周六可以办理吗?', self.history, provider)
|
|
messages = call.call_args.args[0]
|
|
self.assertEqual([m['role'] for m in messages], ['system', 'user', 'assistant', 'user'])
|
|
self.assertIn('KB-742', messages[0]['content'])
|
|
self.assertEqual(messages[1:3], [{k: m[k] for k in ('role', 'content')} for m in self.history])
|
|
self.assertIn('周六', messages[-1]['content'])
|
|
self.assertIs(call.call_args.kwargs['provider'], provider)
|
|
|
|
def test_text_history_does_not_mutate_or_leak_between_calls(self):
|
|
before = copy.deepcopy(self.history)
|
|
with mock.patch('ai_chat._chat_completion', return_value={'content': '可以办理。'}) as call:
|
|
ai_chat.call_ai_text('继续', self.history, self.provider())
|
|
ai_chat.call_ai_text('另一位客户', [], self.provider())
|
|
self.assertEqual(self.history, before)
|
|
self.assertNotIn('蓝色方案', str(call.call_args.args[0]))
|
|
|
|
def test_disabled_context_omits_history(self):
|
|
with mock.patch.object(ai_chat.ai_config, 'AI_CONTEXT_ENABLED', False):
|
|
self.assertEqual(ai_chat._history_messages(self.history), [])
|
|
|
|
def test_history_is_limited_to_recent_rounds(self):
|
|
history = [{'role': r, 'content': f'{i}-{r}'} for i in range(20) for r in ('user', 'assistant')]
|
|
result = ai_chat._history_messages(history)
|
|
self.assertEqual(len(result), 10)
|
|
self.assertTrue(result[0]['content'].startswith('15-'))
|
|
|
|
def test_string_round_limit_from_local_settings_is_accepted(self):
|
|
with mock.patch.object(ai_chat.ai_config, 'AI_CONTEXT_MAX_ROUNDS', '1'):
|
|
self.assertEqual(ai_chat._history_messages(self.history * 3), self.history_as_messages())
|
|
|
|
def history_as_messages(self):
|
|
return [{k: m[k] for k in ('role', 'content')} for m in self.history]
|
|
|
|
def test_invalid_round_limits_are_bounded_and_do_not_crash(self):
|
|
for value in (None, 'invalid', 0, -2, 1000000):
|
|
with self.subTest(value=value), mock.patch.object(ai_chat.ai_config, 'AI_CONTEXT_MAX_ROUNDS', value):
|
|
messages = ai_chat._history_messages(self.history * 100)
|
|
self.assertGreater(len(messages), 0)
|
|
self.assertLessEqual(len(messages), 100)
|
|
|
|
def test_malformed_history_rows_do_not_break_reply(self):
|
|
result = ai_chat._history_messages([None, 'bad', {'role': 'system', 'content': 'untrusted'}, *self.history])
|
|
self.assertEqual(result, self.history_as_messages())
|
|
|
|
def test_dify_direct_text_keeps_history(self):
|
|
with mock.patch('ai_chat._call_dify', return_value='可以办理。') as call:
|
|
ai_chat.call_ai_text('那周六呢?', self.history, self.provider('dify'))
|
|
self.assertIn('蓝色方案', call.call_args.args[0])
|
|
self.assertIn('每月三百元', call.call_args.args[0])
|
|
self.assertIn('那周六呢', call.call_args.args[0])
|
|
|
|
def test_dify_direct_vision_keeps_history_and_image(self):
|
|
with mock.patch('ai_chat._call_dify_with_image', return_value='{}') as call, mock.patch('ai_chat._finalize_vision_reply', return_value='reply'):
|
|
ai_chat.call_ai_vision(b'audit-image', self.history, '这个方案呢?', provider=self.provider('dify'))
|
|
self.assertIn('蓝色方案', call.call_args.args[0])
|
|
self.assertEqual(call.call_args.args[1], b'audit-image')
|
|
|
|
def test_vision_payload_keeps_history_for_openai_claude_gateway(self):
|
|
for kind in ('openai', 'claude', 'gateway'):
|
|
with self.subTest(kind=kind):
|
|
response = mock.Mock()
|
|
response.json.return_value = {'choices': [{'message': {'content': '{}'}}]}
|
|
with mock.patch('ai_chat._post_with_retry', return_value=response) as post, mock.patch('ai_chat._claude_completion', return_value={'content': '{}'}) as claude, mock.patch('ai_chat._gateway_completion', return_value={'content': '{}'}) as gateway, mock.patch('ai_chat._finalize_vision_reply', return_value='reply'):
|
|
ai_chat.call_ai_vision(b'audit-image', self.history, '这个方案呢?', provider=self.provider(kind))
|
|
messages = post.call_args.kwargs['json']['messages'] if kind == 'openai' else (claude if kind == 'claude' else gateway).call_args.args[0]
|
|
self.assertEqual(messages[1:3], self.history_as_messages())
|
|
self.assertEqual(messages[-1]['content'][1]['type'], 'image_url')
|
|
|
|
def test_dify_shared_payload_preserves_system_knowledge_and_turns(self):
|
|
messages = [{'role': 'system', 'content': '参考知识KB-742:周六开放。'}, *self.history_as_messages(), {'role': 'user', 'content': '那周六呢?'}]
|
|
body = mp.chat_payload('dify', model='', messages=messages, max_tokens=100, temperature=.2)
|
|
query = body['query']
|
|
for text in ('KB-742', '蓝色方案', '每月三百元', '那周六呢'):
|
|
self.assertIn(text, query)
|
|
self.assertLess(query.index('我问的是'), query.index('每月三百元'))
|
|
self.assertLess(query.index('每月三百元'), query.index('那周六呢'))
|
|
self.assertNotIn('conversation_id', body)
|
|
|
|
def test_dify_generic_completion_preserves_judge_system_and_history(self):
|
|
messages = [{'role': 'system', 'content': '裁判必须输出JSON,参考知识KB-742。'}, *self.history_as_messages(), {'role': 'user', 'content': '那周六呢?'}]
|
|
provider = self.provider('dify')
|
|
with mock.patch('ai_chat._call_dify', return_value='{"score": 90}') as call:
|
|
ai_chat._chat_completion(messages, provider=provider)
|
|
self.assertIn('必须输出JSON', call.call_args.args[0])
|
|
self.assertIn('KB-742', call.call_args.args[0])
|
|
self.assertIn('蓝色方案', call.call_args.args[0])
|
|
self.assertIs(call.call_args.kwargs['provider'], provider)
|
|
|
|
def test_dify_single_user_query_stays_unchanged(self):
|
|
body = mp.chat_payload('dify', model='', messages=[{'role': 'user', 'content': '请报价'}], max_tokens=100, temperature=.2)
|
|
self.assertEqual(body['query'], '请报价')
|
|
|
|
def test_dify_multimodal_omits_image_bytes_but_preserves_files(self):
|
|
files = [{'type': 'image', 'transfer_method': 'local_file', 'upload_file_id': 'audit'}]
|
|
messages = [{'role': 'system', 'content': 'KB-742'}, {'role': 'user', 'content': [{'type': 'text', 'text': '这张图是什么?'}, {'type': 'image_url', 'image_url': {'url': 'data:image/png;base64,PRIVATE_IMAGE'}}]}]
|
|
body = mp.chat_payload('dify', model='', messages=messages, max_tokens=100, temperature=.2, dify_files=files)
|
|
self.assertIn('KB-742', body['query'])
|
|
self.assertIn('这张图', body['query'])
|
|
self.assertNotIn('PRIVATE_IMAGE', body['query'])
|
|
self.assertEqual(body['files'], files)
|
|
|
|
def test_gateway_customer_text_does_not_contain_image_data(self):
|
|
content = [{'type': 'text', 'text': '周六是否开放?'}, {'type': 'image_url', 'image_url': {'url': 'data:image/png;base64,PRIVATE_IMAGE'}}]
|
|
response = mock.MagicMock()
|
|
response.__enter__.return_value.read.return_value = json.dumps({'reply': '开放。', 'knowledge': {'used': True}}).encode()
|
|
with mock.patch('urllib.request.urlopen', return_value=response) as send, mock.patch('backend_client.device_id', return_value='audit-device'):
|
|
ai_chat._gateway_completion([{'role': 'user', 'content': content}], None, self.provider('gateway'))
|
|
body = json.loads(send.call_args.args[0].data)
|
|
self.assertEqual(body['customer_text'], '周六是否开放?')
|
|
self.assertEqual(body['messages'][-1]['content'], content)
|
|
self.assertEqual(ai_chat.take_last_gateway_trace()['knowledge'], {'used': True})
|
|
|
|
def _fake_mcp(self, tools):
|
|
class Hub:
|
|
tool_count = len(tools)
|
|
server_names = ['audit']
|
|
async def __aenter__(self): return self
|
|
async def __aexit__(self, *args): return False
|
|
def openai_tools(self): return tools
|
|
async def call_tool(self, name, arguments): return '已查询:周六开放。'
|
|
return types.SimpleNamespace(McpHub=Hub, run_coro=asyncio.run)
|
|
|
|
def test_mcp_without_tools_preserves_explicit_provider(self):
|
|
provider = self.provider(model='task-selected')
|
|
with mock.patch.object(ai_chat.ai_config, 'AI_MCP_ENABLED', True), mock.patch.dict(sys.modules, {'mcp_bridge': self._fake_mcp([])}), mock.patch('ai_chat._chat_completion', return_value={'content': '可以办理。'}) as call:
|
|
ai_chat.call_ai_text('那周六呢?', self.history, provider)
|
|
self.assertIs(call.call_args.kwargs.get('provider'), provider)
|
|
self.assertEqual(call.call_args.args[0][1:3], self.history_as_messages())
|
|
|
|
def test_mcp_tool_rounds_preserve_provider_history_and_tool_result(self):
|
|
provider = self.provider(model='task-selected')
|
|
tool = {'type': 'function', 'function': {'name': 'audit_lookup', 'parameters': {'type': 'object'}}}
|
|
responses = [{'content': '', 'tool_calls': [{'id': 'audit-call', 'type': 'function', 'function': {'name': 'audit_lookup', 'arguments': '{}'}}]}, {'content': '周六开放。'}]
|
|
with mock.patch.object(ai_chat.ai_config, 'AI_MCP_ENABLED', True), mock.patch.dict(sys.modules, {'mcp_bridge': self._fake_mcp([tool])}), mock.patch('ai_chat._chat_completion', side_effect=responses) as call:
|
|
ai_chat.call_ai_text('那周六呢?', self.history, provider)
|
|
self.assertEqual(call.call_count, 2)
|
|
for request in call.call_args_list:
|
|
self.assertIs(request.kwargs.get('provider'), provider)
|
|
final_messages = call.call_args.args[0]
|
|
self.assertEqual(final_messages[-1], {'role': 'tool', 'tool_call_id': 'audit-call', 'content': '已查询:周六开放。'})
|
|
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main()
|