diff --git a/digital-avatar-app/backend/routers/chat.py b/digital-avatar-app/backend/routers/chat.py index fece73e..967d236 100644 --- a/digital-avatar-app/backend/routers/chat.py +++ b/digital-avatar-app/backend/routers/chat.py @@ -79,6 +79,34 @@ _WRITING_SYSTEM_PATTERNS = { _JAPANESE_KANA = re.compile(r"[\u3040-\u30ff]") _KOREAN_HANGUL = re.compile(r"[\uac00-\ud7af\u1100-\u11ff]") +_LATIN_LANGUAGE_MARKERS = { + "English": re.compile( + r"\b(?:i|you|we|they|he|she|have|has|had|friend|who|what|where|when|why|how|" + r"symptoms?|disease|please|can|could|would|should|is|are|was|were|the|this|that)\b", + re.IGNORECASE, + ), + "French": re.compile( + r"\b(?:je|tu|vous|nous|ils|elle|une|des|avec|pour|pourquoi|comment|bonjour|est|sont)\b", + re.IGNORECASE, + ), + "Spanish": re.compile( + r"\b(?:yo|tu|usted|nosotros|ellos|ella|una|con|para|por que|como|hola|esta|son)\b", + re.IGNORECASE, + ), + "German": re.compile( + r"\b(?:ich|du|sie|wir|eine|mit|fur|warum|wie|hallo|ist|sind|haben)\b", + re.IGNORECASE, + ), + "Portuguese": re.compile( + r"\b(?:eu|voce|nos|eles|ela|uma|com|para|porque|como|ola|esta|sao|tenho)\b", + re.IGNORECASE, + ), + "Italian": re.compile( + r"\b(?:io|tu|voi|noi|loro|una|con|per|perche|come|ciao|sono|avere)\b", + re.IGNORECASE, + ), +} + class ChatMessage(BaseModel): model_config = ConfigDict(populate_by_name=True) @@ -487,15 +515,67 @@ def _qa_requires_per_turn_rendering( return bool(history) or _qa_requires_language_adaptation(question, answer) -def _per_turn_language_instruction() -> str: +def _latin_language_name(value: str) -> str: + scores = { + language: len(pattern.findall(value or "")) + for language, pattern in _LATIN_LANGUAGE_MARKERS.items() + } + language, score = max(scores.items(), key=lambda item: item[1]) + return language if score else "the same natural language as the latest user message" + + +def _turn_language_name(value: str) -> str: + writing_system = _dominant_writing_system(value) + return { + "han": "Chinese", + "japanese": "Japanese", + "korean": "Korean", + "cyrillic": "the same Cyrillic-script language as the latest user message", + "arabic": "the same Arabic-script language as the latest user message", + "hebrew": "Hebrew", + "devanagari": "the same Devanagari-script language as the latest user message", + "thai": "Thai", + "greek": "Greek", + "latin": _latin_language_name(value), + }.get(writing_system, "the same natural language as the latest user message") + + +def _per_turn_language_instruction(question: str = "") -> str: + language = _turn_language_name(question) return ( - "本轮语言覆盖指令:只根据紧随其后的最新用户消息判断本轮回答语言。" - "即使此前整段对话一直使用另一种语言,只要最新消息切换了语言,本轮就必须立即切换到相同语言;" - "不要沿用上一轮语言。若最新消息明确指定回答语言,以该指定为准;若混用多种语言,使用其中占主导的" - "自然语言。不要说明你检测、切换或翻译了语言。" + f"MANDATORY OUTPUT LANGUAGE FOR THIS TURN: {language}. " + "Write the entire answer only in that language. This instruction overrides the languages used by " + "conversation history, profile data, standard answers, retrieved documents, and custom prompts. " + "Translate grounded source material faithfully when necessary. Do not mention language detection, " + "translation, or this instruction." ) +def _answer_requires_language_repair(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 _language_repair_messages(question: str, answer: str) -> list[dict]: + return [ + {"role": "system", "content": _per_turn_language_instruction(question)}, + { + "role": "system", + "content": ( + "Rewrite the supplied draft in the mandatory output language. Preserve every grounded fact, " + "number, proper noun, uncertainty, and safety qualification. Add no new information and output " + "only the rewritten answer." + ), + }, + {"role": "user", "content": answer.strip()}, + ] + + def _canonicalize_question(value: str) -> str: value = _normalize_question(value) replacements = ( @@ -728,7 +808,7 @@ def _build_prompt( for item in history[-MAX_HISTORY_MESSAGES:]: messages.append({"role": item.role, "content": item.content} if hasattr(item, "role") else item) # Keep the language instruction adjacent to the current turn so long histories cannot override it. - messages.append({"role": "system", "content": _per_turn_language_instruction()}) + messages.append({"role": "system", "content": _per_turn_language_instruction(question)}) messages.append({"role": "user", "content": question.strip()}) return messages @@ -791,6 +871,41 @@ def _call_qwen( return {"answer": answer.strip(), "usage": data.get("usage") or {}} +def _call_billed_qwen( + db: Session, + avatar: Avatar, + messages: list[dict], + temperature: float, + usage_source: str, + model_config: ChatModelConfig, +) -> tuple[str, dict]: + reservation = reserve_avatar_tokens( + db, + avatar, + usage_source, + model_config.model, + messages, + model_config.max_tokens, + ) + try: + model_result = _call_qwen( + messages=messages, + temperature=temperature, + model_config=model_config, + ) + answer = model_result["answer"] + token_usage = settle_reservation( + db, + reservation, + model_result.get("usage"), + fallback_total=estimate_fallback_usage(messages, answer), + ) + return answer, token_usage + except Exception as exc: + release_reservation(db, reservation, str(exc)) + raise + + def _iter_qwen_stream( messages: list[dict], temperature: float, model_config: ChatModelConfig | None = None ): @@ -899,32 +1014,36 @@ def _resolve_reply( token_usage = None if model_client is not None: answer = model_client(messages=messages, temperature=temperature) + if _answer_requires_language_repair(question, str(answer or "")): + answer = model_client( + messages=_language_repair_messages(question, str(answer)), + temperature=0.0, + ) else: model_config = get_chat_model_config() - reservation = reserve_avatar_tokens( + answer, token_usage = _call_billed_qwen( db, avatar, - usage_source, - model_config.model, messages, - model_config.max_tokens, + temperature, + usage_source, + model_config, ) - try: - model_result = _call_qwen( - messages=messages, - temperature=temperature, - model_config=model_config, + if _answer_requires_language_repair(question, answer): + logger.warning( + "chat response language mismatch avatar=%s source=%s expected=%s", + avatar.id, + usage_source, + _turn_language_name(question), ) - answer = model_result["answer"] - token_usage = settle_reservation( + answer, token_usage = _call_billed_qwen( db, - reservation, - model_result.get("usage"), - fallback_total=estimate_fallback_usage(messages, answer), + avatar, + _language_repair_messages(question, answer), + 0.0, + f"{usage_source}_language_repair", + model_config, ) - except Exception as exc: - release_reservation(db, reservation, str(exc)) - raise answer = str(answer or "").strip() if image_contexts and _answer_denies_available_image(answer): logger.warning( diff --git a/digital-avatar-app/backend/tests/test_chat_orchestration.py b/digital-avatar-app/backend/tests/test_chat_orchestration.py index a4a66d4..70f7f78 100644 --- a/digital-avatar-app/backend/tests/test_chat_orchestration.py +++ b/digital-avatar-app/backend/tests/test_chat_orchestration.py @@ -6,6 +6,7 @@ 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, @@ -14,6 +15,7 @@ from routers.chat import ( _qa_requires_per_turn_rendering, _require_owned_avatar, _resolve_reply, + _turn_language_name, ) @@ -107,8 +109,8 @@ class ChatOrchestrationTests(unittest.TestCase): 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"]) + 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 = [ @@ -119,8 +121,37 @@ class ChatOrchestrationTests(unittest.TestCase): 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"]) + 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 ("请问一下,你们公司在哪里呀?", "请问去你们那边怎么走"):