| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161 |
- from dataclasses import replace
- from datetime import UTC, datetime
- from zbt.domains.knowledge.index import KnowledgeIndex
- from zbt.domains.knowledge.models import KnowledgeChunk, KnowledgeSearchHit
- from zbt.domains.knowledge.repository import InMemoryKnowledgeRepository
- from zbt.domains.knowledge.retrieval import KnowledgeReranker, rank_bm25
- from zbt.domains.knowledge.service import KnowledgeService
- class ReverseDenseIndex(KnowledgeIndex):
- def __init__(self) -> None:
- self.chunks: list[KnowledgeChunk] = []
- self.search_limit = 0
- def replace_document(
- self,
- document_id: str,
- chunks: list[KnowledgeChunk],
- ) -> None:
- self.chunks = [
- chunk for chunk in self.chunks if chunk.document_id != document_id
- ] + chunks
- def search(
- self,
- query: str,
- *,
- document_ids: set[str],
- limit: int,
- ) -> list[KnowledgeSearchHit]:
- self.search_limit = limit
- candidates = [
- KnowledgeSearchHit(
- chunk_id=chunk.id,
- document_id=chunk.document_id,
- document_version=chunk.document_version,
- ordinal=chunk.ordinal,
- title=chunk.title,
- content=chunk.content,
- document_type=chunk.document_type,
- source_name=chunk.source_name,
- product_code=chunk.product_code,
- score=0.7,
- )
- for chunk in reversed(self.chunks)
- if chunk.document_id in document_ids
- ]
- return candidates[:limit]
- class ExactAnswerReranker(KnowledgeReranker):
- def rerank(
- self,
- query: str,
- hits: list[KnowledgeSearchHit],
- *,
- limit: int,
- ) -> list[KnowledgeSearchHit]:
- ranked = sorted(hits, key=lambda hit: "30天" in hit.content, reverse=True)
- return [
- replace(hit, score=0.99 - index * 0.01)
- for index, hit in enumerate(ranked[:limit])
- ]
- class FailingReranker(KnowledgeReranker):
- def rerank(
- self,
- query: str,
- hits: list[KnowledgeSearchHit],
- *,
- limit: int,
- ) -> list[KnowledgeSearchHit]:
- raise RuntimeError("reranker unavailable")
- def test_bm25_prefers_chunk_containing_exact_insurance_term() -> None:
- chunks = [
- _chunk("a", "本产品提供基础医疗保障和住院费用补偿。"),
- _chunk("b", "疾病医疗等待期为30天,意外伤害不受等待期限制。"),
- ]
- hits = rank_bm25("医疗险等待期是多少天", chunks, limit=2)
- assert [hit.chunk_id for hit in hits] == ["b", "a"]
- assert hits[0].score > hits[1].score
- def test_service_uses_dense_bm25_fusion_then_reranker() -> None:
- repository = InMemoryKnowledgeRepository()
- dense_index = ReverseDenseIndex()
- service = KnowledgeService(
- repository,
- lambda: datetime(2026, 7, 27, 4, 0, tzinfo=UTC),
- dense_index,
- reranker=ExactAnswerReranker(),
- candidate_multiplier=4,
- )
- document = service.create_document(
- title="安心医疗险条款",
- document_type="INSURANCE_TERMS",
- source_name="安心医疗险条款.md",
- product_code="MED-BASIC",
- content=(
- "疾病医疗等待期为30天,意外伤害不受等待期限制。\n\n"
- "本产品提供住院医疗费用保障,具体以条款为准。"
- ),
- created_by="admin-01",
- )
- service.index_document(document["document_id"])
- service.test_search(document["document_id"], "医疗险等待期是多少天")
- service.publish_document(document["document_id"])
- result = service.search("医疗险等待期是多少天", limit=1)
- assert dense_index.search_limit == 4
- assert result["total"] == 1
- assert "30天" in result["items"][0]["content"]
- assert result["items"][0]["score"] == 0.99
- def test_service_falls_back_to_fused_results_when_reranker_fails() -> None:
- repository = InMemoryKnowledgeRepository()
- dense_index = ReverseDenseIndex()
- service = KnowledgeService(
- repository,
- lambda: datetime(2026, 7, 27, 4, 0, tzinfo=UTC),
- dense_index,
- reranker=FailingReranker(),
- )
- document = service.create_document(
- title="fallback",
- document_type="FAQ",
- source_name="fallback.md",
- product_code=None,
- content="医疗险疾病等待期为30天。\n\n意外伤害不受等待期限制。",
- created_by="admin-01",
- )
- service.index_document(document["document_id"])
- service.test_search(document["document_id"], "医疗险等待期")
- service.publish_document(document["document_id"])
- result = service.search("医疗险等待期", limit=1)
- assert result["total"] == 1
- assert result["items"][0]["content"]
- def _chunk(chunk_id: str, content: str) -> KnowledgeChunk:
- return KnowledgeChunk(
- id=chunk_id,
- document_id="document-01",
- document_version=1,
- ordinal=1,
- title="测试知识",
- content=content,
- document_type="FAQ",
- source_name="测试知识.md",
- product_code=None,
- )
|