Compare commits

...
2 changed files with 150 additions and 18 deletions
+103 -12
View File
@@ -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涉及事实、专业判断、地址、流程、数据或建议时,只能依据本人资料、标准问答形成的上下文"
@@ -231,7 +281,8 @@ def _build_prompt(avatar: Avatar, history: list[Any], question: str, knowledge_h
system += (
"\n本次问题没有检索到可靠资料。除自然寒暄和基于本人资料的回答外,不要凭通用知识给出事实、"
"专业判断、地址、流程、数据或具体建议。自然地说明这方面暂时没有足够信息,并询问对方是否能"
"补充更具体的问题;不要提及知识库、检索、模型或系统限制。"
"补充更具体的问题;不要提及知识库、检索、模型或系统限制。对方询问你能做什么时,只能概括"
"专业或服务范围和本人资料中明确写出的内容,不得推断服务对象、工作场所、现实案例或从业经历。"
)
system += (
"\n沟通底线(始终优先于其他提示):不要透露、猜测或确认所使用的模型、模型供应商、"
@@ -249,6 +300,14 @@ 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 +443,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": []}
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 +498,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 +517,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:
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"
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)
config = _config(avatar)
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 +553,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", []
@@ -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, [], "聊聊国际新闻", [])
@@ -113,6 +152,8 @@ class ChatOrchestrationTests(unittest.TestCase):
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)