from __future__ import annotations import pytest from script_build_host.application.retrieval_pipeline import ( ProviderAttempt, ProviderRetrievalResult, RetrievalLease, RetrievalPipeline, RetrievalPipelineResult, RetrievalStatus, ) from script_build_host.domain.errors import ProtocolViolation from script_build_host.domain.evidence_requirements import ( EvidenceProvider, EvidencePurpose, EvidenceRequirementV1, ) class MemoryLeases: def __init__(self) -> None: self.epochs: dict[str, int] = {} self.active: dict[str, RetrievalLease] = {} self.terminals: dict[str, RetrievalPipelineResult] = {} self.attempts: list[ProviderAttempt] = [] self.provider_results: list[ tuple[EvidenceProvider, int, ProviderRetrievalResult] ] = [] self.renewals = 0 async def reserve( self, retrieval_key: str, *, owner: str, lease_seconds: int ) -> RetrievalLease: if retrieval_key in self.terminals: return RetrievalLease( retrieval_key, owner, self.epochs[retrieval_key], False, self.terminals[retrieval_key], ) if retrieval_key in self.active: return RetrievalLease( retrieval_key, owner, self.active[retrieval_key].epoch, False ) epoch = self.epochs.get(retrieval_key, 0) + 1 self.epochs[retrieval_key] = epoch lease = RetrievalLease(retrieval_key, owner, epoch, True) self.active[retrieval_key] = lease return lease async def append_attempt( self, lease: RetrievalLease, attempt: ProviderAttempt ) -> None: assert self.active[lease.retrieval_key] == lease self.attempts.append(attempt) async def append_provider_result( self, lease: RetrievalLease, provider: EvidenceProvider, number: int, result: ProviderRetrievalResult, ) -> None: assert self.active[lease.retrieval_key] == lease self.provider_results.append((provider, number, result)) async def renew(self, lease: RetrievalLease, *, lease_seconds: int) -> None: assert self.active[lease.retrieval_key] == lease assert lease_seconds > 0 self.renewals += 1 async def complete( self, lease: RetrievalLease, result: RetrievalPipelineResult ) -> RetrievalPipelineResult: assert self.active.pop(lease.retrieval_key) == lease self.terminals[lease.retrieval_key] = result return result class Empty: def __init__(self) -> None: self.calls = 0 async def retrieve(self, **_: object) -> ProviderRetrievalResult: self.calls += 1 return ProviderRetrievalResult((), "empty") class Success: def __init__(self, *, partial: bool = False) -> None: self.calls = 0 self.partial = partial async def retrieve(self, **_: object) -> ProviderRetrievalResult: self.calls += 1 return ProviderRetrievalResult(("note:1",), "ok", partial=self.partial) class ResumedLeaseRepository(MemoryLeases): def __init__(self, journal: tuple[dict[str, object], ...]) -> None: super().__init__() self.journal = journal async def reserve( self, retrieval_key: str, *, owner: str, lease_seconds: int ) -> RetrievalLease: lease = RetrievalLease( retrieval_key, owner, 2, True, journal=self.journal, ) self.active[retrieval_key] = lease return lease @pytest.mark.asyncio async def test_empty_knowledge_falls_through_once_to_partial_xhs() -> None: knowledge, xhs = Empty(), Success(partial=True) repository = MemoryLeases() pipeline = RetrievalPipeline( { EvidenceProvider.KNOWLEDGE: knowledge, EvidenceProvider.EXTERNAL_XHS: xhs, }, repository, owner="host-a", ) requirement = EvidenceRequirementV1( "persona", EvidencePurpose.PERSONA, True, 1, source_chain=(EvidenceProvider.KNOWLEDGE, EvidenceProvider.EXTERNAL_XHS), ) result = await pipeline.require(requirement, query={"account": "a"}, snapshot={}) replay = await pipeline.require(requirement, query={"account": "a"}, snapshot={}) assert result.status is RetrievalStatus.PARTIAL_SUCCESS assert replay == result assert knowledge.calls == xhs.calls == 1 @pytest.mark.asyncio async def test_input_satisfied_never_reserves_or_calls_provider() -> None: provider = Success() repository = MemoryLeases() requirement = EvidenceRequirementV1( "persona", EvidencePurpose.PERSONA, True, 1, satisfied_by_input_ref="input:persona", ) result = await RetrievalPipeline( {EvidenceProvider.EXTERNAL_XHS: provider}, repository, owner="host-a", ).execute(requirement, query={}, snapshot={}) assert result.status is RetrievalStatus.SATISFIED_BY_INPUT assert provider.calls == 0 assert repository.epochs == {} def test_provider_result_rejects_empty_evidence_identity() -> None: with pytest.raises(ProtocolViolation, match="empty evidence"): ProviderRetrievalResult(("",), "invalid") def test_persisted_null_evidence_identity_is_rejected() -> None: payload = RetrievalPipelineResult( "persona", RetrievalStatus.SUCCESS, EvidenceProvider.EXTERNAL_XHS, (), (), ).to_payload() payload["evidence_items"] = [ { "requirement_id": "persona", "evidence_ref": None, "provider": "external_xhs", "concepts": [], } ] with pytest.raises(ProtocolViolation, match="no evidence reference"): RetrievalPipelineResult.from_payload(payload) @pytest.mark.asyncio async def test_expired_lease_reuses_durable_provider_result_without_external_call() -> None: provider = Success() repository = ResumedLeaseRepository( ( { "event_type": "provider_result", "provider": "external_xhs", "number": 1, "result": { "source_refs": ["note:durable"], "summary": "persisted", "concepts": [], "limitations": [], "partial": False, "metadata": {}, }, }, ) ) requirement = EvidenceRequirementV1( "persona", EvidencePurpose.PERSONA, True, 1, source_chain=(EvidenceProvider.EXTERNAL_XHS,), ) result = await RetrievalPipeline( {EvidenceProvider.EXTERNAL_XHS: provider}, repository, owner="host-b", ).require(requirement, query={"account": "a"}, snapshot={}) assert result.evidence_items[0].evidence_ref == "note:durable" assert provider.calls == 0 @pytest.mark.asyncio @pytest.mark.parametrize("invalid_ref", [None, 123]) async def test_durable_provider_journal_rejects_invalid_evidence_identity( invalid_ref: object, ) -> None: provider = Success() repository = ResumedLeaseRepository( ( { "event_type": "provider_result", "provider": "external_xhs", "number": 1, "result": { "source_refs": [invalid_ref], "summary": "invalid persisted result", }, }, ) ) requirement = EvidenceRequirementV1( "persona", EvidencePurpose.PERSONA, True, 1, source_chain=(EvidenceProvider.EXTERNAL_XHS,), ) with pytest.raises(ProtocolViolation, match="no evidence reference"): await RetrievalPipeline( {EvidenceProvider.EXTERNAL_XHS: provider}, repository, owner="host-b", ).require(requirement, query={"account": "a"}, snapshot={}) assert provider.calls == 0