nex_docus/backend/tests/test_chat_citations.py

209 lines
8.8 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

import unittest
from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
from app.api.v1.chat import _canonicalize_message_citations, _compact_cited_refs
from app.services.rag_service import RAGService
class ChatCitationTest(unittest.TestCase):
def test_legacy_interrupt_marker_is_read_only_compat(self):
from app.api.v1.chat import (
_is_interrupted_message,
_legacy_interrupted_content,
_strip_legacy_interrupt_marker,
)
self.assertTrue(_legacy_interrupted_content("interrupt"))
self.assertTrue(_legacy_interrupted_content("内容\n\ninterrupt"))
self.assertFalse(_legacy_interrupted_content("正常内容"))
self.assertEqual(_strip_legacy_interrupt_marker("interrupt"), "")
self.assertEqual(_strip_legacy_interrupt_marker("内容\n\ninterrupt"), "内容")
self.assertEqual(_strip_legacy_interrupt_marker("正常内容"), "正常内容")
# 旧数据status 缺失/pending + 旧标记)按已中断处理,但不修改数据库
legacy = SimpleNamespace(role="assistant", status="pending", content="内容\n\ninterrupt")
self.assertTrue(_is_interrupted_message(legacy))
# 新数据一律以 status 为准,不再比对内容
fresh = SimpleNamespace(role="assistant", status="completed", content="内容\n\ninterrupt")
self.assertFalse(_is_interrupted_message(fresh))
interrupted = SimpleNamespace(role="assistant", status="interrupted", content="")
self.assertTrue(_is_interrupted_message(interrupted))
def test_duplicate_file_citations_are_renumbered_once(self):
content, refs = _canonicalize_message_citations(
"泰坦属于土星的卫星[2][4][5]。",
[
{"citation_id": 2, "file_path": "内容导航.md", "excerpt": "片段一"},
{"citation_id": 4, "file_path": "内容导航.md", "excerpt": "片段二"},
{"citation_id": 5, "file_path": "内容导航.md", "excerpt": "片段三"},
],
)
self.assertEqual(content, "泰坦属于土星的卫星[1]。")
self.assertEqual(len(refs), 1)
self.assertEqual(refs[0]["citation_id"], 1)
self.assertEqual(refs[0]["excerpt"], "片段一\n\n片段二\n\n片段三")
def test_overlapping_blocks_are_not_duplicated(self):
self.assertEqual(
RAGService._merge_text_blocks("泰坦属于土星", "泰坦属于土星"),
"泰坦属于土星",
)
def test_single_used_reference_is_compacted_to_one(self):
content, refs = _compact_cited_refs(
"泰坦是土星的卫星[3],土卫二也属于土星[3]。",
[
{"citation_id": 1, "file_path": "其他一.md"},
{"citation_id": 2, "file_path": "其他二.md"},
{"citation_id": 3, "file_path": "内容导航.md"},
],
)
self.assertEqual(content, "泰坦是土星的卫星[1],土卫二也属于土星[1]。")
self.assertEqual(refs, [{
"citation_id": 1,
"file_path": "内容导航.md",
"anchor_text": "",
"excerpt": "",
"content": "",
"chunk_index": None,
"hit_terms": [],
}])
def test_used_references_follow_first_appearance_order(self):
content, refs = _compact_cited_refs(
"先使用第三份[3],再使用第一份[1]。",
[
{"citation_id": 1, "file_path": "第一份.md"},
{"citation_id": 2, "file_path": "第二份.md"},
{"citation_id": 3, "file_path": "第三份.md"},
],
)
self.assertEqual(content, "先使用第三份[1],再使用第一份[2]。")
self.assertEqual(
[(ref["citation_id"], ref["file_path"]) for ref in refs],
[(1, "第三份.md"), (2, "第一份.md")],
)
class RAGConversationContextTest(unittest.IsolatedAsyncioTestCase):
async def test_previous_assistant_answer_is_not_sent_to_model(self):
db = AsyncMock()
db.execute.return_value = SimpleNamespace(
scalar_one_or_none=lambda: SimpleNamespace(config_id=1)
)
history = [
{"role": "user", "content": "包含了哪几次阿波罗计划?"},
{"role": "assistant", "content": "上一轮完整答案不应再次发送给模型。"},
]
_, system_prompt, messages = await RAGService._prepare_generation(
db,
"项目中的文档包含哪些土星的卫星?",
1,
[],
history,
)
self.assertEqual(messages, [{
"role": "user",
"content": "项目中的文档包含哪些土星的卫星?",
}])
self.assertIn("包含了哪几次阿波罗计划?", system_prompt)
self.assertNotIn("上一轮完整答案不应再次发送给模型。", system_prompt)
self.assertIn("禁止复述、总结或继续回答先前问题", system_prompt)
class CitationQuoteExtractionTest(unittest.TestCase):
def test_split_sentences(self):
self.assertEqual(
RAGService._split_sentences("第一句。第二句!第三句?\n第四句;第五句话。"),
["第一句。", "第二句!", "第三句?", "第四句;", "第五句话。"],
)
def test_extract_citation_claims(self):
answer = (
"土星拥有众多卫星[1]。泰坦是其中最大的一颗[1]。\n"
"关键参数:直径约 120536 公里[1]。\n"
"该产品型号为 X200[1],电池容量 5000mAh[1]。"
)
claims = RAGService._extract_citation_claims(answer)
self.assertEqual(claims[1], [
"土星拥有众多卫星",
"泰坦是其中最大的一颗",
"关键参数:直径约 120536 公里",
"该产品型号为 X200",
"电池容量 5000mAh",
])
def test_sentence_candidates_prefer_term_overlap(self):
candidates = RAGService._sentence_candidates(
"土星是太阳系第二大行星。\n土卫六又称泰坦,是土星最大的卫星。\n"
"木星拥有最多卫星。",
["泰坦是最大的卫星"],
max_candidates=3,
)
self.assertIn("土卫六又称泰坦,是土星最大的卫星。", candidates[:1])
def test_cosine_similarity(self):
self.assertAlmostEqual(RAGService._cosine_similarity([1.0, 0.0], [1.0, 0.0]), 1.0)
self.assertAlmostEqual(RAGService._cosine_similarity([1.0, 0.0], [0.0, 1.0]), 0.0)
self.assertEqual(RAGService._cosine_similarity([], [1.0]), 0.0)
class CitationQuoteAlignmentTest(unittest.IsolatedAsyncioTestCase):
async def test_align_citation_quotes_picks_similar_sentence(self):
refs = [{
"citation_id": 1,
"file_path": "产品资料.md",
"anchor_text": "型号 X200",
"excerpt": "型号 X200 支持 5G。\n电池容量为 5000mAh。\n屏幕尺寸 6.7 英寸。",
"content": "",
}]
answer = "该产品型号为 X200[1],电池续航出色[1]。"
async def fake_embeddings(db, texts, config=None):
return [
[1.0, 0.0, 0.0] if ("电池" in text or "5000mAh" in text) else [0.0, 1.0, 0.0]
for text in texts
]
with patch("app.services.rag_service.zvec_service") as mock_zvec:
mock_zvec.generate_embeddings = fake_embeddings
result = await RAGService.align_citation_quotes(None, answer, refs)
self.assertIn("quotes", result[0])
self.assertTrue(any(
"5000mAh" in q["text"] or "电池" in q["text"]
for q in result[0]["quotes"]
))
# 按出现顺序回填:第一处对应型号,第二处对应电池
occurrences = result[0]["quote_occurrences"]
self.assertEqual(len(occurrences), 2)
self.assertEqual(occurrences[0]["claim"], "该产品型号为 X200")
self.assertEqual(occurrences[1]["claim"], "电池续航出色")
self.assertTrue(any("电池" in q["text"] for q in occurrences[1]["quotes"]))
async def test_align_citation_quotes_falls_back_when_embedding_unavailable(self):
refs = [{
"citation_id": 1,
"file_path": "a.md",
"excerpt": "某句支撑内容。",
"content": "",
}]
with patch("app.services.rag_service.zvec_service") as mock_zvec:
mock_zvec.generate_embeddings = AsyncMock(return_value=None)
result = await RAGService.align_citation_quotes(None, "这是论断[1]。", refs)
self.assertEqual(result, refs)
self.assertNotIn("quotes", result[0])
if __name__ == "__main__":
unittest.main()