diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index e41a6ef..bca3e31 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -62,7 +62,11 @@ jobs: python-version: ${{ matrix.python-version }} - name: Install dependencies - run: uv sync --all-extras --dev + run: uv sync --locked --all-extras --dev + + - name: Verify released Core API + run: | + uv run --locked python -c "from importlib.metadata import version; from amplifier_core.llm_errors import ContextLengthError; assert version('amplifier-core') == '1.6.1'" # No API key placeholder here, unlike provider-openai's CI. Measured: # with ANTHROPIC_API_KEY/OPENAI_API_KEY unset, the suite is @@ -72,4 +76,4 @@ jobs: # No `-m "not live"` deselection either: this repo registers no `live` # marker and has no network-dependent tests. - name: Run test suite - run: uv run pytest -q + run: uv run --locked pytest -q diff --git a/README.md b/README.md index 7fca880..4282ddd 100644 --- a/README.md +++ b/README.md @@ -99,11 +99,29 @@ Compaction triggers when token usage reaches the configured threshold (default: ### Protected Messages (Never Removed) - **System messages**: All system messages are always preserved -- **First user message**: The original task/request is always protected (prevents losing context about what was originally asked) -- **Last user message**: The most recent user input is always preserved +- **First human prompt**: The original human task/request is protected, using message metadata rather than treating every user-role message as human input +- **Last human prompt**: The most recent human input is protected by the same metadata-based classification - **Recent messages**: Last N% of messages (configurable via `protected_recent`) +- **Recent tool results**: The last `protected_tool_results` results (default 5) are protected from both truncation and removal; a protected sibling also prevents removal of its owning call group - **Tool pairs**: Tool_use and tool_result messages are treated as atomic units +### Request-scoped retention + +The optional `context.request_retention` capability lets an orchestrator name +the exact persisted reminder bodies required for its next request. The newest +matching admitted `ephemeral=True, persisted=True` user-role envelope is kept +complete through compaction, along with the first/latest human prompts. +Quoting reminder XML in an ordinary human prompt does not change its identity. + +This protects delivery in the request view; it does not change message roles, +pin every historical reminder, or rewrite canonical history. A missing required +body or an irreducible required set that cannot fit raises `ContextLengthError` +instead of silently dropping instructions. Failed assembly restores the prior +compaction state. + +Protection takes precedence over the compaction target. A protected tool cohort +can leave a view above that target; it is not a strict native-token ceiling. + ### Compaction Phases 1. **Phase 1 - Tool Result Truncation**: Older tool results are truncated to reduce token usage @@ -279,7 +297,18 @@ a corresponding quality regression. ## Dependencies -- `amplifier-core>=1.0.0` +- Host-provided `amplifier-core>=1.6.1`, including + `amplifier_core.llm_errors.ContextLengthError` for fail-loud request retention. + Core remains provided by the host, not installed as a runtime module dependency. + +Development and CI pin the released `amplifier-core==1.6.1` package in `uv.lock`. +A Git `main` source mapping can retain an older commit in the lockfile; the +release pin ensures tests exercise the Core API required by this module. + +```bash +uv sync --locked --all-extras --dev +uv run --locked pytest -q +``` ## Contributing diff --git a/amplifier_module_context_simple/__init__.py b/amplifier_module_context_simple/__init__.py index 4435d13..cfc8830 100644 --- a/amplifier_module_context_simple/__init__.py +++ b/amplifier_module_context_simple/__init__.py @@ -46,6 +46,7 @@ from typing import Any from amplifier_core import ModuleCoordinator, TextBlock +from amplifier_core.llm_errors import ContextLengthError logger = logging.getLogger(__name__) @@ -107,6 +108,16 @@ def _carries_loaded_tool_state(msg: dict[str, Any]) -> bool: return any(meta.get(key) for key in LOADED_TOOL_STATE_METADATA_KEYS) +def _is_human_message(msg: dict[str, Any]) -> bool: + """Wire role alone does not distinguish a prompt from an injection.""" + meta = msg.get("metadata") or {} + return ( + msg.get("role") == "user" + and not meta.get("ephemeral") + and meta.get("source") not in ("hook", "context-compaction") + ) + + async def mount(coordinator: ModuleCoordinator, config: dict[str, Any] | None = None): """ Mount the simple context manager. @@ -190,6 +201,12 @@ async def mount(coordinator: ModuleCoordinator, config: dict[str, Any] | None = ) await coordinator.mount("context", context) + # Optional module capability; the core Context protocol is unchanged. + register_capability = getattr(coordinator, "register_capability", None) + if callable(register_capability): + register_capability( + "context.request_retention", context.get_messages_for_request_retaining + ) logger.info(f"Mounted SimpleContextManager (token_meter={token_meter!r})") async def cleanup() -> None: @@ -224,14 +241,14 @@ class SimpleContextManager: Level 5: Remove more messages (60% of configured protection) Level 6: Truncate remaining tool results (except last N) Level 7: Remove more messages (30% of configured protection - last resort) - Level 8: Stub first user message + remove old stubs (extreme pressure) + Level 8: Stub unprotected machine prefix + remove old stubs (extreme pressure) This interleaved approach ensures minimal data loss by: - Preferring truncation (preserves structure) over removal (loses context) - Progressively relaxing protection as pressure increases - Respecting configured protected_recent as baseline, only relaxing under pressure - - Always protecting: system messages, last user message, last N tool results, tool pairs - - First user message: stubbable at Level 8, but never fully removed + - Always protecting: system messages, first/last human prompts, last N tool results, tool pairs + - Requested active persisted injections remain complete through all levels """ def __init__( @@ -328,6 +345,8 @@ def __init__( self._last_measured_prompt_tokens: int | None = None self._last_token_meter_stats: dict[str, Any] | None = None self._system_prompt_factory: Callable[[], Awaitable[str]] | None = None + self._request_retained_contents: frozenset[str] = frozenset() + self._request_protected_seqs: set[int] = set() # --- Sticky compaction decision state --- # Compaction decisions (remove / truncate / stub) are keyed by a @@ -553,6 +572,82 @@ async def set_system_prompt_factory( self._system_prompt_factory = factory logger.info("System prompt factory registered - will refresh on each request") + async def get_messages_for_request_retaining( + self, + *, + retain_contents: list[str], + provider: Any | None = None, + token_budget: int | None = None, + ) -> list[dict[str, Any]]: + """Optional capability: retain current persisted injections for one view. + + Contents must exactly match admitted user-role messages marked + ephemeral=True and persisted=True. Only the newest matching copy is + protected. This neither admits messages nor changes their lifetime; + the caller supplies the current delivery requirements on every call. + """ + previous_contents = self._request_retained_contents + previous_seqs = self._request_protected_seqs + decisions = ( + self._removed_seqs.copy(), + self._truncated_seqs.copy(), + self._stubbed_seqs.copy(), + self._sticky_level, + self._last_compaction_stats, + self._last_token_meter_stats, + ) + try: + self._request_retained_contents = frozenset(retain_contents) + return await self.get_messages_for_request(token_budget, provider) + except BaseException: + # A failed/cancelled assembly must not commit a reduction that was + # never delivered. Canonical history is unchanged throughout. + ( + self._removed_seqs, + self._truncated_seqs, + self._stubbed_seqs, + self._sticky_level, + self._last_compaction_stats, + self._last_token_meter_stats, + ) = decisions + raise + finally: + self._request_retained_contents = previous_contents + self._request_protected_seqs = previous_seqs + + def _protected_sequences(self, messages: list[dict[str, Any]]) -> set[int]: + humans = [msg for msg in messages if _is_human_message(msg)] + protected = [humans[0], humans[-1]] if humans else [] + remaining = set(self._request_retained_contents) + for msg in reversed(messages): + meta = msg.get("metadata") or {} + content = msg.get("content") + if ( + msg.get("role") == "user" + and meta.get("ephemeral") is True + and meta.get("persisted") is True + and isinstance(content, str) + and content in remaining + ): + protected.append(msg) + remaining.remove(content) + if remaining: + raise ValueError("Requested retained injection is not in admitted history") + return {seq for msg in protected if (seq := self._extract_seq(msg)) is not None} + + def _is_request_protected(self, msg: dict[str, Any]) -> bool: + return self._extract_seq(msg) in self._request_protected_seqs + + def _check_retained_budget( + self, messages: list[dict[str, Any]], budget: int + ) -> None: + if self._request_retained_contents and self._estimate_tokens(messages) > budget: + raise ContextLengthError( + "Context cannot fit the current injections and protected conversation " + "within the estimated input budget; shorten the active instructions " + "or use a larger context window. Required content was not discarded." + ) + async def get_messages_for_request( self, token_budget: int | None = None, @@ -629,6 +724,16 @@ async def get_messages_for_request( # Static mode: use messages as-is (may include stored system messages) working_messages = list(self.messages) + self._request_protected_seqs = self._protected_sequences(working_messages) + self._check_retained_budget( + [ + m + for m in working_messages + if m.get("role") == "system" or self._is_request_protected(m) + ], + effective_budget, + ) + token_count, meter_source, estimated_tokens = self._measure_working_tokens( working_messages ) @@ -643,7 +748,11 @@ async def get_messages_for_request( } # Check if compaction needed (using effective budget with notice reserve deducted) - if self._should_compact(token_count, effective_budget): + retained_view_over_budget = ( + bool(self._request_retained_contents) + and estimated_tokens > effective_budget + ) + if self._should_compact(token_count, effective_budget) or retained_view_over_budget: # Compact EPHEMERALLY - returns new list, working_messages unchanged compacted = await self._compact_ephemeral( effective_budget, working_messages @@ -731,8 +840,10 @@ async def get_messages_for_request( # Strip internal bookkeeping at the module boundary -- everything # above this point (sticky decisions, token accounting) still runs # on messages carrying `_seq`; only what leaves has it removed. + self._check_retained_budget(compacted, budget) return self._strip_internal_metadata(compacted) + self._check_retained_budget(working_messages, budget) return self._strip_internal_metadata(working_messages) # Metadata keys that are internal bookkeeping only and must never cross @@ -876,6 +987,8 @@ def _exceeds_threshold(self, estimated_tokens: int, budget: int) -> bool: """ if budget <= 0: return False + if self._request_retained_contents and estimated_tokens > budget: + return True if ( self.token_meter == TOKEN_METER_ACTUAL and self._last_measured_prompt_tokens is not None @@ -1037,6 +1150,9 @@ def _apply_sticky_decisions( result: list[dict[str, Any]] = [] for msg in messages: seq = self._extract_seq(msg) + if self._is_request_protected(msg): + result.append(dict(msg)) + continue if seq is not None and seq in self._removed_seqs: continue if seq is not None and seq in self._truncated_seqs: @@ -1079,6 +1195,7 @@ async def _compact_ephemeral( messages_to_compact = ( source_messages if source_messages is not None else self.messages ) + self._request_protected_seqs = self._protected_sequences(messages_to_compact) target_tokens = int(budget * self.target_usage) old_count = len(messages_to_compact) old_tokens = self._estimate_tokens(messages_to_compact) @@ -1420,10 +1537,10 @@ async def _compact_ephemeral( # Check if we still need more space if current_tokens > target_tokens: - # === LEVEL 8: Stub first user message + remove old stubs (extreme pressure) === + # === LEVEL 8: Stub unprotected machine prefix + remove old stubs (extreme pressure) === max_level_reached = 8 - # Find first user message and stub it if not already stubbed + # Find first and last user messages. first_user_idx = None last_user_idx = None for i, msg in enumerate(working_messages): @@ -1432,9 +1549,14 @@ async def _compact_ephemeral( first_user_idx = i last_user_idx = i - # Stub first user message (previously protected) - but NEVER if it's also the last - # The last user message is the current intent and must always be preserved - if first_user_idx is not None and first_user_idx != last_user_idx: + # An unprotected machine prefix can still be reduced at extreme pressure. + # Human boundary prompts and current retained injections stay complete. + if ( + first_user_idx is not None + and first_user_idx != last_user_idx + and not self._is_request_protected(working_messages[first_user_idx]) + and not _is_human_message(working_messages[first_user_idx]) + ): first_msg = working_messages[first_user_idx] if not first_msg.get("_stubbed"): content = first_msg.get("content", "") @@ -1462,6 +1584,7 @@ async def _compact_ephemeral( if msg.get("_stubbed") and i < protected_boundary # Outside protected recent zone and i != last_user_idx # Never remove last user message + and not self._is_request_protected(msg) ] stubs_removed = 0 @@ -1625,11 +1748,11 @@ def _remove_messages_with_protection( i for i, msg in enumerate(messages) if msg.get("role") == "user" } - # Find first and last user message indices (always fully protected from stubbing too) + # Human boundaries must not be displaced by machine user-role messages. first_user_idx = None last_user_idx = None for i, msg in enumerate(messages): - if msg.get("role") == "user": + if _is_human_message(msg): if first_user_idx is None: first_user_idx = i last_user_idx = i @@ -1654,11 +1777,11 @@ def _remove_messages_with_protection( len(loaded_tool_state_indices), ) - # First user message is stubbable at extreme pressure (Level 8), but never fully removed - # (It's excluded from removal_candidates via user_message_indices, but can be stubbed) - # We don't add it to protected_indices so it can be stubbed at Level 8 - - # Always protect the LAST user message (current context) + protected_indices.update( + i for i, msg in enumerate(messages) if self._is_request_protected(msg) + ) + if first_user_idx is not None: + protected_indices.add(first_user_idx) if last_user_idx is not None: protected_indices.add(last_user_idx) @@ -1667,6 +1790,13 @@ def _remove_messages_with_protection( for i in range(protected_boundary, len(messages)): protected_indices.add(i) + # The last N tool results are protected from removal as well as truncation. + # Their owning assistant and sibling results are vetoed atomically below. + tool_result_indices = [ + i for i, msg in enumerate(messages) if msg.get("role") == "tool" + ] + protected_indices |= self._protected_tool_indices(tool_result_indices) + # Removal candidates exclude ALL user messages (they can only be stubbed, not removed) removal_candidates = [ i diff --git a/pyproject.toml b/pyproject.toml index 50d7020..8a7dd4f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -30,11 +30,8 @@ packages = [ [tool.hatch.metadata] allow-direct-references = true -[tool.uv.sources] -amplifier-core = { git = "https://github.com/microsoft/amplifier-core", branch = "main" } - [dependency-groups] -dev = ["amplifier-core", "pytest>=9.0.3", "pytest-asyncio>=0.23.0"] +dev = ["amplifier-core==1.6.1", "pytest>=9.0.3", "pytest-asyncio>=0.23.0"] [tool.pytest.ini_options] testpaths = ["tests"] diff --git a/tests/test_progressive_compaction.py b/tests/test_progressive_compaction.py index ee433f5..3d3186f 100644 --- a/tests/test_progressive_compaction.py +++ b/tests/test_progressive_compaction.py @@ -15,18 +15,15 @@ @pytest.mark.asyncio async def test_tool_result_truncation_phase1(): """Phase 1 truncates old tool results while preserving structure.""" - # Configure for easy testing: low tokens, aggressive truncation - # Use lower max_tokens to ensure compaction triggers with our message sizes + # Four results make the first 25% truncation wave nonempty. Its large + # result is unprotected; the three recent results must remain intact. context = SimpleContextManager( - max_tokens=300, # Lower to ensure compaction triggers - compact_threshold=0.5, - target_usage=0.3, + max_tokens=2000, + compact_threshold=0.9, + target_usage=0.5, truncate_chars=50, - protected_recent=0.3, # Protect recent messages but allow truncation of old tool - protected_tool_results=2, # Only protect last 2 tool results - # Without this, the default 800-token notice reserve makes - # effective_budget = 300 - 800 = -700, which silently disables - # compaction entirely (see test_budget_guard.py). + protected_recent=0, + protected_tool_results=3, compaction_notice_enabled=False, ) @@ -46,19 +43,22 @@ async def test_tool_result_truncation_phase1(): } ) # Add a large tool result - this will push us over the threshold - large_content = "x" * 500 # 500 chars = ~125 tokens + large_content = "x" * 10000 await context.add_message( {"role": "tool", "tool_call_id": "toolu_early", "content": large_content} ) - # Add more messages to push tool pair into truncate zone (first 50%) - for i in range(10): + for i in range(3): + call_id = f"recent-call-{i}" await context.add_message( - {"role": "user", "content": f"message {i} with extra padding"} + {"role": "assistant", "content": "", "tool_calls": [ + {"id": call_id, "type": "function", "function": {"name": "read_file"}} + ]} ) await context.add_message( - {"role": "assistant", "content": f"response {i} with extra padding"} + {"role": "tool", "tool_call_id": call_id, "content": "small result"} ) + await context.add_message({"role": "user", "content": "continue"}) # Trigger compaction messages = await context.get_messages_for_request() @@ -70,14 +70,18 @@ async def test_tool_result_truncation_phase1(): tool_result = msg break - # Verify truncation occurred - if tool_result: - content = tool_result.get("content", "") - assert "[truncated:" in content, "Tool result should be truncated" - assert tool_result.get("_truncated") is True, "Should have _truncated marker" - assert len(content) < 200, ( - f"Truncated content should be small, got {len(content)}" - ) + # Removal is not a substitute for the truncation this test promises. + assert tool_result is not None + content = tool_result.get("content", "") + assert "[truncated:" in content, "Tool result should be truncated" + assert tool_result.get("_truncated") is True + assert len(content) < 200 + assert context._last_compaction_stats["strategy_level"] == 1 + assert context._last_compaction_stats["messages_removed"] == 0 + assert [ + m["content"] for m in messages + if m.get("tool_call_id", "").startswith("recent-call-") + ] == ["small result"] * 3 @pytest.mark.asyncio diff --git a/tests/test_protected_tool_results.py b/tests/test_protected_tool_results.py index 2312b4b..0bb9eda 100644 --- a/tests/test_protected_tool_results.py +++ b/tests/test_protected_tool_results.py @@ -164,8 +164,8 @@ async def test_protecting_every_tool_result_still_works(): result = await _run_workload(_make_context(protected_tool_results=N_TOOL_PAIRS)) assert result["truncated_tool_ids"] == [], result - assert result["level"] == 3, result - assert result["messages_removed"] > 0, result + assert result["level"] >= 3, result + assert result["messages_removed"] == 0, result # -------------------------------------------------------------------------- @@ -201,3 +201,55 @@ def test_protected_tool_indices_uses_real_positions_not_ordinals(): """The returned indices are positions in the message list, not 0..N-1.""" ctx = SimpleContextManager(protected_tool_results=2) assert ctx._protected_tool_indices([3, 11, 40, 57]) == {40, 57} + + +@pytest.mark.parametrize("protected", [0, 1, 5]) +def test_one_protected_sibling_vetoes_whole_batch_removal(protected): + """Removal honors the same last-N floor as truncation, including N=0.""" + from copy import deepcopy + + ctx = SimpleContextManager(protected_tool_results=protected) + messages = [ + {"role": "user", "content": "first human"}, + {"role": "assistant", "content": "", "thinking": {"opaque": "unchanged"}, + "tool_calls": [ + {"id": "old", "type": "function", "function": {"name": "read_file"}}, + {"id": "recent", "type": "function", "function": {"name": "read_file"}}, + ]}, + {"role": "tool", "tool_call_id": "old", "content": "older sibling"}, + {"role": "tool", "tool_call_id": "recent", "content": "protected sibling"}, + {"role": "user", "content": "last human"}, + ] + before = deepcopy(messages) + view, removed, stubbed, _ = ctx._remove_messages_with_protection( + messages, target_tokens=1, protected_recent=0, system_tokens=0, + ) + assert messages == before + assert stubbed == 0 + assert removed == (0 if protected else 3) + assert view == (before if protected else [before[0], before[-1]]) + + +@pytest.mark.asyncio +async def test_default_tool_floor_survives_progressive_removal(): + from copy import deepcopy + + ctx = SimpleContextManager(max_tokens=120_000) + for index in range(4): + await ctx.add_message({"role": "user", "content": f"human-{index}:" + "x" * 100_000}) + await ctx.add_message({ + "role": "assistant", "content": "", "tool_calls": [ + {"id": f"call-{index}", "type": "function", "function": {"name": "read_file"}} + ], + }) + await ctx.add_message({ + "role": "tool", "tool_call_id": f"call-{index}", "content": f"result-{index}", + }) + await ctx.add_message({"role": "user", "content": "latest:" + "y" * 100_000}) + canonical = deepcopy(ctx.messages) + view = await ctx.get_messages_for_request() + assert ctx._last_compaction_stats["strategy_level"] >= 4 + assert [m["content"] for m in view if m.get("role") == "tool"] == [ + f"result-{index}" for index in range(4) + ] + assert ctx.messages == canonical diff --git a/tests/test_request_retention.py b/tests/test_request_retention.py new file mode 100644 index 0000000..8a41413 --- /dev/null +++ b/tests/test_request_retention.py @@ -0,0 +1,330 @@ +"""Required request content must survive actual compaction, not just storage.""" + +import copy + +import pytest +from amplifier_core.llm_errors import ContextLengthError + +from amplifier_module_context_simple import SimpleContextManager, mount + + +def reminder(content): + return { + "role": "user", + "content": content, + "metadata": {"ephemeral": True, "persisted": True}, + } + + +def human(content): + return { + "role": "user", + "content": content, + "metadata": {"source": "human"}, + } + + +def public_envelope(message): + """The request view strips only the manager's internal sequence number.""" + return { + **message, + "metadata": { + key: value + for key, value in message["metadata"].items() + if key != "_seq" + }, + } + + +async def pressured_context(): + context = SimpleContextManager(max_tokens=1800, compaction_notice_enabled=False) + await context.add_message({"role": "user", "content": "Original task " + "a" * 150}) + body = "" + "policy " * 180 + " REQUIRED_FACT" + await context.add_message(reminder(body)) + for i in range(40): + await context.add_message( + {"role": "assistant", "content": f"Output {i} " + "x" * 1200} + ) + await context.add_message( + {"role": "user", "content": "Current correction " + "b" * 150} + ) + return context, body + + +@pytest.mark.asyncio +async def test_retention_restores_a_sticky_stub_without_changing_history(): + context, body = await pressured_context() + canonical = copy.deepcopy(await context.get_messages()) + ordinary = await context.get_messages_for_request() + assert not any(m["content"] == body for m in ordinary) + + retained = await context.get_messages_for_request_retaining(retain_contents=[body]) + assert sum(m["content"] == body for m in retained) == 1 + assert context._estimate_tokens(retained) <= context.max_tokens + assert await context.get_messages() == canonical + assert ( + await context.get_messages_for_request_retaining(retain_contents=[body]) + == retained + ) + + # Retention is per request, not a permanent pin that outlives its producer. + released = await context.get_messages_for_request_retaining(retain_contents=[]) + assert not any(m["content"] == body for m in released) + + +@pytest.mark.asyncio +async def test_machine_messages_cannot_take_human_boundary_protection(): + context = SimpleContextManager(max_tokens=1400, compaction_notice_enabled=False) + first = "Original human request " + "details " * 60 + "ORIGINAL_FACT" + last = "Latest human correction " + "details " * 60 + "CURRENT_FACT" + await context.add_message(reminder("Earlier machine state " + "m" * 1200)) + await context.add_message({"role": "user", "content": first}) + for i in range(35): + await context.add_message( + {"role": "assistant", "content": f"Work {i} " + "x" * 1000} + ) + await context.add_message({"role": "user", "content": last}) + await context.add_message(reminder("Later machine state " + "m" * 1200)) + view = await context.get_messages_for_request() + assert any(m["content"] == first for m in view) + assert any(m["content"] == last for m in view) + assert await context.get_messages_for_request() == view + + +@pytest.mark.asyncio +async def test_only_newest_matching_copy_is_retained(): + context, body = await pressured_context() + await context.add_message(reminder(body)) + view = await context.get_messages_for_request_retaining(retain_contents=[body]) + assert sum(m["content"] == body for m in view) == 1 + assert sum(m["content"] == body for m in await context.get_messages()) == 2 + + +@pytest.mark.asyncio +async def test_retention_survives_transcript_round_trip(): + context, body = await pressured_context() + transcript = await context.get_messages() + resumed = SimpleContextManager(max_tokens=1800, compaction_notice_enabled=False) + await resumed.set_messages(transcript) + view = await resumed.get_messages_for_request_retaining(retain_contents=[body]) + assert any(m["content"] == body for m in view) + + +@pytest.mark.asyncio +async def test_missing_content_fails_without_leaking_a_retention_requirement(): + context, body = await pressured_context() + with pytest.raises(ValueError, match="not in admitted history"): + await context.get_messages_for_request_retaining( + retain_contents=["not admitted"] + ) + assert not any( + m["content"] == body for m in await context.get_messages_for_request() + ) + + +@pytest.mark.asyncio +async def test_impossible_retention_fails_before_returning_an_overfull_view(): + context, body = await pressured_context() + huge = "" + "x" * 12000 + "" + await context.add_message(reminder(huge)) + with pytest.raises(ContextLengthError, match="Required content was not discarded"): + await context.get_messages_for_request_retaining(retain_contents=[huge]) + assert any(m["content"] == huge for m in await context.get_messages()) + assert not context._stubbed_seqs and not context._removed_seqs + # Let the withdrawn content age out of the ordinary recent-message window. + for i in range(30): + await context.add_message( + {"role": "assistant", "content": f"Later work {i} " + "x" * 200} + ) + # A later request with a feasible requirement recovers normally. + view = await context.get_messages_for_request_retaining(retain_contents=[body]) + assert any(m["content"] == body for m in view) + assert context._estimate_tokens(view) <= context.max_tokens + + +@pytest.mark.asyncio +async def test_post_compaction_failure_rolls_back_request_state_and_decisions(): + context = SimpleContextManager(max_tokens=1000, compaction_notice_enabled=False) + required = "Active retention requirement." + await context.add_message(human("First human request.")) + await context.add_message(reminder(required)) + await context.add_message( + { + "role": "assistant", + "content": "Unremovable loaded tool state: " + "x" * 5000, + "metadata": {"openai:tool_search_items": ["read_file"]}, + } + ) + await context.add_message(human("Latest human correction.")) + + with pytest.raises(ContextLengthError, match="Required content was not discarded"): + await context.get_messages_for_request_retaining(retain_contents=[required]) + + assert context._last_compaction_stats is None + assert context._last_token_meter_stats is None + assert not context._removed_seqs + assert not context._truncated_seqs + assert not context._stubbed_seqs + assert context._sticky_level == 0 + + +@pytest.mark.asyncio +async def test_mount_advertises_additive_capability_and_keeps_old_getter(): + class Coordinator: + def __init__(self): + self.capabilities = {} + + async def mount(self, name, module): + self.context = module + + def register_capability(self, name, capability): + self.capabilities[name] = capability + + coordinator = Coordinator() + cleanup = await mount(coordinator) + assert callable(coordinator.capabilities["context.request_retention"]) + assert await coordinator.context.get_messages_for_request() == [] + assert ( + await coordinator.capabilities["context.request_retention"](retain_contents=[]) + == [] + ) + await cleanup() + + +@pytest.mark.asyncio +async def test_stale_actual_meter_cannot_prevent_feasible_retention_compaction(): + context, body = await pressured_context() + context.token_meter = "actual" + context._last_measured_prompt_tokens = 100 + view = await context.get_messages_for_request_retaining(retain_contents=[body]) + assert any(m["content"] == body for m in view) + assert context._estimate_tokens(view) <= context.max_tokens + + +@pytest.mark.asyncio +async def test_changed_then_withdrawn_requirements_do_not_leave_a_stale_pin(): + context = SimpleContextManager(max_tokens=1800, compaction_notice_enabled=False) + original = "" + "original policy " * 180 + "" + replacement = ( + "" + "replacement policy " * 180 + "" + ) + await context.add_message(human("First human request: preserve current requirements.")) + await context.add_message(reminder(original)) + for i in range(40): + await context.add_message( + {"role": "assistant", "content": f"Earlier work {i}: " + "x" * 1200} + ) + await context.add_message(human("Latest human correction: use the replacement.")) + await context.add_message(reminder(replacement)) + for i in range(30): + await context.add_message( + {"role": "assistant", "content": f"Later work {i}: " + "z" * 500} + ) + + changed_view = await context.get_messages_for_request_retaining( + retain_contents=[replacement] + ) + assert sum(m["content"] == original for m in changed_view) == 0 + assert sum(m["content"] == replacement for m in changed_view) == 1 + + # A new, empty requirement set must permit the formerly selected envelope + # to age out on a real subsequent compaction. + for i in range(30): + await context.add_message( + {"role": "assistant", "content": f"Even later work {i}: " + "z" * 500} + ) + withdrawn_view = await context.get_messages_for_request_retaining( + retain_contents=[] + ) + assert not any(m["content"] == original for m in withdrawn_view) + assert not any(m["content"] == replacement for m in withdrawn_view) + + +@pytest.mark.asyncio +async def test_level_5_retains_complete_current_reminder_and_human_boundaries(): + """A real Level 5 pass keeps the selected envelope and both human bodies.""" + context = SimpleContextManager( + max_tokens=1800, + compact_threshold=0.92, + target_usage=0.50, + protected_recent=0.30, + compaction_notice_enabled=False, + ) + first = human("First human request: preserve the complete acceptance criteria.") + active_reminder = reminder( + "" + "Keep the current deployment region unchanged. " + "Do not use an unpublished credential. " + "Return the complete migration plan." + "" + ) + latest = human("Latest human correction: the migration must remain reversible.") + await context.add_message(first) + await context.add_message(active_reminder) + for i in range(14): + await context.add_message( + {"role": "assistant", "content": f"Machine output {i}: " + "x" * 800} + ) + await context.add_message(latest) + + canonical = copy.deepcopy(await context.get_messages()) + first_envelope, reminder_envelope, *_, latest_envelope = map( + public_envelope, canonical + ) + view = await context.get_messages_for_request_retaining( + retain_contents=[active_reminder["content"]] + ) + + assert context._last_compaction_stats is not None + assert context._last_compaction_stats["strategy_level"] == 5 + assert first_envelope in view + assert reminder_envelope in view + assert latest_envelope in view + assert await context.get_messages() == canonical + + +@pytest.mark.asyncio +async def test_level_8_keeps_quoted_xml_human_prompt_and_reminder_complete(): + """A quoted system-reminder is human content, not injected provenance.""" + context = SimpleContextManager( + max_tokens=1800, + compact_threshold=0.92, + target_usage=0.50, + protected_recent=0.30, + compaction_notice_enabled=False, + ) + first = human( + "" + "A human quoted this XML while asking whether it was safe to publish." + "" + ) + active_reminder = reminder( + "" + + "Active retention policy. " * 130 + + "" + ) + latest = human("Latest human correction: do not remove the quoted evidence.") + await context.add_message(first) + await context.add_message(active_reminder) + for i in range(8): + await context.add_message( + {"role": "assistant", "content": f"Machine output {i}: " + "x" * 800} + ) + await context.add_message(latest) + + canonical = await context.get_messages() + first_envelope, reminder_envelope, *_, latest_envelope = map( + public_envelope, canonical + ) + view = await context.get_messages_for_request_retaining( + retain_contents=[active_reminder["content"]] + ) + + assert context._last_compaction_stats is not None + assert context._last_compaction_stats["strategy_level"] == 8 + assert first_envelope in view + assert reminder_envelope in view + assert latest_envelope in view + assert context._last_compaction_stats["after_tokens"] > context._last_compaction_stats[ + "target_tokens" + ] \ No newline at end of file diff --git a/uv.lock b/uv.lock index 0f2e81c..3223a56 100644 --- a/uv.lock +++ b/uv.lock @@ -4,8 +4,8 @@ requires-python = ">=3.11" [[package]] name = "amplifier-core" -version = "1.0.0" -source = { git = "https://github.com/microsoft/amplifier-core?branch=main#976fb87335ffe398cf0c1bd1bcd3c2f2c154fc3c" } +version = "1.6.1" +source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "click" }, { name = "pydantic" }, @@ -13,6 +13,14 @@ dependencies = [ { name = "tomli" }, { name = "typing-extensions" }, ] +wheels = [ + { url = "https://files.pythonhosted.org/packages/03/cd/8b0b520bf0de741ea73e069aaf64aca28c9f4ce91a7b8b9239193a6c4c1b/amplifier_core-1.6.1-cp311-abi3-macosx_10_12_x86_64.whl", hash = "sha256:c0f711d8408de78e53e5deddcb38b7240c5c1c497ca51eeaaeff23559b3d3c48", size = 8281633, upload-time = "2026-08-10T02:38:11.98Z" }, + { url = "https://files.pythonhosted.org/packages/14/83/f4fb297d87d35b9d74058da02bb153e12f7891ab62b3aaf7e0857f877798/amplifier_core-1.6.1-cp311-abi3-macosx_11_0_arm64.whl", hash = "sha256:b08f37e2c0b1611349a0e25d5bf9bfdfae3afcee35488f8e26bba1cdd400503b", size = 7366930, upload-time = "2026-08-10T02:38:14.105Z" }, + { url = "https://files.pythonhosted.org/packages/ff/ba/5eb9cecf92d8053c5e6d46ad9668c3ed3558d5423845c1dced1f266b2a38/amplifier_core-1.6.1-cp311-abi3-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:6ebf7e3993c76ea506e70ac7844b286c3ba2e9127b3bcb350fa4fcd2dcdbd38d", size = 7659512, upload-time = "2026-08-10T02:38:16.314Z" }, + { url = "https://files.pythonhosted.org/packages/22/31/121f054e3d079dc33d83f3d8ba9af50fd9f7694c3e2ba3d7d23d7c157d48/amplifier_core-1.6.1-cp311-abi3-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:3c957cd0671d2a003f2c8f7d6a41bd6e808f97d183c57b97e7700bf4c912621d", size = 8678425, upload-time = "2026-08-10T02:38:18.243Z" }, + { url = "https://files.pythonhosted.org/packages/35/25/bfc217f4a9ed2d033995fc59847f1fee2e1b17130632fcb0e0981a1a311b/amplifier_core-1.6.1-cp311-abi3-win_amd64.whl", hash = "sha256:50c80bcfa1f6efe769b19e7af18c925024c7553d4db08880727241709dd44eae", size = 8976601, upload-time = "2026-08-10T02:38:20.505Z" }, + { url = "https://files.pythonhosted.org/packages/a5/14/5f330452c92c6c5d35c51ad5311301949ce5db4d1a1a901456f3ee43eaac/amplifier_core-1.6.1-cp311-abi3-win_arm64.whl", hash = "sha256:cd8b617f132cf5d1ca3e5187d5f831d1f2a508bb40d07b2ab1085961bcb9e1a9", size = 7744837, upload-time = "2026-08-10T02:38:22.562Z" }, +] [[package]] name = "amplifier-module-context-simple" @@ -30,7 +38,7 @@ dev = [ [package.metadata.requires-dev] dev = [ - { name = "amplifier-core", git = "https://github.com/microsoft/amplifier-core?branch=main" }, + { name = "amplifier-core", specifier = "==1.6.1" }, { name = "pytest", specifier = ">=9.0.3" }, { name = "pytest-asyncio", specifier = ">=0.23.0" }, ]