test_retrieval.py 3.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104
  1. from __future__ import annotations
  2. def _evidence(key: str, score: float = 0.8, domain: str = "domain-a"):
  3. from app.core.knowledge.retrieval.contracts import KnowledgeEvidence
  4. return KnowledgeEvidence(
  5. chunk_id=f"chunk-{key}",
  6. content=f"content-{key}",
  7. score=score,
  8. retriever="test",
  9. object_uid=f"object-{key}",
  10. object_type="DataFlow",
  11. object_version=2,
  12. business_domain_uid=domain,
  13. point_keys=(f"DataFlow/object-{key}/purpose",),
  14. point_revisions=(3,),
  15. generation=1,
  16. source_updated_at="2026-07-23T00:00:00+00:00",
  17. )
  18. def test_router_is_deterministic_for_exact_relationship_and_semantic_queries():
  19. from app.core.knowledge.retrieval.router import route_query
  20. assert route_query('字段 "customer_id"') == "exact"
  21. assert route_query("customer_sync 的上游是什么") == "relationship"
  22. assert route_query("客户同步的用途是什么") == "semantic"
  23. def test_rrf_fuses_duplicate_stable_chunks_and_keeps_provenance():
  24. from app.core.knowledge.retrieval.fusion import reciprocal_rank_fusion
  25. lexical = [_evidence("a", 0.8), _evidence("b", 0.7)]
  26. vector = [_evidence("b", 0.9), _evidence("c", 0.6)]
  27. fused = reciprocal_rank_fusion({"lexical": lexical, "vector": vector}, limit=3)
  28. assert [item.chunk_id for item in fused] == ["chunk-b", "chunk-a", "chunk-c"]
  29. assert fused[0].retriever == "lexical+vector"
  30. def test_pipeline_reauthorizes_all_candidates_and_degrades_failed_retriever():
  31. from app.core.knowledge.access import KnowledgeAccessContext
  32. from app.core.knowledge.retrieval.pipeline import KnowledgeRetrievalPipeline
  33. class Retriever:
  34. def __init__(self, values=None, error=None):
  35. self.values = values or []
  36. self.error = error
  37. def retrieve(self, _query, _context, _limit):
  38. if self.error:
  39. raise self.error
  40. return self.values
  41. context = KnowledgeAccessContext(
  42. subject_id="user-1",
  43. roles=frozenset({"viewer"}),
  44. permissions=frozenset({"governance:read"}),
  45. business_domain_uids=frozenset({"domain-a"}),
  46. correlation_id="correlation-1",
  47. )
  48. pipeline = KnowledgeRetrievalPipeline(
  49. lexical=Retriever([_evidence("allowed"), _evidence("denied", 0.9, "domain-b")]),
  50. vector=Retriever(error=RuntimeError("pgvector unavailable")),
  51. )
  52. result = pipeline.search("用途", context=context, mode="semantic")
  53. assert [item.chunk_id for item in result.evidence] == ["chunk-allowed"]
  54. assert result.degraded_components == ("vector",)
  55. def test_pipeline_fuses_device_evidence_and_reauthorizes_it_after_retrieval():
  56. from app.core.knowledge.access import KnowledgeAccessContext
  57. from app.core.knowledge.retrieval.pipeline import KnowledgeRetrievalPipeline
  58. class Retriever:
  59. def __init__(self, values=()):
  60. self.values = values
  61. def retrieve(self, _query, _context, _limit):
  62. return self.values
  63. context = KnowledgeAccessContext(
  64. subject_id="user-1",
  65. roles=frozenset({"viewer"}),
  66. permissions=frozenset({"governance:read"}),
  67. business_domain_uids=frozenset({"domain-a"}),
  68. correlation_id="correlation-1",
  69. )
  70. allowed = _evidence("device-a", 0.95, "domain-a")
  71. denied = _evidence("device-b", 0.99, "domain-b")
  72. pipeline = KnowledgeRetrievalPipeline(
  73. lexical=Retriever(),
  74. vector=Retriever(),
  75. device=Retriever((denied, allowed)),
  76. )
  77. result = pipeline.search("循环泵", context=context, mode="exact")
  78. assert [item.chunk_id for item in result.evidence] == [
  79. "chunk-device-a"
  80. ]