diff --git a/README.md b/README.md index 61bd438..4ae2125 100644 --- a/README.md +++ b/README.md @@ -97,8 +97,9 @@ Per-iteration ephemeral tail messages -- `hooks-status-context`, `inject_context` hook result -- are, by default (`"persist"`), written into canonical context via `context.add_message(...)`, and only when the text differs from the last text this orchestrator persisted. When unchanged, -nothing is added, so request N is a true, append-only prefix of request -N+1. This matters because OpenAI's (and most providers') implicit/explicit +no duplicate is added; between compactions and other request rewrites, this +preserves the opportunity for an append-only request prefix. This matters +because OpenAI's (and most providers') implicit/explicit prompt cache reuses only the longest true prefix of a prior request: the original `"tail"` behavior re-generates and re-appends these messages at the tail of every request, positionally displacing the assistant/tool turn @@ -108,6 +109,10 @@ transcript as a fresh cache write on every call. **Evidence for the default:** +The historical observations below are not a guarantee of current long-history +cache reuse. Request retention protects instruction delivery, not cache +performance; measure reuse from the provider's raw total/read/write counters. + - **OpenAI**: a pre-registered 9-arm live probe found only the persist design (change-gated, canonical-context write) heals prefix reuse (98.9%); byte-stable tails and folding into the tool result do not heal @@ -141,6 +146,20 @@ supported and is byte-identical to the module's original, pre-this-feature behavior. An unknown value falls back to `"persist"` (the current default) with a logged warning. +**Retention capability:** when the active context exposes the callable +`context.request_retention` capability, persist mode requires the current +complete reminder envelope in every request view, including unchanged-body +suppression, pending tool feedback, and bounded finalization. A fresh +orchestrator also reuses an exact matching admitted persisted envelope after +resume instead of adding a duplicate. Required content that cannot fit fails +visibly before a provider request. + +If the capability is unavailable, persist mode keeps legacy assembly and logs +one warning per `execute()` that complete reminder retention is unavailable. +Explicit `tail` mode is unchanged and does not receive this retention +guarantee. None of these changes gives user-carried reminders native +system/developer authority. + ## System-reminder envelope and placement (reminder-redesign-spec.md, W1) ```toml diff --git a/amplifier_module_loop_streaming/__init__.py b/amplifier_module_loop_streaming/__init__.py index 2d29ba1..ad070a2 100644 --- a/amplifier_module_loop_streaming/__init__.py +++ b/amplifier_module_loop_streaming/__init__.py @@ -1176,6 +1176,12 @@ def __init__(self, config: dict[str, Any]): _ephemeral_injection_mode = "persist" self._ephemeral_injection_mode: str = _ephemeral_injection_mode self._last_persisted_injection: str | None = None + self._retention_capability_warned: bool = False + # The exact enveloped message admitted to canonical context. Keeping + # both the raw producer body and this wire representation lets a + # resumed orchestrator recognize an already-admitted reminder without + # confusing a user message that happens to quote the same XML. + self._last_persisted_injection_content: str | None = None # D2 fix (rr wave 20260831 -- envelope accumulation). True once a # PERSISTED reminder envelope carrying the full descriptive header # has been written for the CURRENT turn; every subsequent persisted @@ -1316,6 +1322,7 @@ async def execute( # turn. self._goal_model_cache = None self._goal_model_basis = None + self._retention_capability_warned = False # Peek at goal state *before* the first turn. Goal state can only be # set (by the app layer's /goal command) before execute() is called, @@ -3152,6 +3159,70 @@ async def _evaluate_goal( reason = lines[1] if len(lines) > 1 else "(evaluator gave no reason)" return verdict == "YES", reason + async def _persist_reminder( + self, context, body: str, *, tail: bool, verify_admitted: bool = False + ) -> tuple[str, bool]: + """Return the current admitted envelope without duplicating it.""" + content = _wrap_reminders( + body, tail=tail, header=not self._turn_header_persisted + ) + + def is_admitted(message: dict[str, Any], expected_content: str) -> bool: + return ( + message.get("role") == "user" + and message.get("content") == expected_content + and (message.get("metadata") or {}).get("ephemeral") is True + and (message.get("metadata") or {}).get("persisted") is True + ) + + # The raw-body gate intentionally spans the pre-user and tail envelope + # variants within one turn. Reuse the admitted pre-user variant instead + # of making a second, tail-wrapped copy just because its header differs. + if ( + body == self._last_persisted_injection + and self._last_persisted_injection_content is not None + ): + if not verify_admitted: + return self._last_persisted_injection_content, False + canonical = await context.get_messages() + if any( + is_admitted(message, self._last_persisted_injection_content) + for message in canonical + ): + return self._last_persisted_injection_content, False + + if verify_admitted: + # Retention-aware contexts may compact their request view. Search + # canonical history for the exact *current* envelope instead: a + # fresh/resumed orchestrator has no local body cache, and content + # equality alone would incorrectly reuse a real user quote. + canonical = await context.get_messages() + admitted = any( + is_admitted(message, content) + for message in canonical + ) + if admitted: + self._last_persisted_injection = body + self._last_persisted_injection_content = content + self._turn_header_persisted = True + return content, False + + await context.add_message( + { + "role": "user", + "content": content, + "metadata": { + "ephemeral": True, + "persisted": True, + "reminder_placement": "tail" if tail else "pre_user", + }, + } + ) + self._last_persisted_injection = body + self._last_persisted_injection_content = content + self._turn_header_persisted = True + return content, True + async def _execute_stream( self, prompt: str, @@ -3165,6 +3236,32 @@ async def _execute_stream( Internal streaming execution. Yields tuples of (token, iteration) as they're generated. """ + retaining_getter = None + get_capability = getattr(coordinator, "get_capability", None) + if self._ephemeral_injection_mode == "persist" and callable(get_capability): + candidate = get_capability("context.request_retention") + if callable(candidate): + retaining_getter = candidate + if ( + self._ephemeral_injection_mode == "persist" + and retaining_getter is None + and not self._retention_capability_warned + ): + logger.warning( + "Persisted reminder retention guarantee unavailable: " + "context.request_retention is not callable; using legacy " + "context request fallback." + ) + self._retention_capability_warned = True + + async def request_messages(retain_contents: list[str]): + if retaining_getter is not None: + return await retaining_getter( + provider=provider, retain_contents=retain_contents + ) + return await context.get_messages_for_request(provider=provider) + + turn_start_retained_contents: list[str] = [] # Emit and process prompt submit (allows hooks to inject context before processing) prompt_submit_result = await hooks.emit(PROMPT_SUBMIT, {"prompt": prompt}) if coordinator: @@ -3311,34 +3408,13 @@ async def _execute_stream( turn_start_body = "\n\n".join(turn_start_parts) if turn_start_body: if self._ephemeral_injection_mode == "persist": - # Same change-gate as the in-loop persist path, and the - # SAME comparison basis: the RAW (pre-envelope) merged - # body, not the enveloped string. This is what lets an - # unchanged body suppress correctly even when this - # turn-start (pre_user-headered) block is compared - # against a LATER mid-loop (tail-headered) block wrapping - # the identical text -- the two envelope strings would - # never be equal (different header), which would - # otherwise force a spurious extra persist on the first - # mid-loop iteration of every multi-iteration turn. - if turn_start_body != self._last_persisted_injection: - await context.add_message( - { - "role": "user", - "content": _wrap_reminders( - turn_start_body, - tail=False, - header=not self._turn_header_persisted, - ), - "metadata": { - "ephemeral": True, - "persisted": True, - "reminder_placement": "pre_user", - }, - } - ) - self._last_persisted_injection = turn_start_body - self._turn_header_persisted = True + content, _ = await self._persist_reminder( + context, + turn_start_body, + tail=False, + verify_admitted=retaining_getter is not None, + ) + turn_start_retained_contents.append(content) else: # tail injection mode: never persisted into canonical # context; spliced into the request VIEW at iteration 1 @@ -3499,9 +3575,28 @@ async def exit_for_cancellation() -> None: yield (f"Operation denied: {result.reason}", iteration) return + # Keep the current producer's admitted reminder through compaction + # on every request. A changed body is admitted below before the + # request is assembled. + retained_contents = ( + list(turn_start_retained_contents) if iteration == 1 else [] + ) + if ( + result.action == "inject_context" + and result.ephemeral + and result.context_injection + and retaining_getter is not None + ): + content, _ = await self._persist_reminder( + context, + result.context_injection, + tail=True, + verify_admitted=True, + ) + retained_contents.append(content) # Get messages for LLM request (context handles compaction internally) # Pass provider for dynamic budget calculation based on model's context window - message_dicts = await context.get_messages_for_request(provider=provider) + message_dicts = await request_messages(retained_contents) message_dicts = list(message_dicts) # Convert to list for modification # Splice the turn-start reminder block into the request VIEW @@ -3566,64 +3661,12 @@ async def exit_for_cancellation() -> None: result.context_injection_role, ) if self._ephemeral_injection_mode == "persist": - # Ephemeral-cache fix (ephemeral-cache-fix-spec.md sec 5.1/5.3): - # write the injection into CANONICAL context via - # context.add_message(...), and ONLY when its RAW text - # (pre-envelope) differs from the last text this - # orchestrator persisted. When unchanged, do nothing -- - # the request is then a pure append of the new - # assistant/tool turn, which is what makes request N a - # true prefix of request N+1. Comparing on the raw body - # (not the enveloped string) keeps the change-gate - # working across the pre-user/tail header-variant - # boundary -- the turn-start block and this mid-loop - # block wrap the SAME body differently (different - # header), so comparing enveloped strings would falsely - # look "changed" the first time a turn transitions from - # its turn-start block to a mid-loop one. - # - # Contract change accepted knowingly (spec sec 5.2): once - # persisted, this message is no longer "removed next - # turn" -- it is real history from here on, still marked - # metadata.ephemeral=True (now meaning "machine-generated - # per-turn scaffolding, not a user turn", not "guaranteed - # absent next turn" -- see models.py's updated docstring). - if result.context_injection != self._last_persisted_injection: - await context.add_message( - { - "role": "user", - "content": _wrap_reminders( - result.context_injection, - tail=True, - header=not self._turn_header_persisted, - ), - "metadata": { - "ephemeral": True, - "persisted": True, - "reminder_placement": "tail", - }, - } - ) - self._last_persisted_injection = result.context_injection - self._turn_header_persisted = True - logger.debug( - "Persisted changed ephemeral injection into canonical context" - ) - # Re-fetch so the newly persisted message is present - # and budgeted on THIS request too, not just the next - # one. get_messages_for_request is pure w.r.t. - # self.messages (context-simple returns a new list; - # it never mutates in place), so this is an extra - # call, not a reordering. - message_dicts = await context.get_messages_for_request( - provider=provider - ) - message_dicts = list(message_dicts) - else: - logger.debug( - "Ephemeral injection text unchanged -- skipping persist " - "(change-gate); request is a pure append this iteration" + if retaining_getter is None: + _, changed = await self._persist_reminder( + context, result.context_injection, tail=True ) + if changed: + message_dicts = list(await request_messages([])) # Check if we should append to last tool result elif result.append_to_last_tool_result and len(message_dicts) > 0: last_msg = message_dicts[-1] @@ -3717,36 +3760,16 @@ async def exit_for_cancellation() -> None: pending_to_tool_result + pending_to_message ) if pending_body: - if pending_body != self._last_persisted_injection: - await context.add_message( - { - "role": "user", - "content": _wrap_reminders( - pending_body, - tail=True, - header=not self._turn_header_persisted, - ), - "metadata": { - "ephemeral": True, - "persisted": True, - "reminder_placement": "tail", - }, - } - ) - self._last_persisted_injection = pending_body - self._turn_header_persisted = True - logger.debug( - "Persisted changed pending ephemeral injection(s) " - "into canonical context" - ) - message_dicts = await context.get_messages_for_request( - provider=provider - ) - message_dicts = list(message_dicts) - else: - logger.debug( - "Pending ephemeral injection(s) text unchanged -- " - "skipping persist (change-gate)" + content, changed = await self._persist_reminder( + context, + pending_body, + tail=True, + verify_admitted=retaining_getter is not None, + ) + retained_contents.append(content) + if changed or retaining_getter is not None: + message_dicts = list( + await request_messages(retained_contents) ) else: if pending_to_tool_result: @@ -4355,9 +4378,35 @@ async def exit_for_cancellation() -> None: # allowed provider call, not discarded at the next turn boundary. await self._drain_steering(context, hooks, iteration) - # Get one final response with the reminder (via _execute_stream helper) - message_dicts = await context.get_messages_for_request(provider=provider) - message_dicts = list(message_dicts) + # Current provider-hook requirements and a queued tool-post + # injection must survive the bounded finalization assembly too. + final_retained_contents: list[str] = [] + if ( + retaining_getter is not None + and finalization_result.action == "inject_context" + and finalization_result.ephemeral + and finalization_result.context_injection + ): + content, _ = await self._persist_reminder( + context, + finalization_result.context_injection, + tail=True, + verify_admitted=True, + ) + final_retained_contents.append(content) + if retaining_getter is not None and self._pending_ephemeral_injections: + pending_body = "\n\n".join( + injection["content"] + for injection in self._pending_ephemeral_injections + if injection.get("content") + ) + if pending_body: + content, _ = await self._persist_reminder( + context, pending_body, tail=True, verify_admitted=True + ) + final_retained_contents.append(content) + self._pending_ephemeral_injections.clear() + message_dicts = list(await request_messages(final_retained_contents)) # The finalization hook and context assembly both await. Check # again before contacting the provider so a concurrent # cancellation cannot buy an unrequested final provider call. diff --git a/tests/test_request_retention.py b/tests/test_request_retention.py new file mode 100644 index 0000000..6bff389 --- /dev/null +++ b/tests/test_request_retention.py @@ -0,0 +1,292 @@ +"""Retention-aware request assembly and resume-deduplication coverage.""" + +from __future__ import annotations + +import logging + +import pytest +from amplifier_core import ContextLengthError + +from amplifier_module_loop_streaming import StreamingOrchestrator, _wrap_reminders +from tests.test_ephemeral_cache_persist_mode import ( + MockContext, + MockCoordinator, + NRoundToolProvider, + OneShotTool, + RequestCapturingProvider, + ScriptedHookResult, + ScriptedHooks, + SequencedHooks, +) + + +class RetainingContext(MockContext): + """A context whose request view compacts persisted messages not retained.""" + + def __init__(self, *, fail_retention: bool = False) -> None: + super().__init__() + self.requirements: list[list[str]] = [] + self.fail_retention = fail_retention + + async def get_messages(self) -> list[dict]: + return list(self._messages) + + async def retaining_view(self, *, provider=None, retain_contents: list[str]) -> list[dict]: + self.requirements.append(list(retain_contents)) + if self.fail_retention: + raise ContextLengthError("retention cannot fit") + return [ + dict(message) + if not (message.get("metadata") or {}).get("persisted") + or message["content"] in retain_contents + else {**message, "content": "[compacted]"} + for message in self._messages + ] + + +def injection(body: str) -> ScriptedHookResult: + return ScriptedHookResult( + action="inject_context", ephemeral=True, context_injection=body + ) + + +def coordinator_for(context: RetainingContext) -> MockCoordinator: + coordinator = MockCoordinator() + coordinator.register_capability("context.request_retention", context.retaining_view) + return coordinator + + +def reminder_contents(request) -> list[str]: + return [ + message.content + for message in request.messages + if isinstance(message.content, str) + and message.content.startswith("") + ] + + +def persisted(context: RetainingContext) -> list[dict]: + return [ + message + for message in context.add_message_calls + if (message.get("metadata") or {}).get("persisted") is True + ] + + +@pytest.mark.asyncio +async def test_changed_content_is_retained_and_sent_as_the_current_envelope() -> None: + old_body = "OLD" + new_body = "NEW" + old = _wrap_reminders(old_body, tail=False) + new = _wrap_reminders(new_body, tail=True, header=False) + context = RetainingContext() + provider = NRoundToolProvider(n_tool_rounds=1) + hooks = SequencedHooks({"provider:request": [injection(old_body), injection(new_body)]}) + + await StreamingOrchestrator({}).execute( + "work", + context, + {"main": provider}, + {"mock_tool": OneShotTool()}, + hooks, + coordinator_for(context), + ) + + assert context.requirements == [[old], [new]] + assert [reminder_contents(request) for request in provider.requests] == [[old], [new]] + assert "OLD" not in "\n".join(message.content for message in provider.requests[-1].messages) + + +@pytest.mark.asyncio +async def test_unchanged_content_reuses_its_admitted_envelope_each_request() -> None: + body = "STABLE" + admitted = _wrap_reminders(body, tail=False) + context = RetainingContext() + provider = NRoundToolProvider(n_tool_rounds=2) + hooks = ScriptedHooks({"provider:request": injection(body)}) + + await StreamingOrchestrator({}).execute( + "work", + context, + {"main": provider}, + {"mock_tool": OneShotTool()}, + hooks, + coordinator_for(context), + ) + + assert context.requirements == [[admitted], [admitted], [admitted]] + assert [reminder_contents(request) for request in provider.requests] == [ + [admitted], + [admitted], + [admitted], + ] + assert [message["content"] for message in persisted(context)] == [admitted] + + +@pytest.mark.asyncio +async def test_pending_tool_injection_drains_into_the_retained_request() -> None: + provider_body = "PROVIDER" + tool_body = "TOOL" + provider_envelope = _wrap_reminders(provider_body, tail=False) + tool_envelope = _wrap_reminders(tool_body, tail=True, header=False) + context = RetainingContext() + provider = NRoundToolProvider(n_tool_rounds=1) + hooks = ScriptedHooks( + { + "provider:request": injection(provider_body), + "tool:post": injection(tool_body), + } + ) + + await StreamingOrchestrator({}).execute( + "work", + context, + {"main": provider}, + {"mock_tool": OneShotTool()}, + hooks, + coordinator_for(context), + ) + + assert context.requirements == [ + [provider_envelope], + [provider_envelope], + [provider_envelope, tool_envelope], + ] + assert reminder_contents(provider.requests[-1]) == [provider_envelope, tool_envelope] + + +@pytest.mark.asyncio +async def test_finalization_retains_current_hook_and_pending_tool_injection() -> None: + initial_body = "INITIAL" + final_body = "FINAL" + tool_body = "TOOL" + initial = _wrap_reminders(initial_body, tail=False) + final = _wrap_reminders(final_body, tail=True, header=False) + tool = _wrap_reminders(tool_body, tail=True, header=False) + context = RetainingContext() + provider = NRoundToolProvider(n_tool_rounds=1) + hooks = SequencedHooks( + { + "provider:request": [injection(initial_body), injection(final_body)], + "tool:post": [injection(tool_body)], + } + ) + + await StreamingOrchestrator({"max_iterations": 1}).execute( + "work", + context, + {"main": provider}, + {"mock_tool": OneShotTool()}, + hooks, + coordinator_for(context), + ) + + assert len(provider.requests) == 2 + assert context.requirements == [[initial], [final, tool]] + final_request_reminders = reminder_contents(provider.requests[-1]) + assert final_request_reminders[:2] == [final, tool] + assert 'source="orchestrator-loop-limit"' in final_request_reminders[-1] + + +@pytest.mark.asyncio +async def test_new_orchestrator_resumes_an_admitted_reminder_without_writing_again() -> None: + body = "RESUME" + admitted = _wrap_reminders(body, tail=False) + context = RetainingContext() + provider = RequestCapturingProvider() + hooks = ScriptedHooks({"provider:request": injection(body)}) + coordinator = coordinator_for(context) + + await StreamingOrchestrator({}).execute( + "first", context, {"main": provider}, {}, hooks, coordinator + ) + await StreamingOrchestrator({}).execute( + "second", context, {"main": provider}, {}, hooks, coordinator + ) + + assert [message["content"] for message in persisted(context)] == [admitted] + assert context.requirements == [[admitted], [admitted]] + assert reminder_contents(provider.requests[-1]) == [admitted] + + +@pytest.mark.asyncio +async def test_context_reset_readmits_but_a_human_xml_quote_does_not_deduplicate() -> None: + body = "QUOTE" + admitted = _wrap_reminders(body, tail=False) + context = RetainingContext() + provider = RequestCapturingProvider() + hooks = ScriptedHooks({"provider:request": injection(body)}) + loop = StreamingOrchestrator({}) + coordinator = coordinator_for(context) + + await loop.execute("first", context, {"main": provider}, {}, hooks, coordinator) + context._messages.clear() + await loop.execute("second", context, {"main": provider}, {}, hooks, coordinator) + assert [message["content"] for message in persisted(context)] == [admitted, admitted] + + quoted_context = RetainingContext() + quoted_context._messages.append({"role": "user", "content": admitted}) + await StreamingOrchestrator({}).execute( + "quoted", quoted_context, {"main": RequestCapturingProvider()}, {}, hooks, + coordinator_for(quoted_context), + ) + assert [message["content"] for message in persisted(quoted_context)] == [admitted] + assert sum(message["content"] == admitted for message in quoted_context._messages) == 2 + + +@pytest.mark.asyncio +async def test_missing_retention_capability_warns_once_and_tail_remains_unmodified( + caplog, +) -> None: + body = "LEGACY" + context = MockContext() + provider = RequestCapturingProvider() + with caplog.at_level(logging.WARNING): + await StreamingOrchestrator({}).execute( + "work", + context, + {"main": provider}, + {}, + ScriptedHooks({"provider:request": injection(body)}), + MockCoordinator(), + ) + warnings = [ + record + for record in caplog.records + if "retention guarantee unavailable" in record.message + ] + assert len(warnings) == 1 + assert "legacy context request fallback" in warnings[0].message + + tail_context = RetainingContext() + tail_provider = RequestCapturingProvider() + await StreamingOrchestrator({"ephemeral_injection_mode": "tail"}).execute( + "work", + tail_context, + {"main": tail_provider}, + {}, + ScriptedHooks({"provider:request": injection(body)}), + coordinator_for(tail_context), + ) + assert tail_context.requirements == [] + assert persisted(tail_context) == [] + assert reminder_contents(tail_provider.requests[0]) == [_wrap_reminders(body, tail=False)] + + +@pytest.mark.asyncio +async def test_context_length_during_retention_makes_no_provider_call() -> None: + context = RetainingContext(fail_retention=True) + provider = RequestCapturingProvider() + + with pytest.raises(ContextLengthError, match="retention cannot fit"): + await StreamingOrchestrator({}).execute( + "work", + context, + {"main": provider}, + {}, + ScriptedHooks({"provider:request": injection("FIT")}), + coordinator_for(context), + ) + + assert context.requirements + assert provider.requests == [] \ No newline at end of file