test_bge_reranker.py 1.4 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253
  1. from typing import Any
  2. from zbt.core.config import Settings
  3. from zbt.domains.knowledge.models import KnowledgeSearchHit
  4. from zbt.infrastructure.embedding.reranker import BgeKnowledgeReranker
  5. class FakeRerankerModel:
  6. def compute_score(
  7. self,
  8. pairs: list[list[str]],
  9. *,
  10. normalize: bool,
  11. ) -> list[float]:
  12. assert normalize is True
  13. assert len(pairs) == 2
  14. return [0.21, 0.93]
  15. def test_bge_reranker_reorders_candidates_by_model_score() -> None:
  16. settings = Settings(
  17. app_env="test",
  18. jwt_access_secret="a" * 32,
  19. jwt_refresh_secret="b" * 32,
  20. field_encryption_key="c" * 32,
  21. )
  22. reranker = BgeKnowledgeReranker(settings, model=FakeRerankerModel())
  23. result = reranker.rerank(
  24. "等待期是多少天",
  25. [_hit("general"), _hit("exact")],
  26. limit=1,
  27. )
  28. assert len(result) == 1
  29. assert result[0].chunk_id == "exact"
  30. assert result[0].score == 0.93
  31. def _hit(chunk_id: str) -> KnowledgeSearchHit:
  32. values: dict[str, Any] = {
  33. "chunk_id": chunk_id,
  34. "document_id": "document-01",
  35. "document_version": 1,
  36. "ordinal": 1,
  37. "title": "保险条款",
  38. "content": f"{chunk_id}内容",
  39. "document_type": "INSURANCE_TERMS",
  40. "source_name": "保险条款.md",
  41. "product_code": "MED-BASIC",
  42. "score": 0.1,
  43. }
  44. return KnowledgeSearchHit(**values)