import unittest from types import SimpleNamespace from unittest.mock import Mock from fastapi import HTTPException from models import Avatar, User from routers.chat import ( _answer_requires_language_repair, _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, _turn_language_name, ) 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("MANDATORY OUTPUT LANGUAGE", messages[-2]["content"]) self.assertIn("French", 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("MANDATORY OUTPUT LANGUAGE", messages[-2]["content"]) self.assertIn("English", messages[-2]["content"]) def test_reported_alzheimer_question_is_explicitly_english(self): question = "I have a friend who has symptoms of Alzheimer's disease" self.assertEqual(_turn_language_name(question), "English") messages = _build_prompt(self.avatar, [], question, []) self.assertIn("MANDATORY OUTPUT LANGUAGE FOR THIS TURN: English", messages[-2]["content"]) def test_non_stream_reply_repairs_a_wrong_writing_system_before_sending(self): question = "I have a friend who has symptoms of Alzheimer's disease" fake_model = Mock(side_effect=["建议尽快就医评估。", "Please arrange a medical assessment soon."]) result = _resolve_reply( None, self.avatar, question, [], qa_pairs=[], search_fn=lambda *_args, **_kwargs: [], model_client=fake_model, usage_source="takeover", ) self.assertEqual(result["answer"], "Please arrange a medical assessment soon.") self.assertEqual(fake_model.call_count, 2) repair_messages = fake_model.call_args.kwargs["messages"] self.assertIn("English", repair_messages[0]["content"]) self.assertIn("建议尽快就医评估", repair_messages[-1]["content"]) self.assertTrue(_answer_requires_language_repair(question, "建议尽快就医评估。")) 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()