| 1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253 |
- from typing import Any
- from zbt.core.config import Settings
- from zbt.domains.knowledge.models import KnowledgeSearchHit
- from zbt.infrastructure.embedding.reranker import BgeKnowledgeReranker
- class FakeRerankerModel:
- def compute_score(
- self,
- pairs: list[list[str]],
- *,
- normalize: bool,
- ) -> list[float]:
- assert normalize is True
- assert len(pairs) == 2
- return [0.21, 0.93]
- def test_bge_reranker_reorders_candidates_by_model_score() -> None:
- settings = Settings(
- app_env="test",
- jwt_access_secret="a" * 32,
- jwt_refresh_secret="b" * 32,
- field_encryption_key="c" * 32,
- )
- reranker = BgeKnowledgeReranker(settings, model=FakeRerankerModel())
- result = reranker.rerank(
- "等待期是多少天",
- [_hit("general"), _hit("exact")],
- limit=1,
- )
- assert len(result) == 1
- assert result[0].chunk_id == "exact"
- assert result[0].score == 0.93
- def _hit(chunk_id: str) -> KnowledgeSearchHit:
- values: dict[str, Any] = {
- "chunk_id": chunk_id,
- "document_id": "document-01",
- "document_version": 1,
- "ordinal": 1,
- "title": "保险条款",
- "content": f"{chunk_id}内容",
- "document_type": "INSURANCE_TERMS",
- "source_name": "保险条款.md",
- "product_code": "MED-BASIC",
- "score": 0.1,
- }
- return KnowledgeSearchHit(**values)
|