239 lines
10 KiB
Python
239 lines
10 KiB
Python
import unittest
|
|
from types import SimpleNamespace
|
|
from unittest.mock import Mock
|
|
|
|
from fastapi import HTTPException
|
|
|
|
from models import Avatar, User
|
|
from routers.chat import (
|
|
_build_prompt,
|
|
_iter_text_chunks,
|
|
_match_standard_qa,
|
|
_public_avatar_payload,
|
|
_qa_requires_language_adaptation,
|
|
_qa_requires_per_turn_rendering,
|
|
_require_owned_avatar,
|
|
_resolve_reply,
|
|
)
|
|
|
|
|
|
class ChatOrchestrationTests(unittest.TestCase):
|
|
def setUp(self):
|
|
self.avatar = SimpleNamespace(
|
|
id="avatar-1",
|
|
owner_id="huihui-user-1",
|
|
name="冯医生",
|
|
display_name="冯医生",
|
|
description="耳鼻喉科领域专家",
|
|
photo_url="https://example.test/avatar.png",
|
|
emoji="👨⚕️",
|
|
status="active",
|
|
config={
|
|
"replyStyle": "professional",
|
|
"creativity": 50,
|
|
"rigor": 80,
|
|
"humor": 20,
|
|
"responseLength": "medium",
|
|
"systemPrompt": "不要编造政策。",
|
|
"profession": "医生",
|
|
"position": "主任医师",
|
|
"organization": "测试医院",
|
|
"organizationAddress": "测试路1号",
|
|
},
|
|
)
|
|
self.qa = SimpleNamespace(question="公司地址?", answer="标准地址", enabled=True)
|
|
self.disabled_qa = SimpleNamespace(question="公司地址?", answer="错误答案", enabled=False)
|
|
|
|
def test_enabled_qa_wins_without_calling_model(self):
|
|
fake_model = Mock()
|
|
result = _resolve_reply(
|
|
None,
|
|
self.avatar,
|
|
" 公司地址? ",
|
|
[],
|
|
qa_pairs=[self.disabled_qa, self.qa],
|
|
search_fn=lambda *_args, **_kwargs: [],
|
|
model_client=fake_model,
|
|
)
|
|
self.assertEqual(result["source"], "qa")
|
|
self.assertEqual(result["answer"], "标准地址")
|
|
fake_model.assert_not_called()
|
|
|
|
def test_cross_language_qa_is_faithfully_adapted_by_model(self):
|
|
fake_model = Mock(return_value="Our address is Test Road 1.")
|
|
fake_search = Mock(return_value=[])
|
|
result = _resolve_reply(
|
|
None,
|
|
self.avatar,
|
|
"Where is your office?",
|
|
[],
|
|
qa_pairs=[SimpleNamespace(question="Where is your office?", answer="地址是测试路1号。", enabled=True)],
|
|
search_fn=fake_search,
|
|
model_client=fake_model,
|
|
)
|
|
|
|
self.assertEqual(result["source"], "qa")
|
|
self.assertEqual(result["answer"], "Our address is Test Road 1.")
|
|
self.assertEqual(fake_model.call_args.kwargs["temperature"], 0.0)
|
|
system = fake_model.call_args.kwargs["messages"][0]["content"]
|
|
self.assertIn("已确认标准答案", system)
|
|
self.assertIn("地址是测试路1号", system)
|
|
self.assertIn("只使用该语言回答", system)
|
|
fake_search.assert_not_called()
|
|
|
|
def test_qa_language_adaptation_detects_common_writing_system_changes(self):
|
|
self.assertTrue(_qa_requires_language_adaptation("Hello", "你好"))
|
|
self.assertTrue(_qa_requires_language_adaptation("こんにちは", "你好"))
|
|
self.assertTrue(_qa_requires_language_adaptation("안녕하세요", "你好"))
|
|
self.assertFalse(_qa_requires_language_adaptation("你好", "您好"))
|
|
|
|
def test_conversation_qa_is_rendered_for_the_current_turn_language(self):
|
|
history = [SimpleNamespace(role="user", content="Please answer in English.")]
|
|
self.assertTrue(_qa_requires_per_turn_rendering("Quelle est votre adresse ?", "Our address is Test Road 1.", history))
|
|
|
|
fake_model = Mock(return_value="Notre adresse est Test Road 1.")
|
|
result = _resolve_reply(
|
|
None,
|
|
self.avatar,
|
|
"Quelle est votre adresse ?",
|
|
history,
|
|
qa_pairs=[SimpleNamespace(question="Quelle est votre adresse ?", answer="Our address is Test Road 1.", enabled=True)],
|
|
search_fn=Mock(),
|
|
model_client=fake_model,
|
|
)
|
|
|
|
self.assertEqual(result["source"], "qa")
|
|
self.assertEqual(result["answer"], "Notre adresse est Test Road 1.")
|
|
messages = fake_model.call_args.kwargs["messages"]
|
|
self.assertEqual(messages[-1], {"role": "user", "content": "Quelle est votre adresse ?"})
|
|
self.assertEqual(messages[-2]["role"], "system")
|
|
self.assertIn("本轮语言覆盖指令", messages[-2]["content"])
|
|
self.assertIn("不要沿用上一轮语言", messages[-2]["content"])
|
|
|
|
def test_latest_user_message_has_an_adjacent_language_override(self):
|
|
history = [
|
|
SimpleNamespace(role="user", content="请用中文回答"),
|
|
SimpleNamespace(role="assistant", content="好的,请问有什么可以帮你?"),
|
|
]
|
|
messages = _build_prompt(self.avatar, history, "What can you help me with?", [])
|
|
|
|
self.assertEqual(messages[-1], {"role": "user", "content": "What can you help me with?"})
|
|
self.assertEqual(messages[-2]["role"], "system")
|
|
self.assertIn("最新用户消息", messages[-2]["content"])
|
|
self.assertIn("立即切换到相同语言", messages[-2]["content"])
|
|
|
|
def test_conversational_paraphrase_matches_standard_qa(self):
|
|
for question in ("请问一下,你们公司在哪里呀?", "请问去你们那边怎么走"):
|
|
with self.subTest(question=question):
|
|
matched = _match_standard_qa(question, [self.disabled_qa, self.qa])
|
|
self.assertIs(matched, self.qa)
|
|
|
|
def test_short_related_question_matches_single_standard_qa(self):
|
|
matched = _match_standard_qa("地址", [self.qa])
|
|
self.assertIs(matched, self.qa)
|
|
|
|
def test_ambiguous_short_question_does_not_pick_arbitrarily(self):
|
|
hospital = SimpleNamespace(question="医院地址", answer="医院地址答案", enabled=True)
|
|
company = SimpleNamespace(question="公司地址", answer="公司地址答案", enabled=True)
|
|
self.assertIsNone(_match_standard_qa("地址", [hospital, company]))
|
|
|
|
def test_unrelated_question_does_not_match_standard_qa(self):
|
|
self.assertIsNone(_match_standard_qa("今天天气怎么样", [self.qa]))
|
|
|
|
def test_knowledge_context_is_sent_to_qwen_after_qa_miss(self):
|
|
fake_model = Mock(return_value="根据知识库内容回答")
|
|
knowledge_hit = {
|
|
"filename": "退款.md",
|
|
"snippet": "知识库内容:七日内可申请退款。",
|
|
"score": 0.92,
|
|
}
|
|
result = _resolve_reply(
|
|
None,
|
|
self.avatar,
|
|
"退款规则",
|
|
[],
|
|
qa_pairs=[],
|
|
search_fn=lambda *_args, **_kwargs: [knowledge_hit],
|
|
model_client=fake_model,
|
|
)
|
|
self.assertEqual(result["source"], "knowledge")
|
|
self.assertIn("知识库内容", fake_model.call_args.kwargs["messages"][0]["content"])
|
|
self.assertIn("只能依据本人资料", fake_model.call_args.kwargs["messages"][0]["content"])
|
|
|
|
def test_prompt_contains_personality_configuration(self):
|
|
messages = _build_prompt(self.avatar, [], "你好", [])
|
|
self.assertIn("严谨度", messages[0]["content"])
|
|
self.assertNotIn("冯医生", messages[0]["content"])
|
|
self.assertIn("耳鼻喉科领域专家", messages[0]["content"])
|
|
self.assertIn("职业:医生", messages[0]["content"])
|
|
self.assertIn("职位:主任医师", messages[0]["content"])
|
|
self.assertIn("单位:测试医院", messages[0]["content"])
|
|
self.assertIn("单位地址:测试路1号", messages[0]["content"])
|
|
self.assertIn("不要编造政策", messages[0]["content"])
|
|
self.assertIn("模型供应商", messages[0]["content"])
|
|
self.assertIn("不要称自己为数字人", messages[0]["content"])
|
|
self.assertIn("输出排版规范", messages[0]["content"])
|
|
self.assertIn("任何回答都不要说出自己的姓名", messages[0]["content"])
|
|
self.assertIn("不要自我介绍", messages[0]["content"])
|
|
self.assertIn("像熟人之间微信聊天一样", messages[0]["content"])
|
|
self.assertIn("不隶属于任何机构", messages[0]["content"])
|
|
self.assertIn("不要连续输出空行", messages[0]["content"])
|
|
self.assertIn("回答语言规则", messages[0]["content"])
|
|
self.assertIn("当前最后一条用户消息", messages[0]["content"])
|
|
self.assertIn("历史消息", messages[0]["content"])
|
|
|
|
def test_prompt_blocks_ungrounded_factual_answers(self):
|
|
messages = _build_prompt(self.avatar, [], "聊聊国际新闻", [])
|
|
system = messages[0]["content"]
|
|
self.assertIn("没有检索到可靠资料", system)
|
|
self.assertIn("不要凭通用知识", system)
|
|
self.assertIn("不要提及知识库", system)
|
|
self.assertIn("不得推断服务对象", system)
|
|
self.assertIn("工作场所", system)
|
|
|
|
def test_public_avatar_payload_excludes_internal_configuration(self):
|
|
payload = _public_avatar_payload(self.avatar)
|
|
self.assertEqual(payload["displayName"], "冯医生")
|
|
self.assertEqual(payload["photoUrl"], "https://example.test/avatar.png")
|
|
self.assertNotIn("config", payload)
|
|
self.assertNotIn("ownerId", payload)
|
|
|
|
def test_unshared_avatars_do_not_reuse_a_unique_share_token(self):
|
|
first = Avatar(name="first")
|
|
second = Avatar(name="second")
|
|
self.assertIsNone(first.share_token)
|
|
self.assertIsNone(second.share_token)
|
|
|
|
def test_standard_answer_can_be_emitted_as_sse_chunks(self):
|
|
self.assertEqual(list(_iter_text_chunks("标准答案内容", size=2)), ["标准", "答案", "内容"])
|
|
|
|
def test_chat_rejects_avatar_owned_by_another_user(self):
|
|
class Query:
|
|
def __init__(self, value):
|
|
self.value = value
|
|
|
|
def filter(self, *_args, **_kwargs):
|
|
return self
|
|
|
|
def first(self):
|
|
return self.value
|
|
|
|
self_avatar = self.avatar
|
|
|
|
class DB:
|
|
avatar = self_avatar
|
|
|
|
def query(self, model):
|
|
return Query(
|
|
self.avatar if model is Avatar else SimpleNamespace(huihui_user_id="huihui-user-2")
|
|
)
|
|
|
|
db = DB()
|
|
with self.assertRaises(HTTPException) as caught:
|
|
_require_owned_avatar(db, self.avatar.id, "Bearer other-token")
|
|
self.assertEqual(caught.exception.status_code, 403)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|