diff --git a/README.md b/README.md index f54d974..35e488a 100644 --- a/README.md +++ b/README.md @@ -109,6 +109,10 @@ Compaction triggers when token usage reaches the configured threshold (default: - **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 +- **Unfinished signed tool turns**: provider thinking remains intact while any + sibling tool result is pending. Context reuses the existing sticky view + without a new compaction decision, or fails closed when that complete view + cannot fit the request budget. ### Request-scoped retention diff --git a/amplifier_module_context_simple/__init__.py b/amplifier_module_context_simple/__init__.py index d23f941..6e5494e 100644 --- a/amplifier_module_context_simple/__init__.py +++ b/amplifier_module_context_simple/__init__.py @@ -128,11 +128,83 @@ def _carries_loaded_tool_state(msg: dict[str, Any]) -> bool: return any(meta.get(key) for key in LOADED_TOOL_STATE_METADATA_KEYS) +_STRUCTURED_TOOL_CALL_TYPES = frozenset({"tool_call", "tool_use"}) +_STRUCTURED_TOOL_RESULT_TYPES = frozenset({"tool_result"}) +_SIGNED_THINKING_KEYS = frozenset( + {"thinking", "redacted_thinking", "thinking_block"} +) + + +def _content_blocks(msg: dict[str, Any]) -> list[dict[str, Any]]: + """Return recognized structured content blocks without trusting other shapes.""" + content = msg.get("content") + if not isinstance(content, list): + return [] + return [block for block in content if isinstance(block, dict)] + + +def _tool_call_ids(msg: dict[str, Any]) -> set[str]: + """Return IDs declared by legacy or structured tool-call representations.""" + ids: set[str] = set() + tool_calls = msg.get("tool_calls") + if isinstance(tool_calls, list): + for tool_call in tool_calls: + if not isinstance(tool_call, dict): + continue + tool_call_id = tool_call.get("id") or tool_call.get("tool_call_id") + if isinstance(tool_call_id, str) and tool_call_id: + ids.add(tool_call_id) + for block in _content_blocks(msg): + if block.get("type") not in _STRUCTURED_TOOL_CALL_TYPES: + continue + tool_call_id = ( + block.get("id") or block.get("tool_call_id") or block.get("tool_use_id") + ) + if isinstance(tool_call_id, str) and tool_call_id: + ids.add(tool_call_id) + return ids + + +def _has_tool_calls(msg: dict[str, Any]) -> bool: + """Whether a message declares legacy or structured tool calls.""" + return bool(msg.get("tool_calls")) or any( + block.get("type") in _STRUCTURED_TOOL_CALL_TYPES + for block in _content_blocks(msg) + ) + + +def _tool_result_ids(msg: dict[str, Any]) -> set[str]: + """Return IDs referenced by legacy or structured tool-result representations.""" + ids: set[str] = set() + for key in ("tool_call_id", "tool_use_id"): + tool_call_id = msg.get(key) + if isinstance(tool_call_id, str) and tool_call_id: + ids.add(tool_call_id) + for block in _content_blocks(msg): + if block.get("type") not in _STRUCTURED_TOOL_RESULT_TYPES: + continue + tool_call_id = ( + block.get("tool_use_id") or block.get("tool_call_id") or block.get("id") + ) + if isinstance(tool_call_id, str) and tool_call_id: + ids.add(tool_call_id) + return ids + + +def _is_tool_result_message(msg: dict[str, Any]) -> bool: + """Whether a message is a legacy or structured tool result.""" + return msg.get("role") == "tool" or any( + block.get("type") in _STRUCTURED_TOOL_RESULT_TYPES + for block in _content_blocks(msg) + ) + + 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 _is_tool_result_message(msg) and not meta.get("ephemeral") and meta.get("source") not in ("hook", "context-compaction") ) @@ -422,6 +494,10 @@ def __init__( # "estimate" mode -- see README "Real-usage token meter". self._last_measured_prompt_tokens: int | None = None self._last_token_meter_stats: dict[str, Any] | None = None + # Per-request count of signed-thinking blocks omitted because compaction + # rewrote an earlier prefix. This is intentionally observable rather + # than silently changing a provider-facing replay. + self._last_signed_thinking_blocks_stripped = 0 # A Loop that understands foreground ownership claims this once, then # writes only its successful foreground response usage through the # recorder. Generic llm:response remains a compatibility fallback @@ -1187,6 +1263,39 @@ async def get_measured_request_view( initial_view, initial_attempt, initial_measurement = await self._count_measured_view( count_view, initial ) + tail_tool_call = self._tail_unanswered_tool_call_message(initial) + if tail_tool_call is not None and self._has_provider_thinking(tail_tool_call): + if self._measured_hard_oversize(initial_attempt): + raise ContextLengthError( + "Context cannot fit protected content within the provider " + "input limit." + ) + if initial_measurement is None: + return { + "base_view": initial_view, + "final_attempt": initial_attempt, + "outcome": "measurement_unavailable", + "measured_before": None, + "measured_after": None, + "policy_budget": effective_budget, + "trigger": trigger, + "target": target, + "count_calls": 1, + "transaction": None, + } + before, _estimated, _limit, _measurement_source = initial_measurement + return { + "base_view": initial_view, + "final_attempt": initial_attempt, + "outcome": "protected_floor", + "measured_before": before, + "measured_after": before, + "policy_budget": effective_budget, + "trigger": trigger, + "target": target, + "count_calls": 1, + "transaction": None, + } if initial_measurement is None: if self._measured_hard_oversize(initial_attempt): raise ContextLengthError( @@ -1510,8 +1619,10 @@ def _check_retained_budget( ) @staticmethod - def _tail_has_unanswered_tool_calls(messages: list[dict[str, Any]]) -> bool: - """Return whether the trailing tool-call group is incomplete. + def _tail_unanswered_tool_call_message( + messages: list[dict[str, Any]], + ) -> dict[str, Any] | None: + """Return the assistant message owning an incomplete trailing call group. A tail can end on the assistant tool-call message itself or on one of several sibling tool results. A notice is safe only after every call @@ -1519,19 +1630,29 @@ def _tail_has_unanswered_tool_calls(messages: list[dict[str, Any]]) -> bool: """ trailing_result_ids: set[str] = set() for message in reversed(messages): - if message.get("role") == "tool": - tool_call_id = message.get("tool_call_id") - if tool_call_id: - trailing_result_ids.add(tool_call_id) + if _is_tool_result_message(message): + trailing_result_ids.update(_tool_result_ids(message)) continue - if message.get("role") != "assistant" or not message.get("tool_calls"): - return False - declared_ids = { - tool_call.get("id") or tool_call.get("tool_call_id") - for tool_call in message["tool_calls"] - } - return bool(declared_ids - trailing_result_ids) - return False + if message.get("role") != "assistant" or not _has_tool_calls(message): + return None + declared_ids = _tool_call_ids(message) + if not declared_ids: + return message + return message if declared_ids - trailing_result_ids else None + return None + + @staticmethod + def _has_provider_thinking(message: dict[str, Any]) -> bool: + """Whether an assistant message carries provider-owned thinking state.""" + return any(key in message for key in _SIGNED_THINKING_KEYS) or any( + block.get("type") in _SIGNED_THINKING_KEYS + for block in _content_blocks(message) + ) + + @classmethod + def _tail_has_unanswered_tool_calls(cls, messages: list[dict[str, Any]]) -> bool: + """Return whether the trailing tool-call group is incomplete.""" + return cls._tail_unanswered_tool_call_message(messages) is not None async def get_messages_for_request( self, @@ -1653,6 +1774,23 @@ async def get_messages_for_request( "ratio": (token_count / effective_budget) if effective_budget > 0 else None, } + # Anthropic requires provider thinking to remain present while its tool + # turn is unfinished. Do not rewrite its prefix and then strip that + # thinking: preserve the existing sticky view, or fail before sending + # an invalid replay. The raw provider budget deliberately excludes the + # optional compaction-notice reserve because no notice is added here. + tail_tool_call = self._tail_unanswered_tool_call_message(sticky_view) + if tail_tool_call is not None and self._has_provider_thinking(tail_tool_call): + if estimated_tokens > 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." + ) + self._record_signed_thinking_strip_count(0) + await self._emit_budget() + return self._strip_internal_metadata(sticky_view) + # Check whether this request needs a new compaction escalation. A # provider-directed hard fit applies at the full forced budget, not at # the ordinary compact_threshold or target_usage fraction of it. @@ -1692,6 +1830,7 @@ async def get_messages_for_request( else: self._check_retained_budget(working_messages, budget) await self._emit_budget() + self._record_signed_thinking_strip_count(0) return self._strip_internal_metadata(working_messages) # Append compaction notice at the TAIL if enabled and level threshold met. @@ -1742,6 +1881,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. + compacted, signed_thinking_blocks_stripped = ( + self._strip_invalidated_signed_thinking(compacted, working_messages) + ) + self._record_signed_thinking_strip_count(signed_thinking_blocks_stripped) self._check_retained_budget(compacted, budget) if deferred_hard_fit_stats is not None and self._hooks is not None: try: @@ -1755,6 +1898,82 @@ async def get_messages_for_request( await self._emit_budget() return self._strip_internal_metadata(compacted) + def _record_signed_thinking_strip_count(self, count: int) -> None: + """Expose signed-thinking stripping on the per-request meter surface.""" + self._last_signed_thinking_blocks_stripped = count + if self._last_token_meter_stats is not None: + self._last_token_meter_stats["signed_thinking_blocks_stripped"] = count + if self._last_compaction_stats is not None: + self._last_compaction_stats["signed_thinking_blocks_stripped"] = count + + def _strip_invalidated_signed_thinking( + self, + view: list[dict[str, Any]], + source_messages: list[dict[str, Any]], + ) -> tuple[list[dict[str, Any]], int]: + """Drop signed thinking only after compaction rewrote its prior prefix. + + Canonical history remains untouched. A replay with an unchanged prefix + remains byte-for-byte identical; after a removal, truncation, or stub + changes an earlier message, a later Anthropic signature is no longer + valid and must not be sent back to the provider. + """ + view_by_seq = { + seq: msg + for msg in view + if (seq := self._extract_seq(msg)) is not None + } + prefix_rewritten = False + rewritten_before: dict[int, bool] = {} + for source in source_messages: + seq = self._extract_seq(source) + if seq is None: + continue + rendered = view_by_seq.get(seq) + if rendered is None or rendered != source: + prefix_rewritten = True + continue + rewritten_before[seq] = prefix_rewritten + + stripped = 0 + safe_view: list[dict[str, Any]] = [] + for msg in view: + seq = self._extract_seq(msg) + if seq is None or not rewritten_before.get(seq, False): + safe_view.append(msg) + continue + safe_message, removed_count = self._without_signed_thinking(msg) + safe_view.append(safe_message) + stripped += removed_count + return safe_view, stripped + + @staticmethod + def _without_signed_thinking(msg: dict[str, Any]) -> tuple[dict[str, Any], int]: + """Return a copied message without provider-signed thinking blocks.""" + removed_count = sum(1 for key in _SIGNED_THINKING_KEYS if key in msg) + safe_message = { + key: value for key, value in msg.items() if key not in _SIGNED_THINKING_KEYS + } + content = msg.get("content") + if isinstance(content, list): + safe_content = [ + block + for block in content + if not ( + isinstance(block, dict) and block.get("type") in _SIGNED_THINKING_KEYS + ) + ] + removed_count += len(content) - len(safe_content) + if len(safe_content) != len(content): + safe_message["content"] = safe_content + elif ( + isinstance(content, dict) + and content.get("type") in _SIGNED_THINKING_KEYS + ): + safe_message["content"] = [] + removed_count += 1 + return safe_message if removed_count else msg, removed_count + # Metadata keys that are internal bookkeeping only and must never cross # the module boundary into a provider-facing view. `_seq` is sticky # compaction identity (see _extract_seq): meaningless to a provider, and -- @@ -1851,6 +2070,7 @@ async def clear(self) -> None: self._last_compaction_stats = None self._last_measured_prompt_tokens = None self._last_token_meter_stats = None + self._last_signed_thinking_blocks_stripped = 0 # Ownership is a capability claim, not a turn-local meter value. A # cleared Context must not permit utility hook traffic to take it back. self._foreground_usage_stale = self._foreground_usage_claimed @@ -2689,7 +2909,9 @@ def _remove_messages_with_protection( # Track user messages for stubbing (NEVER removal) user_message_indices = { - i for i, msg in enumerate(messages) if msg.get("role") == "user" + i + for i, msg in enumerate(messages) + if msg.get("role") == "user" and not _is_tool_result_message(msg) } # Human boundaries must not be displaced by machine user-role messages. @@ -2739,7 +2961,7 @@ def _remove_messages_with_protection( # 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" + i for i, msg in enumerate(messages) if _is_tool_result_message(msg) ] protected_indices |= self._protected_tool_indices(tool_result_indices) @@ -2768,8 +2990,7 @@ def _remove_messages_with_protection( token_lens = [self._estimate_tokens([msg]) for msg in messages] tool_call_id_to_indices: dict[str, list[int]] = {} for idx, m in enumerate(messages): - tcid = m.get("tool_call_id") - if tcid: + for tcid in _tool_result_ids(m): tool_call_id_to_indices.setdefault(tcid, []).append(idx) # Remove messages until under target, preserving tool pairs @@ -2789,7 +3010,7 @@ def _remove_messages_with_protection( newly_removed: list[int] = [] # Handle tool result - must remove with its tool_use pair - if msg.get("role") == "tool": + if _is_tool_result_message(msg): pair_removed, newly_removed = self._try_remove_tool_pair_from_result( messages, i, protected_indices, tool_call_id_to_indices ) @@ -2797,7 +3018,7 @@ def _remove_messages_with_protection( continue # Can't remove this one, skip # Handle assistant with tool_calls - must remove with all its tool results - elif msg.get("role") == "assistant" and msg.get("tool_calls"): + elif msg.get("role") == "assistant" and _has_tool_calls(msg): pair_removed, tool_result_indices = ( self._try_remove_tool_pair_from_assistant( msg, protected_indices, tool_call_id_to_indices @@ -2888,7 +3109,7 @@ def _try_remove_tool_pair_from_result( # Find the assistant with tool_calls for j in range(result_idx - 1, -1, -1): check_msg = messages[j] - if check_msg.get("role") == "assistant" and check_msg.get("tool_calls"): + if check_msg.get("role") == "assistant" and _has_tool_calls(check_msg): if j in protected_indices: return False, [] # Can't remove protected assistant @@ -2900,7 +3121,7 @@ def _try_remove_tool_pair_from_result( if all_removable: return True, [j, *tool_result_indices] return False, [] - if check_msg.get("role") != "tool": + if not _is_tool_result_message(check_msg): break return False, [] @@ -2940,20 +3161,21 @@ def _check_tool_pair_removable( all_removable = True tool_result_indices = [] - for tc in assistant_msg.get("tool_calls", []): - tc_id = tc.get("id") or tc.get("tool_call_id") - if tc_id: - result_indices = tool_call_id_to_indices.get(tc_id, []) - # An unanswered call must remain until its result arrives. - # Removing it now would make a subsequently admitted result - # orphaned in the provider-facing view. - if not result_indices: + declared_ids = _tool_call_ids(assistant_msg) + if not declared_ids: + return False, [] + for tc_id in declared_ids: + result_indices = tool_call_id_to_indices.get(tc_id, []) + # An unanswered call must remain until its result arrives. + # Removing it now would make a subsequently admitted result + # orphaned in the provider-facing view. + if not result_indices: + all_removable = False + for k in result_indices: + if k in protected_indices: all_removable = False - for k in result_indices: - if k in protected_indices: - all_removable = False - else: - tool_result_indices.append(k) + else: + tool_result_indices.append(k) return all_removable, tool_result_indices @@ -2981,8 +3203,10 @@ async def _finalize_compaction_with_stats( final_tokens = self._estimate_tokens(final_messages) system_count = len(system_messages) - tool_use_count = sum(1 for m in final_messages if m.get("tool_calls")) - tool_result_count = sum(1 for m in final_messages if m.get("role") == "tool") + tool_use_count = sum(1 for m in final_messages if _has_tool_calls(m)) + tool_result_count = sum( + 1 for m in final_messages if _is_tool_result_message(m) + ) logger.info( f"Compaction complete: {old_count} → {len(final_messages)} messages, " diff --git a/tests/test_measured_request_view.py b/tests/test_measured_request_view.py index f688568..4e2ad20 100644 --- a/tests/test_measured_request_view.py +++ b/tests/test_measured_request_view.py @@ -59,6 +59,29 @@ async def count_view(view): return context, result["transaction"] +async def _add_partial_signed_tool_turn(context): + """Add two calls but return a structured result for only the first.""" + await context.add_message({"role": "user", "content": "first"}) + await context.add_message({"role": "assistant", "content": "old removable prefix"}) + await context.add_message( + { + "role": "assistant", + "content": [ + {"type": "thinking", "thinking": "private", "signature": "sig"}, + {"type": "tool_use", "id": "call-1", "name": "read", "input": {}}, + {"type": "tool_use", "id": "call-2", "name": "list", "input": {}}, + ], + } + ) + await context.add_message( + { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": "call-1", "content": "one"}], + } + ) + context._record_removed(context.messages[1]) + + @pytest.mark.asyncio async def test_measured_capability_is_actual_mode_only_and_foreground_is_always_additive(): estimate = _Coordinator() @@ -141,6 +164,89 @@ async def count_view(view): assert context._last_compaction_stats is None +@pytest.mark.asyncio +async def test_unfinished_signed_tool_turn_is_a_measured_protected_floor(): + """A provider-count fit must not reduce an incomplete signed tool turn.""" + context = SimpleContextManager( + max_tokens=1_000, + compact_threshold=0.1, + compaction_notice_enabled=False, + token_meter="actual", + ) + await _add_partial_signed_tool_turn(context) + calls = 0 + + async def count_view(view): + nonlocal calls + calls += 1 + return {"dispatch": object(), "budget_decision": _decision(900)} + + result = await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=count_view + ) + + assert result["outcome"] == "protected_floor" + assert result["count_calls"] == calls == 1 + assert result["transaction"] is None + assert not context._truncated_seqs + replayed = next(message for message in result["base_view"] if message.get("role") == "assistant") + assert replayed["content"][0]["type"] == "thinking" + + +@pytest.mark.asyncio +async def test_unfinished_signed_tool_turn_measured_hard_overlimit_rolls_back(): + """A known hard overflow fails before staging a reduction.""" + context = SimpleContextManager( + max_tokens=1_000, + compact_threshold=0.1, + compaction_notice_enabled=False, + token_meter="actual", + ) + await _add_partial_signed_tool_turn(context) + + async def count_view(_view): + return { + "dispatch": object(), + "budget_decision": _decision(900, estimated=1_001, limit=1_000), + } + + with pytest.raises(ContextLengthError, match="cannot fit protected content"): + await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=count_view + ) + + assert not context._truncated_seqs + assert context._last_compaction_stats is None + + +@pytest.mark.asyncio +async def test_unfinished_signed_tool_turn_keeps_unavailable_measurement_unreduced(): + """An unavailable count returns the existing outcome without legacy compaction.""" + context = SimpleContextManager( + max_tokens=1_000, + compact_threshold=0.1, + compaction_notice_enabled=False, + token_meter="actual", + ) + await _add_partial_signed_tool_turn(context) + calls = 0 + + async def count_view(_view): + nonlocal calls + calls += 1 + return {"dispatch": object(), "budget_decision": None} + + result = await context.get_measured_request_view( + provider=None, retain_contents=[], count_view=count_view + ) + + assert result["outcome"] == "measurement_unavailable" + assert result["count_calls"] == calls == 1 + assert result["transaction"] is None + assert not context._truncated_seqs + assert context._last_compaction_stats is None + + @pytest.mark.asyncio async def test_unavailable_recount_rolls_back_to_original_counted_view(): context = SimpleContextManager( diff --git a/tests/test_signed_thinking_compaction.py b/tests/test_signed_thinking_compaction.py new file mode 100644 index 0000000..8e4a34e --- /dev/null +++ b/tests/test_signed_thinking_compaction.py @@ -0,0 +1,295 @@ +"""Regression coverage for signed-thinking replay after ephemeral compaction.""" + +from copy import deepcopy + +import pytest + +from amplifier_module_context_simple import SimpleContextManager +from amplifier_core.llm_errors import ContextLengthError + + +async def _add_partial_signed_tool_turn(context): + """Add two structured calls but return a result for only the first.""" + await context.add_message({"role": "user", "content": "Original task"}) + await context.add_message({"role": "assistant", "content": "old removable prefix"}) + await context.add_message( + { + "role": "assistant", + "content": [ + {"type": "thinking", "thinking": "private", "signature": "sig"}, + {"type": "tool_use", "id": "use-1", "name": "read_file", "input": {}}, + {"type": "tool_use", "id": "use-2", "name": "list_files", "input": {}}, + ], + "thinking": {"text": "private", "signature": "sig"}, + } + ) + await context.add_message( + { + "role": "user", + "content": [ + {"type": "tool_result", "tool_use_id": "use-1", "content": "read"}, + ], + } + ) + + +@pytest.mark.asyncio +async def test_partial_signed_tool_turn_preserves_sticky_view_above_threshold(): + """An unfinished signed turn cannot be stripped after its prefix changes.""" + context = SimpleContextManager( + max_tokens=1_000, + compact_threshold=0.01, + compaction_notice_enabled=False, + ) + await _add_partial_signed_tool_turn(context) + canonical = deepcopy(context.messages) + context._record_removed(context.messages[1]) + + view = await context.get_messages_for_request() + + assert context.messages == canonical + assert all(message.get("content") != "old removable prefix" for message in view) + replayed = next(message for message in view if message.get("role") == "assistant") + assert replayed["thinking"] == {"text": "private", "signature": "sig"} + assert replayed["content"][0]["type"] == "thinking" + assert context._last_signed_thinking_blocks_stripped == 0 + assert context._last_token_meter_stats["signed_thinking_blocks_stripped"] == 0 + assert ( + context._last_token_meter_stats["estimated_tokens"] + > context._last_token_meter_stats["budget"] * context.compact_threshold + ) + assert not context._truncated_seqs + + +@pytest.mark.asyncio +async def test_partial_signed_tool_turn_fails_closed_when_raw_budget_cannot_fit(): + """Retention rolls back rather than rewriting an unfinished signed turn.""" + context = SimpleContextManager( + max_tokens=1_000, + compact_threshold=0.01, + compaction_notice_enabled=False, + ) + await context.add_message( + { + "role": "user", + "content": "retained request context", + "metadata": {"ephemeral": True, "persisted": True}, + } + ) + await _add_partial_signed_tool_turn(context) + context._record_removed(context.messages[2]) + decisions = context._removed_seqs.copy() + protected_budget = context._estimate_tokens(context.messages[:2]) + + with pytest.raises(ContextLengthError, match="cannot fit the current injections"): + await context.get_messages_for_request_retaining( + retain_contents=["retained request context"], + token_budget=protected_budget, + hard_fit=True, + ) + + assert context._removed_seqs == decisions + assert not context._truncated_seqs + assert context.messages[3]["thinking"] == {"text": "private", "signature": "sig"} + + +@pytest.mark.asyncio +async def test_completed_signed_tool_turn_resumes_invalidated_thinking_stripping(): + """Once all siblings finish, an invalidated prefix strips signed thinking.""" + context = SimpleContextManager( + max_tokens=100_000, + compact_threshold=1.0, + compaction_notice_enabled=False, + ) + await _add_partial_signed_tool_turn(context) + await context.add_message( + { + "role": "user", + "content": [ + {"type": "tool_result", "tool_use_id": "use-2", "content": "listed"}, + ], + } + ) + context._record_removed(context.messages[1]) + + view = await context.get_messages_for_request() + + replayed = next(message for message in view if message.get("role") == "assistant") + assert "thinking" not in replayed + assert replayed["content"] == [ + {"type": "tool_use", "id": "use-1", "name": "read_file", "input": {}}, + {"type": "tool_use", "id": "use-2", "name": "list_files", "input": {}}, + ] + assert context._last_signed_thinking_blocks_stripped == 2 + + +@pytest.mark.asyncio +async def test_unchanged_replay_preserves_signed_thinking_exactly(): + """Without a rewritten prefix, signed provider history leaves unchanged.""" + context = SimpleContextManager( + max_tokens=100_000, + compact_threshold=1.0, + compaction_notice_enabled=False, + ) + await context.add_message({"role": "user", "content": "Keep this history intact."}) + await context.add_message( + { + "role": "assistant", + "content": [{"type": "thinking", "thinking": "private", "signature": "sig"}], + "thinking": {"text": "private", "signature": "sig"}, + "redacted_thinking": {"data": "opaque"}, + "thinking_block": {"legacy": True}, + } + ) + + first_view = await context.get_messages_for_request() + second_view = await context.get_messages_for_request() + + assert second_view == first_view + replayed = second_view[-1] + assert replayed["thinking"] == {"text": "private", "signature": "sig"} + assert replayed["redacted_thinking"] == {"data": "opaque"} + assert replayed["thinking_block"] == {"legacy": True} + assert replayed["content"] == [ + {"type": "thinking", "thinking": "private", "signature": "sig"} + ] + assert context._last_token_meter_stats["signed_thinking_blocks_stripped"] == 0 + + +@pytest.mark.asyncio +async def test_rewritten_prefix_strips_signed_thinking_without_mutating_history(): + """A removed earlier turn invalidates later signatures only in the wire view.""" + context = SimpleContextManager( + max_tokens=100_000, + compact_threshold=1.0, + compaction_notice_enabled=False, + ) + await context.add_message({"role": "user", "content": "Original task"}) + await context.add_message({"role": "assistant", "content": "old removable prefix"}) + await context.add_message( + { + "role": "assistant", + "content": [ + {"type": "thinking", "thinking": "private", "signature": "sig"}, + {"type": "redacted_thinking", "data": "opaque"}, + {"type": "text", "text": "Visible response"}, + ], + "thinking": {"text": "private", "signature": "sig"}, + "redacted_thinking": {"data": "opaque"}, + "thinking_block": {"legacy": True}, + } + ) + canonical = deepcopy(context.messages) + + # This is the sticky state recorded by the removal phase of compaction. + context._record_removed(context.messages[1]) + + view = await context.get_messages_for_request() + + assert context.messages == canonical + assert all(message.get("content") != "old removable prefix" for message in view) + replayed = view[-1] + assert "thinking" not in replayed + assert "redacted_thinking" not in replayed + assert "thinking_block" not in replayed + assert replayed["content"] == [{"type": "text", "text": "Visible response"}] + assert context._last_signed_thinking_blocks_stripped == 5 + assert context._last_token_meter_stats["signed_thinking_blocks_stripped"] == 5 + + +@pytest.mark.asyncio +async def test_structured_tool_turn_is_retained_atomically_after_prefix_rewrite(): + """Structured tool_use/tool_call blocks pair with structured results atomically.""" + context = SimpleContextManager( + protected_tool_results=1, + max_tokens=100_000, + compact_threshold=1.0, + compaction_notice_enabled=False, + ) + await context.add_message({"role": "user", "content": "Original task"}) + await context.add_message({"role": "assistant", "content": "old removable prefix"}) + await context.add_message( + { + "role": "assistant", + "content": [ + {"type": "thinking", "thinking": "private", "signature": "sig"}, + {"type": "tool_use", "id": "use-1", "name": "read_file", "input": {}}, + {"type": "tool_call", "id": "call-1", "name": "list_files", "input": {}}, + ], + "thinking": {"text": "private", "signature": "sig"}, + } + ) + await context.add_message( + { + "role": "user", + "content": [ + {"type": "tool_result", "tool_use_id": "use-1", "content": "read"}, + {"type": "tool_result", "tool_call_id": "call-1", "content": "listed"}, + ], + } + ) + await context.add_message({"role": "user", "content": "Continue"}) + canonical = deepcopy(context.messages) + + compacted, removed, stubbed, _ = context._remove_messages_with_protection( + context.messages, + target_tokens=1, + protected_recent=0, + system_tokens=0, + ) + + assert context.messages == canonical + assert removed == 1 + assert stubbed == 0 + assert any(message.get("content") == "Continue" for message in compacted) + retained_assistant = next( + message + for message in compacted + if message.get("role") == "assistant" + and any( + isinstance(block, dict) and block.get("type") == "tool_use" + for block in message.get("content", []) + ) + ) + retained_result = next( + message + for message in compacted + if any( + isinstance(block, dict) and block.get("type") == "tool_result" + for block in message.get("content", []) + ) + ) + assert retained_assistant["thinking"] == {"text": "private", "signature": "sig"} + assert {block["tool_use_id"] for block in retained_result["content"] if "tool_use_id" in block} == { + "use-1" + } + assert { + block["tool_call_id"] for block in retained_result["content"] if "tool_call_id" in block + } == {"call-1"} + + view = await context.get_messages_for_request() + + replayed_assistant = next( + message + for message in view + if message.get("role") == "assistant" + and any( + isinstance(block, dict) and block.get("type") == "tool_use" + for block in message.get("content", []) + ) + ) + assert "thinking" not in replayed_assistant + assert [block["type"] for block in replayed_assistant["content"]] == [ + "tool_use", + "tool_call", + ] + assert any( + any( + isinstance(block, dict) + and block.get("type") == "tool_result" + and block.get("tool_use_id") == "use-1" + for block in message.get("content", []) + ) + for message in view + ) + assert context.messages == canonical \ No newline at end of file