test_architecture_retrieval.py 7.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260
  1. from __future__ import annotations
  2. import pytest
  3. from script_build_host.application.retrieval_pipeline import (
  4. ProviderAttempt,
  5. ProviderRetrievalResult,
  6. RetrievalLease,
  7. RetrievalPipeline,
  8. RetrievalPipelineResult,
  9. RetrievalStatus,
  10. )
  11. from script_build_host.domain.errors import ProtocolViolation
  12. from script_build_host.domain.evidence_requirements import (
  13. EvidenceProvider,
  14. EvidencePurpose,
  15. EvidenceRequirementV1,
  16. )
  17. class MemoryLeases:
  18. def __init__(self) -> None:
  19. self.epochs: dict[str, int] = {}
  20. self.active: dict[str, RetrievalLease] = {}
  21. self.terminals: dict[str, RetrievalPipelineResult] = {}
  22. self.attempts: list[ProviderAttempt] = []
  23. self.provider_results: list[
  24. tuple[EvidenceProvider, int, ProviderRetrievalResult]
  25. ] = []
  26. self.renewals = 0
  27. async def reserve(
  28. self, retrieval_key: str, *, owner: str, lease_seconds: int
  29. ) -> RetrievalLease:
  30. if retrieval_key in self.terminals:
  31. return RetrievalLease(
  32. retrieval_key, owner, self.epochs[retrieval_key], False,
  33. self.terminals[retrieval_key],
  34. )
  35. if retrieval_key in self.active:
  36. return RetrievalLease(
  37. retrieval_key, owner, self.active[retrieval_key].epoch, False
  38. )
  39. epoch = self.epochs.get(retrieval_key, 0) + 1
  40. self.epochs[retrieval_key] = epoch
  41. lease = RetrievalLease(retrieval_key, owner, epoch, True)
  42. self.active[retrieval_key] = lease
  43. return lease
  44. async def append_attempt(
  45. self, lease: RetrievalLease, attempt: ProviderAttempt
  46. ) -> None:
  47. assert self.active[lease.retrieval_key] == lease
  48. self.attempts.append(attempt)
  49. async def append_provider_result(
  50. self,
  51. lease: RetrievalLease,
  52. provider: EvidenceProvider,
  53. number: int,
  54. result: ProviderRetrievalResult,
  55. ) -> None:
  56. assert self.active[lease.retrieval_key] == lease
  57. self.provider_results.append((provider, number, result))
  58. async def renew(self, lease: RetrievalLease, *, lease_seconds: int) -> None:
  59. assert self.active[lease.retrieval_key] == lease
  60. assert lease_seconds > 0
  61. self.renewals += 1
  62. async def complete(
  63. self, lease: RetrievalLease, result: RetrievalPipelineResult
  64. ) -> RetrievalPipelineResult:
  65. assert self.active.pop(lease.retrieval_key) == lease
  66. self.terminals[lease.retrieval_key] = result
  67. return result
  68. class Empty:
  69. def __init__(self) -> None:
  70. self.calls = 0
  71. async def retrieve(self, **_: object) -> ProviderRetrievalResult:
  72. self.calls += 1
  73. return ProviderRetrievalResult((), "empty")
  74. class Success:
  75. def __init__(self, *, partial: bool = False) -> None:
  76. self.calls = 0
  77. self.partial = partial
  78. async def retrieve(self, **_: object) -> ProviderRetrievalResult:
  79. self.calls += 1
  80. return ProviderRetrievalResult(("note:1",), "ok", partial=self.partial)
  81. class ResumedLeaseRepository(MemoryLeases):
  82. def __init__(self, journal: tuple[dict[str, object], ...]) -> None:
  83. super().__init__()
  84. self.journal = journal
  85. async def reserve(
  86. self, retrieval_key: str, *, owner: str, lease_seconds: int
  87. ) -> RetrievalLease:
  88. lease = RetrievalLease(
  89. retrieval_key,
  90. owner,
  91. 2,
  92. True,
  93. journal=self.journal,
  94. )
  95. self.active[retrieval_key] = lease
  96. return lease
  97. @pytest.mark.asyncio
  98. async def test_empty_knowledge_falls_through_once_to_partial_xhs() -> None:
  99. knowledge, xhs = Empty(), Success(partial=True)
  100. repository = MemoryLeases()
  101. pipeline = RetrievalPipeline(
  102. {
  103. EvidenceProvider.KNOWLEDGE: knowledge,
  104. EvidenceProvider.EXTERNAL_XHS: xhs,
  105. },
  106. repository,
  107. owner="host-a",
  108. )
  109. requirement = EvidenceRequirementV1(
  110. "persona",
  111. EvidencePurpose.PERSONA,
  112. True,
  113. 1,
  114. source_chain=(EvidenceProvider.KNOWLEDGE, EvidenceProvider.EXTERNAL_XHS),
  115. )
  116. result = await pipeline.require(requirement, query={"account": "a"}, snapshot={})
  117. replay = await pipeline.require(requirement, query={"account": "a"}, snapshot={})
  118. assert result.status is RetrievalStatus.PARTIAL_SUCCESS
  119. assert replay == result
  120. assert knowledge.calls == xhs.calls == 1
  121. @pytest.mark.asyncio
  122. async def test_input_satisfied_never_reserves_or_calls_provider() -> None:
  123. provider = Success()
  124. repository = MemoryLeases()
  125. requirement = EvidenceRequirementV1(
  126. "persona",
  127. EvidencePurpose.PERSONA,
  128. True,
  129. 1,
  130. satisfied_by_input_ref="input:persona",
  131. )
  132. result = await RetrievalPipeline(
  133. {EvidenceProvider.EXTERNAL_XHS: provider},
  134. repository,
  135. owner="host-a",
  136. ).execute(requirement, query={}, snapshot={})
  137. assert result.status is RetrievalStatus.SATISFIED_BY_INPUT
  138. assert provider.calls == 0
  139. assert repository.epochs == {}
  140. def test_provider_result_rejects_empty_evidence_identity() -> None:
  141. with pytest.raises(ProtocolViolation, match="empty evidence"):
  142. ProviderRetrievalResult(("",), "invalid")
  143. def test_persisted_null_evidence_identity_is_rejected() -> None:
  144. payload = RetrievalPipelineResult(
  145. "persona",
  146. RetrievalStatus.SUCCESS,
  147. EvidenceProvider.EXTERNAL_XHS,
  148. (),
  149. (),
  150. ).to_payload()
  151. payload["evidence_items"] = [
  152. {
  153. "requirement_id": "persona",
  154. "evidence_ref": None,
  155. "provider": "external_xhs",
  156. "concepts": [],
  157. }
  158. ]
  159. with pytest.raises(ProtocolViolation, match="no evidence reference"):
  160. RetrievalPipelineResult.from_payload(payload)
  161. @pytest.mark.asyncio
  162. async def test_expired_lease_reuses_durable_provider_result_without_external_call() -> None:
  163. provider = Success()
  164. repository = ResumedLeaseRepository(
  165. (
  166. {
  167. "event_type": "provider_result",
  168. "provider": "external_xhs",
  169. "number": 1,
  170. "result": {
  171. "source_refs": ["note:durable"],
  172. "summary": "persisted",
  173. "concepts": [],
  174. "limitations": [],
  175. "partial": False,
  176. "metadata": {},
  177. },
  178. },
  179. )
  180. )
  181. requirement = EvidenceRequirementV1(
  182. "persona",
  183. EvidencePurpose.PERSONA,
  184. True,
  185. 1,
  186. source_chain=(EvidenceProvider.EXTERNAL_XHS,),
  187. )
  188. result = await RetrievalPipeline(
  189. {EvidenceProvider.EXTERNAL_XHS: provider},
  190. repository,
  191. owner="host-b",
  192. ).require(requirement, query={"account": "a"}, snapshot={})
  193. assert result.evidence_items[0].evidence_ref == "note:durable"
  194. assert provider.calls == 0
  195. @pytest.mark.asyncio
  196. @pytest.mark.parametrize("invalid_ref", [None, 123])
  197. async def test_durable_provider_journal_rejects_invalid_evidence_identity(
  198. invalid_ref: object,
  199. ) -> None:
  200. provider = Success()
  201. repository = ResumedLeaseRepository(
  202. (
  203. {
  204. "event_type": "provider_result",
  205. "provider": "external_xhs",
  206. "number": 1,
  207. "result": {
  208. "source_refs": [invalid_ref],
  209. "summary": "invalid persisted result",
  210. },
  211. },
  212. )
  213. )
  214. requirement = EvidenceRequirementV1(
  215. "persona",
  216. EvidencePurpose.PERSONA,
  217. True,
  218. 1,
  219. source_chain=(EvidenceProvider.EXTERNAL_XHS,),
  220. )
  221. with pytest.raises(ProtocolViolation, match="no evidence reference"):
  222. await RetrievalPipeline(
  223. {EvidenceProvider.EXTERNAL_XHS: provider},
  224. repository,
  225. owner="host-b",
  226. ).require(requirement, query={"account": "a"}, snapshot={})
  227. assert provider.calls == 0