test_retrieval.py 2.6 KB

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