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, )