fix(avatar): follow language changes each turn
This commit is contained in:
@@ -478,6 +478,24 @@ def _qa_requires_language_adaptation(question: str, answer: str) -> bool:
|
||||
)
|
||||
|
||||
|
||||
def _qa_requires_per_turn_rendering(
|
||||
question: str,
|
||||
answer: str,
|
||||
history: list[Any],
|
||||
) -> bool:
|
||||
"""Keep the direct QA fast path only when no conversation can bias language."""
|
||||
return bool(history) or _qa_requires_language_adaptation(question, answer)
|
||||
|
||||
|
||||
def _per_turn_language_instruction() -> str:
|
||||
return (
|
||||
"本轮语言覆盖指令:只根据紧随其后的最新用户消息判断本轮回答语言。"
|
||||
"即使此前整段对话一直使用另一种语言,只要最新消息切换了语言,本轮就必须立即切换到相同语言;"
|
||||
"不要沿用上一轮语言。若最新消息明确指定回答语言,以该指定为准;若混用多种语言,使用其中占主导的"
|
||||
"自然语言。不要说明你检测、切换或翻译了语言。"
|
||||
)
|
||||
|
||||
|
||||
def _canonicalize_question(value: str) -> str:
|
||||
value = _normalize_question(value)
|
||||
replacements = (
|
||||
@@ -709,6 +727,8 @@ def _build_prompt(
|
||||
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)
|
||||
# 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": "user", "content": question.strip()})
|
||||
return messages
|
||||
|
||||
@@ -845,7 +865,7 @@ def _resolve_reply(
|
||||
qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all()
|
||||
matched = _match_standard_qa(question, qa_pairs)
|
||||
adapt_qa_language = bool(
|
||||
matched and _qa_requires_language_adaptation(question, matched.answer)
|
||||
matched and _qa_requires_per_turn_rendering(question, matched.answer, history)
|
||||
)
|
||||
if matched and not adapt_qa_language and not image_contexts:
|
||||
return {"answer": matched.answer, "source": "qa", "references": []}
|
||||
@@ -940,7 +960,7 @@ def _stream_reply(
|
||||
qa_pairs = db.query(QAPair).filter(QAPair.avatar_id == avatar.id).all()
|
||||
matched = _match_standard_qa(question, qa_pairs)
|
||||
adapt_qa_language = bool(
|
||||
matched and _qa_requires_language_adaptation(question, matched.answer)
|
||||
matched and _qa_requires_per_turn_rendering(question, matched.answer, history)
|
||||
)
|
||||
messages, reservation = [], None
|
||||
if matched and not adapt_qa_language and not image_contexts:
|
||||
|
||||
@@ -11,6 +11,7 @@ from routers.chat import (
|
||||
_match_standard_qa,
|
||||
_public_avatar_payload,
|
||||
_qa_requires_language_adaptation,
|
||||
_qa_requires_per_turn_rendering,
|
||||
_require_owned_avatar,
|
||||
_resolve_reply,
|
||||
)
|
||||
@@ -86,6 +87,41 @@ class ChatOrchestrationTests(unittest.TestCase):
|
||||
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):
|
||||
|
||||
Reference in New Issue
Block a user