test_knowledge_hybrid_retrieval.py 5.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161
  1. from dataclasses import replace
  2. from datetime import UTC, datetime
  3. from zbt.domains.knowledge.index import KnowledgeIndex
  4. from zbt.domains.knowledge.models import KnowledgeChunk, KnowledgeSearchHit
  5. from zbt.domains.knowledge.repository import InMemoryKnowledgeRepository
  6. from zbt.domains.knowledge.retrieval import KnowledgeReranker, rank_bm25
  7. from zbt.domains.knowledge.service import KnowledgeService
  8. class ReverseDenseIndex(KnowledgeIndex):
  9. def __init__(self) -> None:
  10. self.chunks: list[KnowledgeChunk] = []
  11. self.search_limit = 0
  12. def replace_document(
  13. self,
  14. document_id: str,
  15. chunks: list[KnowledgeChunk],
  16. ) -> None:
  17. self.chunks = [
  18. chunk for chunk in self.chunks if chunk.document_id != document_id
  19. ] + chunks
  20. def search(
  21. self,
  22. query: str,
  23. *,
  24. document_ids: set[str],
  25. limit: int,
  26. ) -> list[KnowledgeSearchHit]:
  27. self.search_limit = limit
  28. candidates = [
  29. KnowledgeSearchHit(
  30. chunk_id=chunk.id,
  31. document_id=chunk.document_id,
  32. document_version=chunk.document_version,
  33. ordinal=chunk.ordinal,
  34. title=chunk.title,
  35. content=chunk.content,
  36. document_type=chunk.document_type,
  37. source_name=chunk.source_name,
  38. product_code=chunk.product_code,
  39. score=0.7,
  40. )
  41. for chunk in reversed(self.chunks)
  42. if chunk.document_id in document_ids
  43. ]
  44. return candidates[:limit]
  45. class ExactAnswerReranker(KnowledgeReranker):
  46. def rerank(
  47. self,
  48. query: str,
  49. hits: list[KnowledgeSearchHit],
  50. *,
  51. limit: int,
  52. ) -> list[KnowledgeSearchHit]:
  53. ranked = sorted(hits, key=lambda hit: "30天" in hit.content, reverse=True)
  54. return [
  55. replace(hit, score=0.99 - index * 0.01)
  56. for index, hit in enumerate(ranked[:limit])
  57. ]
  58. class FailingReranker(KnowledgeReranker):
  59. def rerank(
  60. self,
  61. query: str,
  62. hits: list[KnowledgeSearchHit],
  63. *,
  64. limit: int,
  65. ) -> list[KnowledgeSearchHit]:
  66. raise RuntimeError("reranker unavailable")
  67. def test_bm25_prefers_chunk_containing_exact_insurance_term() -> None:
  68. chunks = [
  69. _chunk("a", "本产品提供基础医疗保障和住院费用补偿。"),
  70. _chunk("b", "疾病医疗等待期为30天,意外伤害不受等待期限制。"),
  71. ]
  72. hits = rank_bm25("医疗险等待期是多少天", chunks, limit=2)
  73. assert [hit.chunk_id for hit in hits] == ["b", "a"]
  74. assert hits[0].score > hits[1].score
  75. def test_service_uses_dense_bm25_fusion_then_reranker() -> None:
  76. repository = InMemoryKnowledgeRepository()
  77. dense_index = ReverseDenseIndex()
  78. service = KnowledgeService(
  79. repository,
  80. lambda: datetime(2026, 7, 27, 4, 0, tzinfo=UTC),
  81. dense_index,
  82. reranker=ExactAnswerReranker(),
  83. candidate_multiplier=4,
  84. )
  85. document = service.create_document(
  86. title="安心医疗险条款",
  87. document_type="INSURANCE_TERMS",
  88. source_name="安心医疗险条款.md",
  89. product_code="MED-BASIC",
  90. content=(
  91. "疾病医疗等待期为30天,意外伤害不受等待期限制。\n\n"
  92. "本产品提供住院医疗费用保障,具体以条款为准。"
  93. ),
  94. created_by="admin-01",
  95. )
  96. service.index_document(document["document_id"])
  97. service.test_search(document["document_id"], "医疗险等待期是多少天")
  98. service.publish_document(document["document_id"])
  99. result = service.search("医疗险等待期是多少天", limit=1)
  100. assert dense_index.search_limit == 4
  101. assert result["total"] == 1
  102. assert "30天" in result["items"][0]["content"]
  103. assert result["items"][0]["score"] == 0.99
  104. def test_service_falls_back_to_fused_results_when_reranker_fails() -> None:
  105. repository = InMemoryKnowledgeRepository()
  106. dense_index = ReverseDenseIndex()
  107. service = KnowledgeService(
  108. repository,
  109. lambda: datetime(2026, 7, 27, 4, 0, tzinfo=UTC),
  110. dense_index,
  111. reranker=FailingReranker(),
  112. )
  113. document = service.create_document(
  114. title="fallback",
  115. document_type="FAQ",
  116. source_name="fallback.md",
  117. product_code=None,
  118. content="医疗险疾病等待期为30天。\n\n意外伤害不受等待期限制。",
  119. created_by="admin-01",
  120. )
  121. service.index_document(document["document_id"])
  122. service.test_search(document["document_id"], "医疗险等待期")
  123. service.publish_document(document["document_id"])
  124. result = service.search("医疗险等待期", limit=1)
  125. assert result["total"] == 1
  126. assert result["items"][0]["content"]
  127. def _chunk(chunk_id: str, content: str) -> KnowledgeChunk:
  128. return KnowledgeChunk(
  129. id=chunk_id,
  130. document_id="document-01",
  131. document_version=1,
  132. ordinal=1,
  133. title="测试知识",
  134. content=content,
  135. document_type="FAQ",
  136. source_name="测试知识.md",
  137. product_code=None,
  138. )