diff --git a/digital-avatar-app/backend/routers/chat.py b/digital-avatar-app/backend/routers/chat.py index 9e5d863..5daff86 100644 --- a/digital-avatar-app/backend/routers/chat.py +++ b/digital-avatar-app/backend/routers/chat.py @@ -34,6 +34,19 @@ QA_SEMANTIC_THRESHOLD = 0.72 QA_MATCH_MARGIN = 0.06 KNOWLEDGE_MIN_SCORE = float(os.getenv("KNOWLEDGE_MIN_SCORE", "0.42")) +_WRITING_SYSTEM_PATTERNS = { + "han": re.compile(r"[\u3400-\u4dbf\u4e00-\u9fff]"), + "latin": re.compile(r"[A-Za-z\u00c0-\u024f]"), + "cyrillic": re.compile(r"[\u0400-\u052f]"), + "arabic": re.compile(r"[\u0600-\u06ff]"), + "hebrew": re.compile(r"[\u0590-\u05ff]"), + "devanagari": re.compile(r"[\u0900-\u097f]"), + "thai": re.compile(r"[\u0e00-\u0e7f]"), + "greek": re.compile(r"[\u0370-\u03ff]"), +} +_JAPANESE_KANA = re.compile(r"[\u3040-\u30ff]") +_KOREAN_HANGUL = re.compile(r"[\uac00-\ud7af\u1100-\u11ff]") + class ChatMessage(BaseModel): role: str = Field(pattern="^(user|assistant)$") @@ -70,6 +83,30 @@ def _normalize_question(value: str) -> str: return value.translate(str.maketrans("", "", string.punctuation + ",。!?;:、()【】「」‘’“”《》")) +def _dominant_writing_system(value: str) -> str: + value = value or "" + if _JAPANESE_KANA.search(value): + return "japanese" + if _KOREAN_HANGUL.search(value): + return "korean" + counts = { + name: len(pattern.findall(value)) + for name, pattern in _WRITING_SYSTEM_PATTERNS.items() + } + name, count = max(counts.items(), key=lambda item: item[1]) + return name if count else "unknown" + + +def _qa_requires_language_adaptation(question: str, answer: str) -> bool: + question_system = _dominant_writing_system(question) + answer_system = _dominant_writing_system(answer) + return ( + question_system != "unknown" + and answer_system != "unknown" + and question_system != answer_system + ) + + def _canonicalize_question(value: str) -> str: value = _normalize_question(value) replacements = ( @@ -189,7 +226,14 @@ def _config(avatar: Avatar) -> dict: } -def _build_prompt(avatar: Avatar, history: list[Any], question: str, knowledge_hits: list[dict]) -> list[dict]: +def _build_prompt( + avatar: Avatar, + history: list[Any], + question: str, + knowledge_hits: list[dict], + *, + standard_answer: str = "", +) -> list[dict]: config = _config(avatar) description = (getattr(avatar, "description", "") or "").strip() knowledge = "\n".join( @@ -210,7 +254,7 @@ def _build_prompt(avatar: Avatar, history: list[Any], question: str, knowledge_h profile = ";".join(f"{label}:{value}" for label, value in profile_items) system = ( f"你的专业或服务范围是:「{description or '未设置'}」。" - "请基于已提供的知识库回答,不要编造事实;" + "请基于已提供的可靠资料回答,不要编造事实;" f"回复风格:{config['replyStyle']};严谨度:{config['rigor']}/100;" f"幽默感:{config['humor']}/100;回复长度:{config['responseLength']}。" ) @@ -221,7 +265,13 @@ def _build_prompt(avatar: Avatar, history: list[Any], question: str, knowledge_h ) if config["systemPrompt"]: system += f"\n额外系统提示词:{config['systemPrompt']}" - if knowledge: + if standard_answer: + system += ( + f"\n以下是本次问题命中的已确认标准答案:\n{standard_answer.strip()}" + "\n必须保持标准答案中的事实、数字、专有名词和结论不变,只允许为匹配用户当前语言进行忠实转换" + "和必要的自然表达,不得补充、删减或改写其含义。不要提及标准答案或转换过程。" + ) + elif knowledge: system += ( f"\n以下是本次问题检索到的可靠资料:\n{knowledge}" "\n涉及事实、专业判断、地址、流程、数据或建议时,只能依据本人资料、标准问答形成的上下文" @@ -249,6 +299,13 @@ def _build_prompt(avatar: Avatar, history: list[Any], question: str, knowledge_h "只有切换独立观点或确实需要列举时才换行;列举使用 1.、2.、3.,每项单独一行。" "不要在行首或行尾留空格,不要连续输出空行。先给结论,再给简短说明;避免重复和冗长铺垫。" ) + system += ( + "\n回答语言规则(始终优先):先识别当前最后一条用户消息所使用的主要自然语言,并且只使用该语言回答。" + "用户用中文就用中文,使用 English 就用 English,使用日语、韩语、法语、西班牙语或其他语言时," + "也必须使用对应语言。消息混用多种语言时跟随占主导的语言;用户明确指定回答语言时服从其指定。" + "历史消息、本人资料、标准答案和知识库使用的语言都不能覆盖当前用户消息的语言。" + "专有名词、品牌、地址、代码和必要缩写可保留原文。不要解释语言识别或翻译过程。" + ) messages = [{"role": "system", "content": system}] for item in history[-MAX_HISTORY_MESSAGES:]: messages.append({"role": item.role, "content": item.content} if hasattr(item, "role") else item) @@ -384,14 +441,30 @@ def _resolve_reply( if qa_pairs is None: qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all() matched = _match_standard_qa(question, qa_pairs) - if matched: + adapt_qa_language = bool( + matched and _qa_requires_language_adaptation(question, matched.answer) + ) + if matched and not adapt_qa_language: return {"answer": matched.answer, "source": "qa", "references": []} - search_fn = search_fn or (lambda query, avatar_id: _search_knowledge(db, avatar_id, query)) - hits = search_fn(question, avatar.id) - messages = _build_prompt(avatar, history, question, hits) + if matched: + hits = [] + messages = _build_prompt( + avatar, + history, + question, + hits, + standard_answer=matched.answer, + ) + else: + search_fn = search_fn or (lambda query, avatar_id: _search_knowledge(db, avatar_id, query)) + hits = search_fn(question, avatar.id) + messages = _build_prompt(avatar, history, question, hits) config = _config(avatar) - temperature = min(0.45 if hits else 0.25, 0.2 + config["creativity"] / 100 * 0.6) + temperature = 0.0 if matched else min( + 0.45 if hits else 0.25, + 0.2 + config["creativity"] / 100 * 0.6, + ) token_usage = None if model_client is not None: answer = model_client(messages=messages, temperature=temperature) @@ -423,7 +496,7 @@ def _resolve_reply( raise result = { "answer": answer, - "source": "knowledge" if hits else "qwen", + "source": "qa" if matched else ("knowledge" if hits else "qwen"), "references": hits, } if token_usage: @@ -442,14 +515,32 @@ def _stream_reply( ): qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all() matched = _match_standard_qa(question, qa_pairs) - if matched: + adapt_qa_language = bool( + matched and _qa_requires_language_adaptation(question, matched.answer) + ) + messages, reservation = [], None + if matched and not adapt_qa_language: source, references, chunks = "qa", [], _iter_text_chunks(matched.answer) else: - references = _search_knowledge(db, avatar.id, question) - source = "knowledge" if references else "qwen" + if matched: + references = [] + source = "qa" + messages = _build_prompt( + avatar, + history, + question, + references, + standard_answer=matched.answer, + ) + else: + references = _search_knowledge(db, avatar.id, question) + source = "knowledge" if references else "qwen" + messages = _build_prompt(avatar, history, question, references) config = _config(avatar) - temperature = min(0.45 if references else 0.25, 0.2 + config["creativity"] / 100 * 0.6) - messages = _build_prompt(avatar, history, question, references) + temperature = 0.0 if matched else min( + 0.45 if references else 0.25, + 0.2 + config["creativity"] / 100 * 0.6, + ) model_config = get_chat_model_config() reservation = reserve_avatar_tokens( db, @@ -460,8 +551,6 @@ def _stream_reply( model_config.max_tokens, ) chunks = _iter_qwen_stream(messages, temperature, model_config) - if matched: - messages, reservation = [], None if public: source, references = "public", [] diff --git a/digital-avatar-app/backend/tests/test_chat_orchestration.py b/digital-avatar-app/backend/tests/test_chat_orchestration.py index 1c4aa33..3732eff 100644 --- a/digital-avatar-app/backend/tests/test_chat_orchestration.py +++ b/digital-avatar-app/backend/tests/test_chat_orchestration.py @@ -5,7 +5,15 @@ 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, _require_owned_avatar, _resolve_reply +from routers.chat import ( + _build_prompt, + _iter_text_chunks, + _match_standard_qa, + _public_avatar_payload, + _qa_requires_language_adaptation, + _require_owned_avatar, + _resolve_reply, +) class ChatOrchestrationTests(unittest.TestCase): @@ -50,6 +58,34 @@ class ChatOrchestrationTests(unittest.TestCase): 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_conversational_paraphrase_matches_standard_qa(self): for question in ("请问一下,你们公司在哪里呀?", "请问去你们那边怎么走"): with self.subTest(question=question): @@ -106,6 +142,9 @@ class ChatOrchestrationTests(unittest.TestCase): 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, [], "聊聊国际新闻", [])